mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-03 03:59:07 +02:00
[management,proxy] Agent network: per-account LLM gateway (policy, metering, multi-provider) (#6555)
* [agent-network] Shared proto, OpenAPI schema, and generated types * [agent-network] Management: store, manager, synthesizer, policy engine, provider catalog, HTTP/gRPC API Adds the account-scoped agent-network module: provider/policy/budget CRUD and store, the reverse-proxy service synthesizer, policy selection + limit enforcement, the provider catalog (incl. Vertex AI and AWS Bedrock entries), and the management HTTP + proxy gRPC surfaces. * [management] Fix agent-network proxy-peer fan-out on affected-peer recompute The affected-peers resolver loaded only persisted reverse-proxy services, but agent-network services are synthesized on demand and never persisted. As a result the embedded proxy peer was never folded into the affected set when a client's group changed, so the proxy received no network-map update for a newly authorised client and rejected its handshake until a full resync (restart). loadProxyServices now merges the synthesized agent-network services (injected via a registration hook to avoid an import cycle), so proxy peers learn newly authorised clients immediately. * [proxy] Reverse-proxy middleware framework, chain, and request plumbing The per-target middleware chain (slots, dispatcher, mutation gate, metadata merger), body capture, access-log terminal sink, and the proxy wiring that builds + runs chains for synthesized agent-network services. * [proxy] LLM parsers, pricing, and builtin middlewares (OpenAI, Anthropic, Vertex AI, AWS Bedrock) Request/response parsers and SSE/event-stream metering, the embedded pricing table, and the builtin middleware set: request parser, router, policy limit-check/record, cost meter, guardrail, identity inject, response parser. Includes the path-routed providers — Google Vertex AI (keyfile:: service-account OAuth minting) and AWS Bedrock (bearer auth, invoke/converse/streaming, optional /bedrock prefix) — plus the Models allowlist and unmeterable-publisher deny. * [proxy] IPv6 in-place apply and TCP accept-loop hardening on netstack listeners * [agent-network] End-to-end test suite, module docs, and deployment preset * [agent-network] Fix codespell typos and exclude false positives - labelgen word pool: vermillion -> vermilion, racoon -> raccoon. - codespell ignore list: add flate (Go compress/flate package), recordin (a test-local identifier), and unparseable (a valid alternative spelling used consistently across identifiers + a metadata-value constant). * [management] Set LastSeen on injected proxy peer in realstack test (MySQL strict-mode) The injected embedded proxy peer had a PeerStatus with a zero LastSeen, which serializes to '0000-00-00' and is rejected by MySQL in strict mode (SQLite tolerates it). Set LastSeen to a valid time so SaveAccount succeeds on both engines. * [agent-network] Remove e2e shell-script suite from this branch The end-to-end shell scripts under scripts/e2e/ are maintained in a separate testing suite and are not part of this change set. * [agent-network] Polish module docs: remove internal review scaffolding, fix links, verify diagrams Strip PR-review framing, commit references, absolute paths, and stale internal references from the agent-network module docs; fix broken relative links; verify all diagrams against the current architecture. Remove the internal AI-reviewer prompt file. * [management] Refine session expiration handling to support 3-state encoding for SSO deadlines * [agent-network] Relocate agentnetwork package to internals/modules Move management/server/agentnetwork (and its catalog/, labelgen/, types/ subpackages) to management/internals/modules/agentnetwork, alongside the reverse-proxy module, and rewrite all importers. Pure relocation: package names, the synthesizer + affectedpeers registration hook, and store access (shared store.Store) are unchanged, so no import cycle is introduced (affectedpeers still depends only on the agentnetwork/types leaf). * [agent-network] Co-locate HTTP handlers in the module (RegisterEndpoints) Move the agent-network HTTP handlers from server/http/handlers/agentnetwork into the module at internals/modules/agentnetwork/handlers (package handlers) and rename the entrypoint AddEndpoints -> RegisterEndpoints, matching the reverse-proxy module convention. Wiring in http/handler.go updated accordingly.
This commit is contained in:
@@ -0,0 +1,88 @@
|
||||
package cost_meter
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/llm/pricing"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware"
|
||||
"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.
|
||||
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"`
|
||||
}
|
||||
|
||||
// Factory builds cost_meter instances from raw config bytes.
|
||||
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.
|
||||
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
|
||||
}
|
||||
|
||||
loader, err := pricing.NewLoader(fctx.DataDir, pricingPath, ID, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("init pricing loader: %w", err)
|
||||
}
|
||||
|
||||
cancel := startReloader(fctx.Context, loader)
|
||||
|
||||
return newMiddleware(loader, cancel), 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
|
||||
}
|
||||
cctx, cancel := context.WithCancel(ctx)
|
||||
go loader.Reload(cctx)
|
||||
return cancel
|
||||
}
|
||||
|
||||
// decodeConfig accepts empty, null, and {} configs, returning a
|
||||
// zero-value Config. Non-empty payloads must parse cleanly.
|
||||
func decodeConfig(rawConfig []byte) (Config, error) {
|
||||
var cfg Config
|
||||
if len(bytes.TrimSpace(rawConfig)) == 0 {
|
||||
return cfg, nil
|
||||
}
|
||||
if err := json.Unmarshal(rawConfig, &cfg); err != nil {
|
||||
return cfg, fmt.Errorf("decode config: %w", err)
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func init() {
|
||||
builtin.Register(Factory{})
|
||||
}
|
||||
@@ -0,0 +1,193 @@
|
||||
// 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.
|
||||
package cost_meter
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strconv"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/llm/pricing"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware"
|
||||
)
|
||||
|
||||
// ID is the registry identifier for this middleware.
|
||||
const ID = "cost_meter"
|
||||
|
||||
// Version is the implementation version emitted via the spec merge.
|
||||
const Version = "1.0.0"
|
||||
|
||||
// Skip reasons emitted under KeyCostSkipped. The set is closed; the
|
||||
// dashboard surfaces these verbatim.
|
||||
const (
|
||||
skipMissingProvider = "missing_provider"
|
||||
skipMissingModel = "missing_model"
|
||||
skipMissingTokens = "missing_tokens"
|
||||
//nolint:gosec // skip-reason label, not a credential
|
||||
skipUnparseableTokens = "unparseable_tokens"
|
||||
skipZeroTokens = "zero_tokens"
|
||||
skipUnknownModel = "unknown_model"
|
||||
)
|
||||
|
||||
var metadataKeys = []string{
|
||||
middleware.KeyCostUSDTotal,
|
||||
middleware.KeyCostSkipped,
|
||||
}
|
||||
|
||||
// Middleware computes a per-response cost estimate from the token
|
||||
// counts emitted upstream by llm_response_parser.
|
||||
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
|
||||
}
|
||||
|
||||
// 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}
|
||||
}
|
||||
|
||||
// ID returns the registry identifier.
|
||||
func (m *Middleware) ID() string { return ID }
|
||||
|
||||
// Version returns the implementation version.
|
||||
func (m *Middleware) Version() string { return Version }
|
||||
|
||||
// Slot reports that the middleware runs after the upstream call.
|
||||
func (m *Middleware) Slot() middleware.Slot { return middleware.SlotOnResponse }
|
||||
|
||||
// AcceptedContentTypes is empty: cost_meter never inspects bodies.
|
||||
func (m *Middleware) AcceptedContentTypes() []string { return []string{} }
|
||||
|
||||
// MetadataKeys returns the closed allowlist of keys this middleware
|
||||
// may emit.
|
||||
func (m *Middleware) MetadataKeys() []string {
|
||||
return append([]string(nil), metadataKeys...)
|
||||
}
|
||||
|
||||
// MutationsSupported reports that this middleware never mutates the
|
||||
// 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
|
||||
}
|
||||
|
||||
// Invoke reads provider, model, and token metadata, looks up pricing,
|
||||
// and emits either KeyCostUSDTotal or KeyCostSkipped. The decision is
|
||||
// always DecisionAllow; cost metering never denies or mutates.
|
||||
func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middleware.Output, error) {
|
||||
out := &middleware.Output{Decision: middleware.DecisionAllow}
|
||||
if in == nil {
|
||||
return out, nil
|
||||
}
|
||||
|
||||
provider := lookupKV(in.Metadata, middleware.KeyLLMProvider)
|
||||
if provider == "" {
|
||||
out.Metadata = skip(skipMissingProvider)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
model := lookupKV(in.Metadata, middleware.KeyLLMModel)
|
||||
if model == "" {
|
||||
out.Metadata = skip(skipMissingModel)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
inRaw, hasIn := lookupKVOK(in.Metadata, middleware.KeyLLMInputTokens)
|
||||
outRaw, hasOut := lookupKVOK(in.Metadata, middleware.KeyLLMOutputTokens)
|
||||
if !hasIn || !hasOut {
|
||||
out.Metadata = skip(skipMissingTokens)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
inTokens, err := strconv.ParseInt(inRaw, 10, 64)
|
||||
if err != nil || inTokens < 0 {
|
||||
// Unparseable or negative tokens are not a runtime error: the
|
||||
// upstream llm_response_parser emitted a non-numeric / invalid
|
||||
// value, so we surface that as cost.skipped and continue with
|
||||
// Allow rather than pricing a negative count.
|
||||
out.Metadata = skip(skipUnparseableTokens)
|
||||
return out, nil //nolint:nilerr // structured skip; not a runtime error
|
||||
}
|
||||
outTokens, err := strconv.ParseInt(outRaw, 10, 64)
|
||||
if err != nil || outTokens < 0 {
|
||||
out.Metadata = skip(skipUnparseableTokens)
|
||||
return out, nil //nolint:nilerr // structured skip; not a runtime error
|
||||
}
|
||||
|
||||
// Cache buckets are optional and silently zeroed on a missing /
|
||||
// malformed value; they're a refinement on top of input cost,
|
||||
// not a precondition. A buggy value falls back to 0, never aborts.
|
||||
cachedTokens := parseOptionalInt64(in.Metadata, middleware.KeyLLMCachedInputTokens)
|
||||
cacheCreationTokens := parseOptionalInt64(in.Metadata, middleware.KeyLLMCacheCreationTokens)
|
||||
|
||||
if inTokens == 0 && outTokens == 0 && cachedTokens == 0 && cacheCreationTokens == 0 {
|
||||
out.Metadata = skip(skipZeroTokens)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
table := m.loader.Get()
|
||||
cost, ok := table.Cost(provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens)
|
||||
if !ok {
|
||||
out.Metadata = skip(skipUnknownModel)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
out.Metadata = []middleware.KV{
|
||||
{Key: middleware.KeyCostUSDTotal, Value: fmt.Sprintf("%.6f", cost)},
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// skip returns a single-entry metadata slice carrying the given skip
|
||||
// reason under KeyCostSkipped.
|
||||
func skip(reason string) []middleware.KV {
|
||||
return []middleware.KV{{Key: middleware.KeyCostSkipped, Value: reason}}
|
||||
}
|
||||
|
||||
// lookupKV returns the value associated with key, or the empty string
|
||||
// when the key is absent.
|
||||
func lookupKV(kvs []middleware.KV, key string) string {
|
||||
v, _ := lookupKVOK(kvs, key)
|
||||
return v
|
||||
}
|
||||
|
||||
// lookupKVOK returns the value associated with key plus a presence
|
||||
// flag so callers can distinguish absent from empty.
|
||||
func lookupKVOK(kvs []middleware.KV, key string) (string, bool) {
|
||||
for _, kv := range kvs {
|
||||
if kv.Key == key {
|
||||
return kv.Value, true
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
// parseOptionalInt64 reads a metadata value and decodes it as int64.
|
||||
// Absent or unparseable values yield 0 — the caller treats absence as
|
||||
// "no cached tokens" rather than an error, since cache buckets are a
|
||||
// refinement, not a precondition.
|
||||
func parseOptionalInt64(kvs []middleware.KV, key string) int64 {
|
||||
raw, ok := lookupKVOK(kvs, key)
|
||||
if !ok {
|
||||
return 0
|
||||
}
|
||||
v, err := strconv.ParseInt(raw, 10, 64)
|
||||
if err != nil || v < 0 {
|
||||
return 0
|
||||
}
|
||||
return v
|
||||
}
|
||||
@@ -0,0 +1,459 @@
|
||||
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/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 {
|
||||
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
|
||||
}
|
||||
|
||||
func metaValue(t *testing.T, kvs []middleware.KV, key string) (string, bool) {
|
||||
t.Helper()
|
||||
for _, kv := range kvs {
|
||||
if kv.Key == key {
|
||||
return kv.Value, true
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
func buildMiddleware(t *testing.T, raw []byte) middleware.Middleware {
|
||||
t.Helper()
|
||||
mw, err := Factory{}.New(raw)
|
||||
require.NoError(t, err, "factory must accept the supplied config")
|
||||
return mw
|
||||
}
|
||||
|
||||
func TestMiddleware_StaticSurface(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
assert.Equal(t, ID, mw.ID(), "ID must match the registered constant")
|
||||
assert.Equal(t, Version, mw.Version(), "Version must match the constant")
|
||||
assert.Equal(t, middleware.SlotOnResponse, mw.Slot(), "must run in the response slot")
|
||||
assert.Empty(t, mw.AcceptedContentTypes(), "cost_meter does not inspect bodies")
|
||||
assert.False(t, mw.MutationsSupported(), "cost_meter never mutates")
|
||||
assert.NoError(t, mw.Close(), "Close on stateless middleware is a no-op")
|
||||
|
||||
keys := mw.MetadataKeys()
|
||||
expected := []string{middleware.KeyCostUSDTotal, middleware.KeyCostSkipped}
|
||||
assert.Equal(t, expected, keys, "metadata key allowlist must match the spec")
|
||||
}
|
||||
|
||||
func TestFactory_AcceptsEmptyAndJSONConfig(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
cases := [][]byte{nil, {}, []byte("null"), []byte("{}"), []byte(" ")}
|
||||
for _, raw := range cases {
|
||||
mw, err := Factory{}.New(raw)
|
||||
require.NoError(t, err, "empty/null/object config must be accepted")
|
||||
require.NotNil(t, mw, "factory must return a middleware instance")
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "openai"},
|
||||
{Key: middleware.KeyLLMModel, Value: "gpt-4o-mini"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "1000"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "1000"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, middleware.DecisionAllow, out.Decision, "cost_meter always allows")
|
||||
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok, "cost.usd_total must be emitted for known model")
|
||||
assert.Equal(t, "0.000750", value, "0.00015 + 0.0006 per 1k tokens, 6-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"})
|
||||
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.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.015000", value, "2*0.0025 + 1*0.01 = 0.015 with 6-decimal format")
|
||||
}
|
||||
|
||||
func TestInvoke_ComputesCostForKnownModel(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
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.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")
|
||||
assert.Equal(t, "0.018000", value, "0.003 + 0.015 = 0.018 with 6-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)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "10"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "10"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostSkipped)
|
||||
require.True(t, ok, "cost.skipped must be set when provider is missing")
|
||||
assert.Equal(t, skipMissingProvider, value, "skip reason matches missing_provider")
|
||||
}
|
||||
|
||||
func TestInvoke_MissingModel(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "openai"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "10"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "10"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostSkipped)
|
||||
require.True(t, ok, "cost.skipped must be set when model is missing")
|
||||
assert.Equal(t, skipMissingModel, value, "skip reason matches missing_model")
|
||||
}
|
||||
|
||||
func TestInvoke_MissingTokens(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
md []middleware.KV
|
||||
}{
|
||||
{
|
||||
name: "input only",
|
||||
md: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "openai"},
|
||||
{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "10"},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "output only",
|
||||
md: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "openai"},
|
||||
{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "10"},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "neither",
|
||||
md: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "openai"},
|
||||
{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{Metadata: tc.md})
|
||||
require.NoError(t, err)
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostSkipped)
|
||||
require.True(t, ok, "cost.skipped must be set when token keys are missing")
|
||||
assert.Equal(t, skipMissingTokens, value, "skip reason matches missing_tokens")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvoke_UnparseableTokens(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
in string
|
||||
out string
|
||||
}{
|
||||
{name: "input non-numeric", in: "abc", out: "10"},
|
||||
{name: "output non-numeric", in: "10", out: "xyz"},
|
||||
{name: "both garbage", in: "??", out: "??"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.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.KeyLLMInputTokens, Value: tc.in},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: tc.out},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostSkipped)
|
||||
require.True(t, ok, "cost.skipped must be set on unparseable tokens")
|
||||
assert.Equal(t, skipUnparseableTokens, value, "skip reason matches unparseable_tokens")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvoke_ZeroTokens(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
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: "0"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "0"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostSkipped)
|
||||
require.True(t, ok, "cost.skipped must be set when both token counts are zero")
|
||||
assert.Equal(t, skipZeroTokens, value, "skip reason matches zero_tokens")
|
||||
_, hasCost := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
assert.False(t, hasCost, "cost.usd_total must not be emitted for zero tokens")
|
||||
}
|
||||
|
||||
func TestInvoke_UnknownModel(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "openai"},
|
||||
{Key: middleware.KeyLLMModel, Value: "fantasy-model-9000"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "10"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "10"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostSkipped)
|
||||
require.True(t, ok, "cost.skipped must be set when pricing entry is absent")
|
||||
assert.Equal(t, skipUnknownModel, value, "skip reason matches unknown_model")
|
||||
}
|
||||
|
||||
func TestInvoke_NilInput(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, out, "output must be returned even on nil input")
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "decision must be allow on nil input")
|
||||
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)
|
||||
|
||||
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: "500"},
|
||||
{Key: middleware.KeyLLMCachedInputTokens, Value: "750"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, middleware.DecisionAllow, out.Decision)
|
||||
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok, "cached subset path must produce a cost — never a skip")
|
||||
// 250 non-cached at 0.0025/1k + 750 cached at 0.00125/1k + 500 output at 0.01/1k.
|
||||
assert.Equal(t, "0.006563", value,
|
||||
"cached subset must be billed at the discount rate, non-cached at the full rate; never double-billed")
|
||||
}
|
||||
|
||||
// TestInvoke_AnthropicCacheBucketsAdditive proves the Anthropic
|
||||
// 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)
|
||||
|
||||
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.KeyLLMInputTokens, Value: "256"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "200"},
|
||||
{Key: middleware.KeyLLMCachedInputTokens, Value: "768"},
|
||||
{Key: middleware.KeyLLMCacheCreationTokens, Value: "512"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, middleware.DecisionAllow, out.Decision)
|
||||
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok)
|
||||
// 256 input * 0.003 + 768 cache_read * 0.0003 + 512 cache_creation * 0.00375 + 200 output * 0.015
|
||||
// = 0.000768 + 0.0002304 + 0.00192 + 0.003 = 0.0059184 → "0.005918" with 6-decimal format.
|
||||
assert.Equal(t, "0.005918", value,
|
||||
"each Anthropic input bucket must bill at its own rate — cache_read cheap, cache_creation expensive, regular input mid")
|
||||
}
|
||||
|
||||
// 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.
|
||||
func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) {
|
||||
configureBuiltinWithCacheRates(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
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: "500"},
|
||||
// No KeyLLMCachedInputTokens — the parser didn't see one.
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
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.007500", value, "no cached metadata = same cost as before the feature landed")
|
||||
}
|
||||
|
||||
// TestInvoke_UnparseableCachedTokensSkippedSilently proves the
|
||||
// optional-bucket contract: a malformed cached_input_tokens metadata
|
||||
// value falls back to 0 (= no cached count) and continues with the
|
||||
// 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)
|
||||
|
||||
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: "500"},
|
||||
{Key: middleware.KeyLLMCachedInputTokens, Value: "not-a-number"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok, "garbage cache metadata must NOT switch the response from a cost to a skip — fall back to 0 cached")
|
||||
assert.Equal(t, "0.007500", 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.
|
||||
func TestMiddleware_CloseNilSafe(t *testing.T) {
|
||||
require.NoError(t, newMiddleware(nil, nil).Close(), "no-reloader Close must be a no-op")
|
||||
var m *Middleware
|
||||
require.NoError(t, m.Close(), "nil-receiver Close must be safe")
|
||||
}
|
||||
Reference in New Issue
Block a user