From 0780a806f2cc2e8a6a51782cfffe0591b7c3fa9c Mon Sep 17 00:00:00 2001 From: Misha Bragin Date: Fri, 31 Jul 2026 20:52:56 +0200 Subject: [PATCH] [management, proxy] Management-owned LLM pricing: file-backed defaults + (#6965) --- combined/cmd/config.go | 10 + combined/cmd/root.go | 52 ++ combined/config.yaml.example | 13 + e2e/agentnetwork/chat_test.go | 67 +- e2e/agentnetwork/custom_pricing_test.go | 633 ++++++++++++++++++ e2e/harness/agentnetwork.go | 21 + e2e/harness/combined.go | 11 +- e2e/harness/config.go | 39 ++ management/cmd/management.go | 24 + .../modules/agentnetwork/catalog/catalog.go | 176 +++-- .../handlers/providers_handler.go | 71 +- .../handlers/providers_handler_test.go | 53 ++ .../modules/agentnetwork/pricing/defaults.go | 156 +++++ .../pricing/defaults_llm_pricing.example.yaml | 494 +++++++------- .../agentnetwork/pricing/defaults_test.go | 148 ++++ .../agentnetwork/pricing/exampleyaml.go | 109 +++ .../modules/agentnetwork/pricing/gen.go | 20 + .../modules/agentnetwork/pricing/override.go | 249 +++++++ .../agentnetwork/pricing/override_test.go | 173 +++++ .../modules/agentnetwork/synthesizer.go | 14 +- .../agentnetwork/synthesizer_pricing.go | 131 ++++ .../agentnetwork/synthesizer_pricing_test.go | 105 +++ .../modules/agentnetwork/synthesizer_test.go | 41 +- .../modules/agentnetwork/types/provider.go | 50 +- .../modules/agentnetwork/wire_shape_test.go | 6 + management/internals/server/config/config.go | 21 + proxy/internal/llm/bedrock_model.go | 38 -- proxy/internal/llm/fixtures/pricing.yaml | 59 -- proxy/internal/llm/model.go | 21 + .../llm/pricing/defaults_coverage_test.go | 65 -- proxy/internal/llm/pricing/pricing.go | 487 ++++---------- proxy/internal/llm/pricing/pricing_other.go | 20 - proxy/internal/llm/pricing/pricing_test.go | 385 ++--------- proxy/internal/llm/pricing/pricing_unix.go | 68 -- proxy/internal/middleware/builtin/builtin.go | 7 +- .../builtin/cost_calculation_matrix_test.go | 14 +- .../middleware/builtin/cost_meter/factory.go | 82 ++- .../builtin/cost_meter/middleware.go | 68 +- .../builtin/cost_meter/middleware_test.go | 304 +++++---- .../llm_request_parser/bedrock_test.go | 17 - .../builtin/llm_request_parser/middleware.go | 39 +- .../agent_network_chain_realstack_test.go | 40 +- proxy/server.go | 6 +- shared/llm/model.go | 58 ++ .../llm/model_test.go | 13 + shared/management/http/api/openapi.yml | 37 + shared/management/http/api/types.gen.go | 21 + 47 files changed, 3196 insertions(+), 1540 deletions(-) create mode 100644 e2e/agentnetwork/custom_pricing_test.go create mode 100644 management/internals/modules/agentnetwork/handlers/providers_handler_test.go create mode 100644 management/internals/modules/agentnetwork/pricing/defaults.go rename proxy/internal/llm/pricing/defaults_pricing.yaml => management/internals/modules/agentnetwork/pricing/defaults_llm_pricing.example.yaml (63%) create mode 100644 management/internals/modules/agentnetwork/pricing/defaults_test.go create mode 100644 management/internals/modules/agentnetwork/pricing/exampleyaml.go create mode 100644 management/internals/modules/agentnetwork/pricing/gen.go create mode 100644 management/internals/modules/agentnetwork/pricing/override.go create mode 100644 management/internals/modules/agentnetwork/pricing/override_test.go create mode 100644 management/internals/modules/agentnetwork/synthesizer_pricing.go create mode 100644 management/internals/modules/agentnetwork/synthesizer_pricing_test.go delete mode 100644 proxy/internal/llm/bedrock_model.go delete mode 100644 proxy/internal/llm/fixtures/pricing.yaml create mode 100644 proxy/internal/llm/model.go delete mode 100644 proxy/internal/llm/pricing/defaults_coverage_test.go delete mode 100644 proxy/internal/llm/pricing/pricing_other.go delete mode 100644 proxy/internal/llm/pricing/pricing_unix.go create mode 100644 shared/llm/model.go rename proxy/internal/llm/bedrock_model_test.go => shared/llm/model_test.go (66%) diff --git a/combined/cmd/config.go b/combined/cmd/config.go index 7f30cd8a8..890c86876 100644 --- a/combined/cmd/config.go +++ b/combined/cmd/config.go @@ -76,6 +76,13 @@ type ServerConfig struct { SupportedSyncMessageVersions *int `yaml:"supportedSyncMessageVersions,omitempty"` PerAccountSupportedSyncMessageVersions map[string]int `yaml:"perAccountSupportedSyncMessageVersions,omitempty"` + + AgentNetwork AgentNetworkConfig `yaml:"agentNetwork"` +} + +// AgentNetworkConfig contains agent-network (LLM gateway) configuration. +type AgentNetworkConfig struct { + PricingDefaultsFile string `yaml:"pricingDefaultsFile"` } // TLSConfig contains TLS/HTTPS settings @@ -723,6 +730,9 @@ func (c *CombinedConfig) ToManagementConfig() (*nbconfig.Config, error) { EmbeddedIdP: embeddedIdP, HighestSupportedSyncMessageVersion: c.Server.SupportedSyncMessageVersions, PerAccountHighestSupportedSyncMessageVersion: c.Server.PerAccountSupportedSyncMessageVersions, + AgentNetwork: nbconfig.AgentNetwork{ + PricingDefaultsFile: c.Server.AgentNetwork.PricingDefaultsFile, + }, }, nil } diff --git a/combined/cmd/root.go b/combined/cmd/root.go index 5f2564e3a..7eac84ce5 100644 --- a/combined/cmd/root.go +++ b/combined/cmd/root.go @@ -10,6 +10,7 @@ import ( "net/http" "os" "os/signal" + "path/filepath" "strconv" "strings" "sync" @@ -24,6 +25,7 @@ import ( "google.golang.org/grpc" "github.com/netbirdio/netbird/encryption" + agentnetworkpricing "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" mgmtServer "github.com/netbirdio/netbird/management/internals/server" nbconfig "github.com/netbirdio/netbird/management/internals/server/config" "github.com/netbirdio/netbird/management/server/telemetry" @@ -288,6 +290,11 @@ func (s *serverInstances) createManagementServer(ctx context.Context, cfg *Combi return fmt.Errorf("failed to ensure encryption key: %w", err) } + if err := loadAgentNetworkPricing(ctx, mgmtConfig); err != nil { + cleanupSTUNListeners(s.stunListeners) + return fmt.Errorf("failed to load agent-network pricing defaults: %w", err) + } + LogConfigInfo(mgmtConfig) s.mgmtSrv, err = createManagementServer(cfg, mgmtConfig) @@ -622,6 +629,32 @@ func handleRelayWebSocket(w http.ResponseWriter, r *http.Request, acceptFn func( acceptFn(conn) } +// loadAgentNetworkPricing loads the management-side LLM pricing defaults +// file for the combined server and starts its periodic reloader. An +// explicitly configured PricingDefaultsFile is required to load (a typo +// must fail startup rather than silently bill with built-ins the operator +// believes they replaced); a relative path is resolved against the data +// directory so a bare filename like "pricing.yaml" lands in the datadir +// alongside the store. With no path configured, / +// is probed and may be absent (compiled-in defaults serve). +func loadAgentNetworkPricing(ctx context.Context, mgmtConfig *nbconfig.Config) error { + pricingPath := mgmtConfig.AgentNetwork.PricingDefaultsFile + required := pricingPath != "" + if !required { + pricingPath = agentnetworkpricing.DefaultFileName + } + if !filepath.IsAbs(pricingPath) { + pricingPath = filepath.Join(mgmtConfig.Datadir, pricingPath) + } + + log.Infof("loading agent-network pricing defaults from %s (required: %v)", pricingPath, required) + if err := agentnetworkpricing.LoadFile(pricingPath, required); err != nil { + return err + } + agentnetworkpricing.StartReloader(ctx, agentnetworkpricing.ReloadInterval) + return nil +} + // logConfig prints all configuration parameters for debugging func logConfig(cfg *CombinedConfig) { log.Info("=== Configuration ===") @@ -698,6 +731,25 @@ func logManagementConfig(cfg *CombinedConfig) { log.Infof(" Relay addresses: %v", cfg.Management.Relays.Addresses) log.Infof(" Relay credentials TTL: %s", cfg.Management.Relays.CredentialsTTL) } + + logAgentNetworkConfig(cfg) +} + +func logAgentNetworkConfig(cfg *CombinedConfig) { + log.Info(" Agent Network:") + pricingPath := cfg.Server.AgentNetwork.PricingDefaultsFile + configured := pricingPath != "" + if !configured { + pricingPath = agentnetworkpricing.DefaultFileName + } + if !filepath.IsAbs(pricingPath) { + pricingPath = filepath.Join(cfg.Management.DataDir, pricingPath) + } + if configured { + log.Infof(" Pricing defaults file: %s", pricingPath) + } else { + log.Infof(" Pricing defaults file: %s (default, optional)", pricingPath) + } } // logEnvVars logs all NB_ environment variables that are currently set diff --git a/combined/config.yaml.example b/combined/config.yaml.example index 66bc71703..085e4344f 100644 --- a/combined/config.yaml.example +++ b/combined/config.yaml.example @@ -134,3 +134,16 @@ server: # trustedPeers: [] # CIDRs of trusted peer networks (e.g. ["100.64.0.0/10"]) # accessLogRetentionDays: 7 # Days to retain HTTP access logs. 0 (or unset) defaults to 7. Negative values disable cleanup (logs kept indefinitely). # accessLogCleanupIntervalHours: 24 # How often (in hours) to run the access-log cleanup job. 0 (or unset) is treated as "not set" and defaults to 24 hours; cleanup remains enabled. To disable cleanup, set accessLogRetentionDays to a negative value. + + # Agent network (LLM gateway) settings (optional) + # agentNetwork: + # # Path to the YAML file holding the default LLM pricing table. A relative + # # path is resolved against dataDir, so a bare filename like "pricing.yaml" + # # lands in the data directory. When empty, {dataDir}/defaults_llm_pricing.yaml + # # is probed; if no file is present the compiled-in defaults are used. + # # Schema: surface ("openai"/"anthropic"/"bedrock") -> model -> rates in USD + # # per 1k tokens (input_per_1k, output_per_1k, and the optional + # # cached_input_per_1k / cache_read_per_1k / cache_creation_per_1k). The file + # # is re-read periodically (mtime poll). An explicitly configured path that + # # fails to load fails startup; runtime reload errors keep the previous table. + # pricingDefaultsFile: "pricing.yaml" diff --git a/e2e/agentnetwork/chat_test.go b/e2e/agentnetwork/chat_test.go index 65c4d813f..ed94623fe 100644 --- a/e2e/agentnetwork/chat_test.go +++ b/e2e/agentnetwork/chat_test.go @@ -23,8 +23,13 @@ import ( type per1k struct{ in, out, read, write float64 } // publishedPer1k hardcodes the vendors' PUBLISHED rates for the models the live matrix can drive, -// keyed by the normalized model id the proxy stamps. Deliberately independent of the proxy's -// pricing table so a wrong embedded rate or a broken normalization fails the run. +// keyed by the normalized model id the proxy stamps. Deliberately independent of NetBird's own +// default pricing table so a wrong default rate or a broken normalization fails the run. +// +// These rates are also what providerRequest registers as the operator's per-model prices. Since +// management now ships operator prices to the cost meter as a per-provider-record table that is +// consulted BEFORE the surface defaults, registering the published rate is what keeps this matrix +// asserting vendor rates — and exercises the per-record path at the same time. var publishedPer1k = map[string]per1k{ "gpt-4o-mini": {0.00015, 0.0006, 0.000075, 0}, "gpt-4o": {0.0025, 0.01, 0.00125, 0}, @@ -35,12 +40,22 @@ var publishedPer1k = map[string]per1k{ "anthropic.claude-haiku-4-5": {0.001, 0.005, 0.0001, 0.00125}, "anthropic.claude-sonnet-4-5": {0.003, 0.015, 0.0003, 0.00375}, "anthropic.claude-sonnet-4-6": {0.003, 0.015, 0.0003, 0.00375}, + // Gateway-prefixed ids (Vercel AI Gateway, OpenRouter). A gateway model is not in + // NetBird's default table, so before operator pricing it could only be recorded at + // cost 0. The operator names it and prices it — at the underlying vendor's published + // rate, which is what the gateway charges through — so these rows are now priced. + "openai/gpt-4o-mini": {0.00015, 0.0006, 0.000075, 0}, + "openai/gpt-4o": {0.0025, 0.01, 0.00125, 0}, } // rawCostVerificationSQL is the operator-facing double-check, run straight against the management // sqlite store: recompute each usage row's expected total and cache cost from its own persisted // token buckets and hardcoded published rates. OpenAI counts cached tokens as a subset of input; // Anthropic-shape providers count cache buckets additively. +// +// The rate rows must stay in sync with publishedPer1k — they are the same vendor rates the matrix +// registers as operator prices. The join is on model, so rows written by other tests in this +// package (which price their own made-up model ids) are simply not covered here. const rawCostVerificationSQL = ` WITH rates(model, in_rate, out_rate, read_rate, write_rate) AS ( VALUES @@ -52,7 +67,9 @@ WITH rates(model, in_rate, out_rate, read_rate, write_rate) AS ( ('kimi-k3', 0.003, 0.015, 0.0003, 0.003), ('anthropic.claude-haiku-4-5', 0.001, 0.005, 0.0001, 0.00125), ('anthropic.claude-sonnet-4-5', 0.003, 0.015, 0.0003, 0.00375), - ('anthropic.claude-sonnet-4-6', 0.003, 0.015, 0.0003, 0.00375) + ('anthropic.claude-sonnet-4-6', 0.003, 0.015, 0.0003, 0.00375), + ('openai/gpt-4o-mini', 0.00015, 0.0006, 0.000075, 0.0), + ('openai/gpt-4o', 0.0025, 0.01, 0.00125, 0.0) ) SELECT u.provider, @@ -146,6 +163,11 @@ func verifyUsageRowsSQL(t *testing.T, srv *harness.Combined) { require.Positive(t, verified, "raw SQL check must cover at least one usage row") t.Logf("[sql] verified %d usage rows in store.db against published rates", verified) + // Gateway-prefixed model ids are absent from NetBird's default pricing table, so they are + // priced only because the operator registered and priced them on the provider record. Assert + // they are priced (not silently 0) — the join above already checked the exact figures for the + // ones this matrix drives. A gateway row at cost 0 means the per-record table never reached + // the cost meter, which is the regression this guards. gwRows, err := db.Raw(`SELECT model, (input_cost_usd + cached_input_cost_usd + cache_creation_cost_usd + output_cost_usd) AS cost_usd FROM agent_network_request_usage WHERE model LIKE '%/%'`).Rows() @@ -155,8 +177,8 @@ func verifyUsageRowsSQL(t *testing.T, srv *harness.Combined) { var model string var cost float64 require.NoError(t, gwRows.Scan(&model, &cost), "scan gateway usage row") - t.Logf("[sql] gateway %s: stored=$%.6f (must be 0 — deliberately unpriced)", model, cost) - assert.Zerof(t, cost, "gateway-prefixed model %q must store cost 0, never a guessed rate", model) + t.Logf("[sql] gateway %s: stored=$%.6f (priced from the operator's per-record rate)", model, cost) + assert.Positivef(t, cost, "gateway-prefixed model %q is priced on the provider record, so its cost must be > 0", model) } require.NoError(t, gwRows.Err(), "iterate gateway usage rows") } @@ -177,10 +199,6 @@ func validateAccessLogCost(t *testing.T, pc providerCase, row api.AgentNetworkAc rates, known := publishedPer1k[model] if !known { - if strings.Contains(model, "/") { - assert.Zerof(t, row.CostUsd, "gateway-prefixed model %q is not priced so the cost meter must skip (cost 0)", model) - return - } t.Logf("[cost] %s: no published rate on file for model %q (env-overridden?); skipping cost validation", pc.name, model) return } @@ -337,8 +355,17 @@ func availableProviders() []providerCase { } // providerRequest builds a create request for a matrix provider: enabled, with -// a uniquely-priced model for body-routed providers and none for the -// path-routed Vertex (whose model lives in the request path). +// its model registered at the vendor's published rates for body-routed +// providers, and no models for the path-routed Vertex (whose model lives in the +// request path, so it prices from the defaults table management ships). +// +// The registered rates matter: management synthesizes them into the cost +// meter's per-provider-record table, which is consulted before the surface +// defaults, so these are the rates the proxy actually bills with. Registering +// the published rate keeps the cost assertions vendor-anchored while covering +// the operator-pricing path. A model with no published rate on file (an +// env-overridden Bedrock profile) falls back to a nominal rate, and +// validateAccessLogCost skips its cost check. func providerRequest(pc providerCase) api.AgentNetworkProviderRequest { req := api.AgentNetworkProviderRequest{ Name: pc.name, @@ -356,9 +383,23 @@ func providerRequest(pc providerCase) api.AgentNetworkProviderRequest { if pc.kind == harness.WireBedrock { modelID = catalogModel(pc) } - req.Models = &[]api.AgentNetworkProviderModel{ - {Id: modelID, InputPer1k: 0.001, OutputPer1k: 0.002}, + model := api.AgentNetworkProviderModel{Id: modelID, InputPer1k: 0.001, OutputPer1k: 0.002} + if rates, known := publishedPer1k[catalogModel(pc)]; known { + model.InputPer1k = rates.in + model.OutputPer1k = rates.out + // Pin the cache rates too, rather than letting them inherit from the + // defaults table: a gateway-prefixed id has no default entry to + // inherit from, and an unset rate bills that bucket at the input + // rate, which would not match the published-rate recompute. + if rates.read > 0 { + model.CachedInputPer1k = ptr(rates.read) // OpenAI shape + model.CacheReadPer1k = ptr(rates.read) // Anthropic / Bedrock shape + } + if rates.write > 0 { + model.CacheCreationPer1k = ptr(rates.write) + } } + req.Models = &[]api.AgentNetworkProviderModel{model} } return req } diff --git a/e2e/agentnetwork/custom_pricing_test.go b/e2e/agentnetwork/custom_pricing_test.go new file mode 100644 index 000000000..33788a16a --- /dev/null +++ b/e2e/agentnetwork/custom_pricing_test.go @@ -0,0 +1,633 @@ +//go:build e2e + +package agentnetwork + +import ( + "context" + "fmt" + "net/url" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/driver/sqlite" + "gorm.io/gorm" + + "github.com/netbirdio/netbird/e2e/harness" + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// The mock vLLM upstream (harness/vllm.go) always answers with this fixed usage +// block, so every request drives deterministic token counts regardless of the +// model the client asks for. The proxy prices off the REQUEST model, not the +// upstream response model, so a made-up model id billed at operator rates lets +// these tests assert exact costs without a real vendor key. +const ( + vllmPromptTokens = 11 + vllmCompletionTokens = 2 +) + +// pricedEnv is a connected single-provider agent-network deployment pointed at +// the mock vLLM upstream, with the proxy and client up and the endpoint resolved +// — ready to drive chat. All containers are torn down via t.Cleanup. +type pricedEnv struct { + providerID string + groupID string // source group of the policy; the client peer's auto-group + policyID string // policy that authorises (and meters) the requests + upstream string // provider upstream URL, needed to re-send on a PUT update + endpoint string + proxyIP string + client *harness.Client + proxy *harness.Proxy +} + +// provisionPricedProvider brings up the full path for a cost test: a mock vLLM +// upstream, a group + reusable setup key, one openai_api provider pointed at the +// mock enumerating exactly the given models (with the operator's per-1k prices), +// a policy whose token limit switches on usage metering, and a connected proxy + +// client. The provider is created with the given models so the router dispatches +// them to this provider and the cost meter bills at these rates. +// +// Passing nil models makes it a gateway-style catch-all: the router claims every +// model, and since the synthesizer ships no per-provider-record pricing entry +// for a provider that enumerates nothing, the shipped defaults table is the only +// thing that can price the request. The policy sets no model guardrail, so the +// proxy's per-provider allowlist backstop stays empty and any model routes. +func provisionPricedProvider(t *testing.T, ctx context.Context, name string, models []api.AgentNetworkProviderModel) pricedEnv { + t.Helper() + + vllm, err := harness.StartVLLM(ctx, srv) + require.NoError(t, err, "start mock vLLM upstream") + t.Cleanup(func() { _ = vllm.Terminate(context.Background()) }) + + grp, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-price-" + name}) + require.NoError(t, err, "create group") + t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grp.Id) }) + + ephemeral := false + sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{ + Name: "e2e-price-" + name + "-client", + Type: "reusable", + ExpiresIn: 86400, + UsageLimit: 0, + AutoGroups: []string{grp.Id}, + Ephemeral: &ephemeral, + }) + require.NoError(t, err, "mint setup key") + require.NotEmpty(t, sk.Key, "setup key plaintext") + + // The mock ignores auth, so a dummy key satisfies the "Bearer ${API_KEY}" + // template. openai_api is a known catalog provider; the enumerated model id + // need NOT be in the catalog — the operator names it and prices it here. + dummyKey := "sk-price-e2e" + prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{ + Name: name, + ProviderId: "openai_api", + UpstreamUrl: vllm.URL, + ApiKey: &dummyKey, + Enabled: ptr(true), + BootstrapCluster: ptr(harness.AgentNetworkCluster), + Models: &models, + }) + require.NoError(t, err, "create provider") + t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) }) + + // Uncapped token limit: never blocks the handful of tokens driven here, but + // switches on usage metering — the switch that makes consumption rows record. + enabled := true + pol, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-price-" + name, + Enabled: &enabled, + SourceGroups: []string{grp.Id}, + DestinationProviderIds: []string{prov.Id}, + Limits: &api.AgentNetworkPolicyLimits{ + TokenLimit: api.AgentNetworkPolicyTokenLimit{ + Enabled: true, + GroupCap: 10_000_000, + UserCap: 10_000_000, + WindowSeconds: 60, + }, + }, + }) + require.NoError(t, err, "create policy") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), pol.Id) }) + + settings, err := srv.GetSettings(ctx) + require.NoError(t, err, "read settings") + require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned") + + proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-price-"+name+"-proxy") + require.NoError(t, err, "mint proxy token") + px, err := harness.StartProxy(ctx, srv, proxyToken) + require.NoError(t, err, "start proxy") + t.Cleanup(func() { _ = px.Terminate(context.Background()) }) + + cl, err := harness.StartClient(ctx, srv, sk.Key) + require.NoError(t, err, "start client") + t.Cleanup(func() { _ = cl.Terminate(context.Background()) }) + + require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management") + // Probe first: the GET resolves the endpoint and its first packet wakes the + // lazy proxy peer, so WaitProxyPeer then observes it connected. + proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint) + require.NoError(t, err, "resolve endpoint to proxy IP") + if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil { + t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background())) + } + + return pricedEnv{ + providerID: prov.Id, + groupID: grp.Id, + policyID: pol.Id, + upstream: vllm.URL, + endpoint: settings.Endpoint, + proxyIP: proxyIP, + client: cl, + proxy: px, + } +} + +// chatOnce drives one OpenAI-shaped chat for model through the tunnel, retrying +// to absorb first-call tunnel/DNS jitter, and returns the response body. +func chatOnce(t *testing.T, ctx context.Context, env pricedEnv, model, sessionID string) string { + t.Helper() + var code int + var body string + deadline := time.Now().Add(90 * time.Second) + for time.Now().Before(deadline) { + c, b, cerr := env.client.Chat(ctx, env.endpoint, env.proxyIP, harness.WireChat, model, "Reply with exactly: pong", sessionID) + if cerr == nil { + code, body = c, b + if code == 200 { + break + } + } + time.Sleep(5 * time.Second) + } + require.Equal(t, 200, code, + "chat for %s must return 200; body: %s\n=== proxy logs ===\n%s", model, body, env.proxy.Logs(context.Background())) + return body +} + +// findAccessLogBySession polls the access-log page for the row carrying sessionID. +func findAccessLogBySession(t *testing.T, ctx context.Context, sessionID string) api.AgentNetworkAccessLog { + t.Helper() + var row api.AgentNetworkAccessLog + require.Eventually(t, func() bool { + logs, lerr := srv.ListAccessLogs(ctx) + if lerr != nil { + return false + } + for _, r := range logs.Data { + if r.SessionId != nil && *r.SessionId == sessionID { + row = r + return true + } + } + return false + }, 30*time.Second, 2*time.Second, "session id %q must be recorded in an access-log row", sessionID) + return row +} + +// assertOpenAICostAtRates asserts an access-log row's token counts and every cost +// bucket match the mock's fixed usage priced at the given operator rates. The +// openai surface has no cache-write bucket and the mock reports no cache tokens, +// so the whole cost is input + output; cache costs must be exactly zero. +func assertOpenAICostAtRates(t *testing.T, row api.AgentNetworkAccessLog, inRate, outRate float64) { + t.Helper() + wantInput := float64(vllmPromptTokens) / 1000 * inRate + wantOutput := float64(vllmCompletionTokens) / 1000 * outRate + wantTotal := wantInput + wantOutput + + model := "" + if row.Model != nil { + model = *row.Model + } + t.Logf("[cost] model=%s in=%d out=%d rates in/out=%.4f/%.4f stored input/output/total=$%.6f/$%.6f/$%.6f expected input/output/total=$%.6f/$%.6f/$%.6f", + model, row.InputTokens, row.OutputTokens, inRate, outRate, + row.InputCostUsd, row.OutputCostUsd, row.CostUsd, wantInput, wantOutput, wantTotal) + + assert.EqualValues(t, vllmPromptTokens, row.InputTokens, "prompt tokens from the mock usage block") + assert.EqualValues(t, vllmCompletionTokens, row.OutputTokens, "completion tokens from the mock usage block") + assert.InDeltaf(t, wantInput, row.InputCostUsd, 1e-6, "input_cost_usd must be prompt tokens at the operator input rate") + assert.InDeltaf(t, wantOutput, row.OutputCostUsd, 1e-6, "output_cost_usd must be completion tokens at the operator output rate") + assert.InDeltaf(t, wantTotal, row.CostUsd, 1e-6, "cost_usd must be the sum of the priced buckets") + assert.Zerof(t, row.CachedInputCostUsd, "no cache-read tokens, so cached_input_cost_usd must be 0") + assert.Zerof(t, row.CacheCreationCostUsd, "openai surface has no cache-write bucket, so cache_creation_cost_usd must be 0") + assert.Zerof(t, row.CacheCostUsd, "no cache usage, so cache_cost_usd must be 0") + assert.InDeltaf(t, row.InputCostUsd+row.OutputCostUsd, row.CostUsd, 1e-9, "stored buckets must sum to cost_usd") +} + +// verifyUsageRowForSession re-checks the persisted usage row for a session +// directly in the management sqlite store — the same audit an operator runs on a +// production store.db — asserting its cost buckets match the operator rates. +func verifyUsageRowForSession(t *testing.T, sessionID string, inRate, outRate float64) { + t.Helper() + dbPath, err := srv.SnapshotStoreDB(t.TempDir()) + require.NoError(t, err, "snapshot management sqlite store") + + db, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{}) + require.NoError(t, err, "open store snapshot") + sqlDB, err := db.DB() + require.NoError(t, err) + defer func() { _ = sqlDB.Close() }() + + var provider, model string + var inTok, outTok, cachedTok, cacheCreateTok int64 + var inCost, cachedInCost, cacheCreateCost, outCost float64 + row := db.Raw(`SELECT provider, model, input_tokens, output_tokens, cached_input_tokens, cache_creation_tokens, + input_cost_usd, cached_input_cost_usd, cache_creation_cost_usd, output_cost_usd + FROM agent_network_request_usage WHERE session_id = ? ORDER BY timestamp DESC LIMIT 1`, sessionID).Row() + require.NoError(t, row.Scan(&provider, &model, &inTok, &outTok, &cachedTok, &cacheCreateTok, + &inCost, &cachedInCost, &cacheCreateCost, &outCost), + "a usage row must exist for session %q", sessionID) + + wantInput := float64(inTok) / 1000 * inRate + wantOutput := float64(outTok) / 1000 * outRate + t.Logf("[sql] session=%s %s/%s in=%d out=%d stored input/cached/create/output=$%.6f/$%.6f/$%.6f/$%.6f", + sessionID, provider, model, inTok, outTok, inCost, cachedInCost, cacheCreateCost, outCost) + assert.EqualValues(t, vllmPromptTokens, inTok, "usage row prompt tokens") + assert.EqualValues(t, vllmCompletionTokens, outTok, "usage row completion tokens") + assert.InDeltaf(t, wantInput, inCost, 1e-6, "usage input_cost_usd must be prompt tokens at the operator input rate") + assert.InDeltaf(t, wantOutput, outCost, 1e-6, "usage output_cost_usd must be completion tokens at the operator output rate") + assert.Zerof(t, cachedInCost, "usage cached_input_cost_usd must be 0 (no cache usage)") + assert.Zerof(t, cacheCreateCost, "usage cache_creation_cost_usd must be 0 (no cache usage)") +} + +// TestCustomModelPricing proves an operator can serve a model that is NOT in +// NetBird's compiled catalog, at prices they type themselves, and that those +// operator prices drive the recorded cost end to end — access log AND usage +// ledger. The provider enumerates one made-up model id at deliberately odd rates +// (no default entry could supply them), the client requests it, and every cost +// bucket must equal the mock's fixed token counts multiplied by those rates. +func TestCustomModelPricing(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + const ( + customModel = "e2e-custom-model" // absent from the compiled catalog + inRate = 0.037 // odd rates so a stray default can't match + outRate = 0.089 + ) + + env := provisionPricedProvider(t, ctx, "custommodel", []api.AgentNetworkProviderModel{ + {Id: customModel, InputPer1k: inRate, OutputPer1k: outRate}, + }) + + sessionID := "e2e-session-custommodel" + body := chatOnce(t, ctx, env, customModel, sessionID) + require.Contains(t, body, "chat.completion", "body should be an OpenAI-compatible completion; got: %s", body) + + row := findAccessLogBySession(t, ctx, sessionID) + require.NotNil(t, row.Model, "access-log row must carry the requested model") + assert.Equal(t, customModel, *row.Model, "the row must be stamped with the requested (custom) model, not the mock's response model") + assertOpenAICostAtRates(t, row, inRate, outRate) + + // Metering: the uncapped token limit switches on usage recording, so the + // request must surface as a consumption row with positive tokens and cost. + require.Eventually(t, func() bool { + rows, lerr := srv.ListConsumption(ctx) + if lerr != nil { + return false + } + for _, r := range rows { + if r.TokensInput > 0 && r.TokensOutput > 0 && r.CostUsd > 0 { + return true + } + } + return false + }, 60*time.Second, 3*time.Second, "custom-model usage must be metered into a consumption row with positive cost") + + // Final raw-SQL audit: bypass the API and re-verify the persisted usage row. + verifyUsageRowForSession(t, sessionID, inRate, outRate) +} + +// TestPriceChangeUpdatesRecordedCost proves that changing a provider's model +// price is reflected in the cost recorded for subsequent requests — in both the +// access log and the usage ledger — while requests already priced at the old +// rate keep their original cost. The update propagates to the connected proxy +// live (a mapping push rebuilds the cost_meter chain with the new table), so no +// reconnect or restart is needed; the test polls a fresh request until the new +// rate lands. +func TestPriceChangeUpdatesRecordedCost(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + const ( + customModel = "e2e-repriced-model" + inRateA = 0.010 + outRateA = 0.020 + inRateB = 0.050 // 5x / 4x the original, so a repriced row is unmistakable + outRateB = 0.080 + ) + + env := provisionPricedProvider(t, ctx, "reprice", []api.AgentNetworkProviderModel{ + {Id: customModel, InputPer1k: inRateA, OutputPer1k: outRateA}, + }) + + // Phase 1 — request priced at the original rate A. + sessionA := "e2e-session-reprice-a" + chatOnce(t, ctx, env, customModel, sessionA) + rowA := findAccessLogBySession(t, ctx, sessionA) + assertOpenAICostAtRates(t, rowA, inRateA, outRateA) + verifyUsageRowForSession(t, sessionA, inRateA, outRateA) + + // Change the model's price. The API key is omitted so the stored one is kept; + // the models array is re-sent with the new rates (PUT replaces the list). + // This reconciles synchronously and pushes a fresh cost_meter table to the + // already-connected proxy — no reconnect. + _, err := srv.UpdateProvider(ctx, env.providerID, api.AgentNetworkProviderRequest{ + Name: "reprice", + ProviderId: "openai_api", + UpstreamUrl: env.upstream, + Enabled: ptr(true), + Models: &[]api.AgentNetworkProviderModel{ + {Id: customModel, InputPer1k: inRateB, OutputPer1k: outRateB}, + }, + }) + require.NoError(t, err, "update provider price") + + // Phase 2 — the push + chain rebuild is async, so drive fresh requests (each + // under its own session) until one is priced at the new rate B. Each iteration + // fires one request and waits for that session's row to be ingested before + // reading its cost, so an un-ingested row is never mistaken for "still rate A". + // The expected new input cost is unmistakably higher than rate A, so a + // lingering old-rate row can't satisfy the check. + wantInputB := float64(vllmPromptTokens) / 1000 * inRateB + var repriced api.AgentNetworkAccessLog + var lastSession string + deadline := time.Now().Add(90 * time.Second) + for time.Now().Before(deadline) { + lastSession = fmt.Sprintf("e2e-session-reprice-b-%d", time.Now().UnixNano()) + code, _, cerr := env.client.Chat(ctx, env.endpoint, env.proxyIP, harness.WireChat, customModel, "Reply with exactly: pong", lastSession) + if cerr != nil || code != 200 { + time.Sleep(5 * time.Second) + continue + } + row := findAccessLogBySession(t, ctx, lastSession) + if inDelta(row.InputCostUsd, wantInputB, 1e-6) { + repriced = row + break + } + // Still priced at the old rate — the push hasn't landed yet; retry. + time.Sleep(5 * time.Second) + } + require.NotEmpty(t, repriced.Id, "a request after the price change must be priced at the new rate B; last input_cost_usd=$%.6f, wanted $%.6f\n=== proxy logs ===\n%s", + repriced.InputCostUsd, wantInputB, env.proxy.Logs(context.Background())) + + assertOpenAICostAtRates(t, repriced, inRateB, outRateB) + verifyUsageRowForSession(t, lastSession, inRateB, outRateB) + + // The original request keeps its original cost: repricing is not retroactive. + rowAStill := findAccessLogBySession(t, ctx, sessionA) + assertOpenAICostAtRates(t, rowAStill, inRateA, outRateA) + verifyUsageRowForSession(t, sessionA, inRateA, outRateA) +} + +// TestPricingDefaultsFileDrivesCost proves the operator-supplied pricing +// defaults file is what the proxy bills with. The harness configures +// server.agentNetwork.pricingDefaultsFile as a BARE FILENAME and writes that +// file into the bind-mounted datadir (see harness.PricingDefaultsFileName), so a +// pass exercises the whole chain: combined yaml → ToManagementConfig → +// pricing.LoadFile (relative path resolved against datadir) → DefaultTable → +// the synthesizer's cost_meter defaults payload → the proxy's lookup. +// +// The provider enumerates NO models, so it is a catch-all route with no +// per-provider-record pricing entry at all — the only rates that can price the +// request are the shipped defaults. The model is a real catalog model whose +// built-in rates the file replaces with deliberately odd values, so billing at +// the compiled-in rates (i.e. the file never loaded) fails the assertions. +func TestPricingDefaultsFileDrivesCost(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + // nil models: a gateway-style provider claiming every model. The synthesizer + // ships no per-record entry for it, so the defaults table is its price list. + env := provisionPricedProvider(t, ctx, "defaultsfile", nil) + + sessionID := "e2e-session-defaultsfile" + body := chatOnce(t, ctx, env, harness.PricedDefaultModel, sessionID) + require.Contains(t, body, "chat.completion", "body should be an OpenAI-compatible completion; got: %s", body) + + row := findAccessLogBySession(t, ctx, sessionID) + require.NotNil(t, row.Model, "access-log row must carry the requested model") + assert.Equal(t, harness.PricedDefaultModel, *row.Model, "the row must be stamped with the requested model") + + // The file's rates, not the compiled-in catalog rates for this model. + assertOpenAICostAtRates(t, row, harness.PricedDefaultInputPer1k, harness.PricedDefaultOutputPer1k) + verifyUsageRowForSession(t, sessionID, harness.PricedDefaultInputPer1k, harness.PricedDefaultOutputPer1k) +} + +// TestPricingDefaultsFileLeavesOtherModelsAlone proves the defaults file merges +// per entry rather than replacing the whole table: the file names exactly one +// model, so a DIFFERENT catalog model must still bill at its compiled-in rates. +// Without this, a file that shipped as a wholesale replacement would silently +// zero-cost every model the operator didn't list, and TestPricingDefaultsFile- +// DrivesCost alone would not notice. +func TestPricingDefaultsFileLeavesOtherModelsAlone(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + // gpt-4o-mini is a catalog model the pricing file does NOT mention, so it must + // keep its built-in rates. Pinned here independently of the catalog source so + // a rate change in either place surfaces as a failure to reconcile rather + // than passing silently. + const ( + untouchedModel = "gpt-4o-mini" + builtinInRate = 0.00015 + builtinOutRate = 0.0006 + ) + + env := provisionPricedProvider(t, ctx, "defaultsfileother", nil) + + sessionID := "e2e-session-defaultsfile-other" + chatOnce(t, ctx, env, untouchedModel, sessionID) + + row := findAccessLogBySession(t, ctx, sessionID) + assertOpenAICostAtRates(t, row, builtinInRate, builtinOutRate) + verifyUsageRowForSession(t, sessionID, builtinInRate, builtinOutRate) +} + +// TestCustomModelAccessLogAttribution proves a custom (non-catalog) model is +// handled correctly in the ACCESS LOG, not just in the cost columns. The other +// tests here assert money; this one asserts the row's identity and attribution +// dimensions — the columns the dashboard filters, groups and drills down on. +// +// A custom model id is the interesting case precisely because nothing in +// NetBird's catalog describes it. Its provider vendor, parser surface, cost +// buckets, and dashboard filterability all have to come from the operator's +// provider record rather than from a compiled-in entry. So this checks: +// +// - the row is stamped with the REQUESTED model id verbatim, not the mock +// upstream's response model (Qwen/Qwen2.5-0.5B-Instruct) and not a +// normalized or catalog-substituted id; +// - provider is the vendor SURFACE ("openai", from the catalog entry's +// ParserID) — a custom model does not change which wire shape was spoken; +// - resolved_provider_id / selected_policy_id / group_ids attribute the row to +// the operator's provider record, the authorising policy, and the caller's +// group, so spend on a custom model is attributable; +// - decision is "allow" with no deny reason, and the request dimensions +// (status 200, POST, the OpenAI chat path, non-stream, source IP, duration) +// are recorded; +// - management's SERVER-SIDE model filter finds the row by its custom id, so +// the model column is genuinely indexed and queryable rather than merely +// stored; +// - prompt/completion capture stays empty, since prompt collection is off by +// default and a custom model must not bypass that gate. +func TestCustomModelAccessLogAttribution(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + // A model id no catalog entry carries, at odd rates so its cost cannot come + // from anywhere but the provider record. + const ( + customModel = "e2e-attribution-model-v9" + inRate = 0.0271 + outRate = 0.0913 + ) + + env := provisionPricedProvider(t, ctx, "attribution", []api.AgentNetworkProviderModel{ + {Id: customModel, InputPer1k: inRate, OutputPer1k: outRate}, + }) + + sessionID := "e2e-session-attribution" + body := chatOnce(t, ctx, env, customModel, sessionID) + require.Contains(t, body, "chat.completion", "body should be an OpenAI-compatible completion; got: %s", body) + + row := findAccessLogBySession(t, ctx, sessionID) + + // Identity: the requested model verbatim. The mock answers with its own + // served model id, so a row carrying that instead means the log is sourced + // from the response body rather than the parsed request. + require.NotNil(t, row.Model, "access-log row must carry the requested model") + assert.Equal(t, customModel, *row.Model, + "the row must be stamped with the requested custom model id verbatim, not the mock upstream's response model (%s)", harness.VLLMModel) + + // Surface: a custom model id does not change the wire shape that was spoken. + // provider is the vendor surface from the catalog entry's parser, which is + // also the key the cost meter's cache formula switches on. + require.NotNil(t, row.Provider, "access-log row must carry the vendor surface") + assert.Equal(t, "openai", *row.Provider, + "openai_api's parser surface is openai, regardless of how exotic the model id is") + + // Attribution: which provider record served it, which policy authorised it, + // and which group the authorisation came through. Without these, spend on a + // custom model can be seen but not attributed. + require.NotNil(t, row.ResolvedProviderId, "row must name the provider record that served the request") + assert.Equal(t, env.providerID, *row.ResolvedProviderId, + "the router stamps the operator's provider record id; a custom model must attribute to the record that enumerated it") + require.NotNil(t, row.SelectedPolicyId, "row must name the policy that authorised the request") + assert.Equal(t, env.policyID, *row.SelectedPolicyId, + "the policy carrying the token limit is the one that paid for the request") + require.NotNil(t, row.GroupIds, "row must carry the authorising group ids") + assert.Contains(t, *row.GroupIds, env.groupID, + "the caller's group is the policy's source group, so it must be the authorising group") + + // Decision + request dimensions. + require.NotNil(t, row.Decision, "row must carry the policy decision") + assert.Equal(t, "allow", *row.Decision, "the uncapped policy allows this request") + if row.DenyReason != nil { + assert.Empty(t, *row.DenyReason, "an allowed request must carry no deny reason") + } + assert.Equal(t, 200, row.StatusCode, "the mock upstream answers 200") + if row.Method != nil { + assert.Equal(t, "POST", *row.Method, "a chat completion is a POST") + } + require.NotNil(t, row.Path, "row must record the request path") + assert.Equal(t, "/v1/chat/completions", *row.Path, + "the OpenAI chat path the client called, as seen by the proxy") + require.NotNil(t, row.Host, "row must record the host the client addressed") + assert.Equal(t, env.endpoint, *row.Host, "the agent-network endpoint the client resolved") + if row.Stream != nil { + assert.False(t, *row.Stream, "the harness sends a non-streaming request") + } + require.NotNil(t, row.SourceIp, "row must record the caller's tunnel IP") + assert.NotEmpty(t, *row.SourceIp, "the request arrived over the tunnel, so a source IP is known") + + // Tokens and cost, so the attribution above is anchored to a real priced row + // rather than an empty shell that happens to carry the right ids. + assertOpenAICostAtRates(t, row, inRate, outRate) + assert.EqualValues(t, vllmPromptTokens+vllmCompletionTokens, row.TotalTokens, + "total_tokens is the mock's reported total") + + // Prompt capture is off by default (account master switch), and a custom + // model must not bypass that gate. + if row.RequestPrompt != nil { + assert.Empty(t, *row.RequestPrompt, "prompt collection is off by default, so no prompt may be stored") + } + if row.ResponseCompletion != nil { + assert.Empty(t, *row.ResponseCompletion, "prompt collection is off by default, so no completion may be stored") + } + + // Queryability: management's SERVER-SIDE model filter must find the row by + // its custom id. findAccessLogBySession above scans a page client-side, so + // this is the check that the model column is actually indexed and filterable + // — the dashboard's per-model drill-down on a custom model depends on it. + filtered, err := srv.ListAccessLogsFiltered(ctx, url.Values{"model": []string{customModel}}) + require.NoError(t, err, "filter access logs by the custom model id") + require.Positive(t, filtered.TotalRecords, "the custom model must be findable via the server-side model filter") + foundSession := false + for _, r := range filtered.Data { + require.NotNil(t, r.Model, "filtered row must carry a model") + assert.Equal(t, customModel, *r.Model, "the model filter must not return rows for other models") + if r.SessionId != nil && *r.SessionId == sessionID { + foundSession = true + } + } + assert.True(t, foundSession, "the filtered page must include this test's request") + + // Final raw-SQL audit of the parallel usage row: the ledger must carry the + // same custom model, surface, and provider-record attribution as the log. + verifyUsageAttributionForSession(t, sessionID, customModel, "openai", env.providerID, env.groupID) +} + +// verifyUsageAttributionForSession checks the usage ledger's attribution columns +// for a session directly in the management sqlite store — including the group +// child row, which the API renders but which only exists if the proxy's +// authorising-group CSV was parsed into normalised rows. The usage table is +// written unconditionally (independent of the log-collection toggle), so this is +// the record that must attribute spend even for accounts with logs off. +func verifyUsageAttributionForSession(t *testing.T, sessionID, wantModel, wantProvider, wantProviderID, wantGroupID string) { + t.Helper() + dbPath, err := srv.SnapshotStoreDB(t.TempDir()) + require.NoError(t, err, "snapshot management sqlite store") + + db, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{}) + require.NoError(t, err, "open store snapshot") + sqlDB, err := db.DB() + require.NoError(t, err) + defer func() { _ = sqlDB.Close() }() + + var id, provider, model, resolvedProviderID, userID string + require.NoError(t, db.Raw( + `SELECT id, provider, model, resolved_provider_id, user_id + FROM agent_network_request_usage WHERE session_id = ? ORDER BY timestamp DESC LIMIT 1`, sessionID). + Row().Scan(&id, &provider, &model, &resolvedProviderID, &userID), + "a usage row must exist for session %q", sessionID) + + t.Logf("[sql] usage attribution session=%s id=%s provider=%s model=%s resolved_provider_id=%s user_id=%s", + sessionID, id, provider, model, resolvedProviderID, userID) + assert.Equal(t, wantModel, model, "usage row must carry the requested custom model") + assert.Equal(t, wantProvider, provider, "usage row must carry the vendor surface") + assert.Equal(t, wantProviderID, resolvedProviderID, "usage row must attribute to the operator's provider record") + assert.NotEmpty(t, userID, "the tunnel peer resolves to a principal, so the usage row must be attributable to it") + + // The authorising group lands in the normalised child table, which is what + // the usage overview joins on to break spend down by group. + var groupIDs []string + require.NoError(t, db.Raw( + `SELECT group_id FROM agent_network_request_usage_group WHERE usage_id = ?`, id). + Scan(&groupIDs).Error, "read usage group child rows") + assert.Contains(t, groupIDs, wantGroupID, + "the authorising group must be normalised into a usage_group row so spend can be grouped by it") +} + +// inDelta reports whether a and b are within tol of each other. +func inDelta(a, b, tol float64) bool { + d := a - b + if d < 0 { + d = -d + } + return d <= tol +} diff --git a/e2e/harness/agentnetwork.go b/e2e/harness/agentnetwork.go index 53aa8e342..078e697af 100644 --- a/e2e/harness/agentnetwork.go +++ b/e2e/harness/agentnetwork.go @@ -9,6 +9,7 @@ import ( "fmt" "io" "net/http" + "net/url" "github.com/netbirdio/netbird/shared/management/http/api" ) @@ -74,6 +75,13 @@ func (c *Combined) DeleteProvider(ctx context.Context, id string) error { return anDelete(ctx, c, "/api/agent-network/providers/"+id) } +// UpdateProvider replaces a provider by id (PUT). The API key may be omitted on +// the request to keep the stored one; Models replaces the enumerated list, so +// this is the path a test uses to change a model's price mid-run. +func (c *Combined) UpdateProvider(ctx context.Context, id string, req api.AgentNetworkProviderRequest) (api.AgentNetworkProvider, error) { + return anRequest[api.AgentNetworkProvider](ctx, c, http.MethodPut, "/api/agent-network/providers/"+id, req) +} + // SetProviderEnabled toggles a provider's enabled flag, preserving its other // fields (the API key is omitted, which keeps the stored one). Used to run one // provider at a time so model→provider routing is unambiguous. @@ -139,3 +147,16 @@ func (c *Combined) ListConsumption(ctx context.Context) ([]api.AgentNetworkConsu func (c *Combined) ListAccessLogs(ctx context.Context) (api.AgentNetworkAccessLogsResponse, error) { return anRequest[api.AgentNetworkAccessLogsResponse](ctx, c, http.MethodGet, "/api/agent-network/access-logs", nil) } + +// ListAccessLogsFiltered returns the access-log page narrowed by the given +// query parameters (e.g. model=..., session_id=..., provider_id=...). This +// exercises management's server-side filtering rather than filtering client +// side, so a row that is ingested but not indexed under the filtered column +// surfaces as an empty page. +func (c *Combined) ListAccessLogsFiltered(ctx context.Context, query url.Values) (api.AgentNetworkAccessLogsResponse, error) { + path := "/api/agent-network/access-logs" + if encoded := query.Encode(); encoded != "" { + path += "?" + encoded + } + return anRequest[api.AgentNetworkAccessLogsResponse](ctx, c, http.MethodGet, path, nil) +} diff --git a/e2e/harness/combined.go b/e2e/harness/combined.go index 5723100ca..b2f0d89d2 100644 --- a/e2e/harness/combined.go +++ b/e2e/harness/combined.go @@ -93,10 +93,19 @@ func StartCombined(ctx context.Context) (*Combined, error) { _ = net.Remove(ctx) return nil, fmt.Errorf("write combined config: %w", err) } - if err := os.MkdirAll(filepath.Join(workDir, "data"), 0o755); err != nil { + dataDir := filepath.Join(workDir, "data") + if err := os.MkdirAll(dataDir, 0o755); err != nil { _ = net.Remove(ctx) return nil, fmt.Errorf("create datadir: %w", err) } + // The config's agentNetwork.pricingDefaultsFile is a bare filename, so the + // server resolves it against the datadir; write it there. It is an explicitly + // configured path, so a failure to load fails the server's startup — which + // surfaces here as the /api/instance readiness wait timing out. + if err := os.WriteFile(filepath.Join(dataDir, PricingDefaultsFileName), []byte(pricingDefaultsYAML), 0o644); err != nil { //nolint:gosec // non-secret config, bind-mounted and read by the container + _ = net.Remove(ctx) + return nil, fmt.Errorf("write pricing defaults: %w", err) + } req := testcontainers.ContainerRequest{ Image: combinedImage, diff --git a/e2e/harness/config.go b/e2e/harness/config.go index b4bed60a2..71b3656c5 100644 --- a/e2e/harness/config.go +++ b/e2e/harness/config.go @@ -8,6 +8,13 @@ package harness // embedded IdP, local signal/relay/STUN, and a sqlite store under the mounted // data dir. exposedAddress is the address peers use to reach this container; it // is overridden per-run so the value matches the container's network alias. +// +// pricingDefaultsFile is deliberately a BARE FILENAME, not an absolute path: it +// must resolve against dataDir (→ /nb/data/), which is the resolution rule +// the combined server applies. It is also an EXPLICITLY configured path, so the +// server is required to load it — a broken path or malformed file fails startup +// rather than silently falling back to the compiled-in rates, and TestMain then +// fails with the container logs. const combinedConfigYAML = `server: listenAddress: ":8080" exposedAddress: "%s" @@ -23,4 +30,36 @@ const combinedConfigYAML = `server: issuer: "%s" store: engine: "sqlite" + agentNetwork: + pricingDefaultsFile: "` + PricingDefaultsFileName + `" +` + +const ( + // PricingDefaultsFileName is the basename of the operator-supplied LLM + // pricing defaults file the combined server is configured to load. Written + // into the bind-mounted datadir by StartCombined. + PricingDefaultsFileName = "e2e_llm_pricing.yaml" + + // PricedDefaultModel is a real catalog model (openai surface) whose rates the + // defaults file below REPLACES. Tests drive it against the mock vLLM upstream + // and assert the file's rates were billed, which is only true if the file + // travelled: config → LoadFile → DefaultTable → synthesizer → the proxy's + // cost_meter defaults table. + PricedDefaultModel = "gpt-4.1-mini" + // PricedDefaultInputPer1k / PricedDefaultOutputPer1k are deliberately odd + // values that no compiled-in catalog entry carries (gpt-4.1-mini ships as + // 0.0004 / 0.0016), so a test asserting them cannot pass on the built-in + // table. + PricedDefaultInputPer1k = 0.0123 + PricedDefaultOutputPer1k = 0.0456 +) + +// pricingDefaultsYAML is the operator-supplied pricing defaults file. Its schema +// is surface -> model -> per-1k rates. Entries replace the compiled-in entry for +// the same surface+model whole; every other model keeps its built-in rates, so +// this file overriding one model must not disturb the rest of the table. +const pricingDefaultsYAML = `openai: + gpt-4.1-mini: + input_per_1k: 0.0123 + output_per_1k: 0.0456 ` diff --git a/management/cmd/management.go b/management/cmd/management.go index 79c838ec4..147985314 100644 --- a/management/cmd/management.go +++ b/management/cmd/management.go @@ -23,6 +23,7 @@ import ( "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/formatter/hook" + agentnetworkpricing "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" "github.com/netbirdio/netbird/management/internals/server" nbconfig "github.com/netbirdio/netbird/management/internals/server/config" nbdomain "github.com/netbirdio/netbird/shared/management/domain" @@ -112,6 +113,29 @@ var ( mgmtSingleAccModeDomain = "" } + // Load the management-side LLM pricing defaults file: an + // explicitly configured path is required to load (a typo must + // fail startup — the operator believes those rates are live); + // otherwise /defaults_llm_pricing.yaml is probed and + // may be absent (compiled-in defaults serve). A relative path + // is resolved against the datadir so a bare filename lands + // alongside the store. Either way the path stays watched: the + // reloader picks up edits — and the file appearing later — + // without a restart. + pricingPath := config.AgentNetwork.PricingDefaultsFile + pricingRequired := pricingPath != "" + if !pricingRequired { + pricingPath = agentnetworkpricing.DefaultFileName + } + if !filepath.IsAbs(pricingPath) { + pricingPath = filepath.Join(config.Datadir, pricingPath) + } + log.Infof("loading agent-network pricing defaults from %s (required: %v)", pricingPath, pricingRequired) + if err := agentnetworkpricing.LoadFile(pricingPath, pricingRequired); err != nil { + return fmt.Errorf("load agent-network pricing defaults: %v", err) + } + agentnetworkpricing.StartReloader(ctx, agentnetworkpricing.ReloadInterval) + srv := newServer(&server.Config{ NbConfig: config, DNSDomain: dnsDomain, diff --git a/management/internals/modules/agentnetwork/catalog/catalog.go b/management/internals/modules/agentnetwork/catalog/catalog.go index f82e94bf3..2c4efd0b4 100644 --- a/management/internals/modules/agentnetwork/catalog/catalog.go +++ b/management/internals/modules/agentnetwork/catalog/catalog.go @@ -7,12 +7,28 @@ package catalog import "github.com/netbirdio/netbird/shared/management/http/api" // Model is the in-memory representation of a catalog model. +// +// The three cache rates mirror the proxy cost meter's Entry semantics +// (USD per 1k tokens; 0 = no rate configured, that bucket bills at +// InputPer1k): +// - CachedInputPer1k: OpenAI-shape rate for cached prompt tokens +// (a SUBSET of input tokens). Typically 0.1-0.5x input. +// - CacheReadPer1k / CacheCreationPer1k: Anthropic-shape rates for +// the two ADDITIVE prompt-cache buckets. Typically 0.1x / 1.25x +// input. +// +// The catalog is the single default-pricing source: the agentnetwork +// pricing package folds these models into per-surface tables that the +// synthesizer ships to the proxy's cost_meter. type Model struct { - ID string - Label string - InputPer1k float64 - OutputPer1k float64 - ContextWindow int + ID string + Label string + InputPer1k float64 + OutputPer1k float64 + CachedInputPer1k float64 + CacheReadPer1k float64 + CacheCreationPer1k float64 + ContextWindow int } // ProviderKind groups catalog entries for UI presentation. The split @@ -65,6 +81,17 @@ type Provider struct { // surface — the proxy middleware then falls back to URL sniffing // or skips request-side enrichment. ParserID string + // PricingSurfaces names the cost-meter pricing surfaces this + // provider's Models are priced under ("openai", "anthropic", + // "bedrock" — the llm.Parser surface the request parser stamps as + // llm.provider at billing time). NOT derivable from ParserID: + // bedrock_api and vertex_ai_api leave ParserID empty (URL-sniffed) + // yet price under "bedrock" / "anthropic", and kimi_api serves two + // body shapes so it prices under both. Nil for gateway/custom + // entries, which declare no models. Same (surface, model) pair + // contributed by two providers must carry identical rates — the + // pricing package's tests enforce that. + PricingSurfaces []string // IdentityInjection, when non-nil, instructs the proxy to stamp // the caller's NetBird identity onto upstream requests under the // configured header names. Used for gateways like LiteLLM that @@ -219,6 +246,7 @@ var providers = []Provider{ DefaultContentType: "application/json", BrandColor: "#10A37F", ParserID: "openai", + PricingSurfaces: []string{"openai"}, // Pricing + context windows cross-checked against LiteLLM's // model_prices_and_context_window.json. Notable corrections from // earlier values: o4-mini repriced from $4/$16 to $1.10/$4.40 @@ -226,20 +254,20 @@ var providers = []Provider{ // family context windows split between 1.05M for full-size // models and 272K for mini/nano/codex variants. Models: []Model{ - {ID: "gpt-5.5", Label: "GPT-5.5", InputPer1k: 0.005, OutputPer1k: 0.030, ContextWindow: 1050000}, - {ID: "gpt-5.5-pro", Label: "GPT-5.5 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, ContextWindow: 1050000}, - {ID: "gpt-5.4", Label: "GPT-5.4", InputPer1k: 0.0025, OutputPer1k: 0.015, ContextWindow: 1050000}, - {ID: "gpt-5.4-pro", Label: "GPT-5.4 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, ContextWindow: 1050000}, - {ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini", InputPer1k: 0.00075, OutputPer1k: 0.0045, ContextWindow: 272000}, - {ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano", InputPer1k: 0.0002, OutputPer1k: 0.00125, ContextWindow: 272000}, - {ID: "gpt-5.3-codex", Label: "GPT-5.3 Codex", InputPer1k: 0.00175, OutputPer1k: 0.014, ContextWindow: 272000}, - {ID: "gpt-5.3-chat-latest", Label: "GPT-5.3 Chat", InputPer1k: 0.00175, OutputPer1k: 0.014, ContextWindow: 128000}, - {ID: "o4-mini", Label: "o4-mini", InputPer1k: 0.0011, OutputPer1k: 0.0044, ContextWindow: 200000}, - {ID: "gpt-4.1", Label: "GPT-4.1", InputPer1k: 0.002, OutputPer1k: 0.008, ContextWindow: 1047576}, - {ID: "gpt-4.1-mini", Label: "GPT-4.1 mini", InputPer1k: 0.0004, OutputPer1k: 0.0016, ContextWindow: 1047576}, - {ID: "gpt-4.1-nano", Label: "GPT-4.1 nano", InputPer1k: 0.0001, OutputPer1k: 0.0004, ContextWindow: 1047576}, - {ID: "gpt-4o", Label: "GPT-4o", InputPer1k: 0.0025, OutputPer1k: 0.010, ContextWindow: 128000}, - {ID: "gpt-4o-mini", Label: "GPT-4o mini", InputPer1k: 0.00015, OutputPer1k: 0.0006, ContextWindow: 128000}, + {ID: "gpt-5.5", Label: "GPT-5.5", InputPer1k: 0.005, OutputPer1k: 0.030, CachedInputPer1k: 0.0005, ContextWindow: 1050000}, + {ID: "gpt-5.5-pro", Label: "GPT-5.5 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, CachedInputPer1k: 0.003, ContextWindow: 1050000}, + {ID: "gpt-5.4", Label: "GPT-5.4", InputPer1k: 0.0025, OutputPer1k: 0.015, CachedInputPer1k: 0.00025, ContextWindow: 1050000}, + {ID: "gpt-5.4-pro", Label: "GPT-5.4 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, CachedInputPer1k: 0.003, ContextWindow: 1050000}, + {ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini", InputPer1k: 0.00075, OutputPer1k: 0.0045, CachedInputPer1k: 0.000075, ContextWindow: 272000}, + {ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano", InputPer1k: 0.0002, OutputPer1k: 0.00125, CachedInputPer1k: 0.00002, ContextWindow: 272000}, + {ID: "gpt-5.3-codex", Label: "GPT-5.3 Codex", InputPer1k: 0.00175, OutputPer1k: 0.014, CachedInputPer1k: 0.000175, ContextWindow: 272000}, + {ID: "gpt-5.3-chat-latest", Label: "GPT-5.3 Chat", InputPer1k: 0.00175, OutputPer1k: 0.014, CachedInputPer1k: 0.000175, ContextWindow: 128000}, + {ID: "o4-mini", Label: "o4-mini", InputPer1k: 0.0011, OutputPer1k: 0.0044, CachedInputPer1k: 0.000275, ContextWindow: 200000}, + {ID: "gpt-4.1", Label: "GPT-4.1", InputPer1k: 0.002, OutputPer1k: 0.008, CachedInputPer1k: 0.0005, ContextWindow: 1047576}, + {ID: "gpt-4.1-mini", Label: "GPT-4.1 mini", InputPer1k: 0.0004, OutputPer1k: 0.0016, CachedInputPer1k: 0.0001, ContextWindow: 1047576}, + {ID: "gpt-4.1-nano", Label: "GPT-4.1 nano", InputPer1k: 0.0001, OutputPer1k: 0.0004, CachedInputPer1k: 0.000025, ContextWindow: 1047576}, + {ID: "gpt-4o", Label: "GPT-4o", InputPer1k: 0.0025, OutputPer1k: 0.010, CachedInputPer1k: 0.00125, ContextWindow: 128000}, + {ID: "gpt-4o-mini", Label: "GPT-4o mini", InputPer1k: 0.00015, OutputPer1k: 0.0006, CachedInputPer1k: 0.000075, ContextWindow: 128000}, {ID: "gpt-4-turbo", Label: "GPT-4 Turbo", InputPer1k: 0.01, OutputPer1k: 0.03, ContextWindow: 128000}, {ID: "gpt-3.5-turbo", Label: "GPT-3.5 Turbo", InputPer1k: 0.0005, OutputPer1k: 0.0015, ContextWindow: 16385}, {ID: "text-embedding-3-large", Label: "text-embedding-3-large", InputPer1k: 0.00013, OutputPer1k: 0, ContextWindow: 8191}, @@ -257,6 +285,7 @@ var providers = []Provider{ DefaultContentType: "application/json", BrandColor: "#D97757", ParserID: "anthropic", + PricingSurfaces: []string{"anthropic"}, // Per Anthropic's current model lineup. Pricing in USD per 1k // tokens. Context windows: 4.6+ family is 1M; Haiku 4.5 stays at // 200K. claude-3-7-sonnet and claude-3-5-haiku retired @@ -267,14 +296,14 @@ var providers = []Provider{ // account to be on >= 30-day data retention or all requests // 400. Models: []Model{ - {ID: "claude-fable-5", Label: "Claude Fable 5", InputPer1k: 0.010, OutputPer1k: 0.050, ContextWindow: 1000000}, - {ID: "claude-opus-4-8", Label: "Claude Opus 4.8", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "claude-opus-4-7", Label: "Claude Opus 4.7", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "claude-opus-4-6", Label: "Claude Opus 4.6", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (deprecated, retires 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, ContextWindow: 200000}, - {ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000}, - {ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 200000}, - {ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5", InputPer1k: 0.001, OutputPer1k: 0.005, ContextWindow: 200000}, + {ID: "claude-fable-5", Label: "Claude Fable 5", InputPer1k: 0.010, OutputPer1k: 0.050, CacheReadPer1k: 0.001, CacheCreationPer1k: 0.0125, ContextWindow: 1000000}, + {ID: "claude-opus-4-8", Label: "Claude Opus 4.8", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "claude-opus-4-7", Label: "Claude Opus 4.7", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "claude-opus-4-6", Label: "Claude Opus 4.6", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (deprecated, retires 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, CacheReadPer1k: 0.0015, CacheCreationPer1k: 0.01875, ContextWindow: 200000}, + {ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000}, + {ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 200000}, + {ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5", InputPer1k: 0.001, OutputPer1k: 0.005, CacheReadPer1k: 0.0001, CacheCreationPer1k: 0.00125, ContextWindow: 200000}, }, }, { @@ -288,18 +317,19 @@ var providers = []Provider{ DefaultContentType: "application/json", BrandColor: "#0078D4", ParserID: "openai", + PricingSurfaces: []string{"openai"}, // Mirrors openai_api pricing — Azure resells OpenAI models at the // same per-token rates, just under different deployment names. Models: []Model{ - {ID: "gpt-5.5", Label: "GPT-5.5 (Azure)", InputPer1k: 0.005, OutputPer1k: 0.030, ContextWindow: 1050000}, - {ID: "gpt-5.4", Label: "GPT-5.4 (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.015, ContextWindow: 1050000}, - {ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini (Azure)", InputPer1k: 0.00075, OutputPer1k: 0.0045, ContextWindow: 272000}, - {ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano (Azure)", InputPer1k: 0.0002, OutputPer1k: 0.00125, ContextWindow: 272000}, - {ID: "o4-mini", Label: "o4-mini (Azure)", InputPer1k: 0.0011, OutputPer1k: 0.0044, ContextWindow: 200000}, - {ID: "gpt-4.1", Label: "GPT-4.1 (Azure)", InputPer1k: 0.002, OutputPer1k: 0.008, ContextWindow: 1047576}, - {ID: "gpt-4.1-mini", Label: "GPT-4.1 mini (Azure)", InputPer1k: 0.0004, OutputPer1k: 0.0016, ContextWindow: 1047576}, - {ID: "gpt-4o", Label: "GPT-4o (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.010, ContextWindow: 128000}, - {ID: "gpt-4o-mini", Label: "GPT-4o mini (Azure)", InputPer1k: 0.00015, OutputPer1k: 0.0006, ContextWindow: 128000}, + {ID: "gpt-5.5", Label: "GPT-5.5 (Azure)", InputPer1k: 0.005, OutputPer1k: 0.030, CachedInputPer1k: 0.0005, ContextWindow: 1050000}, + {ID: "gpt-5.4", Label: "GPT-5.4 (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.015, CachedInputPer1k: 0.00025, ContextWindow: 1050000}, + {ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini (Azure)", InputPer1k: 0.00075, OutputPer1k: 0.0045, CachedInputPer1k: 0.000075, ContextWindow: 272000}, + {ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano (Azure)", InputPer1k: 0.0002, OutputPer1k: 0.00125, CachedInputPer1k: 0.00002, ContextWindow: 272000}, + {ID: "o4-mini", Label: "o4-mini (Azure)", InputPer1k: 0.0011, OutputPer1k: 0.0044, CachedInputPer1k: 0.000275, ContextWindow: 200000}, + {ID: "gpt-4.1", Label: "GPT-4.1 (Azure)", InputPer1k: 0.002, OutputPer1k: 0.008, CachedInputPer1k: 0.0005, ContextWindow: 1047576}, + {ID: "gpt-4.1-mini", Label: "GPT-4.1 mini (Azure)", InputPer1k: 0.0004, OutputPer1k: 0.0016, CachedInputPer1k: 0.0001, ContextWindow: 1047576}, + {ID: "gpt-4o", Label: "GPT-4o (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.010, CachedInputPer1k: 0.00125, ContextWindow: 128000}, + {ID: "gpt-4o-mini", Label: "GPT-4o mini (Azure)", InputPer1k: 0.00015, OutputPer1k: 0.0006, CachedInputPer1k: 0.000075, ContextWindow: 128000}, {ID: "gpt-35-turbo", Label: "GPT-3.5 Turbo (Azure)", InputPer1k: 0.0005, OutputPer1k: 0.0015, ContextWindow: 16385}, }, }, @@ -313,6 +343,9 @@ var providers = []Provider{ AuthHeaderTemplate: "Bearer ${API_KEY}", DefaultContentType: "application/json", BrandColor: "#FF9900", + // ParserID stays empty (path-style dispatch via IsBedrockPathStyle); + // the request parser meters these under the "bedrock" surface. + PricingSurfaces: []string{"bedrock"}, // Anthropic models on Bedrock take the anthropic.* prefix and // follow the same lineup / pricing as the first-party Anthropic // catalog entry above. claude-3-7-sonnet and claude-3-5-haiku @@ -322,13 +355,13 @@ var providers = []Provider{ // Llama 3.3 70B entry kept unchanged — LiteLLM tracks only // per-region Llama 3 entries; standalone 3.3 not yet listed. Models: []Model{ - {ID: "anthropic.claude-opus-4-8", Label: "Claude Opus 4.8 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "anthropic.claude-opus-4-7", Label: "Claude Opus 4.7 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "anthropic.claude-opus-4-6", Label: "Claude Opus 4.6 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "anthropic.claude-opus-4-1", Label: "Claude Opus 4.1 (Bedrock, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, ContextWindow: 200000}, - {ID: "anthropic.claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000}, - {ID: "anthropic.claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 200000}, - {ID: "anthropic.claude-haiku-4-5", Label: "Claude Haiku 4.5 (Bedrock)", InputPer1k: 0.001, OutputPer1k: 0.005, ContextWindow: 200000}, + {ID: "anthropic.claude-opus-4-8", Label: "Claude Opus 4.8 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "anthropic.claude-opus-4-7", Label: "Claude Opus 4.7 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "anthropic.claude-opus-4-6", Label: "Claude Opus 4.6 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "anthropic.claude-opus-4-1", Label: "Claude Opus 4.1 (Bedrock, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, CacheReadPer1k: 0.0015, CacheCreationPer1k: 0.01875, ContextWindow: 200000}, + {ID: "anthropic.claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000}, + {ID: "anthropic.claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 200000}, + {ID: "anthropic.claude-haiku-4-5", Label: "Claude Haiku 4.5 (Bedrock)", InputPer1k: 0.001, OutputPer1k: 0.005, CacheReadPer1k: 0.0001, CacheCreationPer1k: 0.00125, ContextWindow: 200000}, {ID: "meta.llama3-3-70b-instruct", Label: "Llama 3.3 70B (Bedrock)", InputPer1k: 0.00072, OutputPer1k: 0.00072, ContextWindow: 128000}, {ID: "amazon.nova-2-lite", Label: "Amazon Nova 2 Lite (Bedrock, preview)", InputPer1k: 0.0003, OutputPer1k: 0.0025, ContextWindow: 1000000}, {ID: "amazon.nova-pro", Label: "Amazon Nova Pro (Bedrock)", InputPer1k: 0.0008, OutputPer1k: 0.0032, ContextWindow: 300000}, @@ -358,6 +391,10 @@ var providers = []Provider{ AuthHeaderTemplate: "Bearer ${API_KEY}", DefaultContentType: "application/json", BrandColor: "#4285F4", + // ParserID stays empty (path-style dispatch via IsVertexPathStyle); + // Anthropic-on-Vertex requests are metered under the "anthropic" + // surface with the bare, unversioned model id. + PricingSurfaces: []string{"anthropic"}, // Vertex carries the model in the URL path and authenticates with a // service-account-minted OAuth token (api_key = "keyfile::"). // Only Anthropic-on-Vertex is metered today: the request parser maps the @@ -369,14 +406,14 @@ var providers = []Provider{ // exists — the router denies unmeterable publishers rather than forward // them uncounted. Models: []Model{ - {ID: "claude-fable-5", Label: "Claude Fable 5 (Vertex)", InputPer1k: 0.010, OutputPer1k: 0.050, ContextWindow: 1000000}, - {ID: "claude-opus-4-8", Label: "Claude Opus 4.8 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "claude-opus-4-7", Label: "Claude Opus 4.7 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "claude-opus-4-6", Label: "Claude Opus 4.6 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (Vertex, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, ContextWindow: 200000}, - {ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000}, - {ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 200000}, - {ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5 (Vertex)", InputPer1k: 0.001, OutputPer1k: 0.005, ContextWindow: 200000}, + {ID: "claude-fable-5", Label: "Claude Fable 5 (Vertex)", InputPer1k: 0.010, OutputPer1k: 0.050, CacheReadPer1k: 0.001, CacheCreationPer1k: 0.0125, ContextWindow: 1000000}, + {ID: "claude-opus-4-8", Label: "Claude Opus 4.8 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "claude-opus-4-7", Label: "Claude Opus 4.7 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "claude-opus-4-6", Label: "Claude Opus 4.6 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (Vertex, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, CacheReadPer1k: 0.0015, CacheCreationPer1k: 0.01875, ContextWindow: 200000}, + {ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000}, + {ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 200000}, + {ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5 (Vertex)", InputPer1k: 0.001, OutputPer1k: 0.005, CacheReadPer1k: 0.0001, CacheCreationPer1k: 0.00125, ContextWindow: 200000}, }, }, { @@ -390,6 +427,7 @@ var providers = []Provider{ DefaultContentType: "application/json", BrandColor: "#FF7000", ParserID: "openai", + PricingSurfaces: []string{"openai"}, // Pricing + context windows cross-checked against LiteLLM. Key // gotchas the marketing page hides: // - `mistral-medium-latest` aliases to Medium 3.1 ($0.40/$2), @@ -448,6 +486,10 @@ var providers = []Provider{ // model id "k3") is account-bound seat licensing rather than a // meterable platform key, so it's deliberately not the default. ParserID: "", + // Both body shapes are metered: /v1/chat/completions under + // "openai", /anthropic/v1/messages under "anthropic" — so the + // K3 entry is priced on both surfaces. + PricingSurfaces: []string{"openai", "anthropic"}, // Pricing per Moonshot's platform rates at K3 launch (July 2026): // $3/$15 per MTok with $0.30 cached input, flat across the 1M-token // window. kimi-k3 is the ONLY model the platform serves newer @@ -458,7 +500,14 @@ var providers = []Provider{ // The consumer app's "K3 Swarm Max" mode is not an API SKU, so it // doesn't appear here. Models: []Model{ - {ID: "kimi-k3", Label: "Kimi K3", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000}, + // Carries both cache shapes: Moonshot reports cache hits + // OpenAI-style on /v1/chat/completions (CachedInputPer1k) + // and Anthropic-style on /anthropic/v1/messages + // (CacheReadPer1k) — $0.30/MTok either way. Each surface's + // cost formula reads only its own field, so the superset + // entry prices both endpoints correctly. No cache-creation + // rate published; writes bill at the input rate. + {ID: "kimi-k3", Label: "Kimi K3", InputPer1k: 0.003, OutputPer1k: 0.015, CachedInputPer1k: 0.0003, CacheReadPer1k: 0.0003, ContextWindow: 1000000}, }, }, { @@ -758,13 +807,28 @@ func IsBedrockPathStyle(providerID string) bool { func (p Provider) ToAPIResponse() api.AgentNetworkCatalogProvider { models := make([]api.AgentNetworkCatalogModel, 0, len(p.Models)) for _, m := range p.Models { - models = append(models, api.AgentNetworkCatalogModel{ + am := api.AgentNetworkCatalogModel{ Id: m.ID, Label: m.Label, InputPer1k: m.InputPer1k, OutputPer1k: m.OutputPer1k, ContextWindow: m.ContextWindow, - }) + } + // Cache rates are emitted only when configured so the dashboard + // can prefill them; 0 stays off the wire (absent = no rate). + if m.CachedInputPer1k > 0 { + v := m.CachedInputPer1k + am.CachedInputPer1k = &v + } + if m.CacheReadPer1k > 0 { + v := m.CacheReadPer1k + am.CacheReadPer1k = &v + } + if m.CacheCreationPer1k > 0 { + v := m.CacheCreationPer1k + am.CacheCreationPer1k = &v + } + models = append(models, am) } kind := api.AgentNetworkCatalogProviderKindProvider switch p.Kind { @@ -784,6 +848,10 @@ func (p Provider) ToAPIResponse() api.AgentNetworkCatalogProvider { BrandColor: p.BrandColor, Models: models, } + if len(p.PricingSurfaces) > 0 { + surfaces := append([]string(nil), p.PricingSurfaces...) + resp.PricingSurfaces = &surfaces + } if len(p.ExtraHeaders) > 0 { extras := make([]api.AgentNetworkCatalogExtraHeader, 0, len(p.ExtraHeaders)) for _, h := range p.ExtraHeaders { diff --git a/management/internals/modules/agentnetwork/handlers/providers_handler.go b/management/internals/modules/agentnetwork/handlers/providers_handler.go index 13da137d5..c05363101 100644 --- a/management/internals/modules/agentnetwork/handlers/providers_handler.go +++ b/management/internals/modules/agentnetwork/handlers/providers_handler.go @@ -7,6 +7,7 @@ package handlers import ( "encoding/json" + "math" "net/http" "net/url" "strings" @@ -15,6 +16,7 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog" + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" nbcontext "github.com/netbirdio/netbird/management/server/context" "github.com/netbirdio/netbird/shared/management/http/api" @@ -52,11 +54,45 @@ func (h *handler) getCatalogProviders(w http.ResponseWriter, r *http.Request) { entries := catalog.All() out := make([]api.AgentNetworkCatalogProvider, 0, len(entries)) for _, e := range entries { - out = append(out, e.ToAPIResponse()) + resp := e.ToAPIResponse() + applyDefaultPricing(e, &resp) + out = append(out, resp) } util.WriteJSONObject(r.Context(), w, out) } +// applyDefaultPricing overwrites the catalog response's model rates with +// the LIVE default pricing table, which may differ from the compiled-in +// catalog rates when the operator provides a defaults_llm_pricing.yaml. +// This keeps the dashboard's model-row prefill identical to what the +// proxy will actually bill — the same table the synthesizer ships. +func applyDefaultPricing(cp catalog.Provider, resp *api.AgentNetworkCatalogProvider) { + if len(cp.PricingSurfaces) == 0 { + return + } + for i := range resp.Models { + m := &resp.Models[i] + e, ok := pricing.LookupDefault(cp.PricingSurfaces, m.Id) + if !ok { + continue + } + m.InputPer1k = e.InputPer1k + m.OutputPer1k = e.OutputPer1k + m.CachedInputPer1k = positiveRatePtr(e.CachedInputPer1k) + m.CacheReadPer1k = positiveRatePtr(e.CacheReadPer1k) + m.CacheCreationPer1k = positiveRatePtr(e.CacheCreationPer1k) + } +} + +// positiveRatePtr renders a cache rate for the API: absent (nil) when +// unset, matching the catalog response convention. +func positiveRatePtr(v float64) *float64 { + if v <= 0 { + return nil + } + return &v +} + func (h *handler) getAllProviders(w http.ResponseWriter, r *http.Request) { userAuth, err := nbcontext.GetUserAuthFromContext(r.Context()) if err != nil { @@ -213,5 +249,38 @@ func validate(req *api.AgentNetworkProviderRequest, requireAPIKey bool) error { if requireAPIKey && (req.ApiKey == nil || strings.TrimSpace(*req.ApiKey) == "") { return status.Errorf(status.InvalidArgument, "api_key is required") } + if req.Models != nil { + for i, m := range *req.Models { + if err := validateModel(i, m); err != nil { + return err + } + } + } + return nil +} + +// validateModel is the single ingress guard for operator-entered pricing: +// these rates are synthesized into the proxy's cost_meter config verbatim, +// and a negative or non-finite rate there would poison every cost the +// proxy records, so reject at the API boundary. +func validateModel(i int, m api.AgentNetworkProviderModel) error { + if strings.TrimSpace(m.Id) == "" { + return status.Errorf(status.InvalidArgument, "models[%d]: id is required", i) + } + rates := map[string]*float64{ + "input_per_1k": &m.InputPer1k, + "output_per_1k": &m.OutputPer1k, + "cached_input_per_1k": m.CachedInputPer1k, + "cache_read_per_1k": m.CacheReadPer1k, + "cache_creation_per_1k": m.CacheCreationPer1k, + } + for field, v := range rates { + if v == nil { + continue + } + if *v < 0 || math.IsNaN(*v) || math.IsInf(*v, 0) { + return status.Errorf(status.InvalidArgument, "models[%d] (%s): %s must be a finite, non-negative USD rate", i, m.Id, field) + } + } return nil } diff --git a/management/internals/modules/agentnetwork/handlers/providers_handler_test.go b/management/internals/modules/agentnetwork/handlers/providers_handler_test.go new file mode 100644 index 000000000..649224c02 --- /dev/null +++ b/management/internals/modules/agentnetwork/handlers/providers_handler_test.go @@ -0,0 +1,53 @@ +package handlers + +import ( + "math" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/management/http/api" +) + +func f(v float64) *float64 { return &v } + +// TestValidate_ModelRates guards the single ingress point for operator-entered +// pricing. These rates flow verbatim into the proxy's cost_meter config at +// synthesis time; the proxy treats a bad rate as a chain-build failure, so +// rejecting here is what keeps an account's gateway from going down. +func TestValidate_ModelRates(t *testing.T) { + base := func(models ...api.AgentNetworkProviderModel) *api.AgentNetworkProviderRequest { + key := "sk-test" + return &api.AgentNetworkProviderRequest{ + ProviderId: "openai_api", + Name: "OpenAI", + UpstreamUrl: "https://api.openai.com", + ApiKey: &key, + Models: &models, + } + } + + valid := api.AgentNetworkProviderModel{ + Id: "gpt-4o", InputPer1k: 0.0025, OutputPer1k: 0.01, + CachedInputPer1k: f(0.00125), + } + require.NoError(t, validate(base(valid), true), "finite non-negative rates must pass") + + zeroRates := api.AgentNetworkProviderModel{Id: "self-hosted-llama", InputPer1k: 0, OutputPer1k: 0} + require.NoError(t, validate(base(zeroRates), true), "explicit zero prices are allowed (free / self-hosted models)") + + cases := map[string]api.AgentNetworkProviderModel{ + "empty id": {Id: " ", InputPer1k: 0.001, OutputPer1k: 0.002}, + "negative input": {Id: "m", InputPer1k: -0.001, OutputPer1k: 0.002}, + "negative output": {Id: "m", InputPer1k: 0.001, OutputPer1k: -0.002}, + "NaN input": {Id: "m", InputPer1k: math.NaN(), OutputPer1k: 0.002}, + "Inf output": {Id: "m", InputPer1k: 0.001, OutputPer1k: math.Inf(1)}, + "negative cached": {Id: "m", InputPer1k: 0.001, OutputPer1k: 0.002, CachedInputPer1k: f(-1)}, + "NaN cache read": {Id: "m", InputPer1k: 0.001, OutputPer1k: 0.002, CacheReadPer1k: f(math.NaN())}, + "Inf cache creation": {Id: "m", InputPer1k: 0.001, OutputPer1k: 0.002, CacheCreationPer1k: f(math.Inf(-1))}, + } + for name, m := range cases { + assert.Error(t, validate(base(m), true), "case %q must be rejected", name) + } +} diff --git a/management/internals/modules/agentnetwork/pricing/defaults.go b/management/internals/modules/agentnetwork/pricing/defaults.go new file mode 100644 index 000000000..c690313bc --- /dev/null +++ b/management/internals/modules/agentnetwork/pricing/defaults.go @@ -0,0 +1,156 @@ +// Package pricing builds the default LLM pricing table the synthesizer +// ships to the proxy's cost_meter middleware. The catalog is the single +// source of default rates: every catalog provider's models are folded +// into the pricing surfaces the provider declares (PricingSurfaces), +// then a small supplemental list adds priced-but-not-operator-selectable +// entries. Management is the sole pricing authority — the proxy carries +// no embedded price list and bills exclusively from the table it is +// sent. +//go:generate go run gen.go + +package pricing + +import ( + "sync" + "sync/atomic" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog" +) + +// Entry is a single model's pricing in USD per 1k tokens. This struct IS +// the wire shape: the synthesizer marshals it verbatim into cost_meter's +// ConfigJSON, and the proxy unmarshals the same field names. +// +// A zero rate means "no rate configured" — the proxy bills that cache +// bucket at InputPer1k (identical semantics to the retired proxy-embedded +// table). CachedInputPer1k is the OpenAI shape (cached prompt tokens are +// a subset of input); CacheReadPer1k / CacheCreationPer1k are the +// Anthropic shape (additive buckets). +type Entry struct { + InputPer1k float64 `json:"input_per_1k"` + OutputPer1k float64 `json:"output_per_1k"` + CachedInputPer1k float64 `json:"cached_input_per_1k,omitempty"` + CacheReadPer1k float64 `json:"cache_read_per_1k,omitempty"` + CacheCreationPer1k float64 `json:"cache_creation_per_1k,omitempty"` +} + +// supplementalDefaults are (surface, model) entries that are priced but +// deliberately not operator-selectable in the catalog. Each carries a +// reason; when one of these models joins a catalog lineup, delete the +// row here — the collision test fails loudly if the rates ever disagree. +var supplementalDefaults = map[string]map[string]Entry{ + "openai": { + // GPT-5 (2025) family — kept for gateway requests using the + // unsuffixed ids; the dashboard offers only the 5.x lineup. + "gpt-5": {InputPer1k: 0.00125, OutputPer1k: 0.01, CachedInputPer1k: 0.000125}, + "gpt-5-mini": {InputPer1k: 0.00025, OutputPer1k: 0.002, CachedInputPer1k: 0.000025}, + "gpt-5-nano": {InputPer1k: 0.00005, OutputPer1k: 0.0004, CachedInputPer1k: 0.000005}, + }, + "anthropic": { + // claude-opus-5 is not yet in the catalog lineup but gateway / + // grandfathered traffic uses it; priced so it isn't skipped. + "claude-opus-5": {InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625}, + // "kimi-k3[1m]" is the 1M-context alias some Claude Code guides + // configure against Moonshot's Anthropic-compatible endpoint; + // priced identically to kimi-k3 so those requests aren't skipped. + "kimi-k3[1m]": {InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003}, + }, + "bedrock": { + "anthropic.claude-opus-5": {InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625}, + }, +} + +var ( + compiledOnce sync.Once + compiledTable map[string]map[string]Entry + // mergedTable holds the current live table when a pricing defaults + // file is loaded: the file merged entry-whole over the compiled-in + // base. Nil while no file is loaded (or after the file is removed), + // in which case the compiled-in table serves. Swapped atomically by + // the file loader/reloader; readers never block. + mergedTable atomic.Pointer[map[string]map[string]Entry] +) + +// DefaultTable returns the current default pricing table keyed +// surface -> model -> Entry: the management-side defaults file (see +// LoadFile / StartReloader) when one is loaded, merged over the +// compiled-in catalog table, which alone serves as the fallback when no +// file exists. The snapshot may change between calls as the file is +// re-read — consumers (the synthesizer on every reconcile, the catalog +// endpoint on every request) pick up fresh rates automatically. Callers +// must not mutate the returned maps. +func DefaultTable() map[string]map[string]Entry { + if t := mergedTable.Load(); t != nil { + return *t + } + return compiledBase() +} + +// compiledBase returns the compiled-in table (catalog + supplementals), +// built once. +func compiledBase() map[string]map[string]Entry { + compiledOnce.Do(func() { + compiledTable = buildDefaultTable() + }) + return compiledTable +} + +func buildDefaultTable() map[string]map[string]Entry { + out := make(map[string]map[string]Entry) + for _, p := range catalog.All() { + for _, surface := range p.PricingSurfaces { + inner, ok := out[surface] + if !ok { + inner = make(map[string]Entry) + out[surface] = inner + } + for _, m := range p.Models { + // First writer wins; providers contributing the same + // (surface, model) must agree on rates — enforced by + // TestDefaultTable_NoConflictingContributions. + if _, dup := inner[m.ID]; dup { + continue + } + inner[m.ID] = entryFromCatalogModel(m) + } + } + } + for surface, models := range supplementalDefaults { + inner, ok := out[surface] + if !ok { + inner = make(map[string]Entry) + out[surface] = inner + } + for id, e := range models { + if _, dup := inner[id]; dup { + continue + } + inner[id] = e + } + } + return out +} + +func entryFromCatalogModel(m catalog.Model) Entry { + return Entry{ + InputPer1k: m.InputPer1k, + OutputPer1k: m.OutputPer1k, + CachedInputPer1k: m.CachedInputPer1k, + CacheReadPer1k: m.CacheReadPer1k, + CacheCreationPer1k: m.CacheCreationPer1k, + } +} + +// LookupDefault returns the default entry for model on the first of the +// given surfaces that prices it. Used by the synthesizer to seed a +// per-provider entry with default cache rates before overlaying the +// operator's stored prices. +func LookupDefault(surfaces []string, model string) (Entry, bool) { + table := DefaultTable() + for _, s := range surfaces { + if e, ok := table[s][model]; ok { + return e, true + } + } + return Entry{}, false +} diff --git a/proxy/internal/llm/pricing/defaults_pricing.yaml b/management/internals/modules/agentnetwork/pricing/defaults_llm_pricing.example.yaml similarity index 63% rename from proxy/internal/llm/pricing/defaults_pricing.yaml rename to management/internals/modules/agentnetwork/pricing/defaults_llm_pricing.example.yaml index 988426105..bb1cb09a8 100644 --- a/proxy/internal/llm/pricing/defaults_pricing.yaml +++ b/management/internals/modules/agentnetwork/pricing/defaults_llm_pricing.example.yaml @@ -1,83 +1,176 @@ -# Embedded default pricing for llm_observability. Compiled into the proxy -# binary via go:embed in pricing.go; cost annotation works out of the box -# without any operator action. +# Default LLM pricing used by NetBird's Agent Network cost metering. +# GENERATED from the management catalog — do not edit this file in the +# repository; regenerate with: # -# Operators override entries by dropping a pricing.yaml into --plugin-data-dir -# (or whichever basename is given via params.pricing_path). The override file -# only needs entries the operator wants to change; missing entries fall -# through to these defaults. +# go generate ./management/internals/modules/agentnetwork/pricing # -# Values are USD per 1_000 tokens. Public list prices drift; ship a fresh -# binary or override individual entries via the override file as needed. +# Operators: copy this file to /defaults_llm_pricing.yaml (or +# any path configured via management.json: # -# Optional cache fields: -# cached_input_per_1k OpenAI: rate for prompt_tokens_details.cached_tokens -# (a SUBSET of prompt_tokens). Typically 0.5x input. -# Absent → cached portion bills at input_per_1k. -# cache_read_per_1k Anthropic: rate for cache_read_input_tokens -# (ADDITIVE to input_tokens). Typically 0.1x input. -# Absent → cache reads bill at input_per_1k. -# cache_creation_per_1k Anthropic: rate for cache_creation_input_tokens -# (ADDITIVE to input_tokens). Typically 1.25x input. -# Absent → cache writes bill at input_per_1k. +# { "AgentNetwork": { "PricingDefaultsFile": "/path/defaults_llm_pricing.yaml" } } +# +# ) and adjust the entries you want to change. Management re-reads the +# file periodically (mtime poll, every minute): the live table feeds the +# proxies' cost metering and the dashboard's model-price prefill, so +# edits apply without a restart. Your file only needs the entries you +# want to change — but each entry REPLACES the built-in entry for that +# surface+model whole, so repeat the cache rates you want to keep. +# Unknown fields and negative or non-finite rates are rejected: at +# startup that fails boot (for an explicitly configured path); at +# runtime the previous table is kept and a warning is logged. Deleting +# the file reverts to the built-in defaults below. +# +# Top-level keys are pricing surfaces — the parser shape requests are +# metered under: "openai" (also Azure, Mistral, and OpenAI-compatible +# gateways), "anthropic" (also Anthropic-on-Vertex), "bedrock" +# (normalized ids, e.g. anthropic.claude-sonnet-4-5). Model keys must be +# the normalized id the proxy meters (version/region suffixes stripped). +# +# Values are USD per 1_000 tokens. Optional cache fields: +# cached_input_per_1k OpenAI shape: rate for cached prompt tokens +# (a SUBSET of input tokens). Absent -> cached +# portion bills at input_per_1k. +# cache_read_per_1k Anthropic shape: rate for cache_read tokens +# (ADDITIVE to input). Absent -> input rate. +# cache_creation_per_1k Anthropic shape: rate for cache_creation +# tokens (ADDITIVE to input). Absent -> input +# rate. + +anthropic: + claude-fable-5: + input_per_1k: 0.01 + output_per_1k: 0.05 + cache_read_per_1k: 0.001 + cache_creation_per_1k: 0.0125 + claude-haiku-4-5: + input_per_1k: 0.001 + output_per_1k: 0.005 + cache_read_per_1k: 0.0001 + cache_creation_per_1k: 0.00125 + claude-opus-4-1: + input_per_1k: 0.015 + output_per_1k: 0.075 + cache_read_per_1k: 0.0015 + cache_creation_per_1k: 0.01875 + claude-opus-4-6: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + claude-opus-4-7: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + claude-opus-4-8: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + claude-opus-5: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + claude-sonnet-4-5: + input_per_1k: 0.003 + output_per_1k: 0.015 + cache_read_per_1k: 0.0003 + cache_creation_per_1k: 0.00375 + claude-sonnet-4-6: + input_per_1k: 0.003 + output_per_1k: 0.015 + cache_read_per_1k: 0.0003 + cache_creation_per_1k: 0.00375 + kimi-k3: + input_per_1k: 0.003 + output_per_1k: 0.015 + cached_input_per_1k: 0.0003 + cache_read_per_1k: 0.0003 + "kimi-k3[1m]": + input_per_1k: 0.003 + output_per_1k: 0.015 + cache_read_per_1k: 0.0003 + +bedrock: + amazon.nova-2-lite: + input_per_1k: 0.0003 + output_per_1k: 0.0025 + amazon.nova-lite: + input_per_1k: 0.00006 + output_per_1k: 0.00024 + amazon.nova-micro: + input_per_1k: 0.000035 + output_per_1k: 0.00014 + amazon.nova-pro: + input_per_1k: 0.0008 + output_per_1k: 0.0032 + anthropic.claude-haiku-4-5: + input_per_1k: 0.001 + output_per_1k: 0.005 + cache_read_per_1k: 0.0001 + cache_creation_per_1k: 0.00125 + anthropic.claude-opus-4-1: + input_per_1k: 0.015 + output_per_1k: 0.075 + cache_read_per_1k: 0.0015 + cache_creation_per_1k: 0.01875 + anthropic.claude-opus-4-6: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + anthropic.claude-opus-4-7: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + anthropic.claude-opus-4-8: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + anthropic.claude-opus-5: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + anthropic.claude-sonnet-4-5: + input_per_1k: 0.003 + output_per_1k: 0.015 + cache_read_per_1k: 0.0003 + cache_creation_per_1k: 0.00375 + anthropic.claude-sonnet-4-6: + input_per_1k: 0.003 + output_per_1k: 0.015 + cache_read_per_1k: 0.0003 + cache_creation_per_1k: 0.00375 + meta.llama3-3-70b-instruct: + input_per_1k: 0.00072 + output_per_1k: 0.00072 openai: - # OpenAI + OpenAI-compatible providers (openai_api, azure_openai_api, - # mistral_api, and the openai-parser gateways) all emit llm.provider="openai", - # so their models are priced here. Kept in sync with the management catalog; - # rates cross-checked against LiteLLM model_prices_and_context_window.json. - - # GPT-5.x family — cache reads 10% of input (0.1x). - gpt-5.5: - input_per_1k: 0.005 - output_per_1k: 0.03 - cached_input_per_1k: 0.0005 - gpt-5.5-pro: - input_per_1k: 0.03 - output_per_1k: 0.18 - cached_input_per_1k: 0.003 - gpt-5.4: - input_per_1k: 0.0025 - output_per_1k: 0.015 - cached_input_per_1k: 0.00025 - gpt-5.4-pro: - input_per_1k: 0.03 - output_per_1k: 0.18 - cached_input_per_1k: 0.003 - gpt-5.4-mini: - input_per_1k: 0.00075 - output_per_1k: 0.0045 - cached_input_per_1k: 0.000075 - gpt-5.4-nano: - input_per_1k: 0.0002 - output_per_1k: 0.00125 - cached_input_per_1k: 0.00002 - gpt-5.3-codex: - input_per_1k: 0.00175 - output_per_1k: 0.014 - cached_input_per_1k: 0.000175 - gpt-5.3-chat-latest: - input_per_1k: 0.00175 - output_per_1k: 0.014 - cached_input_per_1k: 0.000175 - # GPT-5 (2025) family — kept for gateway requests using the unsuffixed ids. - gpt-5: - input_per_1k: 0.00125 - output_per_1k: 0.01 - cached_input_per_1k: 0.000125 - gpt-5-mini: - input_per_1k: 0.00025 + codestral-2508: + input_per_1k: 0.0003 + output_per_1k: 0.0009 + codestral-latest: + input_per_1k: 0.001 + output_per_1k: 0.003 + devstral-medium-latest: + input_per_1k: 0.0004 output_per_1k: 0.002 - cached_input_per_1k: 0.000025 - gpt-5-nano: - input_per_1k: 0.00005 - output_per_1k: 0.0004 - cached_input_per_1k: 0.000005 - o4-mini: - input_per_1k: 0.0011 - output_per_1k: 0.0044 - cached_input_per_1k: 0.000275 - # GPT-4.1 family — cache reads 25% of input. + devstral-small-latest: + input_per_1k: 0.0001 + output_per_1k: 0.0003 + gpt-3.5-turbo: + input_per_1k: 0.0005 + output_per_1k: 0.0015 + gpt-35-turbo: + input_per_1k: 0.0005 + output_per_1k: 0.0015 + gpt-4-turbo: + input_per_1k: 0.01 + output_per_1k: 0.03 gpt-4.1: input_per_1k: 0.002 output_per_1k: 0.008 @@ -90,7 +183,6 @@ openai: input_per_1k: 0.0001 output_per_1k: 0.0004 cached_input_per_1k: 0.000025 - # GPT-4o family — cache reads 50% of input (0.5x). gpt-4o: input_per_1k: 0.0025 output_per_1k: 0.01 @@ -99,200 +191,92 @@ openai: input_per_1k: 0.00015 output_per_1k: 0.0006 cached_input_per_1k: 0.000075 - # Older GPT — no prompt caching. - gpt-4-turbo: - input_per_1k: 0.01 - output_per_1k: 0.03 - gpt-3.5-turbo: - input_per_1k: 0.0005 - output_per_1k: 0.0015 - gpt-35-turbo: # Azure deployment alias of gpt-3.5-turbo - input_per_1k: 0.0005 - output_per_1k: 0.0015 - # Embeddings — no caching, no output tokens. - text-embedding-3-large: - input_per_1k: 0.00013 - output_per_1k: 0 - text-embedding-3-small: - input_per_1k: 0.00002 - output_per_1k: 0 - - # Mistral (mistral_api) — routed via the openai parser; no prompt caching. - mistral-large-latest: - input_per_1k: 0.0005 - output_per_1k: 0.0015 - mistral-medium-latest: - input_per_1k: 0.0004 + gpt-5: + input_per_1k: 0.00125 + output_per_1k: 0.01 + cached_input_per_1k: 0.000125 + gpt-5-mini: + input_per_1k: 0.00025 output_per_1k: 0.002 - mistral-medium-3-5: - input_per_1k: 0.0015 - output_per_1k: 0.0075 - mistral-small-latest: - input_per_1k: 0.00006 - output_per_1k: 0.00018 + cached_input_per_1k: 0.000025 + gpt-5-nano: + input_per_1k: 0.00005 + output_per_1k: 0.0004 + cached_input_per_1k: 0.000005 + gpt-5.3-chat-latest: + input_per_1k: 0.00175 + output_per_1k: 0.014 + cached_input_per_1k: 0.000175 + gpt-5.3-codex: + input_per_1k: 0.00175 + output_per_1k: 0.014 + cached_input_per_1k: 0.000175 + gpt-5.4: + input_per_1k: 0.0025 + output_per_1k: 0.015 + cached_input_per_1k: 0.00025 + gpt-5.4-mini: + input_per_1k: 0.00075 + output_per_1k: 0.0045 + cached_input_per_1k: 0.000075 + gpt-5.4-nano: + input_per_1k: 0.0002 + output_per_1k: 0.00125 + cached_input_per_1k: 0.00002 + gpt-5.4-pro: + input_per_1k: 0.03 + output_per_1k: 0.18 + cached_input_per_1k: 0.003 + gpt-5.5: + input_per_1k: 0.005 + output_per_1k: 0.03 + cached_input_per_1k: 0.0005 + gpt-5.5-pro: + input_per_1k: 0.03 + output_per_1k: 0.18 + cached_input_per_1k: 0.003 + kimi-k3: + input_per_1k: 0.003 + output_per_1k: 0.015 + cached_input_per_1k: 0.0003 + cache_read_per_1k: 0.0003 magistral-medium-latest: input_per_1k: 0.002 output_per_1k: 0.005 magistral-small-latest: input_per_1k: 0.0005 output_per_1k: 0.0015 - devstral-medium-latest: - input_per_1k: 0.0004 - output_per_1k: 0.002 - devstral-small-latest: - input_per_1k: 0.0001 - output_per_1k: 0.0003 - codestral-2508: - input_per_1k: 0.0003 - output_per_1k: 0.0009 - codestral-latest: - input_per_1k: 0.001 - output_per_1k: 0.003 ministral-3-14b-2512: input_per_1k: 0.0002 output_per_1k: 0.0002 - ministral-8b-latest: - input_per_1k: 0.00015 - output_per_1k: 0.00015 ministral-3-3b-2512: input_per_1k: 0.0001 output_per_1k: 0.0001 + ministral-8b-latest: + input_per_1k: 0.00015 + output_per_1k: 0.00015 mistral-embed: input_per_1k: 0.0001 output_per_1k: 0 - - # Kimi / Moonshot AI (kimi_api) — OpenAI-compatible /v1 endpoint. Moonshot - # reports cache hits OpenAI-style when present; cached input is 10% of - # input ($0.30 vs $3.00 per MTok). kimi-k3 is the only model the platform - # serves newer accounts (K2-era ids and kimi-latest 404), matching the - # management catalog. - kimi-k3: - input_per_1k: 0.003 - output_per_1k: 0.015 - cached_input_per_1k: 0.0003 - -anthropic: - # Claude 4.x family — cache reads ≈10% of input, cache writes ≈125% of input. - # Pricing source: Anthropic's current published rates per million tokens, - # divided by 1000 for the per-1k figures stored here. - claude-fable-5: - input_per_1k: 0.010 - output_per_1k: 0.050 - cache_read_per_1k: 0.001 - cache_creation_per_1k: 0.0125 - claude-opus-5: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - claude-opus-4-8: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - claude-opus-4-7: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - claude-opus-4-6: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - claude-opus-4-1: - input_per_1k: 0.015 - output_per_1k: 0.075 - cache_read_per_1k: 0.0015 - cache_creation_per_1k: 0.01875 - claude-sonnet-4-6: - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - cache_creation_per_1k: 0.00375 - claude-sonnet-4-5: - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - cache_creation_per_1k: 0.00375 - claude-haiku-4-5: - input_per_1k: 0.001 - output_per_1k: 0.005 - cache_read_per_1k: 0.0001 - cache_creation_per_1k: 0.00125 - - # Kimi / Moonshot AI (kimi_api) via the Anthropic-compatible endpoint - # (/anthropic/v1/messages — the official Claude Code setup). Same rates - # as the OpenAI-shape entry above. "kimi-k3[1m]" is the model id some - # Claude Code guides set for the 1M-context alias; priced identically so - # cost metering doesn't silently skip those requests. - kimi-k3: - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - "kimi-k3[1m]": - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - -bedrock: - # AWS Bedrock model ids, normalised by the request parser (cross-region - # inference-profile prefix + version/throughput suffix stripped), e.g. - # eu.anthropic.claude-sonnet-4-5-20250929-v1:0 -> anthropic.claude-sonnet-4-5. - # Anthropic-on-Bedrock keeps the additive cache buckets (read ≈0.1x input, - # write ≈1.25x input); Nova / Llama report no cache, so cost is input+output. - anthropic.claude-opus-5: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - anthropic.claude-opus-4-8: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - anthropic.claude-opus-4-7: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - anthropic.claude-opus-4-6: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - anthropic.claude-opus-4-1: - input_per_1k: 0.015 - output_per_1k: 0.075 - cache_read_per_1k: 0.0015 - cache_creation_per_1k: 0.01875 - anthropic.claude-sonnet-4-6: - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - cache_creation_per_1k: 0.00375 - anthropic.claude-sonnet-4-5: - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - cache_creation_per_1k: 0.00375 - anthropic.claude-haiku-4-5: - input_per_1k: 0.001 - output_per_1k: 0.005 - cache_read_per_1k: 0.0001 - cache_creation_per_1k: 0.00125 - meta.llama3-3-70b-instruct: - input_per_1k: 0.00072 - output_per_1k: 0.00072 - amazon.nova-2-lite: - input_per_1k: 0.0003 - output_per_1k: 0.0025 - amazon.nova-pro: - input_per_1k: 0.0008 - output_per_1k: 0.0032 - amazon.nova-lite: + mistral-large-latest: + input_per_1k: 0.0005 + output_per_1k: 0.0015 + mistral-medium-3-5: + input_per_1k: 0.0015 + output_per_1k: 0.0075 + mistral-medium-latest: + input_per_1k: 0.0004 + output_per_1k: 0.002 + mistral-small-latest: input_per_1k: 0.00006 - output_per_1k: 0.00024 - amazon.nova-micro: - input_per_1k: 0.000035 - output_per_1k: 0.00014 + output_per_1k: 0.00018 + o4-mini: + input_per_1k: 0.0011 + output_per_1k: 0.0044 + cached_input_per_1k: 0.000275 + text-embedding-3-large: + input_per_1k: 0.00013 + output_per_1k: 0 + text-embedding-3-small: + input_per_1k: 0.00002 + output_per_1k: 0 diff --git a/management/internals/modules/agentnetwork/pricing/defaults_test.go b/management/internals/modules/agentnetwork/pricing/defaults_test.go new file mode 100644 index 000000000..99c965687 --- /dev/null +++ b/management/internals/modules/agentnetwork/pricing/defaults_test.go @@ -0,0 +1,148 @@ +package pricing + +import ( + "math" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog" +) + +// TestDefaultTable_CoversEveryCatalogModel replaces the proxy's old +// hand-maintained coverage list: because the table is built FROM the +// catalog, drift is impossible by construction — this test guards the +// fold itself (every catalog model of every surfaced provider resolves, +// with exactly the catalog's rates). +func TestDefaultTable_CoversEveryCatalogModel(t *testing.T) { + table := DefaultTable() + for _, p := range catalog.All() { + if len(p.PricingSurfaces) == 0 { + assert.Empty(t, p.Models, "catalog entry %s declares models but no pricing surfaces — those models would never be priced", p.ID) + continue + } + for _, surface := range p.PricingSurfaces { + byModel, ok := table[surface] + require.True(t, ok, "surface %q (provider %s) missing from default table", surface, p.ID) + for _, m := range p.Models { + e, ok := byModel[m.ID] + require.True(t, ok, "%s/%s (provider %s) missing from default table", surface, m.ID, p.ID) + assert.Equal(t, m.InputPer1k, e.InputPer1k, "%s/%s input rate", surface, m.ID) + assert.Equal(t, m.OutputPer1k, e.OutputPer1k, "%s/%s output rate", surface, m.ID) + } + } + } +} + +// TestDefaultTable_NoConflictingContributions enforces the collision rule +// documented on catalog.Provider.PricingSurfaces: when two catalog +// providers contribute the same (surface, model) pair — azure/vertex +// mirroring openai/anthropic, kimi on both surfaces — their rates must be +// identical, because the surface-keyed table can only hold one entry. +// If a provider ever diverges (e.g. Azure reprices a model), this fails +// and the divergence must move to per-provider-record pricing. +func TestDefaultTable_NoConflictingContributions(t *testing.T) { + type contribution struct { + providerID string + entry Entry + } + seen := map[string]map[string]contribution{} + for _, p := range catalog.All() { + for _, surface := range p.PricingSurfaces { + if seen[surface] == nil { + seen[surface] = map[string]contribution{} + } + for _, m := range p.Models { + e := entryFromCatalogModel(m) + if prev, dup := seen[surface][m.ID]; dup { + assert.Equal(t, prev.entry, e, + "%s/%s: %s and %s contribute different rates", surface, m.ID, prev.providerID, p.ID) + continue + } + seen[surface][m.ID] = contribution{providerID: p.ID, entry: e} + } + } + } + // Supplemental entries must never shadow a catalog-contributed model — + // they exist precisely because the catalog does NOT list them. + for surface, models := range supplementalDefaults { + for id := range models { + _, fromCatalog := seen[surface][id] + assert.False(t, fromCatalog, "supplemental %s/%s is now in the catalog — delete the supplemental row", surface, id) + } + } +} + +// TestDefaultTable_AllRatesFiniteNonNegative mirrors the proxy-side +// NewTable validation so a bad catalog edit is caught here, at unit-test +// time, rather than as a chain-build failure on every proxy. +func TestDefaultTable_AllRatesFiniteNonNegative(t *testing.T) { + for surface, models := range DefaultTable() { + for id, e := range models { + for field, v := range map[string]float64{ + "input": e.InputPer1k, + "output": e.OutputPer1k, + "cached_input": e.CachedInputPer1k, + "cache_read": e.CacheReadPer1k, + "cache_creation": e.CacheCreationPer1k, + } { + assert.False(t, v < 0 || math.IsNaN(v) || math.IsInf(v, 0), + "%s/%s: %s rate %v must be finite and non-negative", surface, id, field, v) + } + } + } +} + +// TestDefaultTable_PinnedRates pins rates that previously drifted or are +// easy to mis-enter (carried over from the proxy's retired +// defaults_coverage_test), plus the supplemental entries. +func TestDefaultTable_PinnedRates(t *testing.T) { + table := DefaultTable() + + gpt54 := table["openai"]["gpt-5.4"] + assert.InDelta(t, 0.0025, gpt54.InputPer1k, 1e-9, "gpt-5.4 input") + assert.InDelta(t, 0.015, gpt54.OutputPer1k, 1e-9, "gpt-5.4 output") + assert.InDelta(t, 0.00025, gpt54.CachedInputPer1k, 1e-9, "gpt-5.4 cached input") + + sonnet := table["bedrock"]["anthropic.claude-sonnet-4-5"] + assert.InDelta(t, 0.003, sonnet.InputPer1k, 1e-9, "bedrock sonnet-4-5 input") + assert.InDelta(t, 0.015, sonnet.OutputPer1k, 1e-9, "bedrock sonnet-4-5 output") + assert.InDelta(t, 0.0003, sonnet.CacheReadPer1k, 1e-9, "bedrock sonnet-4-5 cache read") + assert.InDelta(t, 0.00375, sonnet.CacheCreationPer1k, 1e-9, "bedrock sonnet-4-5 cache creation") + + // Vertex Claude prices under "anthropic" with the bare id. + fable := table["anthropic"]["claude-fable-5"] + assert.InDelta(t, 0.010, fable.InputPer1k, 1e-9, "claude-fable-5 input") + assert.InDelta(t, 0.0125, fable.CacheCreationPer1k, 1e-9, "claude-fable-5 cache creation") + + // Supplementals present on their surfaces. + for surface, ids := range map[string][]string{ + "openai": {"gpt-5", "gpt-5-mini", "gpt-5-nano"}, + "anthropic": {"claude-opus-5", "kimi-k3[1m]", "kimi-k3"}, + "bedrock": {"anthropic.claude-opus-5"}, + } { + for _, id := range ids { + _, ok := table[surface][id] + assert.True(t, ok, "%s/%s must be priced", surface, id) + } + } + + // Embeddings bill input-only — output stays zero. + emb := table["openai"]["text-embedding-3-large"] + assert.Zero(t, emb.OutputPer1k, "embedding output rate must be zero") + assert.Positive(t, emb.InputPer1k, "embedding input rate must be set") +} + +func TestLookupDefault_SurfaceOrder(t *testing.T) { + // kimi-k3 exists on both surfaces; first surface in the slice wins. + e, ok := LookupDefault([]string{"openai", "anthropic"}, "kimi-k3") + require.True(t, ok) + assert.InDelta(t, 0.003, e.InputPer1k, 1e-9) + + _, ok = LookupDefault([]string{"bedrock"}, "gpt-4o") + assert.False(t, ok, "gpt-4o is not a bedrock model") + + _, ok = LookupDefault(nil, "gpt-4o") + assert.False(t, ok, "no surfaces, no match") +} diff --git a/management/internals/modules/agentnetwork/pricing/exampleyaml.go b/management/internals/modules/agentnetwork/pricing/exampleyaml.go new file mode 100644 index 000000000..6217928d6 --- /dev/null +++ b/management/internals/modules/agentnetwork/pricing/exampleyaml.go @@ -0,0 +1,109 @@ +package pricing + +import ( + "bytes" + "fmt" + "sort" + "strconv" +) + +// MarshalDefaultsYAML renders the built-in default pricing table (catalog +// + supplementals, WITHOUT any operator override) as the YAML schema +// LoadOverrideFile consumes. It backs the generated +// defaults_llm_pricing.example.yaml so operators start from a file that +// matches the compiled-in rates exactly; a golden test keeps the two in +// sync. Output is deterministic (sorted surfaces and models). +func MarshalDefaultsYAML() []byte { + var b bytes.Buffer + b.WriteString(exampleHeader) + + table := buildDefaultTable() + surfaces := make([]string, 0, len(table)) + for s := range table { + surfaces = append(surfaces, s) + } + sort.Strings(surfaces) + + for _, surface := range surfaces { + fmt.Fprintf(&b, "\n%s:\n", surface) + models := table[surface] + ids := make([]string, 0, len(models)) + for id := range models { + ids = append(ids, id) + } + sort.Strings(ids) + for _, id := range ids { + e := models[id] + fmt.Fprintf(&b, " %s:\n", yamlKey(id)) + fmt.Fprintf(&b, " input_per_1k: %s\n", rate(e.InputPer1k)) + fmt.Fprintf(&b, " output_per_1k: %s\n", rate(e.OutputPer1k)) + if e.CachedInputPer1k > 0 { + fmt.Fprintf(&b, " cached_input_per_1k: %s\n", rate(e.CachedInputPer1k)) + } + if e.CacheReadPer1k > 0 { + fmt.Fprintf(&b, " cache_read_per_1k: %s\n", rate(e.CacheReadPer1k)) + } + if e.CacheCreationPer1k > 0 { + fmt.Fprintf(&b, " cache_creation_per_1k: %s\n", rate(e.CacheCreationPer1k)) + } + } + } + return b.Bytes() +} + +// rate renders a USD-per-1k rate without float noise ("0.00015", not +// "0.000150000000..."). +func rate(v float64) string { + return strconv.FormatFloat(v, 'f', -1, 64) +} + +// yamlKey quotes model ids that YAML would otherwise misparse (e.g. +// "kimi-k3[1m]" starts a flow sequence unquoted). +func yamlKey(id string) string { + for _, r := range id { + switch r { + case '[', ']', '{', '}', ':', '#', ',', '&', '*', '!', '|', '>', '\'', '"', '%', '@', '`': + return strconv.Quote(id) + } + } + return id +} + +const exampleHeader = `# Default LLM pricing used by NetBird's Agent Network cost metering. +# GENERATED from the management catalog — do not edit this file in the +# repository; regenerate with: +# +# go generate ./management/internals/modules/agentnetwork/pricing +# +# Operators: copy this file to /defaults_llm_pricing.yaml (or +# any path configured via management.json: +# +# { "AgentNetwork": { "PricingDefaultsFile": "/path/defaults_llm_pricing.yaml" } } +# +# ) and adjust the entries you want to change. Management re-reads the +# file periodically (mtime poll, every minute): the live table feeds the +# proxies' cost metering and the dashboard's model-price prefill, so +# edits apply without a restart. Your file only needs the entries you +# want to change — but each entry REPLACES the built-in entry for that +# surface+model whole, so repeat the cache rates you want to keep. +# Unknown fields and negative or non-finite rates are rejected: at +# startup that fails boot (for an explicitly configured path); at +# runtime the previous table is kept and a warning is logged. Deleting +# the file reverts to the built-in defaults below. +# +# Top-level keys are pricing surfaces — the parser shape requests are +# metered under: "openai" (also Azure, Mistral, and OpenAI-compatible +# gateways), "anthropic" (also Anthropic-on-Vertex), "bedrock" +# (normalized ids, e.g. anthropic.claude-sonnet-4-5). Model keys must be +# the normalized id the proxy meters (version/region suffixes stripped). +# +# Values are USD per 1_000 tokens. Optional cache fields: +# cached_input_per_1k OpenAI shape: rate for cached prompt tokens +# (a SUBSET of input tokens). Absent -> cached +# portion bills at input_per_1k. +# cache_read_per_1k Anthropic shape: rate for cache_read tokens +# (ADDITIVE to input). Absent -> input rate. +# cache_creation_per_1k Anthropic shape: rate for cache_creation +# tokens (ADDITIVE to input). Absent -> input +# rate. +` diff --git a/management/internals/modules/agentnetwork/pricing/gen.go b/management/internals/modules/agentnetwork/pricing/gen.go new file mode 100644 index 000000000..c799ab87c --- /dev/null +++ b/management/internals/modules/agentnetwork/pricing/gen.go @@ -0,0 +1,20 @@ +//go:build ignore + +// Regenerates defaults_llm_pricing.example.yaml from the compiled-in +// default pricing table. Run via: +// +// go generate ./management/internals/modules/agentnetwork/pricing +package main + +import ( + "log" + "os" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" +) + +func main() { + if err := os.WriteFile("defaults_llm_pricing.example.yaml", pricing.MarshalDefaultsYAML(), 0o644); err != nil { + log.Fatalf("write defaults_llm_pricing.example.yaml: %v", err) + } +} diff --git a/management/internals/modules/agentnetwork/pricing/override.go b/management/internals/modules/agentnetwork/pricing/override.go new file mode 100644 index 000000000..07677dbaa --- /dev/null +++ b/management/internals/modules/agentnetwork/pricing/override.go @@ -0,0 +1,249 @@ +package pricing + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "io/fs" + "math" + "os" + "sync" + "time" + + log "github.com/sirupsen/logrus" + "gopkg.in/yaml.v3" +) + +// DefaultFileName is the basename probed under management's datadir when +// AgentNetwork.PricingDefaultsFile doesn't configure an explicit path. +const DefaultFileName = "defaults_llm_pricing.yaml" + +// ReloadInterval is the cadence at which the pricing file's mtime is +// polled for changes. +const ReloadInterval = time.Minute + +// maxFileBytes bounds the pricing file read so a misconfigured path +// (pointed at a huge file) cannot exhaust process memory. +const maxFileBytes = 4 << 20 + +// pricingFile mirrors the on-disk YAML schema — the same schema the +// proxy's retired embedded defaults_pricing.yaml used, so files written +// for it keep working. Keys are pricing surfaces ("openai", "anthropic", +// "bedrock"); nested keys are normalized model ids. +type pricingFile map[string]map[string]struct { + InputPer1k float64 `yaml:"input_per_1k"` + OutputPer1k float64 `yaml:"output_per_1k"` + CachedInputPer1k float64 `yaml:"cached_input_per_1k"` + CacheReadPer1k float64 `yaml:"cache_read_per_1k"` + CacheCreationPer1k float64 `yaml:"cache_creation_per_1k"` +} + +// fileState tracks the watched pricing file across reloads. +var fileState struct { + mu sync.Mutex + path string + mtime int64 +} + +// LoadFile loads the management-side pricing defaults file and makes it +// the live table (merged entry-whole over the compiled-in fallback; see +// DefaultTable). The path stays registered for the periodic reloader, so +// later edits — or the file (re)appearing after deletion — are picked up +// without a restart. +// +// required governs the missing-file case: true for an explicitly +// configured path (a typo must fail startup rather than silently bill +// with built-ins the operator believes they replaced), false for the +// conventional /defaults_llm_pricing.yaml probe (absent file = +// compiled-in defaults, still watched in case it appears). A file that +// exists but is malformed is always an error at load time. +func LoadFile(path string, required bool) error { + if path == "" { + return nil + } + fileState.mu.Lock() + fileState.path = path + fileState.mu.Unlock() + + table, mtime, err := readFile(path) + if err != nil { + if errors.Is(err, fs.ErrNotExist) && !required { + log.Infof("agent-network pricing defaults file %s not present; serving built-in defaults", path) + return nil + } + return err + } + storeFileTable(table, mtime) + log.Infof("agent-network pricing defaults loaded from %s", path) + return nil +} + +// StartReloader launches the periodic mtime-poll goroutine for the file +// registered by LoadFile. Runtime failures are lenient — a parse error +// keeps the previously loaded table and logs a warning; a deleted file +// reverts to the compiled-in defaults — so a mid-edit save can never +// take pricing down. Returns immediately when no path was registered. +func StartReloader(ctx context.Context, interval time.Duration) { + fileState.mu.Lock() + path := fileState.path + fileState.mu.Unlock() + if path == "" { + return + } + if interval <= 0 { + interval = ReloadInterval + } + go func() { + t := time.NewTicker(interval) + defer t.Stop() + for { + select { + case <-ctx.Done(): + return + case <-t.C: + reload() + } + } + }() +} + +// reload performs one mtime check + reload cycle. +func reload() { + fileState.mu.Lock() + path, lastMtime := fileState.path, fileState.mtime + fileState.mu.Unlock() + + log.Debugf("agent-network pricing defaults reload: checking %s for changes", path) + + st, err := os.Stat(path) + if err != nil { + if errors.Is(err, fs.ErrNotExist) { + // File removed (or not yet created): serve compiled-in + // defaults and reset mtime so a future (re)appearance loads. + if mergedTable.Swap(nil) != nil { + log.Warnf("agent-network pricing defaults file %s removed; reverting to built-in defaults", path) + } + setMtime(0) + return + } + log.Warnf("agent-network pricing defaults reload: stat %s: %v", path, err) + return + } + if st.ModTime().UnixNano() == lastMtime { + log.Debugf("agent-network pricing defaults %s unchanged since last check", path) + return + } + + table, mtime, err := readFile(path) + if err != nil { + // Keep the previously loaded table — never blank prices because + // an operator saved mid-edit. + log.Warnf("agent-network pricing defaults reload failed for %s (keeping previous table): %v", path, err) + return + } + storeFileTable(table, mtime) + log.Infof("agent-network pricing defaults reloaded from %s", path) +} + +func readFile(path string) (map[string]map[string]Entry, int64, error) { + f, err := os.Open(path) + if err != nil { + return nil, 0, fmt.Errorf("open pricing defaults %s: %w", path, err) + } + defer func() { _ = f.Close() }() + + st, err := f.Stat() + if err != nil { + return nil, 0, fmt.Errorf("stat pricing defaults %s: %w", path, err) + } + data, err := io.ReadAll(io.LimitReader(f, maxFileBytes+1)) + if err != nil { + return nil, 0, fmt.Errorf("read pricing defaults %s: %w", path, err) + } + if len(data) > maxFileBytes { + return nil, 0, fmt.Errorf("pricing defaults %s exceeds %d bytes", path, maxFileBytes) + } + table, err := parsePricingYAML(data) + if err != nil { + return nil, 0, fmt.Errorf("parse pricing defaults %s: %w", path, err) + } + return table, st.ModTime().UnixNano(), nil +} + +// storeFileTable merges the parsed file over the compiled-in base and +// publishes the result as the live table. File entries replace the +// built-in entry for the same (surface, model) whole — they are not +// field-merged — and surfaces/models the file doesn't mention keep the +// built-in rates, so a partial file only needs the entries it changes. +func storeFileTable(table map[string]map[string]Entry, mtime int64) { + base := compiledBase() + merged := make(map[string]map[string]Entry, len(base)+len(table)) + for surface, models := range base { + inner := make(map[string]Entry, len(models)) + for id, e := range models { + inner[id] = e + } + merged[surface] = inner + } + for surface, models := range table { + inner, ok := merged[surface] + if !ok { + inner = make(map[string]Entry, len(models)) + merged[surface] = inner + } + for id, e := range models { + inner[id] = e + } + } + mergedTable.Store(&merged) + setMtime(mtime) +} + +func setMtime(v int64) { + fileState.mu.Lock() + fileState.mtime = v + fileState.mu.Unlock() +} + +// parsePricingYAML decodes and validates the pricing YAML. Unknown +// fields are rejected (typos surface instead of silently pricing at 0) +// and every rate must be a finite, non-negative USD amount — the same +// constraints the HTTP API enforces on operator per-provider prices. +func parsePricingYAML(data []byte) (map[string]map[string]Entry, error) { + dec := yaml.NewDecoder(bytes.NewReader(data)) + dec.KnownFields(true) + + var raw pricingFile + if err := dec.Decode(&raw); err != nil && !errors.Is(err, io.EOF) { + return nil, fmt.Errorf("decode yaml: %w", err) + } + + out := make(map[string]map[string]Entry, len(raw)) + for surface, models := range raw { + inner := make(map[string]Entry, len(models)) + for model, e := range models { + for field, v := range map[string]float64{ + "input_per_1k": e.InputPer1k, + "output_per_1k": e.OutputPer1k, + "cached_input_per_1k": e.CachedInputPer1k, + "cache_read_per_1k": e.CacheReadPer1k, + "cache_creation_per_1k": e.CacheCreationPer1k, + } { + if v < 0 || math.IsNaN(v) || math.IsInf(v, 0) { + return nil, fmt.Errorf("%s/%s: %s must be a finite, non-negative rate, got %v", surface, model, field, v) + } + } + inner[model] = Entry{ + InputPer1k: e.InputPer1k, + OutputPer1k: e.OutputPer1k, + CachedInputPer1k: e.CachedInputPer1k, + CacheReadPer1k: e.CacheReadPer1k, + CacheCreationPer1k: e.CacheCreationPer1k, + } + } + out[surface] = inner + } + return out, nil +} diff --git a/management/internals/modules/agentnetwork/pricing/override_test.go b/management/internals/modules/agentnetwork/pricing/override_test.go new file mode 100644 index 000000000..86031fc58 --- /dev/null +++ b/management/internals/modules/agentnetwork/pricing/override_test.go @@ -0,0 +1,173 @@ +package pricing + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// resetFileState snapshots and restores the package-level file state so +// tests stay order-independent. +func resetFileState(t *testing.T) { + t.Helper() + prevMerged := mergedTable.Load() + fileState.mu.Lock() + prevPath, prevMtime := fileState.path, fileState.mtime + fileState.mu.Unlock() + t.Cleanup(func() { + mergedTable.Store(prevMerged) + fileState.mu.Lock() + fileState.path, fileState.mtime = prevPath, prevMtime + fileState.mu.Unlock() + }) +} + +func writePricing(t *testing.T, path, yml string) { + t.Helper() + require.NoError(t, os.WriteFile(path, []byte(yml), 0o600)) +} + +func TestLoadFile_MergesOverCompiledDefaults(t *testing.T) { + resetFileState(t) + path := filepath.Join(t.TempDir(), DefaultFileName) + writePricing(t, path, ` +openai: + # Reprice a built-in model. The entry replaces the built-in WHOLE: + # omitting the cache rate here drops the built-in 0.00125 discount. + gpt-4o: + input_per_1k: 0.9 + output_per_1k: 1.8 + # A model NetBird doesn't know at all. + my-private-ft: + input_per_1k: 0.01 + output_per_1k: 0.02 + cached_input_per_1k: 0.005 +gemini: + gemini-pro: + input_per_1k: 0.00125 + output_per_1k: 0.005 +`) + require.NoError(t, LoadFile(path, true)) + table := DefaultTable() + + gpt4o := table["openai"]["gpt-4o"] + assert.InDelta(t, 0.9, gpt4o.InputPer1k, 1e-9, "file rate replaces the compiled-in rate") + assert.Zero(t, gpt4o.CachedInputPer1k, "entries replace whole — omitted cache rate is dropped, not inherited") + + ft := table["openai"]["my-private-ft"] + assert.InDelta(t, 0.005, ft.CachedInputPer1k, 1e-9, "unknown models are added to the surface") + _, ok := table["gemini"]["gemini-pro"] + assert.True(t, ok, "a surface the catalog doesn't declare can be added") + + // Untouched entries keep compiled-in rates (catalog, other surface, + // supplemental). + assert.InDelta(t, 0.00015, table["openai"]["gpt-4o-mini"].InputPer1k, 1e-9, "unlisted model keeps compiled rate") + assert.InDelta(t, 0.003, table["anthropic"]["claude-sonnet-4-5"].InputPer1k, 1e-9, "unlisted surface untouched") + assert.InDelta(t, 0.00125, table["openai"]["gpt-5"].InputPer1k, 1e-9, "supplemental entries untouched") + + // The synthesizer-facing lookup reads the live table too. + e, ok := LookupDefault([]string{"openai"}, "gpt-4o") + require.True(t, ok) + assert.InDelta(t, 0.9, e.InputPer1k, 1e-9, "LookupDefault serves the file-backed rate") +} + +func TestLoadFile_MissingPath(t *testing.T) { + resetFileState(t) + missing := filepath.Join(t.TempDir(), DefaultFileName) + + require.Error(t, LoadFile(missing, true), + "explicitly configured path that doesn't exist must fail startup") + + require.NoError(t, LoadFile(missing, false), + "conventional datadir probe tolerates an absent file (compiled-in defaults serve)") + assert.Nil(t, mergedTable.Load(), "no file, no merged table") + fileState.mu.Lock() + path := fileState.path + fileState.mu.Unlock() + assert.Equal(t, missing, path, "the path stays registered so the reloader picks the file up when it appears") +} + +func TestLoadFile_RejectsInvalid(t *testing.T) { + resetFileState(t) + dir := t.TempDir() + cases := map[string]string{ + "unknown field (typo)": "openai:\n gpt-4o:\n input_per1k: 0.1\n", + "negative rate": "openai:\n gpt-4o:\n input_per_1k: -0.1\n", + "non-numeric rate": "openai:\n gpt-4o:\n input_per_1k: cheap\n", + "not a mapping": "- just\n- a\n- list\n", + } + for name, yml := range cases { + path := filepath.Join(dir, DefaultFileName) + writePricing(t, path, yml) + assert.Error(t, LoadFile(path, true), "case %q must be rejected", name) + } +} + +// TestReload_LifeCycle drives the reloader's single-shot reload through +// its full lifecycle: file edit picked up on mtime change, a broken save +// keeps the previous table, and file removal reverts to the compiled-in +// defaults (then a re-created file loads again). +func TestReload_LifeCycle(t *testing.T) { + resetFileState(t) + path := filepath.Join(t.TempDir(), DefaultFileName) + writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: 0.5\n output_per_1k: 1\n") + require.NoError(t, LoadFile(path, true)) + require.InDelta(t, 0.5, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9) + + // Edit: new mtime, new rates. + writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: 0.7\n output_per_1k: 1.4\n") + bumpMtime(t, path) + reload() + assert.InDelta(t, 0.7, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9, "edit must be picked up") + + // Broken save: previous table survives. + writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: -1\n") + bumpMtime(t, path) + reload() + assert.InDelta(t, 0.7, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9, + "a malformed save must keep the previously loaded table, never blank prices") + + // Removal: compiled-in defaults serve again. + require.NoError(t, os.Remove(path)) + reload() + assert.InDelta(t, 0.0025, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9, + "file removal reverts to the compiled-in rate") + + // Re-created file loads without a restart. + writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: 0.9\n output_per_1k: 1.8\n") + bumpMtime(t, path) + reload() + assert.InDelta(t, 0.9, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9, + "a file appearing after removal (or after a missing-probe boot) must load") +} + +// bumpMtime pushes the file's mtime forward past the previously recorded +// value — timestamps can otherwise collide within the test's timescale. +func bumpMtime(t *testing.T, path string) { + t.Helper() + st, err := os.Stat(path) + require.NoError(t, err) + next := st.ModTime().Add(2 * 1e9) + require.NoError(t, os.Chtimes(path, next, next)) +} + +// TestExampleYAML_InSyncWithBuiltins is the golden guard for +// defaults_llm_pricing.example.yaml: the shipped example must stay +// byte-identical to what the compiled-in table renders (catalog edits +// require `go generate ./management/internals/modules/agentnetwork/pricing`) +// and must round-trip through the same parser operators' files go +// through, reproducing the compiled-in table exactly. +func TestExampleYAML_InSyncWithBuiltins(t *testing.T) { + onDisk, err := os.ReadFile("defaults_llm_pricing.example.yaml") + require.NoError(t, err, "example file must exist next to the package") + require.Equal(t, string(MarshalDefaultsYAML()), string(onDisk), + "defaults_llm_pricing.example.yaml is stale — run: go generate ./management/internals/modules/agentnetwork/pricing") + + parsed, err := parsePricingYAML(onDisk) + require.NoError(t, err, "the example must be a valid pricing defaults file") + assert.Equal(t, buildDefaultTable(), parsed, + "parsing the example must reproduce the compiled-in table exactly") +} diff --git a/management/internals/modules/agentnetwork/synthesizer.go b/management/internals/modules/agentnetwork/synthesizer.go index 169bdd4fd..64711387b 100644 --- a/management/internals/modules/agentnetwork/synthesizer.go +++ b/management/internals/modules/agentnetwork/synthesizer.go @@ -233,6 +233,11 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([ return nil, err } + costMeterJSON, err := buildCostMeterConfigJSON(enabledProviders, groupIndex) + if err != nil { + return nil, err + } + mergedGuardrails := mergeGuardrails(enabledPolicies, guardrailsByID) applyAccountCollectionControls(&mergedGuardrails, settings) // The proxy guardrail is a per-provider fail-closed backstop; the @@ -248,7 +253,7 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([ // Use the merged decision (account settings OR policy-required redaction), // not the raw account flag, so a policy that mandates PII redaction is // honored by the capture parsers even when the account toggle is off. - middlewares := buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, mergedGuardrails.PromptCapture.RedactPii, mergedGuardrails.PromptCapture.Enabled) + middlewares := buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, costMeterJSON, mergedGuardrails.PromptCapture.RedactPii, mergedGuardrails.PromptCapture.Enabled) priv, pub, err := pickServiceSessionKeys(enabledProviders) if err != nil { @@ -700,7 +705,7 @@ func buildIdentityExtraHeaders(p *types.Provider, extras []catalog.ExtraHeader) // requests bound for gateways like LiteLLM that key budgets and // attribution off request headers. CanMutate is required so its // HeadersAdd / HeadersRemove pass the framework's mutation gate. -func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON []byte, redactPii, capturePromptContent bool) []rpservice.MiddlewareConfig { +func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, costMeterJSON []byte, redactPii, capturePromptContent bool) []rpservice.MiddlewareConfig { // Both parsers receive an explicit capture flag derived from the account's // enable_prompt_collection toggle; nil/unset would default to the legacy // "always emit" behavior in the middleware, which is precisely what we @@ -769,10 +774,13 @@ func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON []byt ConfigJSON: []byte("{}"), }, { + // Carries the full pricing table (defaults + per-provider + // operator prices) so the proxy bills without an embedded + // price list; see buildCostMeterConfigJSON. ID: middlewareIDCostMeter, Enabled: true, Slot: rpservice.MiddlewareSlotOnResponse, - ConfigJSON: []byte("{}"), + ConfigJSON: costMeterJSON, }, { ID: middlewareIDLLMResponseParser, diff --git a/management/internals/modules/agentnetwork/synthesizer_pricing.go b/management/internals/modules/agentnetwork/synthesizer_pricing.go new file mode 100644 index 000000000..6d57ffaef --- /dev/null +++ b/management/internals/modules/agentnetwork/synthesizer_pricing.go @@ -0,0 +1,131 @@ +package agentnetwork + +import ( + "encoding/json" + "fmt" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog" + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + sharedllm "github.com/netbirdio/netbird/shared/llm" +) + +// costMeterConfig is the JSON shape the proxy-side cost_meter middleware +// expects (mirror-type pattern, same as routerConfig). The top-level +// "pricing" wrapper is the feature-detection signal: an old proxy's config +// struct ignores it as an unknown field, and a new proxy treats its +// absence as "old management" (skips every cost computation and warns). +type costMeterConfig struct { + Pricing *costMeterPricing `json:"pricing,omitempty"` +} + +// costMeterPricing carries the full pricing table: +// - Defaults: surface ("openai"/"anthropic"/"bedrock") -> normalized +// model id -> rates. The full default table ships to every account — +// it is small (~10 KB) and keeps gateway-style providers (which +// enumerate no models) priced for every catalog model. +// - Providers: provider record id (matched against the +// llm.resolved_provider_id metadata llm_router stamps) -> normalized +// model id -> rates. Entries are fully materialized here at synth +// time — default cache rates already folded in — so the proxy does +// two map lookups and no merging. +type costMeterPricing struct { + Defaults map[string]map[string]pricing.Entry `json:"defaults,omitempty"` + Providers map[string]map[string]pricing.Entry `json:"providers,omitempty"` +} + +// buildCostMeterConfigJSON assembles the cost_meter middleware config +// from the default pricing table plus the operator's stored per-provider +// model prices. Same orphan rule as the router: a provider no enabled +// policy authorises is unreachable, so its prices are not shipped. +// +// Overlay semantics per model row: +// - The row's model id is normalized exactly the way the proxy's +// request parser normalizes the ids it meters (bedrock ARN/region/ +// version stripping, vertex "@version" stripping), so the per-record +// lookup key compares equal to llm.model at billing time. +// - The entry starts from the default entry for that model (when one +// exists) to inherit cache rates the operator didn't state. +// - Operator input/output overlay verbatim — including an explicit 0, +// which prices the model as free (self-hosted / internal endpoints) +// rather than silently reverting to list price. +// - Cache-rate pointers overlay only when non-nil: nil means "inherit +// the default", an explicit 0 means "no discount, bill this bucket +// at the input rate". +func buildCostMeterConfigJSON(providers []*types.Provider, groupIndex map[string][]string) ([]byte, error) { + cfg := costMeterConfig{Pricing: &costMeterPricing{ + Defaults: pricing.DefaultTable(), + }} + + perRecord := make(map[string]map[string]pricing.Entry) + for _, p := range providers { + if _, hasPolicy := groupIndex[p.ID]; !hasPolicy { + // Orphan: unreachable via the router, so unpriceable. + continue + } + if len(p.Models) == 0 { + // Gateway-style "claim every model" provider: the defaults + // table is its price list. + continue + } + entry, _ := catalog.Lookup(p.ProviderID) + models := make(map[string]pricing.Entry, len(p.Models)) + for _, m := range p.Models { + id := normalizePricingModelID(p.ProviderID, m.ID) + if id == "" { + continue + } + if _, dup := models[id]; dup { + // First occurrence wins on post-normalization duplicates, + // matching providerModelIDs' dedup order for routing. + continue + } + models[id] = materializeEntry(entry.PricingSurfaces, id, m) + } + if len(models) > 0 { + perRecord[p.ID] = models + } + } + if len(perRecord) > 0 { + cfg.Pricing.Providers = perRecord + } + + out, err := json.Marshal(cfg) + if err != nil { + return nil, fmt.Errorf("marshal cost_meter middleware config: %w", err) + } + return out, nil +} + +// normalizePricingModelID maps an operator-entered model id onto the +// normalized id the proxy's request parser emits as llm.model — the key +// the cost meter looks up at billing time. +func normalizePricingModelID(catalogProviderID, modelID string) string { + switch { + case catalog.IsBedrockPathStyle(catalogProviderID): + return sharedllm.NormalizeBedrockModel(modelID) + case catalog.IsVertexPathStyle(catalogProviderID): + return sharedllm.NormalizeVertexModel(modelID) + default: + return modelID + } +} + +// materializeEntry folds the default entry for (surfaces, model) — when +// one exists — under the operator's stored prices, producing the fully +// materialized wire entry. +func materializeEntry(surfaces []string, normalizedID string, m types.ProviderModel) pricing.Entry { + e, _ := pricing.LookupDefault(surfaces, normalizedID) // zero Entry on miss + e.InputPer1k = m.InputPer1k + e.OutputPer1k = m.OutputPer1k + if m.CachedInputPer1k != nil { + e.CachedInputPer1k = *m.CachedInputPer1k + } + if m.CacheReadPer1k != nil { + e.CacheReadPer1k = *m.CacheReadPer1k + } + if m.CacheCreationPer1k != nil { + e.CacheCreationPer1k = *m.CacheCreationPer1k + } + return e +} diff --git a/management/internals/modules/agentnetwork/synthesizer_pricing_test.go b/management/internals/modules/agentnetwork/synthesizer_pricing_test.go new file mode 100644 index 000000000..83961878a --- /dev/null +++ b/management/internals/modules/agentnetwork/synthesizer_pricing_test.go @@ -0,0 +1,105 @@ +package agentnetwork + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" +) + +func fptr(v float64) *float64 { return &v } + +func decodeCostMeterConfig(t *testing.T, raw []byte) costMeterConfig { + t.Helper() + var cfg costMeterConfig + require.NoError(t, json.Unmarshal(raw, &cfg), "cost meter config must round-trip") + require.NotNil(t, cfg.Pricing, "pricing wrapper must be present") + return cfg +} + +// TestBuildCostMeterConfig_BedrockModelNormalization: the operator may +// paste region-prefixed, versioned, or ARN-wrapped Bedrock ids; the +// per-record map must be keyed by the normalized id the request parser +// emits as llm.model, or the lookup never hits at billing time. +func TestBuildCostMeterConfig_BedrockModelNormalization(t *testing.T) { + bedrock := &types.Provider{ + ID: "prov-bedrock", + ProviderID: "bedrock_api", + Enabled: true, + Models: []types.ProviderModel{ + {ID: "eu.anthropic.claude-sonnet-4-5-20250929-v1:0", InputPer1k: 0.0033, OutputPer1k: 0.0165}, + // Post-normalization duplicate of the row above under a + // different regional spelling — first occurrence wins. + {ID: "us.anthropic.claude-sonnet-4-5-20250929-v1:0", InputPer1k: 9.9, OutputPer1k: 9.9}, + }, + } + raw, err := buildCostMeterConfigJSON([]*types.Provider{bedrock}, map[string][]string{"prov-bedrock": {"grp"}}) + require.NoError(t, err) + cfg := decodeCostMeterConfig(t, raw) + + models := cfg.Pricing.Providers["prov-bedrock"] + require.Len(t, models, 1, "both spellings normalize to one model; first row wins") + e, ok := models["anthropic.claude-sonnet-4-5"] + require.True(t, ok, "key must be the normalized id the parser emits, not the operator's raw spelling") + assert.InDelta(t, 0.0033, e.InputPer1k, 1e-9, "first row's rate wins the dedup") + assert.InDelta(t, 0.0003, e.CacheReadPer1k, 1e-9, "cache read inherited from the bedrock default entry") + assert.InDelta(t, 0.00375, e.CacheCreationPer1k, 1e-9, "cache creation inherited from the bedrock default entry") +} + +// TestBuildCostMeterConfig_CacheRateNilVsZero pins the pointer semantics: +// nil inherits the default cache rate, explicit 0 clears it (that bucket +// bills at the input rate on the proxy). +func TestBuildCostMeterConfig_CacheRateNilVsZero(t *testing.T) { + p := &types.Provider{ + ID: "prov-oai", + ProviderID: "openai_api", + Enabled: true, + Models: []types.ProviderModel{ + {ID: "gpt-4o", InputPer1k: 0.002, OutputPer1k: 0.008}, // nil → inherit 0.00125 + {ID: "gpt-4o-mini", InputPer1k: 0.0001, OutputPer1k: 0.0005, CachedInputPer1k: fptr(0)}, // explicit 0 → no discount + {ID: "my-custom-ft", InputPer1k: 0.01, OutputPer1k: 0.02, CachedInputPer1k: fptr(0.005)}, // unknown model, explicit rate + }, + } + raw, err := buildCostMeterConfigJSON([]*types.Provider{p}, map[string][]string{"prov-oai": {"grp"}}) + require.NoError(t, err) + cfg := decodeCostMeterConfig(t, raw) + models := cfg.Pricing.Providers["prov-oai"] + + assert.InDelta(t, 0.00125, models["gpt-4o"].CachedInputPer1k, 1e-9, "nil cache pointer inherits the default rate") + assert.Zero(t, models["gpt-4o-mini"].CachedInputPer1k, "explicit 0 overrides the default (0.000075) — bucket bills at input rate") + custom := models["my-custom-ft"] + assert.InDelta(t, 0.005, custom.CachedInputPer1k, 1e-9, "unknown model keeps the operator's explicit cache rate") + assert.Zero(t, custom.CacheReadPer1k, "no default to inherit for a model outside the catalog") +} + +// TestBuildCostMeterConfig_OrphanAndGatewayProviders: an orphan (no +// authorising policy) is unreachable so its prices must not ship; a +// gateway with no model rows relies on the defaults table and gets no +// per-record entry. +func TestBuildCostMeterConfig_OrphanAndGatewayProviders(t *testing.T) { + orphan := &types.Provider{ + ID: "prov-orphan", + ProviderID: "openai_api", + Enabled: true, + Models: []types.ProviderModel{{ID: "gpt-4o", InputPer1k: 1, OutputPer1k: 1}}, + } + gateway := &types.Provider{ + ID: "prov-litellm", + ProviderID: "litellm_proxy", + Enabled: true, + Models: []types.ProviderModel{}, + } + raw, err := buildCostMeterConfigJSON( + []*types.Provider{orphan, gateway}, + map[string][]string{"prov-litellm": {"grp"}}, // orphan has no policy + ) + require.NoError(t, err) + cfg := decodeCostMeterConfig(t, raw) + + assert.NotContains(t, cfg.Pricing.Providers, "prov-orphan", "orphan provider prices must not ship") + assert.NotContains(t, cfg.Pricing.Providers, "prov-litellm", "empty-models gateway needs no per-record entry") + assert.NotEmpty(t, cfg.Pricing.Defaults["openai"], "defaults still ship so the gateway's catalog-model traffic is priced") +} diff --git a/management/internals/modules/agentnetwork/synthesizer_test.go b/management/internals/modules/agentnetwork/synthesizer_test.go index 8a18a9b59..7b14f8209 100644 --- a/management/internals/modules/agentnetwork/synthesizer_test.go +++ b/management/internals/modules/agentnetwork/synthesizer_test.go @@ -33,14 +33,17 @@ func newSynthTestSettings() *types.Settings { func newSynthTestProvider() *types.Provider { return &types.Provider{ - ID: "prov-1", - AccountID: testAccountID, - ProviderID: "openai_api", - Name: "OpenAI", - UpstreamURL: "https://api.openai.com", - APIKey: "sk-test-key", - Enabled: true, - Models: []types.ProviderModel{{ID: "gpt-5.4", InputPer1k: 0.0025, OutputPer1k: 0.015}}, + ID: "prov-1", + AccountID: testAccountID, + ProviderID: "openai_api", + Name: "OpenAI", + UpstreamURL: "https://api.openai.com", + APIKey: "sk-test-key", + Enabled: true, + // Prices deliberately differ from the catalog's gpt-5.4 rates + // (0.0025/0.015) so pricing tests can prove the operator's + // stored price overlays the catalog default. + Models: []types.ProviderModel{{ID: "gpt-5.4", InputPer1k: 0.004, OutputPer1k: 0.02}}, SessionPrivateKey: "test-priv-key", SessionPublicKey: "test-pub-key", CreatedAt: time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC), @@ -214,7 +217,27 @@ func TestSynthesizeServices_HappyPath(t *testing.T) { assert.Equal(t, middlewareIDCostMeter, mws[6].ID, "seventh middleware is the cost meter") assert.Equal(t, rpservice.MiddlewareSlotOnResponse, mws[6].Slot, "cost meter runs on_response") - assert.Equal(t, []byte("{}"), mws[6].ConfigJSON, "cost meter carries an explicit empty config") + + var costCfg costMeterConfig + require.NoError(t, json.Unmarshal(mws[6].ConfigJSON, &costCfg), "cost meter config must unmarshal") + require.NotNil(t, costCfg.Pricing, "cost meter config must carry the pricing table — its absence tells the proxy management predates config-delivered pricing") + + gpt4o, ok := costCfg.Pricing.Defaults["openai"]["gpt-4o"] + require.True(t, ok, "the full default table ships regardless of the account's providers") + assert.InDelta(t, 0.0025, gpt4o.InputPer1k, 1e-9, "default gpt-4o input rate comes from the catalog") + + openaiPrices, ok := costCfg.Pricing.Providers[openai.ID] + require.True(t, ok, "operator-priced provider must have a per-record entry") + gpt54, ok := openaiPrices["gpt-5.4"] + require.True(t, ok, "operator's model row keys the per-record map") + assert.InDelta(t, 0.004, gpt54.InputPer1k, 1e-9, "operator input price overlays the catalog default (0.0025)") + assert.InDelta(t, 0.02, gpt54.OutputPer1k, 1e-9, "operator output price overlays the catalog default (0.015)") + assert.InDelta(t, 0.00025, gpt54.CachedInputPer1k, 1e-9, "cache rate the operator didn't state is inherited from the default entry") + + opus, ok := costCfg.Pricing.Providers[anthropic.ID]["claude-opus-4-7"] + require.True(t, ok, "anthropic's model row keys its per-record map") + assert.Zero(t, opus.InputPer1k, "operator-stored zero prices ship verbatim — an explicit $0 model bills as free, it does not revert to list price") + assert.InDelta(t, 0.0005, opus.CacheReadPer1k, 1e-9, "cache rates still inherit from the default entry") assert.Equal(t, middlewareIDLLMResponseParser, mws[7].ID, "eighth middleware is the response parser") assert.Equal(t, rpservice.MiddlewareSlotOnResponse, mws[7].Slot, "response parser runs on_response") diff --git a/management/internals/modules/agentnetwork/types/provider.go b/management/internals/modules/agentnetwork/types/provider.go index b3287168e..96242f45f 100644 --- a/management/internals/modules/agentnetwork/types/provider.go +++ b/management/internals/modules/agentnetwork/types/provider.go @@ -14,10 +14,24 @@ import ( // ProviderModel is one row in the provider's models list. The operator // pins the per-1k input/output price for cost tracking; ID is the // model identifier the upstream provider expects on the wire. +// +// The three cache rates are pointers because absence is meaningful: nil +// means "inherit NetBird's default rate for this model" (folded in at +// synthesis time), while an explicit 0 means "no discount — bill this +// cache bucket at the input rate". type ProviderModel struct { ID string `json:"id"` InputPer1k float64 `json:"input_per_1k"` OutputPer1k float64 `json:"output_per_1k"` + // CachedInputPer1k is the OpenAI-shape rate for cached prompt tokens + // (a subset of input tokens). + CachedInputPer1k *float64 `json:"cached_input_per_1k,omitempty"` + // CacheReadPer1k is the Anthropic-shape rate for cache-read tokens + // (additive to input tokens). + CacheReadPer1k *float64 `json:"cache_read_per_1k,omitempty"` + // CacheCreationPer1k is the Anthropic-shape rate for cache-creation + // tokens (additive to input tokens). + CacheCreationPer1k *float64 `json:"cache_creation_per_1k,omitempty"` } // Provider is an Agent Network AI provider record persisted per account. @@ -128,9 +142,12 @@ func (p *Provider) FromAPIRequest(req *api.AgentNetworkProviderRequest) { if req.Models != nil { for _, m := range *req.Models { p.Models = append(p.Models, ProviderModel{ - ID: m.Id, - InputPer1k: m.InputPer1k, - OutputPer1k: m.OutputPer1k, + ID: m.Id, + InputPer1k: m.InputPer1k, + OutputPer1k: m.OutputPer1k, + CachedInputPer1k: copyFloatPtr(m.CachedInputPer1k), + CacheReadPer1k: copyFloatPtr(m.CacheReadPer1k), + CacheCreationPer1k: copyFloatPtr(m.CacheCreationPer1k), }) } } @@ -164,9 +181,12 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider { models := make([]api.AgentNetworkProviderModel, 0, len(p.Models)) for _, m := range p.Models { models = append(models, api.AgentNetworkProviderModel{ - Id: m.ID, - InputPer1k: m.InputPer1k, - OutputPer1k: m.OutputPer1k, + Id: m.ID, + InputPer1k: m.InputPer1k, + OutputPer1k: m.OutputPer1k, + CachedInputPer1k: copyFloatPtr(m.CachedInputPer1k), + CacheReadPer1k: copyFloatPtr(m.CacheReadPer1k), + CacheCreationPer1k: copyFloatPtr(m.CacheCreationPer1k), }) } created := p.CreatedAt @@ -201,11 +221,27 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider { return resp } +// copyFloatPtr returns a fresh pointer to the same value, or nil. Keeps +// stored models and API payloads from aliasing each other's rate fields. +func copyFloatPtr(v *float64) *float64 { + if v == nil { + return nil + } + out := *v + return &out +} + // Copy returns a deep copy of the provider. func (p *Provider) Copy() *Provider { clone := *p if p.Models != nil { - clone.Models = append([]ProviderModel(nil), p.Models...) + clone.Models = make([]ProviderModel, len(p.Models)) + for i, m := range p.Models { + m.CachedInputPer1k = copyFloatPtr(m.CachedInputPer1k) + m.CacheReadPer1k = copyFloatPtr(m.CacheReadPer1k) + m.CacheCreationPer1k = copyFloatPtr(m.CacheCreationPer1k) + clone.Models[i] = m + } } if p.ExtraValues != nil { clone.ExtraValues = make(map[string]string, len(p.ExtraValues)) diff --git a/management/internals/modules/agentnetwork/wire_shape_test.go b/management/internals/modules/agentnetwork/wire_shape_test.go index b574ab3e1..779dd77f9 100644 --- a/management/internals/modules/agentnetwork/wire_shape_test.go +++ b/management/internals/modules/agentnetwork/wire_shape_test.go @@ -103,6 +103,12 @@ func TestSynthesizedService_WireShape(t *testing.T) { assert.Equal(t, middlewareIDCostMeter, mws[6].GetId(), "seventh middleware id") assert.Equal(t, proto.MiddlewareSlot_MIDDLEWARE_SLOT_ON_RESPONSE, mws[6].GetSlot(), "cost meter slot") + var costCfg costMeterConfig + require.NoError(t, json.Unmarshal(mws[6].GetConfigJson(), &costCfg), "cost meter config JSON must decode from the wire") + require.NotNil(t, costCfg.Pricing, "the pricing table must travel on the wire — the proxy has no embedded price list to fall back to") + assert.NotEmpty(t, costCfg.Pricing.Defaults["openai"], "default table rides in every mapping") + assert.NotEmpty(t, costCfg.Pricing.Defaults["anthropic"], "default table covers all surfaces") + assert.NotEmpty(t, costCfg.Pricing.Defaults["bedrock"], "default table covers all surfaces") assert.Equal(t, middlewareIDLLMResponseParser, mws[7].GetId(), "eighth middleware id") assert.Equal(t, proto.MiddlewareSlot_MIDDLEWARE_SLOT_ON_RESPONSE, mws[7].GetSlot(), "response parser slot") diff --git a/management/internals/server/config/config.go b/management/internals/server/config/config.go index a77d5c19b..dc60ed822 100644 --- a/management/internals/server/config/config.go +++ b/management/internals/server/config/config.go @@ -55,6 +55,8 @@ type Config struct { ReverseProxy ReverseProxy + AgentNetwork AgentNetwork + // disable default all-to-all policy DisableDefaultPolicy bool @@ -185,6 +187,25 @@ type StoreConfig struct { Engine types.Engine } +// AgentNetwork contains agent-network (LLM gateway) configuration. +type AgentNetwork struct { + // PricingDefaultsFile is the path to the YAML file holding the default + // LLM pricing table (defaults_llm_pricing.yaml). A relative path is + // resolved against , so a bare filename lands alongside the + // store. Empty falls back to probing /defaults_llm_pricing.yaml; + // with no file present the compiled-in defaults serve. Schema: surface ("openai"/"anthropic"/ + // "bedrock") -> model -> rates in USD per 1k tokens (input_per_1k, + // output_per_1k, and the optional cached_input_per_1k / + // cache_read_per_1k / cache_creation_per_1k). File entries replace the + // compiled-in entry for the same surface+model whole; everything else + // keeps the compiled-in rates. The file is re-read periodically (mtime + // poll), and the live table feeds both the synthesizer (what proxies + // bill with) and the dashboard's catalog endpoint (what model rows + // prefill with). An explicitly configured path that fails to load + // fails startup; runtime reload errors keep the previous table. + PricingDefaultsFile string +} + // ReverseProxy contains reverse proxy configuration in front of management. type ReverseProxy struct { // TrustedHTTPProxies represents a list of trusted HTTP proxies by their IP prefixes. diff --git a/proxy/internal/llm/bedrock_model.go b/proxy/internal/llm/bedrock_model.go deleted file mode 100644 index a4c4704f7..000000000 --- a/proxy/internal/llm/bedrock_model.go +++ /dev/null @@ -1,38 +0,0 @@ -package llm - -import ( - "regexp" - "strings" -) - -// bedrockRegionPrefixes are the cross-region inference-profile prefixes that -// front a Bedrock model id (e.g. "eu.anthropic.claude-..."). -var bedrockRegionPrefixes = []string{"us.", "eu.", "apac.", "global."} - -// bedrockVersionSuffix matches the trailing "-vN[:N]" or "-YYYYMMDD-vN[:N]" -// version/throughput suffix of a Bedrock model id. -var bedrockVersionSuffix = regexp.MustCompile(`-(\d{8}-)?v\d+(:\d+)?$`) - -// NormalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile -// prefix, and the version/throughput suffix from a Bedrock model id so it -// matches the catalog/pricing key, e.g. -// "eu.anthropic.claude-sonnet-4-5-20250929-v1:0" -> "anthropic.claude-sonnet-4-5" -// and the inference-profile ARN's last segment likewise. It is the single -// source of truth shared by the request parser (which normalizes the request -// model from the URL path) and the router (which normalizes the operator's -// registered Bedrock model ids so both sides compare equal). -func NormalizeBedrockModel(modelID string) string { - m := modelID - if strings.HasPrefix(m, "arn:") { - if i := strings.LastIndex(m, "/"); i >= 0 { - m = m[i+1:] - } - } - for _, p := range bedrockRegionPrefixes { - if strings.HasPrefix(m, p) { - m = m[len(p):] - break - } - } - return bedrockVersionSuffix.ReplaceAllString(m, "") -} diff --git a/proxy/internal/llm/fixtures/pricing.yaml b/proxy/internal/llm/fixtures/pricing.yaml deleted file mode 100644 index 3d26ff803..000000000 --- a/proxy/internal/llm/fixtures/pricing.yaml +++ /dev/null @@ -1,59 +0,0 @@ -# Realistic-pricing starter for llm_observability. Drop this into the -# directory you point the proxy at via --plugin-data-dir, then reference it -# from the target's plugin config: -# -# plugins: -# - id: llm_observability -# enabled: true -# params: -# pricing_path: pricing.yaml -# -# Values are USD per 1_000 tokens. Public list prices drift; treat this as a -# starting point and keep your production copy current. - -openai: - # GPT-5 family - gpt-5: - input_per_1k: 0.00125 - output_per_1k: 0.01 - gpt-5-mini: - input_per_1k: 0.00025 - output_per_1k: 0.002 - gpt-5-nano: - input_per_1k: 0.00005 - output_per_1k: 0.0004 - gpt-5.4: - input_per_1k: 0.00125 - output_per_1k: 0.01 - # GPT-4o family - gpt-4o: - input_per_1k: 0.0025 - output_per_1k: 0.01 - gpt-4o-mini: - input_per_1k: 0.00015 - output_per_1k: 0.0006 - # Embeddings - text-embedding-3-large: - input_per_1k: 0.00013 - output_per_1k: 0 - text-embedding-3-small: - input_per_1k: 0.00002 - output_per_1k: 0 - -anthropic: - # Claude 4.x family - claude-opus-4-7: - input_per_1k: 0.015 - output_per_1k: 0.075 - claude-sonnet-4-7: - input_per_1k: 0.003 - output_per_1k: 0.015 - claude-sonnet-4-6: - input_per_1k: 0.003 - output_per_1k: 0.015 - claude-sonnet-4-5: - input_per_1k: 0.003 - output_per_1k: 0.015 - claude-haiku-4-5: - input_per_1k: 0.0008 - output_per_1k: 0.004 diff --git a/proxy/internal/llm/model.go b/proxy/internal/llm/model.go new file mode 100644 index 000000000..76ccfeccf --- /dev/null +++ b/proxy/internal/llm/model.go @@ -0,0 +1,21 @@ +package llm + +import ( + sharedllm "github.com/netbirdio/netbird/shared/llm" +) + +// NormalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile +// prefix, and the version/throughput suffix from a Bedrock model id so it +// matches the catalog/pricing key. Thin delegate to the shared implementation +// (shared/llm), which management also uses at synthesis time so both sides of +// the pricing / routing contract normalize identically. +func NormalizeBedrockModel(modelID string) string { + return sharedllm.NormalizeBedrockModel(modelID) +} + +// NormalizeVertexModel strips the "@version" suffix from a Vertex AI model id +// so it matches the catalog/pricing key. Thin delegate to shared/llm, kept +// beside NormalizeBedrockModel for the same contract reason. +func NormalizeVertexModel(modelID string) string { + return sharedllm.NormalizeVertexModel(modelID) +} diff --git a/proxy/internal/llm/pricing/defaults_coverage_test.go b/proxy/internal/llm/pricing/defaults_coverage_test.go deleted file mode 100644 index 8df1557ea..000000000 --- a/proxy/internal/llm/pricing/defaults_coverage_test.go +++ /dev/null @@ -1,65 +0,0 @@ -package pricing - -import ( - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -// TestDefaultTable_FirstPartyModelCoverage guards the embedded defaults against -// silent drift/gaps: every metered first-party model the management catalog -// enumerates must resolve to a price, and a few rates that previously drifted -// are pinned to their LiteLLM-validated values. Keep this list in step with the -// catalog (management/server/agentnetwork/catalog) when adding models. -func TestDefaultTable_FirstPartyModelCoverage(t *testing.T) { - tbl := DefaultTable() - require.NotNil(t, tbl, "embedded default pricing table must load") - - mustPrice := map[string][]string{ - // openai parser covers openai_api, azure_openai_api, and mistral_api. - "openai": { - "gpt-5.5", "gpt-5.5-pro", "gpt-5.4", "gpt-5.4-mini", "gpt-5.4-nano", - "gpt-5.3-codex", "gpt-5.3-chat-latest", "o4-mini", - "gpt-4.1", "gpt-4.1-mini", "gpt-4.1-nano", "gpt-4o", "gpt-4o-mini", - "gpt-4-turbo", "gpt-3.5-turbo", "gpt-35-turbo", - "text-embedding-3-large", "text-embedding-3-small", - "mistral-large-latest", "mistral-medium-3-5", "codestral-2508", - "ministral-8b-latest", "mistral-embed", - }, - "anthropic": { - "claude-fable-5", "claude-opus-5", "claude-opus-4-8", "claude-opus-4-7", "claude-opus-4-6", - "claude-opus-4-1", "claude-sonnet-4-6", "claude-sonnet-4-5", "claude-haiku-4-5", - }, - // bedrock keys are the normalized ids the request parser emits. - "bedrock": { - "anthropic.claude-opus-5", "anthropic.claude-opus-4-8", "anthropic.claude-opus-4-7", "anthropic.claude-opus-4-6", - "anthropic.claude-opus-4-1", "anthropic.claude-sonnet-4-6", "anthropic.claude-sonnet-4-5", - "anthropic.claude-haiku-4-5", "meta.llama3-3-70b-instruct", - "amazon.nova-pro", "amazon.nova-lite", "amazon.nova-micro", "amazon.nova-2-lite", - }, - } - for provider, models := range mustPrice { - for _, m := range models { - _, ok := tbl.Cost(provider, m, 1000, 1000, 0, 0) - assert.True(t, ok, "%s/%s must be priced in the embedded defaults", provider, m) - } - } - - // Pin per-direction rates independently (input-only then output-only) so a - // swap or skew of input<->output that preserves the combined total is still - // caught — these are rates that previously drifted or are easy to mis-enter. - in, ok := tbl.Cost("openai", "gpt-5.4", 1000, 0, 0, 0) - require.True(t, ok) - assert.InDelta(t, 0.0025, in, 1e-9, "gpt-5.4 input = 0.0025 per 1k") - out, ok := tbl.Cost("openai", "gpt-5.4", 0, 1000, 0, 0) - require.True(t, ok) - assert.InDelta(t, 0.015, out, 1e-9, "gpt-5.4 output = 0.015 per 1k") - - in, ok = tbl.Cost("bedrock", "anthropic.claude-sonnet-4-5", 1000, 0, 0, 0) - require.True(t, ok) - assert.InDelta(t, 0.003, in, 1e-9, "bedrock sonnet-4-5 input = 0.003 per 1k") - out, ok = tbl.Cost("bedrock", "anthropic.claude-sonnet-4-5", 0, 1000, 0, 0) - require.True(t, ok) - assert.InDelta(t, 0.015, out, 1e-9, "bedrock sonnet-4-5 output = 0.015 per 1k") -} diff --git a/proxy/internal/llm/pricing/pricing.go b/proxy/internal/llm/pricing/pricing.go index b77000000..ce6e636cf 100644 --- a/proxy/internal/llm/pricing/pricing.go +++ b/proxy/internal/llm/pricing/pricing.go @@ -1,102 +1,30 @@ -// Package pricing implements the embedded-default + override pricing table -// shared by middleware that converts LLM token usage into a USD cost -// estimate. The table is hot-reloadable from a basename under the proxy -// data directory; missing override files keep the embedded defaults so -// cost annotation works without operator action. +// Package pricing implements the pricing table and cost formula the +// cost_meter middleware uses to convert LLM token usage into a USD cost +// estimate. The table's content arrives from the management server inside +// cost_meter's middleware config (synthesized from the catalog plus the +// operator's stored per-provider prices) — the proxy carries no embedded +// price list. Price updates ride the ordinary mapping push: a chain +// rebuild constructs a fresh table, so there is nothing to reload. package pricing import ( - "bytes" - "context" - _ "embed" - "errors" "fmt" - "io" - "io/fs" "math" - "path/filepath" - "regexp" - "strings" - "sync" - "sync/atomic" - "time" - - log "github.com/sirupsen/logrus" - "go.opentelemetry.io/otel/attribute" - "go.opentelemetry.io/otel/metric" - "gopkg.in/yaml.v3" ) -//go:embed defaults_pricing.yaml -var defaultPricingYAML []byte - -var ( - defaultTableOnce sync.Once - defaultTablePtr *Table -) - -// DefaultTable returns the pricing table embedded in the binary. The result -// is parsed once and shared; callers must not mutate the returned value. -// Cost annotation works without any operator action because every loader -// starts with this table. -func DefaultTable() *Table { - defaultTableOnce.Do(func() { - t, err := parsePricingBytes(defaultPricingYAML) - if err != nil { - panic(fmt.Sprintf("llmobs: embedded default pricing failed to parse: %v", err)) - } - defaultTablePtr = t - }) - return defaultTablePtr -} - -// mergeOver returns a new Table containing every entry from base, with any -// matching entry from overlay replacing the base value. Either argument may -// be nil. Result is a fresh allocation so callers can mutate / Store safely. -func mergeOver(base, overlay *Table) *Table { - if overlay == nil || len(overlay.entries) == 0 { - return base - } - if base == nil || len(base.entries) == 0 { - return overlay - } - out := make(map[string]map[string]Entry, len(base.entries)) - for provider, models := range base.entries { - inner := make(map[string]Entry, len(models)) - for model, e := range models { - inner[model] = e - } - out[provider] = inner - } - for provider, models := range overlay.entries { - inner, ok := out[provider] - if !ok { - inner = make(map[string]Entry, len(models)) - out[provider] = inner - } - for model, e := range models { - inner[model] = e - } - } - return &Table{entries: out} -} - // Entry is a single model's input and output pricing, expressed in USD per // 1000 tokens. // // CachedInputPer1K applies to OpenAI's cached prompt tokens, which are a // subset of input_tokens — when set, the cached portion is billed at this // rate and the non-cached remainder at InputPer1K. Zero means "no discount -// configured", and cached tokens are billed at InputPer1K (matches current -// behaviour where cached counts weren't extracted at all). +// configured", and cached tokens are billed at InputPer1K. // // CacheReadPer1K and CacheCreationPer1K apply to Anthropic's two prompt- // cache fields, which are additive to input_tokens: cache_read is the // cheaper read-from-cache rate, cache_creation is the more expensive // write-to-cache rate. Zero means "no rate configured" and the -// corresponding token bucket is billed at InputPer1K. This is more -// accurate than today's behaviour, where Anthropic's cache tokens are -// ignored and not charged at all. +// corresponding token bucket is billed at InputPer1K. type Entry struct { InputPer1K float64 OutputPer1K float64 @@ -105,33 +33,102 @@ type Entry struct { CacheCreationPer1K float64 } -// Table is a provider-to-model pricing lookup. Instances are immutable once -// built and are swapped atomically by Loader. +// EntryJSON is the wire shape of a pricing entry inside cost_meter's +// middleware config. Field names are the management→proxy contract; the +// management synthesizer marshals the same names (its pricing.Entry). +type EntryJSON struct { + InputPer1K float64 `json:"input_per_1k"` + OutputPer1K float64 `json:"output_per_1k"` + CachedInputPer1K float64 `json:"cached_input_per_1k"` + CacheReadPer1K float64 `json:"cache_read_per_1k"` + CacheCreationPer1K float64 `json:"cache_creation_per_1k"` +} + +// Table is a provider-surface-to-model pricing lookup. Instances are +// immutable once built; a mapping update builds a whole new middleware +// instance (and with it a new table) rather than mutating this one. type Table struct { entries map[string]map[string]Entry } +// NewEntries validates and converts a wire-shape map (surface-or-record -> +// model -> rates) into the internal representation. Every rate must be a +// finite, non-negative USD amount; a violation is returned as an error so +// a corrupt config fails the chain build loudly instead of mispricing. +// Management validates the same constraints at its API boundary, so this +// is defense-in-depth. Nil input yields an empty (never-matching) map. +func NewEntries(raw map[string]map[string]EntryJSON) (map[string]map[string]Entry, error) { + out := make(map[string]map[string]Entry, len(raw)) + for outer, models := range raw { + inner := make(map[string]Entry, len(models)) + for model, e := range models { + for field, v := range map[string]float64{ + "input_per_1k": e.InputPer1K, + "output_per_1k": e.OutputPer1K, + "cached_input_per_1k": e.CachedInputPer1K, + "cache_read_per_1k": e.CacheReadPer1K, + "cache_creation_per_1k": e.CacheCreationPer1K, + } { + if v < 0 || math.IsNaN(v) || math.IsInf(v, 0) { + return nil, fmt.Errorf("pricing %s/%s: %s must be a finite, non-negative rate, got %v", outer, model, field, v) + } + } + // EntryJSON and Entry are field-identical (tags aside), so a + // direct conversion carries all five rates. + inner[model] = Entry(e) + } + out[outer] = inner + } + return out, nil +} + +// NewTable builds an immutable Table from the wire-shape defaults map. +// See NewEntries for validation semantics. +func NewTable(raw map[string]map[string]EntryJSON) (*Table, error) { + entries, err := NewEntries(raw) + if err != nil { + return nil, err + } + return &Table{entries: entries}, nil +} + +// Lookup returns the entry for the given provider surface and model. +func (t *Table) Lookup(provider, model string) (Entry, bool) { + if t == nil { + return Entry{}, false + } + byModel, ok := t.entries[provider] + if !ok { + return Entry{}, false + } + e, ok := byModel[model] + return e, ok +} + +// Has reports whether the provider/model pair is present in the table. +func (t *Table) Has(provider, model string) bool { + _, ok := t.Lookup(provider, model) + return ok +} + // Cost returns the estimated USD cost for the given token counts. ok is // false when the provider or model is not present in the table; the caller // can still emit token metrics with a model=unknown label. -// -// Provider-shape semantics for cached / cache-creation counts: -// -// - OpenAI: cachedInput is a SUBSET of inTokens. The cached portion is -// billed at CachedInputPer1K (or InputPer1K when no override), and the -// non-cached remainder of inTokens at InputPer1K. cacheCreation is -// ignored (OpenAI has no analogue). -// - Anthropic: cachedInput (cache_read) and cacheCreation are ADDITIVE to -// inTokens. The three buckets are billed at CacheReadPer1K, -// CacheCreationPer1K, and InputPer1K respectively, each falling back -// to InputPer1K when the corresponding rate is zero. -// - Other providers: cached and cacheCreation are ignored; cost is -// inTokens*InputPer1K + outTokens*OutputPer1K. func (t *Table) Cost(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (float64, bool) { c, ok := t.Costs(provider, model, inTokens, outTokens, cachedInput, cacheCreation) return c.TotalUSD, ok } +// Costs returns the estimated USD cost split for the given token counts. +// The provider surface selects the cache formula; see EntryCosts. +func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (Costs, bool) { + entry, ok := t.Lookup(provider, model) + if !ok { + return Costs{}, false + } + return EntryCosts(entry, provider, inTokens, outTokens, cachedInput, cacheCreation), true +} + // Costs is a per-request cost split. The four per-bucket fields are the base // of the breakdown — one per token bucket the provider bills separately — and // the two aggregates are derived from them: @@ -165,9 +162,25 @@ func newCosts(input, cachedInput, cacheCreation, output float64) Costs { } } -// Costs returns the estimated USD cost split for the given token counts, with -// the same semantics as Cost. -func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (Costs, bool) { +// EntryCosts computes the USD cost split for the given entry and token +// counts. The surface (the llm.provider value the request parser stamps) +// selects the cache formula; the entry may come from the surface-keyed +// defaults table or from a per-provider-record override — the math is +// identical either way. +// +// Provider-shape semantics for cached / cache-creation counts: +// +// - "openai": cachedInput is a SUBSET of inTokens. The cached portion is +// billed at CachedInputPer1K (or InputPer1K when no override), and the +// non-cached remainder of inTokens at InputPer1K. cacheCreation is +// ignored (OpenAI has no analogue). +// - "anthropic", "bedrock": cachedInput (cache_read) and cacheCreation are +// ADDITIVE to inTokens. The three buckets are billed at CacheReadPer1K, +// CacheCreationPer1K, and InputPer1K respectively, each falling back +// to InputPer1K when the corresponding rate is zero. +// - Other surfaces: cached and cacheCreation are ignored; cost is +// inTokens*InputPer1K + outTokens*OutputPer1K. +func EntryCosts(entry Entry, surface string, inTokens, outTokens, cachedInput, cacheCreation int64) Costs { // Clamp negatives to zero before any pricing math so a malformed // upstream count can never produce a negative cost. if inTokens < 0 { @@ -182,19 +195,8 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, if cacheCreation < 0 { cacheCreation = 0 } - if t == nil { - return Costs{}, false - } - byModel, ok := t.entries[provider] - if !ok { - return Costs{}, false - } - entry, ok := byModel[model] - if !ok { - return Costs{}, false - } output := (float64(outTokens) / 1000.0) * entry.OutputPer1K - switch provider { + switch surface { case "openai": // cachedInput is a subset of inTokens; clamp so a malformed // upstream (cached > total) can't produce a negative remainder. @@ -208,7 +210,7 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, } nonCached := float64(inTokens-clamped) / 1000.0 * entry.InputPer1K cached := float64(clamped) / 1000.0 * cachedRate - return newCosts(nonCached, cached, 0, output), true + return newCosts(nonCached, cached, 0, output) case "anthropic", "bedrock": // Bedrock-Anthropic returns the same additive cache buckets as // first-party Anthropic; non-Anthropic Bedrock models simply report @@ -224,266 +226,9 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, input := float64(inTokens) / 1000.0 * entry.InputPer1K read := float64(cachedInput) / 1000.0 * readRate create := float64(cacheCreation) / 1000.0 * createRate - return newCosts(input, read, create, output), true + return newCosts(input, read, create, output) default: input := float64(inTokens) / 1000.0 * entry.InputPer1K - return newCosts(input, 0, 0, output), true + return newCosts(input, 0, 0, output) } } - -// Has reports whether the provider/model pair is present in the table. -func (t *Table) Has(provider, model string) bool { - if t == nil { - return false - } - byModel, ok := t.entries[provider] - if !ok { - return false - } - _, ok = byModel[model] - return ok -} - -// pricingFile mirrors the on-disk YAML schema. Keys are provider names; the -// nested map keys are model names. -type pricingFile map[string]map[string]struct { - InputPer1K float64 `yaml:"input_per_1k"` - OutputPer1K float64 `yaml:"output_per_1k"` - CachedInputPer1K float64 `yaml:"cached_input_per_1k"` - CacheReadPer1K float64 `yaml:"cache_read_per_1k"` - CacheCreationPer1K float64 `yaml:"cache_creation_per_1k"` -} - -const ( - // ReloadInterval is the mtime-poll cadence for the background reloader. - ReloadInterval = 30 * time.Second - - // errorBackoff bounds how often the loader logs a repeated parse error. - errorBackoff = 5 * time.Minute -) - -var basenameRegex = regexp.MustCompile(`^[a-zA-Z0-9._-]+$`) - -// Loader is a confined, hot-reloadable pricing table reader. Construction -// must succeed against the target file; subsequent reload failures keep the -// previously-loaded table so callers never observe a blank price list. -type Loader struct { - baseDir string - fullPath string - pluginID string - table atomic.Pointer[Table] - mtime atomic.Int64 - failures metric.Int64Counter - interval time.Duration -} - -// NewLoader returns a pricing loader that overlays an optional file-based -// table on top of the embedded defaults. Missing override file, baseDir, or -// relPath is not an error: the loader keeps the embedded defaults so cost -// metadata is still emitted for known models. -// -// Errors: -// - bad basename, traversal segment, or absolute relPath are rejected so a -// misconfigured target surfaces immediately. -// - permission errors and YAML parse errors keep the defaults but log a -// warning; cost annotation does not silently break. -// -// failures is optional; pass nil in tests that do not care about -// reload-failure telemetry. -func NewLoader(baseDir, relPath, pluginID string, failures metric.Int64Counter) (*Loader, error) { - defaults := DefaultTable() - l := &Loader{ - baseDir: baseDir, - pluginID: pluginID, - failures: failures, - } - l.table.Store(defaults) - - if strings.TrimSpace(baseDir) == "" || strings.TrimSpace(relPath) == "" { - return l, nil - } - - full, err := resolveMiddlewareDataPath(baseDir, relPath) - if err != nil { - return nil, err - } - l.fullPath = full - - overlay, mtime, err := loadPricing(full) - if err != nil { - if errors.Is(err, fs.ErrNotExist) { - // Override file is optional. Defaults already stored. - return l, nil - } - // Symlink rejection, oversize file, parse failure, permission errors - // — surface so a misconfigured operator sees the problem instead of - // silently running with stale defaults. - return nil, fmt.Errorf("load pricing %s: %w", full, err) - } - l.table.Store(mergeOver(defaults, overlay)) - l.mtime.Store(mtime.UnixNano()) - return l, nil -} - -// Get returns the current pricing table. The returned pointer is immutable; -// callers must not mutate its contents. -func (l *Loader) Get() *Table { - if l == nil { - return nil - } - return l.table.Load() -} - -// WatchesFile reports whether this loader is bound to an override file on -// disk. False for defaults-only loaders (no operator override given). -// Callers use this to decide whether to spawn the mtime-poll goroutine. -func (l *Loader) WatchesFile() bool { - if l == nil { - return false - } - return l.fullPath != "" -} - -// SetReloadInterval overrides the mtime-poll cadence used by Reload. Calls -// after Reload has started have no effect on the running loop. Intended for -// tests; production code uses the default ReloadInterval. -func (l *Loader) SetReloadInterval(d time.Duration) { - if l == nil || d <= 0 { - return - } - l.interval = d -} - -// Reload runs a polling loop that checks the pricing file mtime every -// ReloadInterval (or the value passed to SetReloadInterval). Returns when -// ctx is cancelled. -func (l *Loader) Reload(ctx context.Context) { - if l == nil { - return - } - interval := l.interval - if interval <= 0 { - interval = ReloadInterval - } - t := time.NewTicker(interval) - defer t.Stop() - - var lastErrAt time.Time - for { - select { - case <-ctx.Done(): - return - case <-t.C: - if err := l.reload(); err != nil { - if l.failures != nil { - l.failures.Add(ctx, 1, metric.WithAttributes( - attribute.String("plugin", l.pluginID), - )) - } - now := time.Now() - if now.Sub(lastErrAt) >= errorBackoff { - log.Warnf("llmobs: pricing reload failed for %s: %v", l.fullPath, err) - lastErrAt = now - } - } - } - } -} - -// reload performs a single-shot mtime check and reload. The reloaded -// override file is merged on top of the embedded defaults; missing override -// (e.g. operator deleted the file) is not an error and reverts to defaults. -func (l *Loader) reload() error { - if l.fullPath == "" { - // Defaults-only loader; nothing on disk to reload. - return nil - } - mtime, err := statMtime(l.fullPath) - if err != nil { - if errors.Is(err, fs.ErrNotExist) { - // File was removed since startup. Drop back to defaults and - // reset mtime so a future re-creation triggers a reload. - l.table.Store(DefaultTable()) - l.mtime.Store(0) - return nil - } - return err - } - if mtime.UnixNano() == l.mtime.Load() { - return nil - } - - overlay, newMtime, err := loadPricing(l.fullPath) - if err != nil { - return err - } - l.table.Store(mergeOver(DefaultTable(), overlay)) - l.mtime.Store(newMtime.UnixNano()) - return nil -} - -// resolveMiddlewareDataPath validates relPath is a safe basename and resolves -// it under baseDir. An additional cleaned-prefix check guards against -// CVE-style edge cases where Join is used with trailing path segments. -func resolveMiddlewareDataPath(baseDir, relPath string) (string, error) { - if strings.TrimSpace(baseDir) == "" { - return "", errors.New("middleware-data-dir is not configured") - } - if relPath == "" { - return "", errors.New("pricing path is empty") - } - if !basenameRegex.MatchString(relPath) { - return "", fmt.Errorf("pricing path %q is not a safe basename", relPath) - } - if filepath.IsAbs(relPath) { - return "", fmt.Errorf("pricing path %q must be a basename, not absolute", relPath) - } - - cleanBase, err := filepath.Abs(filepath.Clean(baseDir)) - if err != nil { - return "", fmt.Errorf("resolve middleware-data-dir: %w", err) - } - full := filepath.Join(cleanBase, relPath) - cleanedFull := filepath.Clean(full) - if !strings.HasPrefix(cleanedFull, cleanBase+string(filepath.Separator)) && cleanedFull != cleanBase { - return "", fmt.Errorf("pricing path %q escapes middleware-data-dir", relPath) - } - return cleanedFull, nil -} - -func parsePricingBytes(data []byte) (*Table, error) { - dec := yaml.NewDecoder(bytes.NewReader(data)) - dec.KnownFields(true) - - var raw pricingFile - if err := dec.Decode(&raw); err != nil && !errors.Is(err, io.EOF) { - return nil, fmt.Errorf("decode pricing yaml: %w", err) - } - - out := make(map[string]map[string]Entry, len(raw)) - for provider, models := range raw { - inner := make(map[string]Entry, len(models)) - for model, entry := range models { - for field, v := range map[string]float64{ - "input_per_1k": entry.InputPer1K, - "output_per_1k": entry.OutputPer1K, - "cached_input_per_1k": entry.CachedInputPer1K, - "cache_read_per_1k": entry.CacheReadPer1K, - "cache_creation_per_1k": entry.CacheCreationPer1K, - } { - if v < 0 || math.IsNaN(v) || math.IsInf(v, 0) { - return nil, fmt.Errorf("pricing %s/%s: %s must be a finite, non-negative rate, got %v", provider, model, field, v) - } - } - inner[model] = Entry{ - InputPer1K: entry.InputPer1K, - OutputPer1K: entry.OutputPer1K, - CachedInputPer1K: entry.CachedInputPer1K, - CacheReadPer1K: entry.CacheReadPer1K, - CacheCreationPer1K: entry.CacheCreationPer1K, - } - } - out[provider] = inner - } - return &Table{entries: out}, nil -} diff --git a/proxy/internal/llm/pricing/pricing_other.go b/proxy/internal/llm/pricing/pricing_other.go deleted file mode 100644 index e65fffff1..000000000 --- a/proxy/internal/llm/pricing/pricing_other.go +++ /dev/null @@ -1,20 +0,0 @@ -//go:build !unix - -package pricing - -import ( - "fmt" - "time" -) - -// loadPricing is unavailable on non-Unix platforms because O_NOFOLLOW and -// fstat-from-FD are required to honour the spec's symlink-safety rules. The -// proxy is only deployed on Linux today; a Windows port would need an -// equivalent path-as-handle implementation. -func loadPricing(path string) (*Table, time.Time, error) { - return nil, time.Time{}, fmt.Errorf("llmobs pricing loader is not supported on this platform: %s", path) -} - -func statMtime(path string) (time.Time, error) { - return time.Time{}, fmt.Errorf("llmobs pricing loader is not supported on this platform: %s", path) -} diff --git a/proxy/internal/llm/pricing/pricing_test.go b/proxy/internal/llm/pricing/pricing_test.go index 7ac2a85dc..b946faa7f 100644 --- a/proxy/internal/llm/pricing/pricing_test.go +++ b/proxy/internal/llm/pricing/pricing_test.go @@ -1,47 +1,13 @@ -//go:build unix - package pricing import ( - "context" - "os" - "path/filepath" + "math" "testing" - "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) -func copyFixture(t *testing.T, src, dst string) { - t.Helper() - data, err := os.ReadFile(src) - require.NoError(t, err, "read source fixture") - require.NoError(t, os.WriteFile(dst, data, 0o600), "write target fixture") -} - -func TestNewLoader_HappyPath(t *testing.T) { - base := t.TempDir() - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml")) - - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err, "NewLoader must succeed with a valid fixture") - table := l.Get() - require.NotNil(t, table, "table populated after load") - - cost, ok := table.Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0) - require.True(t, ok, "known provider/model resolves") - assert.InDelta(t, 0.00075, cost, 1e-9, "cost = 0.00015 + 0.0006 per 1k tokens") - - cost, ok = table.Cost("openai", "gpt-4o", 2000, 1000, 0, 0) - require.True(t, ok, "second known model resolves") - assert.InDelta(t, 0.015, cost, 1e-9, "cost for gpt-4o: 2*0.0025 + 1*0.01") - - cost, ok = table.Cost("anthropic", "claude-sonnet-4-5", 1000, 1000, 0, 0) - require.True(t, ok, "anthropic model resolves") - assert.InDelta(t, 0.018, cost, 1e-9, "cost for claude-sonnet-4-5: 0.003 + 0.015") -} - // TestCost_OpenAICachedSubsetDiscount proves OpenAI's cached input // tokens are billed at the configured cached_input_per_1k rate while // the non-cached remainder of input_tokens is billed at the regular @@ -65,11 +31,9 @@ func TestCost_OpenAICachedSubsetDiscount(t *testing.T) { "cached subset must bill at the discount rate; non-cached remainder at regular rate") } -// TestCost_OpenAICachedFallsBackToInputRate covers the operator -// opt-in contract: when CachedInputPer1K is unset (zero), cached -// tokens bill at the regular input rate. This matches today's -// behaviour (cached counts weren't extracted at all so they -// implicitly billed at the input rate via prompt_tokens). +// TestCost_OpenAICachedFallsBackToInputRate covers the fallback +// contract: when CachedInputPer1K is unset (zero), cached tokens bill +// at the regular input rate. func TestCost_OpenAICachedFallsBackToInputRate(t *testing.T) { tbl := &Table{entries: map[string]map[string]Entry{ "openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}}, @@ -78,7 +42,7 @@ func TestCost_OpenAICachedFallsBackToInputRate(t *testing.T) { require.True(t, ok) want := 0.0025 + (500.0/1000.0)*0.01 assert.InDelta(t, want, cost, 1e-12, - "absent cached_input_per_1k rate must fall back to input_per_1k — same as pre-feature behaviour") + "absent cached_input_per_1k rate must fall back to input_per_1k") } // TestCost_OpenAIClampsCachedToInputCount is the defensive guard @@ -100,10 +64,7 @@ func TestCost_OpenAIClampsCachedToInputCount(t *testing.T) { // TestCost_AnthropicCacheReadAndCreationAreAdditive proves the // Anthropic shape: cache_read and cache_creation tokens are // ADDITIVE to input_tokens (not subset), each billed at its own -// configured rate. The two rates pull in opposite directions — -// cache_read is the cheaper read-from-cache rate (≈0.1× input), -// cache_creation is the more expensive write-to-cache rate -// (≈1.25× input). +// configured rate. func TestCost_AnthropicCacheReadAndCreationAreAdditive(t *testing.T) { tbl := &Table{entries: map[string]map[string]Entry{ "anthropic": {"claude-sonnet": { @@ -125,11 +86,9 @@ func TestCost_AnthropicCacheReadAndCreationAreAdditive(t *testing.T) { "each Anthropic input bucket must bill at its own configured rate") } -// TestCost_AnthropicCacheRatesFallBackToInput covers the no-opt-in +// TestCost_AnthropicCacheRatesFallBackToInput covers the no-rate // path: when neither CacheReadPer1K nor CacheCreationPer1K is set, -// cache tokens bill at the regular input rate. This is more -// accurate than today's behaviour (cache tokens ignored entirely) -// without requiring operators to opt in via YAML. +// cache tokens bill at the regular input rate. func TestCost_AnthropicCacheRatesFallBackToInput(t *testing.T) { tbl := &Table{entries: map[string]map[string]Entry{ "anthropic": {"claude-sonnet": {InputPer1K: 0.003, OutputPer1K: 0.015}}, @@ -139,259 +98,39 @@ func TestCost_AnthropicCacheRatesFallBackToInput(t *testing.T) { // Without overrides: every input bucket at input_per_1k. want := ((256.0+768.0+512.0)/1000.0)*0.003 + (200.0/1000.0)*0.015 assert.InDelta(t, want, cost, 1e-12, - "absent cache rates must fall back to input_per_1k — Anthropic cache tokens were ignored before this change, billing at input rate is more accurate as a default") + "absent cache rates must fall back to input_per_1k") } -func TestNewLoader_UnknownModel(t *testing.T) { - base := t.TempDir() - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml")) +// TestEntryCosts_SurfaceSelectsFormula pins that the formula branches on +// the SURFACE, not on which table the entry came from: the same entry +// bills a subset carve-out on "openai", additive buckets on +// "anthropic"/"bedrock", and ignores cache counts everywhere else. This +// is what keeps per-provider-record entries (looked up by record id) +// mathematically identical to defaults-table entries. +func TestEntryCosts_SurfaceSelectsFormula(t *testing.T) { + e := Entry{InputPer1K: 0.002, OutputPer1K: 0.01, CachedInputPer1K: 0.001, CacheReadPer1K: 0.0002, CacheCreationPer1K: 0.0025} - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err) + openai := EntryCosts(e, "openai", 1000, 0, 400, 300) + assert.InDelta(t, (600.0/1000.0)*0.002+(400.0/1000.0)*0.001, openai.TotalUSD, 1e-12, + "openai: cached is a subset, cacheCreation ignored") - _, ok := l.Get().Cost("openai", "fantasy-model", 10, 10, 0, 0) - assert.False(t, ok, "unknown model returns ok=false") + anthropic := EntryCosts(e, "anthropic", 1000, 0, 400, 300) + assert.InDelta(t, 0.002+(400.0/1000.0)*0.0002+(300.0/1000.0)*0.0025, anthropic.TotalUSD, 1e-12, + "anthropic: cache buckets are additive") - _, ok = l.Get().Cost("cohere", "anything", 10, 10, 0, 0) - assert.False(t, ok, "unknown provider returns ok=false") + bedrock := EntryCosts(e, "bedrock", 1000, 0, 400, 300) + assert.InDelta(t, anthropic.TotalUSD, bedrock.TotalUSD, 1e-12, "bedrock shares the anthropic formula") + + other := EntryCosts(e, "gemini", 1000, 0, 400, 300) + assert.InDelta(t, 0.002, other.TotalUSD, 1e-12, "unknown surface: cache counts ignored") } -func TestNewLoader_InvalidYAMLRejected(t *testing.T) { - base := t.TempDir() - require.NoError(t, os.WriteFile(filepath.Join(base, "pricing.yaml"), []byte("\t- this is not: valid: yaml: :["), 0o600)) - - _, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.Error(t, err, "invalid YAML must surface as construction error") -} - -func TestLoader_ReloadKeepsPreviousOnParseError(t *testing.T) { - base := t.TempDir() - target := filepath.Join(base, "pricing.yaml") - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target) - - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err) - require.NotNil(t, l.Get(), "initial table populated") - - // Overwrite with content that violates the strict schema (extra field) - // plus a bumped mtime to trigger reload. - require.NoError(t, os.WriteFile(target, []byte("openai:\n gpt-4o:\n input_per_1k: 1.0\n output_per_1k: 2.0\n bogus_field: nope\n"), 0o600)) - future := time.Now().Add(time.Hour) - require.NoError(t, os.Chtimes(target, future, future)) - - err = l.reload() - require.Error(t, err, "parse error surfaced by reload()") - - cost, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0) - require.True(t, ok, "previous table still available after parse failure") - assert.InDelta(t, 0.00075, cost, 1e-9, "previous cost preserved") -} - -func TestLoader_ReloadNoChangeIsNoOp(t *testing.T) { - base := t.TempDir() - target := filepath.Join(base, "pricing.yaml") - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target) - - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err) - ptrBefore := l.Get() - - require.NoError(t, l.reload(), "no-change reload must not error") - ptrAfter := l.Get() - assert.Same(t, ptrBefore, ptrAfter, "table pointer unchanged when mtime unchanged") -} - -func TestLoader_ReloadDetectsChange(t *testing.T) { - base := t.TempDir() - target := filepath.Join(base, "pricing.yaml") - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target) - - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err) - - updated := []byte("openai:\n gpt-4o-mini:\n input_per_1k: 1.00\n output_per_1k: 2.00\n") - require.NoError(t, os.WriteFile(target, updated, 0o600)) - future := time.Now().Add(time.Hour) - require.NoError(t, os.Chtimes(target, future, future)) - - require.NoError(t, l.reload(), "reload must succeed on valid new content") - - cost, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0) - require.True(t, ok, "updated model still present") - assert.InDelta(t, 3.0, cost, 0.0001, "new prices are applied: 1 + 2 per 1k") -} - -// TestLoader_ReloadGoroutinePicksUpChanges proves the background goroutine -// started via Reload actually swaps the pricing table when the file changes -// on disk. Without that goroutine running, pricing edits would never reach -// requests until a proxy restart. -func TestLoader_ReloadGoroutinePicksUpChanges(t *testing.T) { - base := t.TempDir() - target := filepath.Join(base, "pricing.yaml") - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target) - - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err) - l.SetReloadInterval(20 * time.Millisecond) - - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - done := make(chan struct{}) - go func() { - l.Reload(ctx) - close(done) - }() - - // Before any rewrite, the loader holds the fixture's prices. - costBefore, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0) - require.True(t, ok, "fixture model must resolve initially") - assert.InDelta(t, 0.00075, costBefore, 1e-9, "fixture prices apply before rewrite") - - updated := []byte("openai:\n gpt-4o-mini:\n input_per_1k: 1.00\n output_per_1k: 2.00\n") - require.NoError(t, os.WriteFile(target, updated, 0o600)) - future := time.Now().Add(time.Hour) - require.NoError(t, os.Chtimes(target, future, future)) - - deadline := time.Now().Add(2 * time.Second) - for { - cost, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0) - if ok && cost > 2.5 { - break - } - if time.Now().After(deadline) { - t.Fatalf("background reloader did not pick up rewrite within deadline") - } - time.Sleep(10 * time.Millisecond) - } - - cancel() - select { - case <-done: - case <-time.After(2 * time.Second): - t.Fatal("Reload loop did not exit after cancel") - } -} - -func TestLoader_ReloadBackgroundLoopCancellation(t *testing.T) { - base := t.TempDir() - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml")) - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err) - - ctx, cancel := context.WithCancel(context.Background()) - done := make(chan struct{}) - go func() { - l.Reload(ctx) - close(done) - }() - cancel() - - select { - case <-done: - case <-time.After(2 * time.Second): - t.Fatal("Reload loop did not exit on context cancel") - } -} - -func TestNewLoader_PathValidation(t *testing.T) { - base := t.TempDir() - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml")) - - cases := []struct { - name string - relPath string - }{ - {"traversal", "../../etc/passwd"}, - {"absolute", "/etc/passwd"}, - {"slash in basename", "sub/pricing.yaml"}, - {"control chars", "pricing\x00.yaml"}, - } - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - _, err := NewLoader(base, tc.relPath, "llm_observability", nil) - require.Error(t, err, "NewLoader must reject %q", tc.relPath) - }) - } - - // Empty relPath is no longer a validation error: the loader treats it - // as "no override file, defaults only" so cost metadata is still - // emitted for the embedded models out of the box. - t.Run("empty falls back to defaults", func(t *testing.T) { - l, err := NewLoader(base, "", "llm_observability", nil) - require.NoError(t, err, "empty relPath should yield a defaults-only loader") - require.NotNil(t, l, "loader must be returned") - require.False(t, l.WatchesFile(), "no file watching when no override is given") - _, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0) - assert.True(t, ok, "embedded defaults should still resolve gpt-4o-mini") - }) -} - -// TestNewLoader_PathValidation_Extended covers the remaining attack shapes -// called out in C2: dot references, embedded traversal segments, and a -// newline in the basename. The basename regex must reject each one even -// though filepath.Clean would otherwise collapse them. -func TestNewLoader_PathValidation_Extended(t *testing.T) { - base := t.TempDir() - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml")) - - cases := []struct { - name string - relPath string - }{ - {"dot", "."}, - {"dotdot", ".."}, - {"relative traversal", "../pricing.yaml"}, - {"embedded slash", "pri/cing.yaml"}, - {"newline", "pricing\n.yaml"}, - } - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - _, err := NewLoader(base, tc.relPath, "llm_observability", nil) - require.Error(t, err, "NewLoader must reject %q", tc.relPath) - }) - } -} - -// TestNewLoader_ValidBasenameLoads proves the allowlist is exclusive: a -// basename containing only safe characters under baseDir loads. Without this -// a regression that over-tightened the regex would silently break valid -// deployments. -func TestNewLoader_ValidBasenameLoads(t *testing.T) { - base := t.TempDir() - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing-v2_prod.yaml")) - - l, err := NewLoader(base, "pricing-v2_prod.yaml", "llm_observability", nil) - require.NoError(t, err, "basename with _, -, . must load") - require.NotNil(t, l.Get(), "table populated") -} - -// TestNewLoader_SymlinkOutsideBaseDirRejected constructs a symlink under -// baseDir that points to a file outside it. O_NOFOLLOW must refuse to open -// the symlink even though the symlink path itself is a valid basename under -// baseDir. -func TestNewLoader_SymlinkOutsideBaseDirRejected(t *testing.T) { - outside := t.TempDir() - target := filepath.Join(outside, "evil.yaml") - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target) - - base := t.TempDir() - link := filepath.Join(base, "pricing.yaml") - require.NoError(t, os.Symlink(target, link), "symlink setup") - - _, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.Error(t, err, "O_NOFOLLOW must reject symlink even when it points outside baseDir") -} - -func TestNewLoader_SymlinkRejected(t *testing.T) { - base := t.TempDir() - concrete := filepath.Join(base, "real.yaml") - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), concrete) - - link := filepath.Join(base, "pricing.yaml") - require.NoError(t, os.Symlink(concrete, link), "symlink setup") - - _, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.Error(t, err, "O_NOFOLLOW must reject symlinked targets") +// TestEntryCosts_ClampsNegativeTokens: malformed upstream counts must +// never produce a negative cost. +func TestEntryCosts_ClampsNegativeTokens(t *testing.T) { + e := Entry{InputPer1K: 0.002, OutputPer1K: 0.01} + c := EntryCosts(e, "openai", -50, -10, -5, -3) + assert.Zero(t, c.TotalUSD, "all-negative counts clamp to zero cost") } func TestTableCost_NilSafe(t *testing.T) { @@ -402,31 +141,37 @@ func TestTableCost_NilSafe(t *testing.T) { assert.False(t, t1.Has("x", "y"), "nil table has nothing") } -func TestLoaderGet_NilSafe(t *testing.T) { - var l *Loader - assert.Nil(t, l.Get(), "nil loader returns nil table") -} - -// TestNewLoader_RejectsOversizedFile_FixesM4 proves the loader bounds reads -// at maxPricingBytes so a hostile file cannot exhaust process memory. -func TestNewLoader_RejectsOversizedFile_FixesM4(t *testing.T) { - base := t.TempDir() - target := filepath.Join(base, "pricing.yaml") - - // Build a YAML payload larger than the cap. We pad with valid YAML - // comments so a partial read would still fail the size check rather - // than the parser. - header := "openai:\n" - bigComment := make([]byte, maxPricingBytes+1024) - for i := range bigComment { - bigComment[i] = ' ' +func TestNewTable_ValidatesRates(t *testing.T) { + good := map[string]map[string]EntryJSON{ + "openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01, CachedInputPer1K: 0.00125}}, } - bigComment[0] = '#' - bigComment[len(bigComment)-1] = '\n' - payload := append([]byte(header), bigComment...) - require.NoError(t, os.WriteFile(target, payload, 0o600)) + tbl, err := NewTable(good) + require.NoError(t, err) + cost, ok := tbl.Cost("openai", "gpt-4o", 1000, 1000, 0, 0) + require.True(t, ok, "entry survives the wire conversion") + assert.InDelta(t, 0.0125, cost, 1e-9) - _, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.Error(t, err, "oversized pricing file must be rejected") - assert.Contains(t, err.Error(), "exceeds", "rejection must reference the byte cap") + _, ok = tbl.Cost("openai", "unknown-model", 1, 1, 0, 0) + assert.False(t, ok, "unknown model misses") + + for name, bad := range map[string]EntryJSON{ + "negative input": {InputPer1K: -1, OutputPer1K: 0.01}, + "NaN output": {InputPer1K: 0.01, OutputPer1K: math.NaN()}, + "Inf cache read": {InputPer1K: 0.01, OutputPer1K: 0.01, CacheReadPer1K: math.Inf(1)}, + "negative cached": {InputPer1K: 0.01, OutputPer1K: 0.01, CachedInputPer1K: -0.001}, + } { + _, err := NewTable(map[string]map[string]EntryJSON{"openai": {"m": bad}}) + assert.Error(t, err, "case %q must be rejected so a corrupt config fails the chain build instead of mispricing", name) + } +} + +func TestNewTable_NilAndEmpty(t *testing.T) { + tbl, err := NewTable(nil) + require.NoError(t, err, "nil map builds an empty (never-matching) table") + _, ok := tbl.Cost("openai", "gpt-4o", 1, 1, 0, 0) + assert.False(t, ok, "empty table prices nothing") + + entries, err := NewEntries(nil) + require.NoError(t, err) + assert.Empty(t, entries, "nil in, empty (never-matching) map out for the per-record map") } diff --git a/proxy/internal/llm/pricing/pricing_unix.go b/proxy/internal/llm/pricing/pricing_unix.go deleted file mode 100644 index 4f3ea33a2..000000000 --- a/proxy/internal/llm/pricing/pricing_unix.go +++ /dev/null @@ -1,68 +0,0 @@ -//go:build unix - -package pricing - -import ( - "fmt" - "io" - "os" - "syscall" - "time" - - log "github.com/sirupsen/logrus" -) - -// maxPricingBytes caps the size of the pricing YAML on read so a hostile or -// runaway file cannot exhaust process memory during reload. 1 MiB is several -// orders of magnitude larger than any reasonable pricing table. -const maxPricingBytes int64 = 1 << 20 - -// loadPricing opens the file with O_NOFOLLOW, fstats the open descriptor, -// and parses from that same descriptor. Never re-opens by path so a -// mid-read rename or symlink swap cannot substitute content. Bytes are -// capped at maxPricingBytes so the loader cannot be coerced into reading an -// unbounded file. -func loadPricing(path string) (*Table, time.Time, error) { - f, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW, 0) - if err != nil { - return nil, time.Time{}, fmt.Errorf("open %s: %w", path, err) - } - defer func() { - if cerr := f.Close(); cerr != nil { - log.Debugf("close pricing file %s: %v", path, cerr) - } - }() - - info, err := f.Stat() - if err != nil { - return nil, time.Time{}, fmt.Errorf("fstat %s: %w", path, err) - } - if !info.Mode().IsRegular() { - return nil, time.Time{}, fmt.Errorf("pricing file %s is not a regular file", path) - } - - data, err := io.ReadAll(io.LimitReader(f, maxPricingBytes+1)) - if err != nil { - return nil, time.Time{}, fmt.Errorf("read %s: %w", path, err) - } - if int64(len(data)) > maxPricingBytes { - return nil, time.Time{}, fmt.Errorf("pricing file %s exceeds %d bytes", path, maxPricingBytes) - } - - table, err := parsePricingBytes(data) - if err != nil { - return nil, time.Time{}, err - } - return table, info.ModTime(), nil -} - -// statMtime returns the mtime of the file at path. It uses lstat semantics -// via os.Lstat so a symlink swap is detected even though O_NOFOLLOW will -// later reject the open. -func statMtime(path string) (time.Time, error) { - info, err := os.Lstat(path) - if err != nil { - return time.Time{}, fmt.Errorf("lstat %s: %w", path, err) - } - return info.ModTime(), nil -} diff --git a/proxy/internal/middleware/builtin/builtin.go b/proxy/internal/middleware/builtin/builtin.go index 9ea4cf89d..9df60dd65 100644 --- a/proxy/internal/middleware/builtin/builtin.go +++ b/proxy/internal/middleware/builtin/builtin.go @@ -36,15 +36,13 @@ var defaultRegistry = middleware.NewRegistry() // FactoryContext is the per-process bag that concrete factories may // consult during construction. It carries the proxy-lifetime context, -// the data directory used for static config files (pricing tables, -// allowlists), the OTel meter, and the proxy logger. +// the OTel meter, and the proxy logger. // // Configure must be called once at boot before any chain build calls // Resolve. Calling it twice overwrites the prior value; tests may rely // on this to reset state. type FactoryContext struct { Context context.Context - DataDir string Meter metric.Meter Logger *log.Logger MgmtClient MgmtClient @@ -58,12 +56,11 @@ var ( // Configure stores the per-process FactoryContext. Concrete factories // reach for it via Context(). mgmt may be nil on tests / standalone // builds with no management server; consumers must guard. -func Configure(ctx context.Context, dataDir string, meter metric.Meter, logger *log.Logger, mgmt MgmtClient) { +func Configure(ctx context.Context, meter metric.Meter, logger *log.Logger, mgmt MgmtClient) { ctxMu.Lock() defer ctxMu.Unlock() ctxStore = FactoryContext{ Context: ctx, - DataDir: dataDir, Meter: meter, Logger: logger, MgmtClient: mgmt, diff --git a/proxy/internal/middleware/builtin/cost_calculation_matrix_test.go b/proxy/internal/middleware/builtin/cost_calculation_matrix_test.go index f479eb563..be3682c0e 100644 --- a/proxy/internal/middleware/builtin/cost_calculation_matrix_test.go +++ b/proxy/internal/middleware/builtin/cost_calculation_matrix_test.go @@ -12,6 +12,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + mgmtpricing "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" "github.com/netbirdio/netbird/proxy/internal/middleware" "github.com/netbirdio/netbird/proxy/internal/middleware/builtin" "github.com/netbirdio/netbird/proxy/internal/middleware/builtin/cost_meter" @@ -19,17 +20,20 @@ import ( "github.com/netbirdio/netbird/proxy/internal/middleware/builtin/llm_response_parser" ) -// Drives the real pipeline (llm_request_parser → llm_response_parser → cost_meter) on the embedded default pricing -// table and asserts exact USD amounts hardcoded from the vendors' published prices, including the cache split. +// Drives the real pipeline (llm_request_parser → llm_response_parser → cost_meter) on the REAL default pricing +// table management ships (mgmtpricing.DefaultTable, catalog-derived) and asserts exact USD amounts hardcoded from +// the vendors' published prices, including the cache split. This is the cross-stack pricing contract test: the +// management-side Entry JSON must decode into the proxy-side table and produce these exact costs. func TestCostCalculation_ProviderMatrix(t *testing.T) { - // Empty data dir → embedded defaults, like a proxy with no pricing override. - builtin.Configure(context.Background(), t.TempDir(), nil, nil, nil) + builtin.Configure(context.Background(), nil, nil, nil) reqMW, err := llm_request_parser.Factory{}.New(nil) require.NoError(t, err, "build llm_request_parser") respMW, err := llm_response_parser.Factory{}.New(nil) require.NoError(t, err, "build llm_response_parser") - costMW, err := cost_meter.Factory{}.New(nil) + costCfgJSON, err := json.Marshal(map[string]any{"pricing": map[string]any{"defaults": mgmtpricing.DefaultTable()}}) + require.NoError(t, err, "marshal management default table into cost_meter config") + costMW, err := cost_meter.Factory{}.New(costCfgJSON) require.NoError(t, err, "build cost_meter") t.Cleanup(func() { _ = costMW.Close() }) diff --git a/proxy/internal/middleware/builtin/cost_meter/factory.go b/proxy/internal/middleware/builtin/cost_meter/factory.go index b8a58d10e..2993ce32e 100644 --- a/proxy/internal/middleware/builtin/cost_meter/factory.go +++ b/proxy/internal/middleware/builtin/cost_meter/factory.go @@ -2,7 +2,6 @@ package cost_meter import ( "bytes" - "context" "encoding/json" "fmt" @@ -11,16 +10,27 @@ import ( "github.com/netbirdio/netbird/proxy/internal/middleware/builtin" ) -// defaultPricingFilename is the basename probed inside the proxy data -// directory when no override is configured. -const defaultPricingFilename = "pricing.yaml" - -// Config is the on-wire configuration for the middleware. +// Config is the on-wire configuration for the middleware, synthesized by +// management (buildCostMeterConfigJSON). The proxy has no embedded price +// list: this payload is the only pricing source, and updates arrive as +// ordinary mapping pushes that rebuild the chain (and with it this +// middleware instance) — no per-request fetches, no reload loops. type Config struct { - // PricingPath optionally overrides the basename of the pricing - // file probed inside the proxy data directory. When empty the - // loader falls back to "pricing.yaml". - PricingPath string `json:"pricing_path"` + Pricing *PricingConfig `json:"pricing"` +} + +// PricingConfig carries the full pricing table: +// - Defaults: parser surface ("openai"/"anthropic"/"bedrock") -> +// normalized model id -> rates, matched against llm.provider + +// llm.model. +// - Providers: provider record id -> normalized model id -> rates, +// matched against the llm.resolved_provider_id metadata llm_router +// stamps. Entries arrive fully materialized (management folds default +// cache rates in at synth time), so lookup order is simply +// per-record first, defaults second. +type PricingConfig struct { + Defaults map[string]map[string]pricing.EntryJSON `json:"defaults"` + Providers map[string]map[string]pricing.EntryJSON `json:"providers"` } // Factory builds cost_meter instances from raw config bytes. @@ -29,45 +39,45 @@ type Factory struct{} // ID returns the registry identifier. func (Factory) ID() string { return ID } -// New constructs a middleware instance. Empty, null, and {} configs -// are accepted; non-empty rawConfig that fails to unmarshal is -// rejected so misconfigurations surface at chain build time. The -// pricing loader is built once per instance and reused across -// invocations. +// New constructs a middleware instance. Empty, null, and {} configs are +// accepted for backward compatibility with a management server that +// predates config-delivered pricing — the instance then skips every cost +// computation (unknown_model) and a warning is logged once at build time. +// Non-empty rawConfig that fails to unmarshal, or a table carrying a +// non-finite / negative rate, is rejected so misconfigurations surface at +// chain build time. func (Factory) New(rawConfig []byte) (middleware.Middleware, error) { cfg, err := decodeConfig(rawConfig) if err != nil { return nil, err } - fctx := builtin.Context() - pricingPath := cfg.PricingPath - if pricingPath == "" { - pricingPath = defaultPricingFilename + if cfg.Pricing == nil { + if logger := builtin.Context().Logger; logger != nil { + logger.Warnf("cost_meter: no pricing table in middleware config; management predates config-delivered pricing — every request will record cost.skipped=unknown_model ($0)") + } + return newMiddleware(mustEmptyTable(), nil), nil } - loader, err := pricing.NewLoader(fctx.DataDir, pricingPath, ID, nil) + defaults, err := pricing.NewTable(cfg.Pricing.Defaults) if err != nil { - return nil, fmt.Errorf("init pricing loader: %w", err) + return nil, fmt.Errorf("cost_meter pricing defaults: %w", err) } - - cancel := startReloader(fctx.Context, loader) - - return newMiddleware(loader, cancel), nil + perRecord, err := pricing.NewEntries(cfg.Pricing.Providers) + if err != nil { + return nil, fmt.Errorf("cost_meter per-provider pricing: %w", err) + } + return newMiddleware(defaults, perRecord), nil } -// startReloader binds the loader's mtime-poll goroutine to a context -// derived from the proxy-lifetime context and returns its cancel func so -// the owning middleware can stop the goroutine on teardown. Returns nil -// when there's nothing to watch (nil context or defaults-only loader), in -// which case the middleware's Close is a no-op. -func startReloader(ctx context.Context, loader *pricing.Loader) context.CancelFunc { - if ctx == nil || !loader.WatchesFile() { - return nil +// mustEmptyTable returns a valid empty table. NewTable on a nil map cannot +// fail; the panic guard documents that invariant. +func mustEmptyTable() *pricing.Table { + t, err := pricing.NewTable(nil) + if err != nil { + panic(fmt.Sprintf("cost_meter: empty pricing table must build: %v", err)) } - cctx, cancel := context.WithCancel(ctx) - go loader.Reload(cctx) - return cancel + return t } // decodeConfig accepts empty, null, and {} configs, returning a diff --git a/proxy/internal/middleware/builtin/cost_meter/middleware.go b/proxy/internal/middleware/builtin/cost_meter/middleware.go index 63da6d17b..2ce706cda 100644 --- a/proxy/internal/middleware/builtin/cost_meter/middleware.go +++ b/proxy/internal/middleware/builtin/cost_meter/middleware.go @@ -1,7 +1,9 @@ // Package cost_meter implements the SlotOnResponse middleware that // converts token-usage metadata emitted by llm_response_parser into a -// per-request USD cost estimate. The middleware uses the shared pricing -// loader so operator pricing overrides apply to the chain. +// per-request USD cost estimate. Pricing arrives from management inside +// the middleware config: a per-provider-record table (the operator's +// stored prices, matched via llm.resolved_provider_id) consulted first, +// then the surface-keyed defaults table. package cost_meter import ( @@ -17,7 +19,9 @@ import ( const ID = "cost_meter" // Version is the implementation version emitted via the spec merge. -const Version = "1.0.0" +// 1.1.0: pricing is config-delivered (defaults + per-provider-record +// entries) instead of proxy-embedded. +const Version = "1.1.0" // Skip reasons emitted under KeyCostSkipped. The set is closed; the // dashboard surfaces these verbatim. @@ -42,19 +46,21 @@ var metadataKeys = []string{ } // Middleware computes a per-response cost estimate from the token -// counts emitted upstream by llm_response_parser. +// counts emitted upstream by llm_response_parser. Both tables are +// immutable — a pricing change arrives as a mapping push that rebuilds +// the chain with a fresh instance. type Middleware struct { - loader *pricing.Loader - // cancel stops this instance's pricing-reload goroutine. Non-nil only - // when the loader watches an override file; Close calls it so a chain - // rebuild doesn't leak a poll goroutine per retired instance. - cancel context.CancelFunc + // defaults is the surface-keyed table (llm.provider x llm.model). + defaults *pricing.Table + // perRecord is keyed by provider record id (llm.resolved_provider_id) + // then normalized model id; entries arrive fully materialized from + // management. Consulted before defaults. May be nil. + perRecord map[string]map[string]pricing.Entry } -// newMiddleware constructs a Middleware bound to the given pricing loader. -// cancel may be nil (defaults-only loader with no reloader to stop). -func newMiddleware(loader *pricing.Loader, cancel context.CancelFunc) *Middleware { - return &Middleware{loader: loader, cancel: cancel} +// newMiddleware constructs a Middleware over the given pricing tables. +func newMiddleware(defaults *pricing.Table, perRecord map[string]map[string]pricing.Entry) *Middleware { + return &Middleware{defaults: defaults, perRecord: perRecord} } // ID returns the registry identifier. @@ -79,16 +85,9 @@ func (m *Middleware) MetadataKeys() []string { // response. func (m *Middleware) MutationsSupported() bool { return false } -// Close stops this instance's pricing-reload goroutine, if any. Called by -// the chain when a rebuild retires the instance, so the mtime-poll loop -// doesn't outlive the chain it belonged to. Safe to call on a nil receiver -// and on an instance with no reloader. -func (m *Middleware) Close() error { - if m != nil && m.cancel != nil { - m.cancel() - } - return nil -} +// Close releases resources owned by the middleware. Stateless — the +// pricing tables are plain maps owned by this instance. +func (m *Middleware) Close() error { return nil } // Invoke reads provider, model, and token metadata, looks up pricing, // and emits either KeyCostUSDTotal or KeyCostSkipped. The decision is @@ -144,8 +143,7 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar return out, nil } - table := m.loader.Get() - costs, ok := table.Costs(provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens) + costs, ok := m.lookupCosts(in.Metadata, provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens) if !ok { out.Metadata = skip(skipUnknownModel) return out, nil @@ -164,6 +162,26 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar return out, nil } +// lookupCosts resolves the price for this request and computes the cost +// split. Resolution order: +// +// 1. Per-provider-record entry: the operator's stored price for the +// provider route that served the request, keyed by the +// llm.resolved_provider_id metadata llm_router stamped on the allow +// path. Absent metadata (e.g. no router in the chain) skips this tier. +// 2. Surface defaults: the catalog-derived table keyed by llm.provider. +// +// The surface always selects the cache formula — a per-record entry for an +// Anthropic route still bills its cache buckets additively. +func (m *Middleware) lookupCosts(md []middleware.KV, surface, model string, inTokens, outTokens, cachedTokens, cacheCreationTokens int64) (pricing.Costs, bool) { + if recordID := lookupKV(md, middleware.KeyLLMResolvedProviderID); recordID != "" { + if entry, ok := m.perRecord[recordID][model]; ok { + return pricing.EntryCosts(entry, surface, inTokens, outTokens, cachedTokens, cacheCreationTokens), true + } + } + return m.defaults.Costs(surface, model, inTokens, outTokens, cachedTokens, cacheCreationTokens) +} + // usd renders a cost as the fixed-precision string every cost.usd_* key // carries, so the per-bucket values and the aggregates round identically. // diff --git a/proxy/internal/middleware/builtin/cost_meter/middleware_test.go b/proxy/internal/middleware/builtin/cost_meter/middleware_test.go index e5d431d77..482061270 100644 --- a/proxy/internal/middleware/builtin/cost_meter/middleware_test.go +++ b/proxy/internal/middleware/builtin/cost_meter/middleware_test.go @@ -3,39 +3,50 @@ package cost_meter import ( "context" "encoding/json" - "os" - "path/filepath" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/netbirdio/netbird/proxy/internal/llm/pricing" "github.com/netbirdio/netbird/proxy/internal/middleware" - "github.com/netbirdio/netbird/proxy/internal/middleware/builtin" ) -const fixturePricing = `openai: - gpt-4o: - input_per_1k: 0.0025 - output_per_1k: 0.01 - gpt-4o-mini: - input_per_1k: 0.00015 - output_per_1k: 0.0006 -anthropic: - claude-sonnet-4-5: - input_per_1k: 0.003 - output_per_1k: 0.015 -` - -// configureBuiltin points the package-level FactoryContext at a tmp -// directory containing the test pricing fixture. Returns the path so -// callers can override files later if needed. -func configureBuiltin(t *testing.T) string { +// fixtureConfig mirrors what management's buildCostMeterConfigJSON ships: +// a surface-keyed defaults table. Rates match the retired YAML fixture so +// every cost assertion below is byte-identical to the pre-feature values. +func fixtureConfig(t *testing.T) []byte { t.Helper() - dir := t.TempDir() - require.NoError(t, os.WriteFile(filepath.Join(dir, "pricing.yaml"), []byte(fixturePricing), 0o600), "write pricing fixture") - builtin.Configure(context.Background(), dir, nil, nil, nil) - return dir + raw, err := json.Marshal(Config{Pricing: &PricingConfig{ + Defaults: map[string]map[string]pricing.EntryJSON{ + "openai": { + "gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}, + "gpt-4o-mini": {InputPer1K: 0.00015, OutputPer1K: 0.0006}, + }, + "anthropic": { + "claude-sonnet-4-5": {InputPer1K: 0.003, OutputPer1K: 0.015}, + }, + }, + }}) + require.NoError(t, err) + return raw +} + +// fixtureConfigWithCache adds the cache-rate fields. +func fixtureConfigWithCache(t *testing.T) []byte { + t.Helper() + raw, err := json.Marshal(Config{Pricing: &PricingConfig{ + Defaults: map[string]map[string]pricing.EntryJSON{ + "openai": { + "gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01, CachedInputPer1K: 0.00125}, + }, + "anthropic": { + "claude-sonnet-4-5": {InputPer1K: 0.003, OutputPer1K: 0.015, CacheReadPer1K: 0.0003, CacheCreationPer1K: 0.00375}, + }, + }, + }}) + require.NoError(t, err) + return raw } func metaValue(t *testing.T, kvs []middleware.KV, key string) (string, bool) { @@ -56,8 +67,7 @@ func buildMiddleware(t *testing.T, raw []byte) middleware.Middleware { } func TestMiddleware_StaticSurface(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) assert.Equal(t, ID, mw.ID(), "ID must match the registered constant") assert.Equal(t, Version, mw.Version(), "Version must match the constant") @@ -79,8 +89,10 @@ func TestMiddleware_StaticSurface(t *testing.T) { assert.Equal(t, expected, keys, "metadata key allowlist must match the spec") } +// TestFactory_AcceptsEmptyAndJSONConfig: empty/null/{} configs are what an +// old management (pre config-delivered pricing) sends — they must build a +// working (all-skip) instance, never fail the chain. func TestFactory_AcceptsEmptyAndJSONConfig(t *testing.T) { - configureBuiltin(t) cases := [][]byte{nil, {}, []byte("null"), []byte("{}"), []byte(" ")} for _, raw := range cases { mw, err := Factory{}.New(raw) @@ -90,15 +102,57 @@ func TestFactory_AcceptsEmptyAndJSONConfig(t *testing.T) { } func TestFactory_RejectsMalformedConfig(t *testing.T) { - configureBuiltin(t) mw, err := Factory{}.New([]byte("{not json")) require.Error(t, err, "malformed config must surface at construction") assert.Nil(t, mw, "no instance is returned on error") } -func TestFactory_DefaultPricingPathLoadsFixture(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) +// TestFactory_RejectsInvalidRates: a non-finite or negative rate anywhere +// in the table fails the chain build (defense-in-depth behind management's +// API validation) rather than silently mispricing. +func TestFactory_RejectsInvalidRates(t *testing.T) { + raw, err := json.Marshal(Config{Pricing: &PricingConfig{ + Defaults: map[string]map[string]pricing.EntryJSON{ + "openai": {"gpt-4o": {InputPer1K: -0.0025, OutputPer1K: 0.01}}, + }, + }}) + require.NoError(t, err) + mw, err := Factory{}.New(raw) + require.Error(t, err, "negative rate must fail the build") + assert.Nil(t, mw) + + raw, err = json.Marshal(Config{Pricing: &PricingConfig{ + Providers: map[string]map[string]pricing.EntryJSON{ + "prov-1": {"m": {InputPer1K: 0.01, OutputPer1K: 0.01, CacheReadPer1K: -1}}, + }, + }}) + require.NoError(t, err) + _, err = Factory{}.New(raw) + require.Error(t, err, "per-record tables validate too") +} + +// TestFactory_NilPricingSkipsEverything is the version-skew contract: a +// new proxy under an old management ({} config) must build, allow, and +// skip with unknown_model — degraded but never broken. +func TestFactory_NilPricingSkipsEverything(t *testing.T) { + mw := buildMiddleware(t, []byte("{}")) + out, err := mw.Invoke(context.Background(), &middleware.Input{ + Metadata: []middleware.KV{ + {Key: middleware.KeyLLMProvider, Value: "openai"}, + {Key: middleware.KeyLLMModel, Value: "gpt-4o"}, + {Key: middleware.KeyLLMInputTokens, Value: "1000"}, + {Key: middleware.KeyLLMOutputTokens, Value: "1000"}, + }, + }) + require.NoError(t, err) + assert.Equal(t, middleware.DecisionAllow, out.Decision, "cost_meter always allows") + value, ok := metaValue(t, out.Metadata, middleware.KeyCostSkipped) + require.True(t, ok, "no pricing table means every request skips") + assert.Equal(t, skipUnknownModel, value) +} + +func TestFactory_ConfigDefaultsPriceRequests(t *testing.T) { + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -116,34 +170,78 @@ func TestFactory_DefaultPricingPathLoadsFixture(t *testing.T) { assert.Equal(t, "0.000750000", value, "0.00015 + 0.0006 per 1k tokens, 9-decimal format") } -func TestFactory_PricingPathOverride(t *testing.T) { - dir := t.TempDir() - require.NoError(t, os.WriteFile(filepath.Join(dir, "custom.yaml"), []byte(fixturePricing), 0o600), "write custom pricing") - builtin.Configure(context.Background(), dir, nil, nil, nil) - - raw, err := json.Marshal(Config{PricingPath: "custom.yaml"}) +// TestInvoke_PerRecordEntryBeatsDefaults: when llm_router resolved a +// provider record whose operator pinned a price for the model, that price +// wins over the surface default. +func TestInvoke_PerRecordEntryBeatsDefaults(t *testing.T) { + raw, err := json.Marshal(Config{Pricing: &PricingConfig{ + Defaults: map[string]map[string]pricing.EntryJSON{ + "openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}}, + }, + Providers: map[string]map[string]pricing.EntryJSON{ + "prov-azure": {"gpt-4o": {InputPer1K: 0.005, OutputPer1K: 0.02}}, + }, + }}) require.NoError(t, err) - mw := buildMiddleware(t, raw) + out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ {Key: middleware.KeyLLMProvider, Value: "openai"}, {Key: middleware.KeyLLMModel, Value: "gpt-4o"}, - {Key: middleware.KeyLLMInputTokens, Value: "2000"}, + {Key: middleware.KeyLLMResolvedProviderID, Value: "prov-azure"}, + {Key: middleware.KeyLLMInputTokens, Value: "1000"}, {Key: middleware.KeyLLMOutputTokens, Value: "1000"}, }, }) require.NoError(t, err) - value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal) - require.True(t, ok, "cost.usd_total must be emitted with custom pricing path") - assert.Equal(t, "0.015000000", value, "2*0.0025 + 1*0.01 = 0.015 with 9-decimal format") + require.True(t, ok) + assert.Equal(t, "0.025000000", value, "operator's per-record price (0.005+0.02) wins over the default (0.0025+0.01)") } -func TestInvoke_ComputesCostForKnownModel(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) +// TestInvoke_PerRecordMissFallsBackToDefaults: a resolved record with no +// entry for this model (or no entries at all) falls through to the +// surface defaults — gateway providers rely on exactly this. +func TestInvoke_PerRecordMissFallsBackToDefaults(t *testing.T) { + raw, err := json.Marshal(Config{Pricing: &PricingConfig{ + Defaults: map[string]map[string]pricing.EntryJSON{ + "openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}}, + }, + Providers: map[string]map[string]pricing.EntryJSON{ + "prov-1": {"some-other-model": {InputPer1K: 1, OutputPer1K: 1}}, + }, + }}) + require.NoError(t, err) + mw := buildMiddleware(t, raw) + for name, recordID := range map[string]string{ + "record with other models": "prov-1", + "record with no entries": "prov-gateway", + } { + t.Run(name, func(t *testing.T) { + out, err := mw.Invoke(context.Background(), &middleware.Input{ + Metadata: []middleware.KV{ + {Key: middleware.KeyLLMProvider, Value: "openai"}, + {Key: middleware.KeyLLMModel, Value: "gpt-4o"}, + {Key: middleware.KeyLLMResolvedProviderID, Value: recordID}, + {Key: middleware.KeyLLMInputTokens, Value: "1000"}, + {Key: middleware.KeyLLMOutputTokens, Value: "1000"}, + }, + }) + require.NoError(t, err) + value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal) + require.True(t, ok, "per-record miss must fall back to the surface default, not skip") + assert.Equal(t, "0.012500000", value, "default rates apply") + }) + } +} + +// TestInvoke_NoResolvedProviderIDUsesDefaults: metadata without a +// resolved provider id (router denied, or a chain without llm_router) +// prices from the defaults table directly. +func TestInvoke_NoResolvedProviderIDUsesDefaults(t *testing.T) { + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ {Key: middleware.KeyLLMProvider, Value: "anthropic"}, @@ -153,17 +251,15 @@ func TestInvoke_ComputesCostForKnownModel(t *testing.T) { }, }) require.NoError(t, err) - value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal) - require.True(t, ok, "cost.usd_total must be emitted") + require.True(t, ok) assert.Equal(t, "0.018000000", value, "0.003 + 0.015 = 0.018 with 9-decimal format") _, skipped := metaValue(t, out.Metadata, middleware.KeyCostSkipped) assert.False(t, skipped, "cost.skipped must not be set when cost is computed") } func TestInvoke_MissingProvider(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -179,8 +275,7 @@ func TestInvoke_MissingProvider(t *testing.T) { } func TestInvoke_MissingModel(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -196,8 +291,7 @@ func TestInvoke_MissingModel(t *testing.T) { } func TestInvoke_MissingTokens(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) cases := []struct { name string @@ -240,8 +334,7 @@ func TestInvoke_MissingTokens(t *testing.T) { } func TestInvoke_UnparseableTokens(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) cases := []struct { name string @@ -271,8 +364,7 @@ func TestInvoke_UnparseableTokens(t *testing.T) { } func TestInvoke_ZeroTokens(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -291,8 +383,7 @@ func TestInvoke_ZeroTokens(t *testing.T) { } func TestInvoke_UnknownModel(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -309,8 +400,7 @@ func TestInvoke_UnknownModel(t *testing.T) { } func TestInvoke_NilInput(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), nil) require.NoError(t, err) @@ -319,36 +409,12 @@ func TestInvoke_NilInput(t *testing.T) { assert.Empty(t, out.Metadata, "no metadata must be emitted on nil input") } -const fixturePricingWithCache = `openai: - gpt-4o: - input_per_1k: 0.0025 - output_per_1k: 0.01 - cached_input_per_1k: 0.00125 -anthropic: - claude-sonnet-4-5: - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - cache_creation_per_1k: 0.00375 -` - -// configureBuiltinWithCacheRates points the package-level -// FactoryContext at a tmp directory containing pricing entries that -// include the cache rate fields. -func configureBuiltinWithCacheRates(t *testing.T) { - t.Helper() - dir := t.TempDir() - require.NoError(t, os.WriteFile(filepath.Join(dir, "pricing.yaml"), []byte(fixturePricingWithCache), 0o600), "write cache-aware pricing fixture") - builtin.Configure(context.Background(), dir, nil, nil, nil) -} - // TestInvoke_OpenAICachedSubsetDiscount proves the OpenAI shape end // to end through the middleware: cached_input_tokens is treated as a // SUBSET of input_tokens and discounted at the configured rate, not // added on top. func TestInvoke_OpenAICachedSubsetDiscount(t *testing.T) { - configureBuiltinWithCacheRates(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfigWithCache(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -390,8 +456,7 @@ func TestInvoke_OpenAICachedSubsetDiscount(t *testing.T) { // shape: cache_read and cache_creation are additive to input_tokens // and each carries its own rate. func TestInvoke_AnthropicCacheBucketsAdditive(t *testing.T) { - configureBuiltinWithCacheRates(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfigWithCache(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -429,8 +494,37 @@ func TestInvoke_AnthropicCacheBucketsAdditive(t *testing.T) { "output bucket bills 200 tokens at 0.015/1k") } +// TestInvoke_PerRecordEntryUsesSurfaceFormula: a per-record entry for an +// anthropic-surface request must bill its cache buckets additively — the +// formula follows llm.provider, not which table the entry came from. +func TestInvoke_PerRecordEntryUsesSurfaceFormula(t *testing.T) { + raw, err := json.Marshal(Config{Pricing: &PricingConfig{ + Providers: map[string]map[string]pricing.EntryJSON{ + "prov-ant": {"claude-sonnet-4-5": {InputPer1K: 0.003, OutputPer1K: 0.015, CacheReadPer1K: 0.0003, CacheCreationPer1K: 0.00375}}, + }, + }}) + require.NoError(t, err) + mw := buildMiddleware(t, raw) + + out, err := mw.Invoke(context.Background(), &middleware.Input{ + Metadata: []middleware.KV{ + {Key: middleware.KeyLLMProvider, Value: "anthropic"}, + {Key: middleware.KeyLLMModel, Value: "claude-sonnet-4-5"}, + {Key: middleware.KeyLLMResolvedProviderID, Value: "prov-ant"}, + {Key: middleware.KeyLLMInputTokens, Value: "256"}, + {Key: middleware.KeyLLMOutputTokens, Value: "200"}, + {Key: middleware.KeyLLMCachedInputTokens, Value: "768"}, + {Key: middleware.KeyLLMCacheCreationTokens, Value: "512"}, + }, + }) + require.NoError(t, err) + value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal) + require.True(t, ok) + assert.Equal(t, "0.005918400", value, "identical math to the defaults-table entry with the same rates") +} + // assertBucket asserts one per-bucket cost key carries the expected -// 6-decimal value. +// 9-decimal value. func assertBucket(t *testing.T, md []middleware.KV, key, want, msg string) { t.Helper() got, ok := metaValue(t, md, key) @@ -439,13 +533,10 @@ func assertBucket(t *testing.T, md []middleware.KV, key, want, msg string) { } // TestInvoke_CachedTokensAbsentFallsBackToBaseFormula covers the -// "operator hasn't opted in" path: with no cached metadata keys -// emitted, the meter must produce exactly the same cost as before -// the feature landed. Critical so operators with the new binary but -// no YAML changes see no behavioural drift on OpenAI requests. +// no-cache-metadata path: with no cached keys emitted, the meter must +// produce exactly the input+output cost. func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) { - configureBuiltinWithCacheRates(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfigWithCache(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -460,7 +551,7 @@ func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) { value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal) require.True(t, ok) // 1000 input * 0.0025 + 500 output * 0.01 = 0.0025 + 0.005 = 0.0075 - assert.Equal(t, "0.007500000", value, "no cached metadata = same cost as before the feature landed") + assert.Equal(t, "0.007500000", value, "no cached metadata = plain input+output cost") } // TestInvoke_UnparseableCachedTokensSkippedSilently proves the @@ -469,8 +560,7 @@ func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) { // regular formula. Cache buckets are a refinement, never a reason to // abort cost computation. func TestInvoke_UnparseableCachedTokensSkippedSilently(t *testing.T) { - configureBuiltinWithCacheRates(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfigWithCache(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -487,22 +577,10 @@ func TestInvoke_UnparseableCachedTokensSkippedSilently(t *testing.T) { assert.Equal(t, "0.007500000", value, "same as the no-cached-metadata path") } -// TestMiddleware_CloseCancelsReloader proves Close stops the per-instance -// pricing-reload goroutine: a chain rebuild retires the old instance and -// calls Close, which must invoke the cancel func startReloader handed it so -// the mtime-poll loop doesn't outlive the chain. -func TestMiddleware_CloseCancelsReloader(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - m := newMiddleware(nil, cancel) - - require.NoError(t, m.Close(), "Close must not error") - require.Error(t, ctx.Err(), "Close must cancel the reloader context so the poll goroutine exits") -} - -// TestMiddleware_CloseNilSafe confirms Close is a no-op (no panic) for an -// instance with no reloader and for a nil receiver. +// TestMiddleware_CloseNilSafe confirms Close is a no-op (no panic) even +// for a nil receiver. func TestMiddleware_CloseNilSafe(t *testing.T) { - require.NoError(t, newMiddleware(nil, nil).Close(), "no-reloader Close must be a no-op") + require.NoError(t, newMiddleware(nil, nil).Close(), "Close must be a no-op") var m *Middleware require.NoError(t, m.Close(), "nil-receiver Close must be safe") } diff --git a/proxy/internal/middleware/builtin/llm_request_parser/bedrock_test.go b/proxy/internal/middleware/builtin/llm_request_parser/bedrock_test.go index 827b81d07..d8cd81437 100644 --- a/proxy/internal/middleware/builtin/llm_request_parser/bedrock_test.go +++ b/proxy/internal/middleware/builtin/llm_request_parser/bedrock_test.go @@ -6,23 +6,6 @@ import ( "github.com/stretchr/testify/require" ) -func TestNormalizeBedrockModel(t *testing.T) { - cases := map[string]string{ - "eu.anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5", - "us.anthropic.claude-opus-4-8-20250101-v1:0": "anthropic.claude-opus-4-8", - "apac.anthropic.claude-haiku-4-5-v1:0": "anthropic.claude-haiku-4-5", - "anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5", - "meta.llama3-3-70b-instruct-v1:0": "meta.llama3-3-70b-instruct", - "amazon.nova-pro-v1:0": "amazon.nova-pro", - "amazon.nova-2-lite-v1:0": "amazon.nova-2-lite", - // Inference-profile ARN — model id lives in the last path segment. - "arn:aws:bedrock:eu-central-1:123456789012:inference-profile/eu.anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5", - } - for in, want := range cases { - require.Equal(t, want, normalizeBedrockModel(in), "normalize %q", in) - } -} - func TestParseBedrockPath(t *testing.T) { tests := []struct { path string diff --git a/proxy/internal/middleware/builtin/llm_request_parser/middleware.go b/proxy/internal/middleware/builtin/llm_request_parser/middleware.go index 64ca04e6a..b4d1e16d4 100644 --- a/proxy/internal/middleware/builtin/llm_request_parser/middleware.go +++ b/proxy/internal/middleware/builtin/llm_request_parser/middleware.go @@ -8,7 +8,6 @@ package llm_request_parser import ( "context" "net/url" - "regexp" "strconv" "strings" "unicode/utf8" @@ -253,9 +252,7 @@ func parseVertexPath(reqPath string) (vertexRequest, bool) { if c := strings.LastIndex(rest, ":"); c >= 0 { model, action = rest[:c], rest[c+1:] } - if at := strings.Index(model, "@"); at >= 0 { - model = model[:at] - } + model = llm.NormalizeVertexModel(model) if model == "" { return vertexRequest{}, false } @@ -343,14 +340,6 @@ func trimBedrockNamespace(reqPath string) string { return reqPath } -// bedrockRegionPrefixes are the cross-region inference-profile prefixes that -// front a Bedrock model id (e.g. "eu.anthropic.claude-..."). -var bedrockRegionPrefixes = []string{"us.", "eu.", "apac.", "global."} - -// bedrockVersionSuffix matches the trailing "-vN[:N]" or "-YYYYMMDD-vN[:N]" -// version/throughput suffix of a Bedrock model id. -var bedrockVersionSuffix = regexp.MustCompile(`-(\d{8}-)?v\d+(:\d+)?$`) - // parseBedrockPath extracts the model and streaming/converse flags from an AWS // Bedrock runtime model endpoint: // @@ -375,7 +364,7 @@ func parseBedrockPath(reqPath string) (bedrockRequest, bool) { if decoded, err := url.PathUnescape(rawModel); err == nil { rawModel = decoded } - model := normalizeBedrockModel(rawModel) + model := llm.NormalizeBedrockModel(rawModel) if model == "" { return bedrockRequest{}, false } @@ -389,30 +378,6 @@ func parseBedrockPath(reqPath string) (bedrockRequest, bool) { } } -// normalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile -// prefix, and the version/throughput suffix from a Bedrock model id so it -// matches the catalog/pricing key, e.g. -// "eu.anthropic.claude-sonnet-4-5-20250929-v1:0" -> "anthropic.claude-sonnet-4-5" -// and "arn:aws:bedrock:eu-central-1:123:inference-profile/eu.anthropic.claude-sonnet-4-5-20250929-v1:0" -// -> "anthropic.claude-sonnet-4-5". -func normalizeBedrockModel(modelID string) string { - m := modelID - // A full ARN (inference-profile / provisioned-throughput / foundation-model) - // carries the model id in its last path segment. - if strings.HasPrefix(m, "arn:") { - if i := strings.LastIndex(m, "/"); i >= 0 { - m = m[i+1:] - } - } - for _, p := range bedrockRegionPrefixes { - if strings.HasPrefix(m, p) { - m = m[len(p):] - break - } - } - return bedrockVersionSuffix.ReplaceAllString(m, "") -} - // invokeBedrock emits the model/provider/session/prompt for an AWS Bedrock // request. Bedrock is metered under the dedicated "bedrock" parser, which reads // both the InvokeModel and Converse response shapes. diff --git a/proxy/internal/proxy/agent_network_chain_realstack_test.go b/proxy/internal/proxy/agent_network_chain_realstack_test.go index bc611fc98..924d37ace 100644 --- a/proxy/internal/proxy/agent_network_chain_realstack_test.go +++ b/proxy/internal/proxy/agent_network_chain_realstack_test.go @@ -18,10 +18,10 @@ import ( "google.golang.org/grpc/credentials/insecure" "google.golang.org/grpc/test/bufconn" - rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" - mgmtgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + mgmtgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/proxy/internal/middleware" "github.com/netbirdio/netbird/proxy/internal/middleware/bodytap" @@ -134,14 +134,18 @@ func TestReverseProxy_AgentNetworkRequest_FullChain(t *testing.T) { RedactPii: true, })) require.NoError(t, st.SaveAgentNetworkProvider(ctx, &agentNetworkTypes.Provider{ - ID: providerID, - AccountID: testAccountID, - ProviderID: "openai_api", - Name: "openai-fullchain-test", - UpstreamURL: upstream.URL, // router rewrites to this - APIKey: "sk-test", - Enabled: true, - Models: []agentNetworkTypes.ProviderModel{{ID: "gpt-5.4"}}, + ID: providerID, + AccountID: testAccountID, + ProviderID: "openai_api", + Name: "openai-fullchain-test", + UpstreamURL: upstream.URL, // router rewrites to this + APIKey: "sk-test", + Enabled: true, + // Operator-pinned prices deliberately differ from the catalog's + // gpt-5.4 rates (0.0025/0.015) so the cost assertion below proves + // the per-provider-record price — not the default table — billed + // this request: stored price → synth → wire → cost_meter. + Models: []agentNetworkTypes.ProviderModel{{ID: "gpt-5.4", InputPer1k: 0.004, OutputPer1k: 0.02}}, SessionPrivateKey: "priv", SessionPublicKey: "pub", })) @@ -176,7 +180,7 @@ func TestReverseProxy_AgentNetworkRequest_FullChain(t *testing.T) { // ---- 5. Wire the middleware framework — same registry the proxy uses // in production, configured with our bufconn-backed management client. - mwbuiltin.Configure(ctx, t.TempDir(), nil, testLogger, mgmtClient) + mwbuiltin.Configure(ctx, nil, testLogger, mgmtClient) registry := mwbuiltin.DefaultRegistry() mwMetrics, err := middleware.NewMetrics(nil) require.NoError(t, err) @@ -283,13 +287,23 @@ func TestReverseProxy_AgentNetworkRequest_FullChain(t *testing.T) { if r.DimensionKind == agentNetworkTypes.DimensionGroup && r.DimensionID == adminGroupID && r.WindowSeconds == 60 && - r.TokensInput+r.TokensOutput > 0 { + r.TokensInput+r.TokensOutput > 0 && + r.CostUSD > 0 { return true } } return false }, 5*time.Second, 50*time.Millisecond, - "Admins group consumption row must increment via the response leg — if this fails the proxy's respInput dropped UserGroups again or the parser/recorder wiring is broken") + "Admins group consumption row must increment via the response leg WITH a non-zero cost — a zero cost means the operator's stored price never reached cost_meter (synth → wire → per-record lookup broken)") + + // 8a-cost. Exact cost from the OPERATOR's stored price, not the catalog + // default: 12 prompt tokens × 0.004/1k + 40 completion tokens × 0.02/1k + // = 0.000048 + 0.0008 = 0.000848. With catalog rates it would be 0.00063 + // — this assertion distinguishes the two, closing the loop on the whole + // dynamic-pricing feature (dashboard save → synth → gRPC wire → + // llm.resolved_provider_id lookup → billing). + assert.Equal(t, "0.000848000", cd.GetMetadata()["cost.usd_total"], + "cost must be computed from the provider record's operator-pinned price") // 8b. Both the captured prompt and the captured completion are // redacted — proves the synth threads redact_pii=true into BOTH parser diff --git a/proxy/server.go b/proxy/server.go index 4f448e4b8..bd70b7e70 100644 --- a/proxy/server.go +++ b/proxy/server.go @@ -246,10 +246,6 @@ type Server struct { // in processMappings before the receive loop reconnects to resync. // Zero uses defaultMappingBatchWatchdog. MappingBatchWatchdog time.Duration - // MiddlewareDataDir is the base directory the middleware system uses to - // resolve file-backed configuration (e.g. the cost_meter pricing table). - // Empty means any middleware that requires a file fails at configure time. - MiddlewareDataDir string // MiddlewareCaptureBudgetBytes overrides the proxy-wide in-flight capture // budget passed to middleware.NewManager. Zero or negative values fall // back to defaultMiddlewareCaptureBudgetBytes (256 MiB). @@ -2093,7 +2089,7 @@ func (s *Server) initMiddlewareManager(ctx context.Context) error { return fmt.Errorf("middleware manager requires metrics bundle") } otelMeter := s.meter.Meter() - mwbuiltin.Configure(ctx, s.MiddlewareDataDir, otelMeter, s.Logger, s.mgmtClient) + mwbuiltin.Configure(ctx, otelMeter, s.Logger, s.mgmtClient) mwMetrics, err := middleware.NewMetrics(otelMeter) if err != nil { diff --git a/shared/llm/model.go b/shared/llm/model.go new file mode 100644 index 000000000..08e42e5a4 --- /dev/null +++ b/shared/llm/model.go @@ -0,0 +1,58 @@ +// Package llm holds LLM model-identifier helpers shared by the proxy and +// the management server. The proxy normalizes model ids parsed off inbound +// requests; management normalizes the operator's registered model ids at +// synthesis time so both sides of the pricing / routing contract compare +// equal. +package llm + +import ( + "regexp" + "strings" +) + +// bedrockRegionPrefixes are the cross-region inference-profile prefixes that +// front a Bedrock model id (e.g. "eu.anthropic.claude-..."). +var bedrockRegionPrefixes = []string{"us.", "eu.", "apac.", "global."} + +// bedrockVersionSuffix matches the trailing "-vN[:N]" or "-YYYYMMDD-vN[:N]" +// version/throughput suffix of a Bedrock model id. +var bedrockVersionSuffix = regexp.MustCompile(`-(\d{8}-)?v\d+(:\d+)?$`) + +// NormalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile +// prefix, and the version/throughput suffix from a Bedrock model id so it +// matches the catalog/pricing key, e.g. +// "eu.anthropic.claude-sonnet-4-5-20250929-v1:0" -> "anthropic.claude-sonnet-4-5" +// and the inference-profile ARN's last segment likewise. It is the single +// source of truth shared by the proxy's request parser (which normalizes the +// request model from the URL path), the proxy's router (which normalizes the +// operator's registered Bedrock model ids so both sides compare equal), and +// the management synthesizer (which keys per-provider pricing entries by the +// normalized id the parser will emit at billing time). +func NormalizeBedrockModel(modelID string) string { + m := modelID + // A full ARN (inference-profile / provisioned-throughput / foundation-model) + // carries the model id in its last path segment. + if strings.HasPrefix(m, "arn:") { + if i := strings.LastIndex(m, "/"); i >= 0 { + m = m[i+1:] + } + } + for _, p := range bedrockRegionPrefixes { + if strings.HasPrefix(m, p) { + m = m[len(p):] + break + } + } + return bedrockVersionSuffix.ReplaceAllString(m, "") +} + +// NormalizeVertexModel strips the "@version" suffix from a Vertex AI model id +// (e.g. "claude-sonnet-4-5@20250929" -> "claude-sonnet-4-5") so it matches +// the catalog/pricing key. Vertex publisher models are priced under their +// vendor surface with the bare, unversioned id. +func NormalizeVertexModel(modelID string) string { + if at := strings.Index(modelID, "@"); at >= 0 { + return modelID[:at] + } + return modelID +} diff --git a/proxy/internal/llm/bedrock_model_test.go b/shared/llm/model_test.go similarity index 66% rename from proxy/internal/llm/bedrock_model_test.go rename to shared/llm/model_test.go index 3bd9662b7..42f2e9ca5 100644 --- a/proxy/internal/llm/bedrock_model_test.go +++ b/shared/llm/model_test.go @@ -11,6 +11,8 @@ func TestNormalizeBedrockModel(t *testing.T) { "eu.anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5", "us.anthropic.claude-haiku-4-5": "anthropic.claude-haiku-4-5", "us.anthropic.claude-opus-4-8-20250101-v1:0": "anthropic.claude-opus-4-8", + "apac.anthropic.claude-haiku-4-5-v1:0": "anthropic.claude-haiku-4-5", + "amazon.nova-2-lite-v1:0": "amazon.nova-2-lite", "anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5", "meta.llama3-3-70b-instruct-v1:0": "meta.llama3-3-70b-instruct", "amazon.nova-pro-v1:0": "amazon.nova-pro", @@ -21,3 +23,14 @@ func TestNormalizeBedrockModel(t *testing.T) { require.Equal(t, want, NormalizeBedrockModel(in), "normalize %q", in) } } + +func TestNormalizeVertexModel(t *testing.T) { + cases := map[string]string{ + "claude-sonnet-4-5@20250929": "claude-sonnet-4-5", + "claude-haiku-4-5": "claude-haiku-4-5", + "gpt-4o@2024-08-06": "gpt-4o", + } + for in, want := range cases { + require.Equal(t, want, NormalizeVertexModel(in), "normalize %q", in) + } +} diff --git a/shared/management/http/api/openapi.yml b/shared/management/http/api/openapi.yml index e3d11227a..8ad3d932c 100644 --- a/shared/management/http/api/openapi.yml +++ b/shared/management/http/api/openapi.yml @@ -5271,6 +5271,21 @@ components: format: double description: Cost per 1k output tokens, in USD. example: 0.0006 + cached_input_per_1k: + type: number + format: double + description: OpenAI-shape cache rate — cost per 1k cached prompt tokens (a subset of input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means no discount (cached tokens bill at input_per_1k). + example: 0.000075 + cache_read_per_1k: + type: number + format: double + description: Anthropic-shape cache rate — cost per 1k cache-read tokens (additive to input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means cache reads bill at input_per_1k. + example: 0.0003 + cache_creation_per_1k: + type: number + format: double + description: Anthropic-shape cache rate — cost per 1k cache-creation tokens (additive to input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means cache writes bill at input_per_1k. + example: 0.00375 required: - id - input_per_1k @@ -5296,6 +5311,21 @@ components: format: double description: Output token price per 1k tokens, in USD. example: 0.015 + cached_input_per_1k: + type: number + format: double + description: OpenAI-shape cache rate — default cost per 1k cached prompt tokens (a subset of input tokens), in USD. Absent when the model has no cached-input discount. + example: 0.000075 + cache_read_per_1k: + type: number + format: double + description: Anthropic-shape cache rate — default cost per 1k cache-read tokens (additive to input tokens), in USD. Absent when the model has no cache-read rate. + example: 0.0003 + cache_creation_per_1k: + type: number + format: double + description: Anthropic-shape cache rate — default cost per 1k cache-creation tokens (additive to input tokens), in USD. Absent when the model has no cache-creation rate. + example: 0.00375 context_window: type: integer description: Maximum context window in tokens. @@ -5354,6 +5384,13 @@ components: $ref: '#/components/schemas/AgentNetworkCatalogExtraHeader' identity_injection: $ref: '#/components/schemas/AgentNetworkCatalogIdentityInjection' + pricing_surfaces: + type: array + description: | + Cost-meter pricing surfaces this provider's traffic is metered under ("openai", "anthropic", "bedrock"). Tells the dashboard which cache-rate fields apply to this provider's models: "openai" → cached_input_per_1k (cached prompt tokens are a subset of input); "anthropic"/"bedrock" → cache_read_per_1k + cache_creation_per_1k (additive buckets). Absent/empty for gateway and custom entries, whose upstream shape NetBird cannot know ahead of time — surface all cache fields for those. + items: + type: string + example: ["openai"] models: type: array description: Catalog models available for this provider. diff --git a/shared/management/http/api/types.gen.go b/shared/management/http/api/types.gen.go index a4de48a09..87dd9ccfc 100644 --- a/shared/management/http/api/types.gen.go +++ b/shared/management/http/api/types.gen.go @@ -2017,6 +2017,15 @@ type AgentNetworkCatalogJSONMetadataInjection struct { // AgentNetworkCatalogModel defines model for AgentNetworkCatalogModel. type AgentNetworkCatalogModel struct { + // CacheCreationPer1k Anthropic-shape cache rate — default cost per 1k cache-creation tokens (additive to input tokens), in USD. Absent when the model has no cache-creation rate. + CacheCreationPer1k *float64 `json:"cache_creation_per_1k,omitempty"` + + // CacheReadPer1k Anthropic-shape cache rate — default cost per 1k cache-read tokens (additive to input tokens), in USD. Absent when the model has no cache-read rate. + CacheReadPer1k *float64 `json:"cache_read_per_1k,omitempty"` + + // CachedInputPer1k OpenAI-shape cache rate — default cost per 1k cached prompt tokens (a subset of input tokens), in USD. Absent when the model has no cached-input discount. + CachedInputPer1k *float64 `json:"cached_input_per_1k,omitempty"` + // ContextWindow Maximum context window in tokens. ContextWindow int `json:"context_window"` @@ -2070,6 +2079,9 @@ type AgentNetworkCatalogProvider struct { // Name Display name for the provider. Name string `json:"name"` + + // PricingSurfaces Cost-meter pricing surfaces this provider's traffic is metered under ("openai", "anthropic", "bedrock"). Tells the dashboard which cache-rate fields apply to this provider's models: "openai" → cached_input_per_1k (cached prompt tokens are a subset of input); "anthropic"/"bedrock" → cache_read_per_1k + cache_creation_per_1k (additive buckets). Absent/empty for gateway and custom entries, whose upstream shape NetBird cannot know ahead of time — surface all cache fields for those. + PricingSurfaces *[]string `json:"pricing_surfaces,omitempty"` } // AgentNetworkCatalogProviderKind Presentation grouping for the provider Select on the dashboard. @@ -2293,6 +2305,15 @@ type AgentNetworkProvider struct { // AgentNetworkProviderModel A model exposed by the provider, with the operator's per-1k input/output prices in USD. type AgentNetworkProviderModel struct { + // CacheCreationPer1k Anthropic-shape cache rate — cost per 1k cache-creation tokens (additive to input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means cache writes bill at input_per_1k. + CacheCreationPer1k *float64 `json:"cache_creation_per_1k,omitempty"` + + // CacheReadPer1k Anthropic-shape cache rate — cost per 1k cache-read tokens (additive to input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means cache reads bill at input_per_1k. + CacheReadPer1k *float64 `json:"cache_read_per_1k,omitempty"` + + // CachedInputPer1k OpenAI-shape cache rate — cost per 1k cached prompt tokens (a subset of input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means no discount (cached tokens bill at input_per_1k). + CachedInputPer1k *float64 `json:"cached_input_per_1k,omitempty"` + // Id Model identifier (e.g. "gpt-4o-mini"). Id string `json:"id"`