Compare commits

...

5 Commits

38 changed files with 2665 additions and 246 deletions

View File

@@ -5,6 +5,13 @@ on:
schedule:
- cron: "0 3 * * *"
workflow_dispatch:
inputs:
bedrock_model:
description: >-
Bedrock inference-profile id to drive the matrix with, exactly as
AWS issues it. Leave empty for the Sonnet 4.6 default.
required: false
default: ""
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
@@ -62,6 +69,8 @@ jobs:
CLOUDFLARE_TOKEN: ${{ secrets.E2E_CLOUDFLARE_TOKEN }}
AWS_BEARER_TOKEN_BEDROCK: ${{ secrets.E2E_AWS_BEARER_TOKEN_BEDROCK }}
AWS_REGION: ${{ secrets.E2E_AWS_REGION }}
# Bedrock model override: dispatch input wins, then the repo variable, else the test default.
AWS_BEDROCK_MODEL: ${{ inputs.bedrock_model || vars.E2E_AWS_BEDROCK_MODEL }}
# Vertex (Anthropic-on-Vertex): SA + project required; region defaults
# to "global", model to a pinned claude snapshot.
GOOGLE_VERTEX_SA_BASE64: ${{ secrets.E2E_GOOGLE_VERTEX_SA_BASE64 }}

View File

@@ -3,6 +3,7 @@ package dns
import (
"context"
"fmt"
"maps"
"math"
"net"
"slices"
@@ -31,6 +32,26 @@ type SubdomainMatcher interface {
MatchSubdomains() bool
}
// responseMeta holds the annotations handlers attach to a request to explain the
// response the chain ends up writing. It survives a deferral, so an answer that
// did not come from the handler that owns the name still says who stepped aside
// and why.
type responseMeta map[resutil.MetaKey]string
// format renders the annotations for the response log line. The order is stable
// so the same event reads the same way every time; a map's own order is not.
func (m responseMeta) format() string {
if len(m) == 0 {
return ""
}
var b strings.Builder
for _, k := range slices.Sorted(maps.Keys(m)) {
b.WriteString(" " + string(k) + "=" + m[k])
}
return b.String()
}
type HandlerEntry struct {
Handler dns.Handler
Priority int
@@ -52,8 +73,19 @@ type ResponseWriterChain struct {
origPattern string
requestID string
shouldContinue bool
response *dns.Msg
meta map[string]string
// softNegative suppresses a poisoning negative verdict for this request. A
// handler that owns the name but cannot answer this query type sets it
// before deferring, and it stays set for every handler that runs after.
softNegative bool
// clientHasEdns records whether the original query advertised EDNS0, taken
// before any handler ran: handlers add EDNS0 to the query they forward
// upstream, so the message itself no longer answers the question later.
clientHasEdns bool
response *dns.Msg
// meta is handed to the next handler when this one defers, so the same map
// outlives the writer that created it. Handlers must only set metadata from
// within ServeDNS, never from a goroutine that outlives the call.
meta responseMeta
}
// RequestID returns the request ID for tracing
@@ -62,26 +94,64 @@ func (w *ResponseWriterChain) RequestID() string {
}
// SetMeta sets a metadata key-value pair for logging
func (w *ResponseWriterChain) SetMeta(key, value string) {
func (w *ResponseWriterChain) SetMeta(key resutil.MetaKey, value string) {
if w.meta == nil {
w.meta = make(map[string]string)
w.meta = make(responseMeta)
}
w.meta[key] = value
}
// RequestSoftNegative marks the request so a downstream NXDOMAIN is turned into
// NODATA before it reaches the client. Set by a handler that defers a query for
// a name it owns but a type it cannot resolve.
func (w *ResponseWriterChain) RequestSoftNegative() {
w.softNegative = true
}
func (w *ResponseWriterChain) WriteMsg(m *dns.Msg) error {
// Check if this is a continue signal (NXDOMAIN with Zero bit set)
if m.Rcode == dns.RcodeNameError && m.MsgHdr.Zero {
w.shouldContinue = true
return nil
}
if w.softNegative && m.Rcode == dns.RcodeNameError {
m = softenNegative(m, w.clientHasEdns)
w.SetMeta(resutil.MetaKeySoftened, "nxdomain->nodata")
}
w.response = m
if m.MsgHdr.Truncated {
w.SetMeta("truncated", "true")
w.SetMeta(resutil.MetaKeyTruncated, "true")
}
return w.ResponseWriter.WriteMsg(m)
}
// softenNegative downgrades an NXDOMAIN to NODATA for a request a handler that
// owns the name deferred. NXDOMAIN is cached for the name and every type below
// it, so a resolver that has never heard of a routed name would take out the
// record types the route does serve; NODATA is cached for this name and type
// only. The authority section goes with it: the negative TTL of a zone we just
// overruled does not apply, and RFC 2308 keeps a negative answer that carries no
// SOA out of caches altogether, so the rewrite cannot outlive the route.
//
// The rewrite is ours, not the answering resolver's, so an EDNS0 client is told
// as much: the reply we hand back travels a path no capture on this host sees,
// and an empty answer is otherwise indistinguishable from a real one. Returns a
// copy so the answering handler keeps whatever it may hold on to.
func softenNegative(m *dns.Msg, clientHasEdns bool) *dns.Msg {
out := m.Copy()
out.Rcode = dns.RcodeSuccess
out.Ns = nil
if clientHasEdns {
resutil.AttachEDE(out, resutil.EDENetbirdSoftenedNegative,
"netbird: name is served locally, NXDOMAIN from the fallthrough resolver suppressed")
} else {
resutil.StripOPT(out)
}
return out
}
func NewHandlerChain() *HandlerChain {
return &HandlerChain{
handlers: make([]HandlerEntry, 0),
@@ -223,6 +293,19 @@ func (c *HandlerChain) dispatch(w dns.ResponseWriter, r *dns.Msg, maxPriority in
handlers := slices.Clone(c.handlers)
c.mu.RUnlock()
// Carried across deferrals: once a handler that owns the name defers, no
// handler after it may answer with a verdict that poisons the name. The
// metadata of a handler that stepped aside is carried too, so the response
// log line explains an answer that did not come from the handler that owns
// the name.
var softNegative bool
var carried responseMeta
// Taken before any handler runs: handlers advertise EDNS0 on the query they
// forward upstream, so afterwards the message no longer tells us whether the
// client did.
clientHasEdns := r.IsEdns0() != nil
// Try handlers in priority order
for _, entry := range handlers {
if entry.Priority > maxPriority {
@@ -245,11 +328,16 @@ func (c *HandlerChain) dispatch(w dns.ResponseWriter, r *dns.Msg, maxPriority in
ResponseWriter: w,
origPattern: entry.OrigPattern,
requestID: requestID,
softNegative: softNegative,
clientHasEdns: clientHasEdns,
meta: carried,
}
entry.Handler.ServeDNS(chainWriter, r)
// If handler wants to continue, try next handler
if chainWriter.shouldContinue {
softNegative = softNegative || chainWriter.softNegative
carried = chainWriter.meta
if entry.Priority != PriorityMgmtCache {
logger.Tracef("handler requested continue for domain=%s", qname)
}
@@ -265,30 +353,59 @@ func (c *HandlerChain) dispatch(w dns.ResponseWriter, r *dns.Msg, maxPriority in
qname, dns.TypeToString[question.Qtype], dns.ClassToString[question.Qclass])
resp := &dns.Msg{}
resp.SetRcode(r, dns.RcodeRefused)
// A handler that owns the name deferred and nothing below it could answer
// (a client with no primary nameserver group). The name exists as far as
// this client is concerned, since the route serves its addresses, so REFUSED
// would contradict the route: a stub that takes it as "not served here" and
// retries another resolver can come back with an NXDOMAIN that takes the
// whole name down. Answer "no records of this type" instead, without an SOA,
// so RFC 2308 keeps it out of negative caches, and tell an EDNS0 client the
// empty answer is ours rather than a resolver's.
if softNegative {
resp.Rcode = dns.RcodeSuccess
if clientHasEdns {
resutil.AttachEDE(resp, resutil.EDENetbirdSoftenedNegative,
"netbird: name is served locally, no fallthrough resolver for this query type")
}
// logResponse never runs on this path, so the carried metadata is
// appended here or the reason for the deferral is lost in exactly the
// case that is hardest to diagnose.
logger.Tracef("no handler below the deferring one for domain=%s type=%s, answering NODATA%s",
qname, dns.TypeToString[question.Qtype], carried.format())
}
if err := w.WriteMsg(resp); err != nil {
logger.Errorf("failed to write DNS response: %v", err)
}
}
func (c *HandlerChain) logResponse(logger *log.Entry, cw *ResponseWriterChain, qname string, startTime time.Time) {
// Runs for every query, and the arguments below are not free: Len() packs
// the message to measure it, and formatting the answers and the metadata
// allocates. None of it is worth doing when the line is discarded.
if !log.IsLevelEnabled(log.TraceLevel) {
return
}
if cw.response == nil {
return
}
var meta string
for k, v := range cw.meta {
meta += " " + k + "=" + v
}
logger.Tracef("response: domain=%s rcode=%s answers=%s size=%dB%s took=%s",
qname, dns.RcodeToString[cw.response.Rcode], resutil.FormatAnswers(cw.response.Answer),
cw.response.Len(), meta, time.Since(startTime))
cw.response.Len(), cw.meta.format(), time.Since(startTime))
}
// ResolveInternal runs an in-process DNS query against the chain, skipping any
// handler with priority > maxPriority. Used by internal callers (e.g. the mgmt
// cache refresher) that must bypass themselves to avoid loops. Honors ctx
// cancellation; on ctx.Done the dispatch goroutine is left to drain on its own
// cache refresher) that must bypass themselves to avoid loops.
//
// "Nothing answered" is read off RcodeRefused, which a request soft-negatived by
// a deferring handler never carries: it ends in an empty NOERROR instead, and
// would look resolved. No caller can reach that today, since every handler that
// defers sits above the maxPriority any caller passes. Lowering one below it
// means this check needs the soft-negative case too.
//
// Honors ctx cancellation; on ctx.Done the dispatch goroutine is left to drain on its own
// (bounded by the invoked handler's internal timeout).
func (c *HandlerChain) ResolveInternal(ctx context.Context, r *dns.Msg, maxPriority int) (*dns.Msg, error) {
if len(r.Question) == 0 {

View File

@@ -3,15 +3,19 @@ package dns_test
import (
"context"
"net"
"strings"
"testing"
"time"
"github.com/miekg/dns"
log "github.com/sirupsen/logrus"
logtest "github.com/sirupsen/logrus/hooks/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"github.com/stretchr/testify/require"
nbdns "github.com/netbirdio/netbird/client/internal/dns"
"github.com/netbirdio/netbird/client/internal/dns/resutil"
"github.com/netbirdio/netbird/client/internal/dns/test"
)
@@ -1238,6 +1242,276 @@ func TestHandlerChain_ResolveInternal_HonorsContextTimeout(t *testing.T) {
assert.Less(t, elapsed, 500*time.Millisecond, "ResolveInternal must return shortly after ctx deadline")
}
// requestSoftNegative asks the chain to soften a negative verdict produced by
// the handlers that run after the caller defers, reporting whether the chain
// supports the signal. Written as a type assertion so these tests compile
// against a chain that does not support it yet.
func requestSoftNegative(w dns.ResponseWriter) bool {
sn, ok := w.(interface{ RequestSoftNegative() })
if ok {
sn.RequestSoftNegative()
}
return ok
}
// deferringHandler defers to the next handler in the chain, optionally asking
// for the negative verdict of whatever answers instead to be softened. This is
// what a DNS route handler does for a record type its routing peer cannot
// resolve.
type deferringHandler struct {
softNegative bool
called bool
supported bool
}
func (h *deferringHandler) ServeDNS(w dns.ResponseWriter, r *dns.Msg) {
h.called = true
if h.softNegative {
h.supported = requestSoftNegative(w)
// The real handler records why it stepped aside, for the response log
// line of whichever handler answers instead.
resutil.SetMeta(w, resutil.MetaKeyDeferredBy, "test handler")
}
resp := new(dns.Msg)
resp.SetRcode(r, dns.RcodeNameError)
resp.MsgHdr.Zero = true
_ = w.WriteMsg(resp)
}
// nxdomainHandler answers an authoritative NXDOMAIN with an SOA in the
// authority section, the way a public resolver answers for a name that only
// exists inside the routed network.
type nxdomainHandler struct{}
func (h *nxdomainHandler) ServeDNS(w dns.ResponseWriter, r *dns.Msg) {
resp := new(dns.Msg)
resp.SetRcode(r, dns.RcodeNameError)
resp.Ns = []dns.RR{&dns.SOA{
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeSOA, Class: dns.ClassINET, Ttl: 3600},
Ns: "ns1.example.com.",
Mbox: "hostmaster.example.com.",
Minttl: 3600,
Expire: 604800,
Refresh: 7200,
Retry: 3600,
}}
_ = w.WriteMsg(resp)
}
// TestHandlerChain_SoftNegative_DowngradesDownstreamNXDOMAIN is the whole point
// of the soft-negative signal: a route handler may only defer a query to the
// public chain if the answer cannot poison the routed name. NXDOMAIN is cached
// for the name and every type under it (RFC 2308, RFC 8020), so it has to be
// rewritten to NODATA, which is cached per name and type only.
func TestHandlerChain_SoftNegative_DowngradesDownstreamNXDOMAIN(t *testing.T) {
chain := nbdns.NewHandlerChain()
route := &deferringHandler{softNegative: true}
chain.AddHandler("*.example.com.", route, nbdns.PriorityDNSRoute)
chain.AddHandler(".", &nxdomainHandler{}, nbdns.PriorityDefault)
r := new(dns.Msg)
r.SetQuestion("_mongodb._tcp.db.example.com.", dns.TypeSRV)
mw := &test.MockResponseWriter{}
chain.ServeDNS(&nbdns.ResponseWriterChain{ResponseWriter: mw}, r)
resp := mw.GetLastResponse()
require.NotNil(t, resp, "a response must reach the client")
assert.True(t, route.called, "the route handler must run first")
require.True(t, route.supported, "the chain writer must accept a soft-negative request")
assert.Equal(t, dns.RcodeSuccess, resp.Rcode, "NXDOMAIN must be softened to NODATA")
assert.Empty(t, resp.Answer, "a softened negative carries no answer")
assert.Empty(t, resp.Ns, "the downstream zone's SOA must not set the negative TTL for a name we overrode")
}
// responseLineFor returns the chain's response log line for one query. Raising
// the level to trace also unmutes whatever else is logging in this package, so
// the line has to be picked by the name it was asked about rather than by being
// the last one seen.
func responseLineFor(hook *logtest.Hook, qname string) string {
for _, e := range hook.AllEntries() {
if strings.HasPrefix(e.Message, "response:") && strings.Contains(e.Message, qname) {
return e.Message
}
}
return ""
}
// TestHandlerChain_SoftNegative_IsVisibleToTheClient covers observability of the
// rewrite. The reply we hand the application travels over loopback, which the
// bundle capture does not see, so a softened verdict has to say so on the wire:
// without it an empty answer is indistinguishable from a real "no such record"
// in a dig output or a capture taken next to the application.
func TestHandlerChain_SoftNegative_IsVisibleToTheClient(t *testing.T) {
// One hook for the whole test: logtest installs it on the standard logger
// and logrus has no way to take it off again, so a hook per subtest would
// leave several behind buffering every later line in the package.
hook := logtest.NewGlobal()
t.Cleanup(hook.Reset)
newChain := func() *nbdns.HandlerChain {
chain := nbdns.NewHandlerChain()
chain.AddHandler("*.example.com.", &deferringHandler{softNegative: true}, nbdns.PriorityDNSRoute)
chain.AddHandler(".", &nxdomainHandler{}, nbdns.PriorityDefault)
return chain
}
t.Run("EDNS0 client gets an extended error", func(t *testing.T) {
r := new(dns.Msg)
r.SetQuestion("_mongodb._tcp.db.example.com.", dns.TypeSRV)
r.SetEdns0(dns.DefaultMsgSize, false)
mw := &test.MockResponseWriter{}
newChain().ServeDNS(&nbdns.ResponseWriterChain{ResponseWriter: mw}, r)
resp := mw.GetLastResponse()
require.NotNil(t, resp)
require.Equal(t, dns.RcodeSuccess, resp.Rcode)
ede, ok := resutil.ExtractEDE(resp)
require.True(t, ok, "a softened verdict must carry an extended DNS error")
assert.Equal(t, resutil.EDENetbirdSoftenedNegative, ede.InfoCode)
assert.Contains(t, ede.ExtraText, "netbird", "the text must name us as the source of the rewrite")
})
t.Run("plain client gets no OPT", func(t *testing.T) {
r := new(dns.Msg)
r.SetQuestion("_mongodb._tcp.db.example.com.", dns.TypeSRV)
mw := &test.MockResponseWriter{}
newChain().ServeDNS(&nbdns.ResponseWriterChain{ResponseWriter: mw}, r)
resp := mw.GetLastResponse()
require.NotNil(t, resp)
assert.Nil(t, resp.IsEdns0(), "RFC 6891 forbids an OPT toward a client that did not advertise EDNS0")
})
// The response line is what support reads out of a debug bundle, and it now
// carries several annotations at once. Built from a map, their order would
// differ on every query, so the same event never looks the same twice.
t.Run("log fields keep a stable order", func(t *testing.T) {
const qname = "_mongodb._tcp.stable.example.com."
prev := log.GetLevel()
log.SetLevel(log.TraceLevel)
t.Cleanup(func() { log.SetLevel(prev) })
lineFor := func() string {
hook.Reset()
r := new(dns.Msg)
r.SetQuestion(qname, dns.TypeSRV)
mw := &test.MockResponseWriter{}
newChain().ServeDNS(&nbdns.ResponseWriterChain{ResponseWriter: mw}, r)
line := responseLineFor(hook, qname)
// The duration differs per query and is not what we compare.
line, _, _ = strings.Cut(line, " took=")
return line
}
first := lineFor()
require.NotEmpty(t, first, "the chain must log the response it wrote")
require.Contains(t, first, "deferred_by=", "the line must carry more than one annotation to be worth ordering")
require.Contains(t, first, "softened=")
for range 20 {
assert.Equal(t, first, lineFor(), "the same event must produce the same line")
}
})
t.Run("logged with the reason it was deferred", func(t *testing.T) {
const qname = "_mongodb._tcp.reason.example.com."
hook.Reset()
prev := log.GetLevel()
log.SetLevel(log.TraceLevel)
t.Cleanup(func() { log.SetLevel(prev) })
r := new(dns.Msg)
r.SetQuestion(qname, dns.TypeSRV)
mw := &test.MockResponseWriter{}
newChain().ServeDNS(&nbdns.ResponseWriterChain{ResponseWriter: mw}, r)
response := responseLineFor(hook, qname)
require.NotEmpty(t, response, "the chain must log the response it wrote")
assert.Contains(t, response, "softened=", "the log must show the verdict was rewritten")
assert.Contains(t, response, "deferred_by=", "the log must name the handler that deferred")
})
}
// TestHandlerChain_SoftNegative_NoHandlerBelow covers a client with no primary
// nameserver group: the deferred query reaches the end of the chain unanswered.
// REFUSED would say the name is not served here while the route serves its
// addresses, and a stub that acts on that by asking elsewhere can bring back an
// NXDOMAIN for the whole name. The answer must be an empty, uncacheable NODATA.
func TestHandlerChain_SoftNegative_NoHandlerBelow(t *testing.T) {
chain := nbdns.NewHandlerChain()
route := &deferringHandler{softNegative: true}
chain.AddHandler("*.example.com.", route, nbdns.PriorityDNSRoute)
r := new(dns.Msg)
r.SetQuestion("db.example.com.", dns.TypeHTTPS)
mw := &test.MockResponseWriter{}
chain.ServeDNS(&nbdns.ResponseWriterChain{ResponseWriter: mw}, r)
resp := mw.GetLastResponse()
require.NotNil(t, resp, "a response must reach the client")
assert.Equal(t, dns.RcodeSuccess, resp.Rcode, "an unanswered soft-negative query must be NODATA, not REFUSED")
assert.Empty(t, resp.Answer)
assert.Empty(t, resp.Ns,
"no SOA, so RFC 2308 keeps the empty answer out of negative caches and it cannot outlive the route")
}
// TestHandlerChain_SoftNegative_KeepsRealAnswers guards the other direction:
// softening applies to negative verdicts only. A real answer from a downstream
// handler must reach the client untouched.
func TestHandlerChain_SoftNegative_KeepsRealAnswers(t *testing.T) {
chain := nbdns.NewHandlerChain()
route := &deferringHandler{softNegative: true}
chain.AddHandler("*.example.com.", route, nbdns.PriorityDNSRoute)
chain.AddHandler(".", &answeringHandler{name: "public", ip: "203.0.113.10"}, nbdns.PriorityDefault)
r := new(dns.Msg)
r.SetQuestion("db.example.com.", dns.TypeA)
mw := &test.MockResponseWriter{}
chain.ServeDNS(&nbdns.ResponseWriterChain{ResponseWriter: mw}, r)
resp := mw.GetLastResponse()
require.NotNil(t, resp, "a response must reach the client")
assert.Equal(t, dns.RcodeSuccess, resp.Rcode)
require.Len(t, resp.Answer, 1, "the downstream answer must pass through")
}
// TestHandlerChain_NXDOMAINPreservedWithoutSoftNegative makes sure the
// softening is opt-in: an ordinary chain continuation still yields NXDOMAIN, so
// genuine non-existence keeps being reported.
func TestHandlerChain_NXDOMAINPreservedWithoutSoftNegative(t *testing.T) {
chain := nbdns.NewHandlerChain()
route := &deferringHandler{}
chain.AddHandler("*.example.com.", route, nbdns.PriorityDNSRoute)
chain.AddHandler(".", &nxdomainHandler{}, nbdns.PriorityDefault)
r := new(dns.Msg)
r.SetQuestion("nope.example.com.", dns.TypeA)
mw := &test.MockResponseWriter{}
chain.ServeDNS(&nbdns.ResponseWriterChain{ResponseWriter: mw}, r)
resp := mw.GetLastResponse()
require.NotNil(t, resp, "a response must reach the client")
assert.Equal(t, dns.RcodeNameError, resp.Rcode, "without the signal a real NXDOMAIN must survive")
}
func TestHandlerChain_HasRootHandlerAtOrBelow(t *testing.T) {
chain := nbdns.NewHandlerChain()
h := &answeringHandler{name: "h", ip: "10.0.0.1"}

View File

@@ -19,6 +19,34 @@ import (
// uses when a resolved host has no addresses of the requested family.
const errNoSuitableAddress = "no suitable address found"
// Extended DNS Error info codes NetBird emits so a client can see why an answer
// looks the way it does without reading this peer's logs. They live in the RFC
// 8914 Private Use range (49152-65535) and are registered here, in one place,
// because nothing else guarantees two NetBird components pick distinct codes.
const (
// EDENetbirdUpstreamTimeout: a DNS forwarder's upstream did not answer.
EDENetbirdUpstreamTimeout uint16 = 49152
// EDENetbirdUpstreamFailure: a DNS forwarder's upstream failed.
EDENetbirdUpstreamFailure uint16 = 49153
// EDENetbirdSoftenedNegative: the empty answer is ours, not the answering
// resolver's. A handler that owns the name deferred this query type, and the
// negative verdict that came back was downgraded so it cannot poison the
// name.
EDENetbirdSoftenedNegative uint16 = 49154
)
// AttachEDE adds an Extended DNS Error (RFC 8914) option to a message, creating
// the OPT pseudo-record if it has none. Callers must only use it toward a client
// that advertised EDNS0: per RFC 6891 an OPT must not appear otherwise.
func AttachEDE(msg *dns.Msg, code uint16, text string) {
opt := msg.IsEdns0()
if opt == nil {
msg.SetEdns0(dns.DefaultMsgSize, false)
opt = msg.IsEdns0()
}
opt.Option = append(opt.Option, &dns.EDNS0_EDE{InfoCode: code, ExtraText: text})
}
// GenerateRequestID creates a random 8-character hex string for request tracing.
func GenerateRequestID() string {
bytes := make([]byte, 4)
@@ -78,10 +106,38 @@ type resolver interface {
LookupNetIP(ctx context.Context, network, host string) ([]netip.Addr, error)
}
// MetaKey names an annotation a handler attaches to a request to explain the
// response the chain ends up writing. The set is closed and rendered onto a
// single log line, so an unrecognized key is a typo inventing a field rather
// than a new fact.
type MetaKey string
const (
// MetaKeyProtocol: the transport a query arrived on.
MetaKeyProtocol MetaKey = "protocol"
// MetaKeyUpstream: the upstream that answered.
MetaKeyUpstream MetaKey = "upstream"
// MetaKeyUpstreamProtocol: the transport used toward that upstream.
MetaKeyUpstreamProtocol MetaKey = "upstream_protocol"
// MetaKeyPeer: the routing peer whose forwarder answered.
MetaKeyPeer MetaKey = "peer"
// MetaKeyEDE: an Extended DNS Error carried by the answer.
MetaKeyEDE MetaKey = "ede"
// MetaKeyTruncated: the answer did not fit and was truncated.
MetaKeyTruncated MetaKey = "truncated"
// MetaKeySoftened: a negative verdict was rewritten so it cannot poison a
// name served locally.
MetaKeySoftened MetaKey = "softened"
// MetaKeyDeferredBy: the handler that owned the name and stepped aside.
MetaKeyDeferredBy MetaKey = "deferred_by"
// MetaKeyDeferredReason: why it stepped aside.
MetaKeyDeferredReason MetaKey = "deferred_reason"
)
// chainedWriter is implemented by ResponseWriters that carry request metadata
type chainedWriter interface {
RequestID() string
SetMeta(key, value string)
SetMeta(key MetaKey, value string)
}
// GetRequestID extracts a request ID from the ResponseWriter if available,
@@ -96,12 +152,30 @@ func GetRequestID(w dns.ResponseWriter) string {
}
// SetMeta sets metadata on the ResponseWriter if it supports it.
func SetMeta(w dns.ResponseWriter, key, value string) {
func SetMeta(w dns.ResponseWriter, key MetaKey, value string) {
if cw, ok := w.(chainedWriter); ok {
cw.SetMeta(key, value)
}
}
// softNegativeRequester is implemented by chain writers that can soften the
// negative verdict of the handlers a deferring handler falls through to.
type softNegativeRequester interface {
RequestSoftNegative()
}
// RequestSoftNegative asks the handler chain to downgrade an NXDOMAIN from the
// handlers that run after the caller defers. A handler that owns a name but
// cannot answer one query type for it needs this: the resolvers it falls
// through to cannot prove the name absent, and an NXDOMAIN from them is cached
// for the name and every type under it (RFC 2308, RFC 8020), taking the
// addresses the handler does serve down with it.
func RequestSoftNegative(w dns.ResponseWriter) {
if sn, ok := w.(softNegativeRequester); ok {
sn.RequestSoftNegative()
}
}
// LookupResult contains the result of an external DNS lookup
type LookupResult struct {
IPs []netip.Addr
@@ -199,11 +273,27 @@ type RecordResolver interface {
LookupAddr(ctx context.Context, addr string) ([]string, error)
}
// SupportedRecordQtype reports whether LookupRecords can resolve qtype. The
// set is bounded by the net.Resolver API, which exposes no way to query an
// arbitrary record type, and going around it with a raw exchange would mean
// picking nameservers ourselves instead of resolving the way the host does.
//
// Both ends read this: a DNS forwarder answers these types for the domains it
// routes, and a client uses it to tell which types are worth forwarding to a
// peer at all.
func SupportedRecordQtype(qtype uint16) bool {
switch qtype {
case dns.TypeMX, dns.TypeTXT, dns.TypeNS, dns.TypeSRV, dns.TypeCNAME, dns.TypePTR:
return true
default:
return false
}
}
// LookupRecords resolves a non-address DNS record type through the host
// resolver and returns the resource records and the DNS rcode. Types the host
// resolver cannot answer (anything not covered by the net.Resolver Lookup*
// methods) yield NODATA so that a routed name is never poisoned with NXDOMAIN
// for an unsupported type.
// resolver and returns the resource records and the DNS rcode. Types outside
// SupportedRecordQtype yield NODATA so that a routed name is never poisoned
// with NXDOMAIN for a type we cannot look up.
func LookupRecords(ctx context.Context, r RecordResolver, name string, qtype uint16, ttl uint32) ([]dns.RR, int) {
fqdn := dns.Fqdn(name)

View File

@@ -295,7 +295,7 @@ func (u *upstreamResolverBase) ServeDNS(w dns.ResponseWriter, r *dns.Msg) {
if addr := w.RemoteAddr(); addr != nil {
network := addr.Network()
ctx = contextWithDNSProtocol(ctx, network)
resutil.SetMeta(w, "protocol", network)
resutil.SetMeta(w, resutil.MetaKeyProtocol, network)
}
ok, failures := u.tryUpstreamServers(ctx, w, r, logger)
@@ -331,7 +331,7 @@ func (u *upstreamResolverBase) tryOnlyRace(ctx context.Context, w dns.ResponseWr
return false, res.failures
}
if res.ede != "" {
resutil.SetMeta(w, "ede", res.ede)
resutil.SetMeta(w, resutil.MetaKeyEDE, res.ede)
}
u.writeSuccessResponse(w, res.msg, res.upstream, r.Question[0].Name, res.protocol, logger)
return true, res.failures
@@ -361,7 +361,7 @@ func (u *upstreamResolverBase) raceAll(ctx context.Context, w dns.ResponseWriter
failures = append(failures, res.failures...)
if res.msg != nil {
if res.ede != "" {
resutil.SetMeta(w, "ede", res.ede)
resutil.SetMeta(w, resutil.MetaKeyEDE, res.ede)
}
u.writeSuccessResponse(w, res.msg, res.upstream, r.Question[0].Name, res.protocol, logger)
return true, failures
@@ -550,9 +550,9 @@ func (u *upstreamResolverBase) debugUpstreamTimeout(upstream netip.AddrPort) str
}
func (u *upstreamResolverBase) writeSuccessResponse(w dns.ResponseWriter, rm *dns.Msg, upstream netip.AddrPort, domain string, proto string, logger *log.Entry) {
resutil.SetMeta(w, "upstream", upstream.String())
resutil.SetMeta(w, resutil.MetaKeyUpstream, upstream.String())
if proto != "" {
resutil.SetMeta(w, "upstream_protocol", proto)
resutil.SetMeta(w, resutil.MetaKeyUpstreamProtocol, proto)
}
// Clear Zero bit from external responses to prevent upstream servers from

View File

@@ -26,15 +26,6 @@ import (
const errResolveFailed = "failed to resolve query for domain=%s: %v"
const upstreamTimeout = 15 * time.Second
// EDE info codes the forwarder emits on upstream failures so the querying
// client can see the reason without inspecting this peer's logs. They live in
// the RFC 8914 Private Use range (49152-65535); the Go resolver never exposes a
// real upstream EDE here, so these cannot collide with a genuine code.
const (
edeNetbirdUpstreamTimeout uint16 = 49152
edeNetbirdUpstreamFailure uint16 = 49153
)
type resolver interface {
LookupNetIP(ctx context.Context, network, host string) ([]netip.Addr, error)
LookupMX(ctx context.Context, name string) ([]*net.MX, error)
@@ -216,6 +207,9 @@ func (f *DNSForwarder) handleDNSQuery(logger *log.Entry, w dns.ResponseWriter, q
qname, dns.TypeToString[question.Qtype], dns.ClassToString[question.Qclass])
resp := query.SetReply(query)
// Every answer here comes from a recursive lookup on this peer. SetReply
// leaves RA unset, which reads to a client as a server that cannot recurse.
resp.RecursionAvailable = true
mostSpecificResId, matchingEntries := f.getMatchingEntries(strings.TrimSuffix(qname, "."))
if mostSpecificResId == "" {
@@ -229,20 +223,22 @@ func (f *DNSForwarder) handleDNSQuery(logger *log.Entry, w dns.ResponseWriter, q
reqHasEdns := query.IsEdns0() != nil
switch question.Qtype {
case dns.TypeA, dns.TypeAAAA:
switch {
case question.Qtype == dns.TypeA || question.Qtype == dns.TypeAAAA:
f.handleAddressQuery(ctx, logger, w, resp, mostSpecificResId, matchingEntries, reqHasEdns, startTime)
case dns.TypeMX, dns.TypeTXT, dns.TypeNS, dns.TypeSRV, dns.TypeCNAME, dns.TypePTR:
case resutil.SupportedRecordQtype(question.Qtype):
f.handleRecordQuery(ctx, logger, w, resp, startTime)
default:
// The domain is routed here, so any other type is answered NODATA
// (NOERROR, empty answer) rather than falling back to a resolver that
// would poison the name with NXDOMAIN. The Extended DNS Error lets a
// client tell this capability-driven NODATA apart from an
// authoritative one. The OPT pseudo-record must not appear unless the
// query advertised EDNS0.
// authoritative one; a current client knows the type is unsupported and
// resolves it through its own handler chain instead of asking at all.
// The OPT pseudo-record must not appear unless the query advertised
// EDNS0.
if reqHasEdns {
attachEDE(resp, dns.ExtendedErrorCodeNotSupported, "netbird forwarder: unsupported query type")
resutil.AttachEDE(resp, dns.ExtendedErrorCodeNotSupported, "netbird forwarder: unsupported query type")
}
f.writeResponse(logger, w, resp, qname, startTime)
}
@@ -441,7 +437,7 @@ func (f *DNSForwarder) handleDNSError(
}
if reqHasEdns {
attachEDE(resp, edeCodeFor(dnsErr), edeText(dnsErr))
resutil.AttachEDE(resp, edeCodeFor(dnsErr), edeText(dnsErr))
}
f.writeResponse(logger, w, resp, domain, startTime)
@@ -488,9 +484,9 @@ func (f *DNSForwarder) getMatchingEntries(domain string) (route.ResID, []*Forwar
// edeCodeFor maps an upstream lookup error to the NetBird EDE info code.
func edeCodeFor(dnsErr *net.DNSError) uint16 {
if dnsErr != nil && dnsErr.IsTimeout {
return edeNetbirdUpstreamTimeout
return resutil.EDENetbirdUpstreamTimeout
}
return edeNetbirdUpstreamFailure
return resutil.EDENetbirdUpstreamFailure
}
// edeText builds the EDE extra-text describing the class of upstream failure.
@@ -503,14 +499,3 @@ func edeText(dnsErr *net.DNSError) string {
}
return "netbird forwarder: upstream failure"
}
// attachEDE adds an Extended DNS Error (RFC 8914) option to the response,
// creating the OPT pseudo-record if the response does not already carry one.
func attachEDE(resp *dns.Msg, code uint16, text string) {
opt := resp.IsEdns0()
if opt == nil {
resp.SetEdns0(dns.DefaultMsgSize, false)
opt = resp.IsEdns0()
}
opt.Option = append(opt.Option, &dns.EDNS0_EDE{InfoCode: code, ExtraText: text})
}

View File

@@ -649,6 +649,39 @@ func TestDNSForwarder_ResponseCodes(t *testing.T) {
}
}
// TestDNSForwarder_RecursionAvailable covers the RA bit. Every answer this
// forwarder produces comes from a recursive lookup on the routing peer, but
// dns.Msg.SetReply leaves RA unset, so clients report "recursion not available"
// and some stub resolvers treat the server as unable to serve the query.
func TestDNSForwarder_RecursionAvailable(t *testing.T) {
t.Run("record answer", func(t *testing.T) {
mockResolver := &MockResolver{}
forwarder := newRecordTestForwarder(t, mockResolver, "example.com")
mockResolver.On("LookupMX", mock.Anything, "example.com.").
Return([]*net.MX{{Host: "mail.example.com.", Pref: 10}}, nil).Once()
resp := runRecordQuery(t, forwarder, "example.com", dns.TypeMX)
assert.True(t, resp.RecursionAvailable, "the forwarder resolves recursively")
})
t.Run("unsupported type NODATA", func(t *testing.T) {
forwarder := newRecordTestForwarder(t, &MockResolver{}, "example.com")
resp := runRecordQuery(t, forwarder, "example.com", dns.TypeCAA)
require.Equal(t, dns.RcodeSuccess, resp.Rcode)
assert.True(t, resp.RecursionAvailable, "RA describes the server, not the query type")
})
t.Run("unauthorized domain", func(t *testing.T) {
forwarder := newRecordTestForwarder(t, &MockResolver{}, "example.com")
resp := runRecordQuery(t, forwarder, "other.com", dns.TypeMX)
require.Equal(t, dns.RcodeRefused, resp.Rcode)
assert.True(t, resp.RecursionAvailable, "RA describes the server, not the verdict")
})
}
func hasEDE(m *dns.Msg, code uint16) bool {
opt := m.IsEdns0()
if opt == nil {
@@ -859,7 +892,7 @@ func TestDNSForwarder_UpstreamFailureEDE(t *testing.T) {
lookupErr: &net.DNSError{Err: "i/o timeout", Server: "10.0.0.53:53", IsTimeout: true},
reqEdns: true,
wantEDE: true,
wantCode: edeNetbirdUpstreamTimeout,
wantCode: resutil.EDENetbirdUpstreamTimeout,
wantTextHas: "netbird forwarder: upstream timeout",
},
{
@@ -867,7 +900,7 @@ func TestDNSForwarder_UpstreamFailureEDE(t *testing.T) {
lookupErr: &net.DNSError{Err: "server misbehaving", Server: "10.0.0.53:53"},
reqEdns: true,
wantEDE: true,
wantCode: edeNetbirdUpstreamFailure,
wantCode: resutil.EDENetbirdUpstreamFailure,
wantTextHas: "netbird forwarder: upstream failure",
},
{

View File

@@ -22,7 +22,6 @@ import (
nbdns "github.com/netbirdio/netbird/client/internal/dns"
"github.com/netbirdio/netbird/client/internal/dns/resutil"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/peerstore"
"github.com/netbirdio/netbird/client/internal/routemanager/common"
"github.com/netbirdio/netbird/client/internal/routemanager/fakeip"
iface "github.com/netbirdio/netbird/client/internal/routemanager/iface"
@@ -40,6 +39,61 @@ type internalDNATer interface {
AddInternalDNATMapping(netip.Addr, netip.Addr) error
}
// peerAllowedIPs reports the tunnel address a peer is reachable on.
type peerAllowedIPs interface {
AllowedIP(pubKey string) (netip.Addr, bool)
}
// disposition says where a query for a given record type has to be answered.
type disposition int
const (
// dispositionPeer: the routing peer owns the answer. Address records are
// what a DNS route exists for, and their answer programs the routes,
// allowed IPs and firewall sets, so they are never resolved anywhere else.
dispositionPeer disposition = iota
// dispositionPeerFirst: the peer's forwarder resolves the type, and the
// records may only exist inside the routed network, so it has to be asked.
// A peer running a client that predates that support cannot answer, which
// only its reply reveals.
dispositionPeerFirst
// dispositionChain: no forwarder resolves the type, so asking would cost a
// tunnel round trip for a reply that says nothing. Public records for the
// name are still better than none, so the query goes to the rest of the
// chain.
dispositionChain
)
func dispositionFor(qtype uint16) disposition {
switch {
case qtype == dns.TypeA || qtype == dns.TypeAAAA:
return dispositionPeer
case resutil.SupportedRecordQtype(qtype):
return dispositionPeerFirst
default:
return dispositionChain
}
}
// peerCannotAnswer reports whether a forwarder reply means the peer is unable to
// resolve this query type, as opposed to having resolved it and found nothing.
// A forwarder older than the one that learned to resolve non-address types
// answers NOTIMP to all of them, before it even looks at the domain; FORMERR
// covers one that cannot parse the EDNS0 we add to the query.
//
// REFUSED is deliberately not here. It says the peer does not hold the name,
// which happens while the client and the peer disagree about the route set, and
// falling through then would hand the name of an internal-only domain to a
// public resolver. No forwarder version reports a missing record type that way.
func peerCannotAnswer(rcode int) bool {
switch rcode {
case dns.RcodeNotImplemented, dns.RcodeFormatError:
return true
default:
return false
}
}
type DnsInterceptor struct {
mu sync.RWMutex
route *route.Route
@@ -50,7 +104,7 @@ type DnsInterceptor struct {
currentPeerKey string
interceptedDomains domainMap
wgInterface iface.WGIface
peerStore *peerstore.Store
peerStore peerAllowedIPs
firewall firewall.Manager
fakeIPManager *fakeip.Manager
forwarderPort *atomic.Uint32
@@ -226,11 +280,14 @@ func (d *DnsInterceptor) ServeDNS(w dns.ResponseWriter, r *dns.Msg) {
return
}
// All query types for an intercepted domain are forwarded to the peer's
// DNS forwarder, which owns the name. Falling through to the system
// resolver would let it answer NXDOMAIN for a name it isn't authoritative
// for, poisoning the whole name (including the A/AAAA records the route
// does serve). The forwarder answers NODATA for types it cannot resolve.
qtype := r.Question[0].Qtype
dispose := dispositionFor(qtype)
if dispose == dispositionChain {
d.deferToChain(w, r, logger, "unsupported-qtype")
return
}
d.mu.RLock()
peerKey := d.currentPeerKey
d.mu.RUnlock()
@@ -246,35 +303,33 @@ func (d *DnsInterceptor) ServeDNS(w dns.ResponseWriter, r *dns.Msg) {
return
}
if r.Extra == nil {
r.MsgHdr.AuthenticatedData = true
}
// Advertise EDNS0 to the forwarder so it may return an Extended DNS Error
// describing why a lookup failed. The OPT is stripped from the reply when
// the original client did not request EDNS0.
hadEdns := r.IsEdns0() != nil
if !hadEdns {
r.SetEdns0(dns.DefaultMsgSize, false)
}
query, hadEdns := peerQuery(r)
upstream := net.JoinHostPort(upstreamIP.String(), strconv.FormatUint(uint64(d.forwarderPort.Load()), 10))
ctx, cancel := context.WithTimeout(context.Background(), dnsTimeout)
defer cancel()
reply := d.queryUpstreamDNS(ctx, w, r, upstream, upstreamIP, peerKey, logger)
reply := d.queryUpstreamDNS(ctx, w, query, upstream, upstreamIP, peerKey, logger)
if reply == nil {
// queryUpstreamDNS already logged the failure and answered the client
return
}
// The peer owns the name but its forwarder cannot resolve this type. No
// capability is announced anywhere, so the round trip above is the probe.
if dispose == dispositionPeerFirst && peerCannotAnswer(reply.Rcode) {
d.deferToChain(w, r, logger, "peer-rcode-"+dns.RcodeToString[reply.Rcode])
return
}
if ede, ok := resutil.ExtractEDE(reply); ok {
resutil.SetMeta(w, "ede", fmt.Sprintf("%d %s", ede.InfoCode, ede.ExtraText))
resutil.SetMeta(w, resutil.MetaKeyEDE, fmt.Sprintf("%d %s", ede.InfoCode, ede.ExtraText))
}
if !hadEdns {
resutil.StripOPT(reply)
}
resutil.SetMeta(w, "peer", peerKey)
resutil.SetMeta(w, resutil.MetaKeyPeer, peerKey)
reply.Id = r.Id
if err := d.writeMsg(w, reply, logger); err != nil {
@@ -282,6 +337,30 @@ func (d *DnsInterceptor) ServeDNS(w dns.ResponseWriter, r *dns.Msg) {
}
}
// deferToChain hands the query to the next handler in the chain and asks the
// chain to soften the negative verdict of whatever answers instead. The
// resolvers below cannot prove a name inside the routed network absent, so their
// NXDOMAIN must not reach the client: it would be cached for the name and every
// type under it, taking the addresses this route does serve with it.
func (d *DnsInterceptor) deferToChain(w dns.ResponseWriter, r *dns.Msg, logger *log.Entry, reason string) {
logger.Tracef("continuing to next handler for domain=%s type=%s reason=%s",
r.Question[0].Name, dns.TypeToString[r.Question[0].Qtype], reason)
resutil.RequestSoftNegative(w)
// Carried to the chain's response log line: without it the answer looks like
// it came from the fallthrough resolver on its own.
resutil.SetMeta(w, resutil.MetaKeyDeferredBy, "dns-route")
resutil.SetMeta(w, resutil.MetaKeyDeferredReason, reason)
resp := new(dns.Msg)
resp.SetRcode(r, dns.RcodeNameError)
// Set Zero bit to signal handler chain to continue
resp.MsgHdr.Zero = true
if err := w.WriteMsg(resp); err != nil {
logger.Errorf("failed writing DNS continue response: %v", err)
}
}
func (d *DnsInterceptor) writeDNSError(w dns.ResponseWriter, r *dns.Msg, logger *log.Entry, reason string) {
logger.Warnf("failed to query upstream for domain=%s: %s", r.Question[0].Name, reason)
@@ -621,3 +700,26 @@ func (d *DnsInterceptor) debugPeerTimeout(peerIP netip.Addr, peerKey string) str
return fmt.Sprintf(" (peer %s)", nbdns.FormatPeerStatus(&peerState))
}
// peerQuery builds the query sent to a peer's DNS forwarder and reports whether
// the client itself advertised EDNS0. EDNS0 is added so the forwarder can return
// an Extended DNS Error describing an upstream failure, and the OPT is stripped
// from the reply again when the client did not ask for it. The client's message
// is left untouched: it may still be handed to the next handler in the chain,
// and neither the OPT nor the AD bit is ours to put on the wire on its behalf.
func peerQuery(r *dns.Msg) (query *dns.Msg, hadEdns bool) {
query = r.Copy()
hadEdns = query.IsEdns0() != nil
// AD tells the forwarder we understand authenticated data. Only set when the
// client sent no additional section of its own, so we never overrule what it
// asked for.
if len(query.Extra) == 0 {
query.MsgHdr.AuthenticatedData = true
}
if !hadEdns {
query.SetEdns0(dns.DefaultMsgSize, false)
}
return query, hadEdns
}

View File

@@ -0,0 +1,399 @@
package dnsinterceptor
import (
"net"
"net/netip"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/miekg/dns"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.zx2c4.com/wireguard/tun/netstack"
"github.com/netbirdio/netbird/client/iface/device"
"github.com/netbirdio/netbird/client/iface/wgaddr"
"github.com/netbirdio/netbird/client/internal/dns/resutil"
"github.com/netbirdio/netbird/client/internal/dns/test"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/route"
)
// softNegativeWriter records what the handler told the chain: whether to soften
// a negative verdict from the handlers it defers to, and the metadata that ends
// up on the chain's response log line.
type softNegativeWriter struct {
test.MockResponseWriter
softNegative bool
meta map[resutil.MetaKey]string
}
func (w *softNegativeWriter) RequestSoftNegative() { w.softNegative = true }
func (w *softNegativeWriter) RequestID() string { return "test" }
func (w *softNegativeWriter) SetMeta(key resutil.MetaKey, value string) {
if w.meta == nil {
w.meta = make(map[resutil.MetaKey]string)
}
w.meta[key] = value
}
// TestServeDNS_QtypeTheForwarderCannotResolve covers the record types no DNS
// forwarder can answer, because the host resolver exposes no API for them.
// Asking the peer only burns a tunnel round trip, so the query goes straight to
// the rest of the chain. No peer key is configured, proving no round trip is
// attempted.
func TestServeDNS_QtypeTheForwarderCannotResolve(t *testing.T) {
qtypes := []uint16{
dns.TypeHTTPS,
dns.TypeSVCB,
dns.TypeCAA,
dns.TypeNAPTR,
dns.TypeTLSA,
dns.TypeSOA,
}
for _, qtype := range qtypes {
t.Run(dns.TypeToString[qtype], func(t *testing.T) {
d := &DnsInterceptor{}
w := &softNegativeWriter{}
r := new(dns.Msg)
r.SetQuestion("db.example.com.", qtype)
d.ServeDNS(w, r)
resp := w.GetLastResponse()
require.NotNil(t, resp, "a response must be written")
assert.Equal(t, dns.RcodeNameError, resp.Rcode, "chain continuation is signalled as NXDOMAIN")
assert.True(t, resp.MsgHdr.Zero, "Zero bit must be set so the chain continues")
assert.True(t, w.softNegative,
"a downstream NXDOMAIN must be softened, or the fallthrough poisons the routed name")
assert.Equal(t, "dns-route", w.meta[resutil.MetaKeyDeferredBy],
"the chain's response log line must name us as the handler that stepped aside")
assert.Equal(t, "unsupported-qtype", w.meta[resutil.MetaKeyDeferredReason],
"the reason must survive to the log line, or the fallthrough is invisible")
})
}
}
// TestServeDNS_AddressQtypeNeverFallsThrough locks in the opposite case: A and
// AAAA are what the route exists for, and the peer is authoritative for them.
// A missing peer is an error the client must see, never a fallthrough to a
// resolver that knows nothing about the routed name.
func TestServeDNS_AddressQtypeNeverFallsThrough(t *testing.T) {
for _, qtype := range []uint16{dns.TypeA, dns.TypeAAAA} {
t.Run(dns.TypeToString[qtype], func(t *testing.T) {
d := &DnsInterceptor{}
w := &softNegativeWriter{}
r := new(dns.Msg)
r.SetQuestion("db.example.com.", qtype)
d.ServeDNS(w, r)
resp := w.GetLastResponse()
require.NotNil(t, resp, "a response must be written")
assert.Equal(t, dns.RcodeServerFailure, resp.Rcode, "an unusable route must fail, not fall through")
assert.False(t, resp.MsgHdr.Zero, "the chain must not continue for address queries")
assert.False(t, w.softNegative, "no fallthrough means nothing to soften")
})
}
}
// TestDispositionFor pins the query-type policy, and ties the types we ask a
// peer about to the types a forwarder can actually resolve, so the two ends
// cannot drift apart.
func TestDispositionFor(t *testing.T) {
tests := []struct {
qtype uint16
want disposition
}{
{dns.TypeA, dispositionPeer},
{dns.TypeAAAA, dispositionPeer},
{dns.TypeMX, dispositionPeerFirst},
{dns.TypeTXT, dispositionPeerFirst},
{dns.TypeNS, dispositionPeerFirst},
{dns.TypeSRV, dispositionPeerFirst},
{dns.TypeCNAME, dispositionPeerFirst},
{dns.TypePTR, dispositionPeerFirst},
{dns.TypeHTTPS, dispositionChain},
{dns.TypeSVCB, dispositionChain},
{dns.TypeCAA, dispositionChain},
{dns.TypeNAPTR, dispositionChain},
{dns.TypeTLSA, dispositionChain},
{dns.TypeDS, dispositionChain},
{dns.TypeDNSKEY, dispositionChain},
{dns.TypeSOA, dispositionChain},
{dns.TypeANY, dispositionChain},
}
for _, tt := range tests {
t.Run(dns.TypeToString[tt.qtype], func(t *testing.T) {
assert.Equal(t, tt.want, dispositionFor(tt.qtype), "disposition for %s", dns.TypeToString[tt.qtype])
// Address types are resolved by the forwarder's address path, not by
// its record path, so they are deliberately outside
// SupportedRecordQtype and their disposition does not depend on it.
if tt.want == dispositionPeer {
assert.False(t, resutil.SupportedRecordQtype(tt.qtype),
"address types take the forwarder's address path")
return
}
assert.Equal(t, tt.want == dispositionPeerFirst, resutil.SupportedRecordQtype(tt.qtype),
"only types a forwarder resolves may be sent to a peer")
})
}
}
// TestPeerCannotAnswer separates a peer that cannot resolve a type from one that
// resolved it and found nothing. Only the former may be retried elsewhere:
// treating NODATA or NXDOMAIN as a capability failure would send the client to a
// public resolver behind the peer's back and prefer a public record over the
// private one the route exists to reach.
func TestPeerCannotAnswer(t *testing.T) {
cannot := []int{dns.RcodeNotImplemented, dns.RcodeFormatError}
for _, rcode := range cannot {
assert.True(t, peerCannotAnswer(rcode), "%s means the peer cannot answer", dns.RcodeToString[rcode])
}
// REFUSED belongs here, not above: it says the peer does not hold the name,
// so retrying elsewhere would leak the name of an internal-only domain to a
// public resolver. No forwarder version reports a missing record type this
// way, so nothing is lost by keeping it out.
can := []int{
dns.RcodeSuccess,
dns.RcodeNameError,
dns.RcodeServerFailure,
dns.RcodeNotAuth,
dns.RcodeRefused,
}
for _, rcode := range can {
assert.False(t, peerCannotAnswer(rcode), "%s is the peer's answer, not a capability failure", dns.RcodeToString[rcode])
}
}
// fakeWGIface is the minimum an interceptor needs to reach a peer's DNS
// forwarder over a plain socket. A nil netstack keeps the exchange on the host
// stack so the test can point it at a loopback listener.
type fakeWGIface struct{}
func (fakeWGIface) AddAllowedIP(string, netip.Prefix) error { return nil }
func (fakeWGIface) RemoveAllowedIP(string, netip.Prefix) error { return nil }
func (fakeWGIface) Name() string { return "wt0" }
func (fakeWGIface) Address() wgaddr.Address { return wgaddr.Address{} }
func (fakeWGIface) ToInterface() *net.Interface { return nil }
func (fakeWGIface) IsUserspaceBind() bool { return false }
func (fakeWGIface) GetFilter() device.PacketFilter { return nil }
func (fakeWGIface) GetDevice() *device.FilteredDevice { return nil }
func (fakeWGIface) GetNet() *netstack.Net { return nil }
// fakePeerIPs resolves every peer key to the loopback address the test's
// forwarder stub listens on.
type fakePeerIPs struct {
addr netip.Addr
}
func (p fakePeerIPs) AllowedIP(string) (netip.Addr, bool) { return p.addr, p.addr.IsValid() }
// forwarderStub stands in for the DNS forwarder of a routing peer, recording
// what it was asked and answering with a canned reply.
type forwarderStub struct {
mu sync.Mutex
queries []*dns.Msg
reply func(q *dns.Msg) *dns.Msg
port uint16
shutdown func()
}
func (s *forwarderStub) received() []*dns.Msg {
s.mu.Lock()
defer s.mu.Unlock()
return append([]*dns.Msg(nil), s.queries...)
}
func newForwarderStub(t *testing.T, reply func(q *dns.Msg) *dns.Msg) *forwarderStub {
t.Helper()
conn, err := net.ListenUDP("udp", net.UDPAddrFromAddrPort(netip.MustParseAddrPort("127.0.0.1:0")))
require.NoError(t, err, "listen for the forwarder stub")
stub := &forwarderStub{reply: reply}
stub.port = uint16(conn.LocalAddr().(*net.UDPAddr).Port)
mux := dns.NewServeMux()
mux.HandleFunc(".", func(w dns.ResponseWriter, q *dns.Msg) {
stub.mu.Lock()
stub.queries = append(stub.queries, q.Copy())
stub.mu.Unlock()
_ = w.WriteMsg(stub.reply(q))
})
srv := &dns.Server{PacketConn: conn, Handler: mux}
started := make(chan struct{})
srv.NotifyStartedFunc = func() { close(started) }
go func() {
_ = srv.ActivateAndServe()
}()
select {
case <-started:
case <-time.After(5 * time.Second):
t.Fatal("forwarder stub did not start")
}
stub.shutdown = func() { _ = srv.Shutdown() }
t.Cleanup(stub.shutdown)
return stub
}
func newTestInterceptor(t *testing.T, stub *forwarderStub) *DnsInterceptor {
t.Helper()
port := new(atomic.Uint32)
port.Store(uint32(stub.port))
return &DnsInterceptor{
route: &route.Route{Domains: nil},
statusRecorder: peer.NewRecorder("https://mgm"),
currentPeerKey: "peer-key",
interceptedDomains: make(domainMap),
wgInterface: fakeWGIface{},
peerStore: fakePeerIPs{addr: netip.MustParseAddr("127.0.0.1")},
forwarderPort: port,
}
}
// TestServeDNS_PeerCannotResolveQtype is the reported regression: a routing peer
// running a client older than the one that taught the forwarder about non-address
// record types answers NOTIMP for every SRV query, and the interceptor used to
// hand that straight to the application, breaking SRV-based service discovery.
// The client cannot know the peer's capability up front, so it must try the peer
// and fall back to the chain when the peer says it cannot answer.
func TestServeDNS_PeerCannotResolveQtype(t *testing.T) {
for _, rcode := range []int{dns.RcodeNotImplemented, dns.RcodeFormatError} {
t.Run(dns.RcodeToString[rcode], func(t *testing.T) {
stub := newForwarderStub(t, func(q *dns.Msg) *dns.Msg {
resp := new(dns.Msg)
resp.SetRcode(q, rcode)
return resp
})
d := newTestInterceptor(t, stub)
w := &softNegativeWriter{}
r := new(dns.Msg)
r.SetQuestion("_mongodb._tcp.db.example.com.", dns.TypeSRV)
d.ServeDNS(w, r)
require.Len(t, stub.received(), 1, "the peer must be asked before giving up on it")
resp := w.GetLastResponse()
require.NotNil(t, resp, "a response must be written")
assert.Equal(t, dns.RcodeNameError, resp.Rcode, "chain continuation is signalled as NXDOMAIN")
assert.True(t, resp.MsgHdr.Zero, "Zero bit must be set so the chain continues")
assert.True(t, w.softNegative, "the fallthrough must not be allowed to poison the routed name")
assert.Equal(t, "dns-route", w.meta[resutil.MetaKeyDeferredBy])
assert.Equal(t, "peer-rcode-"+dns.RcodeToString[rcode], w.meta[resutil.MetaKeyDeferredReason],
"the peer's verdict must be visible in the log, it is the only trace of the probe")
})
}
}
// TestServeDNS_PeerVerdictPassesThrough covers the replies that are the peer's
// answer rather than a statement about its capability. The peer is authoritative
// for a name routed to it, so its verdict reaches the client as it is: retrying
// these elsewhere would send the name of an internal-only domain to a public
// resolver and prefer whatever that resolver says over the route.
func TestServeDNS_PeerVerdictPassesThrough(t *testing.T) {
// REFUSED means the peer does not consider the name routed to it, which
// happens while the client and the peer disagree about the route set. It is
// never how a peer reports a record type it cannot resolve: a forwarder too
// old for non-address types answers NOTIMP before it looks at the domain.
for _, rcode := range []int{dns.RcodeRefused, dns.RcodeNameError, dns.RcodeSuccess} {
t.Run(dns.RcodeToString[rcode], func(t *testing.T) {
stub := newForwarderStub(t, func(q *dns.Msg) *dns.Msg {
resp := new(dns.Msg)
resp.SetRcode(q, rcode)
return resp
})
d := newTestInterceptor(t, stub)
w := &softNegativeWriter{}
r := new(dns.Msg)
r.SetQuestion("_mongodb._tcp.db.example.com.", dns.TypeSRV)
d.ServeDNS(w, r)
require.Len(t, stub.received(), 1, "the peer must be asked exactly once")
resp := w.GetLastResponse()
require.NotNil(t, resp, "a response must be written")
assert.Equal(t, rcode, resp.Rcode, "the peer's verdict must reach the client unchanged")
assert.False(t, resp.MsgHdr.Zero, "the chain must not continue past an answered query")
assert.False(t, w.softNegative, "nothing was deferred, so nothing may be softened")
})
}
}
// TestServeDNS_PeerResolvesQtype guards the common case: a peer that can answer
// the type stays authoritative, and its records reach the client unchanged.
func TestServeDNS_PeerResolvesQtype(t *testing.T) {
stub := newForwarderStub(t, func(q *dns.Msg) *dns.Msg {
resp := new(dns.Msg)
resp.SetReply(q)
resp.Answer = []dns.RR{&dns.SRV{
Hdr: dns.RR_Header{Name: q.Question[0].Name, Rrtype: dns.TypeSRV, Class: dns.ClassINET, Ttl: 60},
Target: "shard-00.db.example.com.",
Port: 27017,
}}
return resp
})
d := newTestInterceptor(t, stub)
w := &softNegativeWriter{}
r := new(dns.Msg)
r.SetQuestion("_mongodb._tcp.db.example.com.", dns.TypeSRV)
d.ServeDNS(w, r)
resp := w.GetLastResponse()
require.NotNil(t, resp, "a response must be written")
assert.Equal(t, dns.RcodeSuccess, resp.Rcode)
require.Len(t, resp.Answer, 1, "the peer's records must pass through")
assert.False(t, resp.MsgHdr.Zero, "an answered query must not continue the chain")
assert.False(t, w.softNegative)
}
// TestServeDNS_DeferredQueryIsPristine covers what the fallthrough hands to the
// next handler. The interceptor advertises EDNS0 to the peer so the forwarder can
// return an Extended DNS Error, and sets AD, but neither belongs to the client's
// query: on the deferred path they would travel to the public resolver and the
// reply could come back carrying an OPT the client never advertised, which
// RFC 6891 forbids us from passing on.
func TestServeDNS_DeferredQueryIsPristine(t *testing.T) {
stub := newForwarderStub(t, func(q *dns.Msg) *dns.Msg {
resp := new(dns.Msg)
resp.SetRcode(q, dns.RcodeNotImplemented)
return resp
})
d := newTestInterceptor(t, stub)
w := &softNegativeWriter{}
r := new(dns.Msg)
r.SetQuestion("_mongodb._tcp.db.example.com.", dns.TypeSRV)
d.ServeDNS(w, r)
sent := stub.received()
require.Len(t, sent, 1)
assert.NotNil(t, sent[0].IsEdns0(), "the peer must be asked with EDNS0 so it can return an EDE")
assert.Nil(t, r.IsEdns0(), "the deferred query must not carry an OPT the client never sent")
assert.False(t, r.AuthenticatedData, "the deferred query must not carry an AD bit the client never set")
assert.Empty(t, r.Extra, "the deferred query must reach the next handler as the client sent it")
}

View File

@@ -0,0 +1,7 @@
//go:build windows
package dnsinterceptor
// GetInterfaceGUIDString completes iface.WGIface on Windows, which requires the
// interface GUID for DNS registration.
func (fakeWGIface) GetInterfaceGUIDString() (string, error) { return "", nil }

View File

@@ -9,12 +9,220 @@ import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"github.com/netbirdio/netbird/e2e/harness"
"github.com/netbirdio/netbird/shared/management/http/api"
)
// per1k is a model's published USD rates per 1k tokens. read is the prompt-cache read rate
// (OpenAI: the cached-input discount rate); write is the cache-creation rate where one exists.
type per1k struct{ in, out, read, write float64 }
// publishedPer1k hardcodes the vendors' PUBLISHED rates for the models the live matrix can drive,
// keyed by the normalized model id the proxy stamps. Deliberately independent of the proxy's
// pricing table so a wrong embedded rate or a broken normalization fails the run.
var publishedPer1k = map[string]per1k{
"gpt-4o-mini": {0.00015, 0.0006, 0.000075, 0},
"gpt-4o": {0.0025, 0.01, 0.00125, 0},
"claude-haiku-4-5": {0.001, 0.005, 0.0001, 0.00125},
"claude-sonnet-4-5": {0.003, 0.015, 0.0003, 0.00375},
"claude-sonnet-4-6": {0.003, 0.015, 0.0003, 0.00375},
"kimi-k3": {0.003, 0.015, 0.0003, 0.003}, // no published write rate: bills at the input rate
"anthropic.claude-haiku-4-5": {0.001, 0.005, 0.0001, 0.00125},
"anthropic.claude-sonnet-4-5": {0.003, 0.015, 0.0003, 0.00375},
"anthropic.claude-sonnet-4-6": {0.003, 0.015, 0.0003, 0.00375},
}
// rawCostVerificationSQL is the operator-facing double-check, run straight against the management
// sqlite store: recompute each usage row's expected total and cache cost from its own persisted
// token buckets and hardcoded published rates. OpenAI counts cached tokens as a subset of input;
// Anthropic-shape providers count cache buckets additively.
const rawCostVerificationSQL = `
WITH rates(model, in_rate, out_rate, read_rate, write_rate) AS (
VALUES
('gpt-4o-mini', 0.00015, 0.0006, 0.000075, 0.0),
('gpt-4o', 0.0025, 0.01, 0.00125, 0.0),
('claude-haiku-4-5', 0.001, 0.005, 0.0001, 0.00125),
('claude-sonnet-4-5', 0.003, 0.015, 0.0003, 0.00375),
('claude-sonnet-4-6', 0.003, 0.015, 0.0003, 0.00375),
('kimi-k3', 0.003, 0.015, 0.0003, 0.003),
('anthropic.claude-haiku-4-5', 0.001, 0.005, 0.0001, 0.00125),
('anthropic.claude-sonnet-4-5', 0.003, 0.015, 0.0003, 0.00375),
('anthropic.claude-sonnet-4-6', 0.003, 0.015, 0.0003, 0.00375)
)
SELECT
u.provider,
u.model,
u.input_tokens,
u.output_tokens,
u.cached_input_tokens,
u.cache_creation_tokens,
u.input_cost_usd,
u.cached_input_cost_usd,
u.cache_creation_cost_usd,
u.output_cost_usd,
-- No cost_usd / cache_cost_usd columns are stored: both are derived from the
-- four per-bucket columns above, exactly as the API renders them.
(u.input_cost_usd + u.cached_input_cost_usd + u.cache_creation_cost_usd + u.output_cost_usd) AS cost_usd,
(u.cached_input_cost_usd + u.cache_creation_cost_usd) AS cache_cost_usd,
CASE WHEN u.provider = 'openai' THEN
(u.input_tokens - MIN(u.cached_input_tokens, u.input_tokens))*r.in_rate/1000.0
ELSE
u.input_tokens*r.in_rate/1000.0
END AS expected_input,
CASE WHEN u.provider = 'openai' THEN
MIN(u.cached_input_tokens, u.input_tokens)*r.read_rate/1000.0
ELSE
u.cached_input_tokens*r.read_rate/1000.0
END AS expected_cached_input,
CASE WHEN u.provider = 'openai' THEN
0.0
ELSE
u.cache_creation_tokens*r.write_rate/1000.0
END AS expected_cache_creation,
u.output_tokens*r.out_rate/1000.0 AS expected_output,
CASE WHEN u.provider = 'openai' THEN
(u.input_tokens - MIN(u.cached_input_tokens, u.input_tokens))*r.in_rate/1000.0
+ MIN(u.cached_input_tokens, u.input_tokens)*r.read_rate/1000.0
+ u.output_tokens*r.out_rate/1000.0
ELSE
u.input_tokens*r.in_rate/1000.0 + u.cached_input_tokens*r.read_rate/1000.0
+ u.cache_creation_tokens*r.write_rate/1000.0 + u.output_tokens*r.out_rate/1000.0
END AS expected_total,
CASE WHEN u.provider = 'openai' THEN
MIN(u.cached_input_tokens, u.input_tokens)*r.read_rate/1000.0
ELSE
u.cached_input_tokens*r.read_rate/1000.0 + u.cache_creation_tokens*r.write_rate/1000.0
END AS expected_cache
FROM agent_network_request_usage u
JOIN rates r ON r.model = u.model
ORDER BY u.timestamp`
// verifyUsageRowsSQL re-checks every persisted usage row directly in the management sqlite store,
// bypassing the API path — the same audit an operator can run on a production store.db.
func verifyUsageRowsSQL(t *testing.T, srv *harness.Combined) {
t.Helper()
dbPath, err := srv.SnapshotStoreDB(t.TempDir())
require.NoError(t, err, "snapshot management sqlite store")
db, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{})
require.NoError(t, err, "open store snapshot")
sqlDB, err := db.DB()
require.NoError(t, err)
defer func() { _ = sqlDB.Close() }()
rows, err := db.Raw(rawCostVerificationSQL).Rows()
require.NoError(t, err, "run raw cost verification query")
defer func() { _ = rows.Close() }()
verified := 0
for rows.Next() {
var provider, model string
var inTok, outTok, readTok, writeTok int64
var inCost, cachedInCost, cacheCreateCost, outCost, cost, cacheCost float64
var wantInput, wantCachedInput, wantCacheCreation, wantOutput, wantTotal, wantCache float64
require.NoError(t, rows.Scan(&provider, &model, &inTok, &outTok, &readTok, &writeTok,
&inCost, &cachedInCost, &cacheCreateCost, &outCost, &cost, &cacheCost,
&wantInput, &wantCachedInput, &wantCacheCreation, &wantOutput, &wantTotal, &wantCache), "scan usage row")
t.Logf("[sql] %s/%s: in=%d out=%d cache_read=%d cache_write=%d stored in/cached/create/out=$%.6f/$%.6f/$%.6f/$%.6f total=$%.6f cache=$%.6f expected total=$%.6f cache=$%.6f",
provider, model, inTok, outTok, readTok, writeTok,
inCost, cachedInCost, cacheCreateCost, outCost, cost, cacheCost, wantTotal, wantCache)
assert.InDeltaf(t, wantInput, inCost, 1e-6, "stored input_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantCachedInput, cachedInCost, 1e-6, "stored cached_input_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantCacheCreation, cacheCreateCost, 1e-6, "stored cache_creation_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantOutput, outCost, 1e-6, "stored output_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantTotal, cost, 1e-6, "derived cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantCache, cacheCost, 1e-6, "derived cache_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, inCost+cachedInCost+cacheCreateCost+outCost, cost, 1e-9,
"stored buckets must sum to the derived cost_usd for %s/%s", provider, model)
verified++
}
require.NoError(t, rows.Err(), "iterate usage rows")
require.Positive(t, verified, "raw SQL check must cover at least one usage row")
t.Logf("[sql] verified %d usage rows in store.db against published rates", verified)
gwRows, err := db.Raw(`SELECT model,
(input_cost_usd + cached_input_cost_usd + cache_creation_cost_usd + output_cost_usd) AS cost_usd
FROM agent_network_request_usage WHERE model LIKE '%/%'`).Rows()
require.NoError(t, err, "query gateway-prefixed usage rows")
defer func() { _ = gwRows.Close() }()
for gwRows.Next() {
var model string
var cost float64
require.NoError(t, gwRows.Scan(&model, &cost), "scan gateway usage row")
t.Logf("[sql] gateway %s: stored=$%.6f (must be 0 — deliberately unpriced)", model, cost)
assert.Zerof(t, cost, "gateway-prefixed model %q must store cost 0, never a guessed rate", model)
}
require.NoError(t, gwRows.Err(), "iterate gateway usage rows")
}
// validateAccessLogCost recomputes a live access-log row's expected total and cache cost from the
// published per-1k rates and the row's persisted token buckets, and asserts both stored values.
// Gateway-prefixed model ids the proxy deliberately does not price must store cost 0.
func validateAccessLogCost(t *testing.T, pc providerCase, row api.AgentNetworkAccessLog) {
t.Helper()
model := catalogModel(pc)
provider := ""
if row.Provider != nil {
provider = *row.Provider
}
t.Logf("[cost] %s: provider=%s model=%s in=%d out=%d total=%d cache_read=%d cache_write=%d cost=$%.6f cache_cost=$%.6f",
pc.name, provider, model, row.InputTokens, row.OutputTokens, row.TotalTokens,
row.CachedInputTokens, row.CacheCreationTokens, row.CostUsd, row.CacheCostUsd)
rates, known := publishedPer1k[model]
if !known {
if strings.Contains(model, "/") {
assert.Zerof(t, row.CostUsd, "gateway-prefixed model %q is not priced so the cost meter must skip (cost 0)", model)
return
}
t.Logf("[cost] %s: no published rate on file for model %q (env-overridden?); skipping cost validation", pc.name, model)
return
}
// input_tokens may legitimately be 0: Moonshot/Kimi reports fully cached prompts under the cache
// buckets only. Output and total must always be present on a priced row.
require.Positive(t, row.OutputTokens, "priced row must carry output tokens")
require.Positive(t, row.TotalTokens, "priced row must carry total tokens")
var wantInput, wantCachedInput, wantCacheCreation float64
if provider == "openai" {
cached := min(row.CachedInputTokens, row.InputTokens) // cached is a subset of input
wantInput = float64(row.InputTokens-cached) / 1000 * rates.in
wantCachedInput = float64(cached) / 1000 * rates.read
// OpenAI has no cache-write bucket; wantCacheCreation stays 0.
} else {
// Anthropic / Bedrock shape: cache buckets are additive to input_tokens.
wantInput = float64(row.InputTokens) / 1000 * rates.in
wantCachedInput = float64(row.CachedInputTokens) / 1000 * rates.read
wantCacheCreation = float64(row.CacheCreationTokens) / 1000 * rates.write
}
wantOutput := float64(row.OutputTokens) / 1000 * rates.out
wantCache := wantCachedInput + wantCacheCreation
wantTotal := wantInput + wantCache + wantOutput
t.Logf("[cost] %s: expecting input=$%.6f cached_input=$%.6f cache_creation=$%.6f output=$%.6f total=$%.6f cache=$%.6f from published rates",
pc.name, wantInput, wantCachedInput, wantCacheCreation, wantOutput, wantTotal, wantCache)
assert.InDeltaf(t, wantInput, row.InputCostUsd, 1e-6, "stored input_cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantCachedInput, row.CachedInputCostUsd, 1e-6, "stored cached_input_cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantCacheCreation, row.CacheCreationCostUsd, 1e-6, "stored cache_creation_cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantOutput, row.OutputCostUsd, 1e-6, "stored output_cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantTotal, row.CostUsd, 1e-6, "derived cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantCache, row.CacheCostUsd, 1e-6, "derived cache_cost_usd for %s (%s)", pc.name, model)
// The aggregates must be exactly the sum of the stored components, not an
// independently-computed figure that could drift from the breakdown.
assert.InDeltaf(t, row.InputCostUsd+row.CachedInputCostUsd+row.CacheCreationCostUsd+row.OutputCostUsd,
row.CostUsd, 1e-9, "stored buckets must sum to the derived cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, row.CachedInputCostUsd+row.CacheCreationCostUsd,
row.CacheCostUsd, 1e-9, "stored cache buckets must sum to the derived cache_cost_usd for %s (%s)", pc.name, model)
}
// providerCase is one entry in the live provider matrix. The same scenario runs
// for every available provider; availability is keyed off env vars so the suite
// covers whatever credentials are present (source ~/.llm-keys locally / set the
@@ -116,12 +324,12 @@ func availableProviders() []providerCase {
if region == "" {
region = "eu-central-1"
}
// A valid Bedrock inference-profile id (region prefix + date + version),
// overridable per account. `global.` profiles can be invoked from any
// region; set AWS_BEDROCK_MODEL to match the enabled profile for the token.
// A valid Bedrock inference-profile id, overridable per account (AWS_BEDROCK_MODEL, also the
// workflow's bedrock_model dispatch input). `global.` profiles work from any region. Defaults to
// Sonnet 4.6, whose id convention dropped the -YYYYMMDD-v1:0 suffix that Haiku 4.5 still carries.
model := os.Getenv("AWS_BEDROCK_MODEL")
if model == "" {
model = "global.anthropic.claude-haiku-4-5-20251001-v1:0"
model = "global.anthropic.claude-sonnet-4-6"
}
ps = append(ps, providerCase{name: "bedrock", catalogID: "bedrock_api", upstream: "https://bedrock-runtime." + region + ".amazonaws.com", apiKey: k, model: model, kind: harness.WireBedrock})
}
@@ -257,6 +465,10 @@ func TestProvidersMatrix(t *testing.T) {
// session id and confirm the marker propagated end-to-end.
sessionID := "e2e-session-" + pc.name
// A long-form prompt so completions carry realistic token counts for cost validation;
// max_tokens in the harness bodies (2048) lets the full answer through.
const matrixPrompt = "explain GitHub workflow in 1000 words"
// Retry briefly to absorb tunnel/DNS jitter on the first call.
var code int
var body string
@@ -267,11 +479,11 @@ func TestProvidersMatrix(t *testing.T) {
var cerr error
switch pc.kind {
case harness.WireVertex:
c, b, cerr = cl.Vertex(ctx, settings.Endpoint, proxyIP, pc.project, pc.region, pc.model, "Reply with exactly: pong", sessionID)
c, b, cerr = cl.Vertex(ctx, settings.Endpoint, proxyIP, pc.project, pc.region, pc.model, matrixPrompt, sessionID)
case harness.WireBedrock:
c, b, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, pc.model, "Reply with exactly: pong", sessionID)
c, b, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, pc.model, matrixPrompt, sessionID)
default:
c, b, cerr = cl.ChatPrefixed(ctx, settings.Endpoint, proxyIP, pc.pathPrefix, pc.kind, pc.model, "Reply with exactly: pong", sessionID)
c, b, cerr = cl.ChatPrefixed(ctx, settings.Endpoint, proxyIP, pc.pathPrefix, pc.kind, pc.model, matrixPrompt, sessionID)
}
if cerr == nil {
code, body = c, b
@@ -290,6 +502,7 @@ func TestProvidersMatrix(t *testing.T) {
// The session id sent as x-session-id must round-trip into the
// access-log row for this provider.
var row api.AgentNetworkAccessLog
require.Eventually(t, func() bool {
logs, lerr := srv.ListAccessLogs(ctx)
if lerr != nil {
@@ -297,11 +510,15 @@ func TestProvidersMatrix(t *testing.T) {
}
for _, r := range logs.Data {
if r.SessionId != nil && *r.SessionId == sessionID {
row = r
return true
}
}
return false
}, 30*time.Second, 2*time.Second, "session id %q must be recorded in an access-log row for %s", sessionID, pc.name)
// Stored total and cache cost must match the published rates applied to the row's buckets.
validateAccessLogCost(t, pc, row)
})
}
@@ -322,4 +539,7 @@ func TestProvidersMatrix(t *testing.T) {
}
return false
}, 60*time.Second, 3*time.Second, "consumption must be recorded with positive token counts after live traffic")
// Final raw-SQL audit: bypass the API and re-verify every persisted usage row in the store.
verifyUsageRowsSQL(t, srv)
}

View File

@@ -256,7 +256,7 @@ func (cl *Client) ChatPrefixed(ctx context.Context, endpoint, proxyIP, pathPrefi
case WireMessages:
path = "/v1/messages"
headers = []string{"anthropic-version: 2023-06-01"}
body = fmt.Sprintf(`{"model":%q,"max_tokens":64,"messages":[{"role":"user","content":%q}]}`, model, prompt)
body = fmt.Sprintf(`{"model":%q,"max_tokens":2048,"messages":[{"role":"user","content":%q}]}`, model, prompt)
default:
path = "/v1/chat/completions"
body = fmt.Sprintf(`{"model":%q,"messages":[{"role":"user","content":%q}]}`, model, prompt)
@@ -271,7 +271,7 @@ func (cl *Client) ChatPrefixed(ctx context.Context, endpoint, proxyIP, pathPrefi
// is sent as the universal x-session-id header the proxy records.
func (cl *Client) Vertex(ctx context.Context, endpoint, proxyIP, project, region, model, prompt, sessionID string) (int, string, error) {
path := fmt.Sprintf("/v1/projects/%s/locations/%s/publishers/anthropic/models/%s:rawPredict", project, region, model)
body := fmt.Sprintf(`{"anthropic_version":"vertex-2023-10-16","max_tokens":64,"messages":[{"role":"user","content":%q}]}`, prompt)
body := fmt.Sprintf(`{"anthropic_version":"vertex-2023-10-16","max_tokens":2048,"messages":[{"role":"user","content":%q}]}`, prompt)
return cl.post(ctx, endpoint, proxyIP, path, body, withSessionID(nil, sessionID))
}
@@ -282,7 +282,7 @@ func (cl *Client) Vertex(ctx context.Context, endpoint, proxyIP, project, region
// header the proxy records.
func (cl *Client) Bedrock(ctx context.Context, endpoint, proxyIP, model, prompt, sessionID string) (int, string, error) {
path := "/model/" + model + "/invoke"
body := fmt.Sprintf(`{"anthropic_version":"bedrock-2023-05-31","max_tokens":64,"messages":[{"role":"user","content":%q}]}`, prompt)
body := fmt.Sprintf(`{"anthropic_version":"bedrock-2023-05-31","max_tokens":2048,"messages":[{"role":"user","content":%q}]}`, prompt)
return cl.post(ctx, endpoint, proxyIP, path, body, withSessionID(nil, sessionID))
}
@@ -341,7 +341,7 @@ func (cl *Client) Terminate(ctx context.Context) error {
return cl.container.Terminate(ctx)
}
// containerLogs reads up to 256 KiB of a container's logs for diagnostics.
// containerLogs reads up to 4 MiB of a container's logs for diagnostics — enough for a whole provider-matrix run.
func containerLogs(ctx context.Context, c testcontainers.Container) string {
if c == nil {
return ""
@@ -351,6 +351,6 @@ func containerLogs(ctx context.Context, c testcontainers.Container) string {
return fmt.Sprintf("<logs error: %v>", err)
}
defer r.Close()
b, _ := io.ReadAll(io.LimitReader(r, 256<<10))
b, _ := io.ReadAll(io.LimitReader(r, 4<<20))
return string(b)
}

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -221,14 +221,21 @@ func (l *Logger) allowDenyLog(serviceID types.ServiceID, reason string) bool {
// proxy/internal/middleware/keys.go — only the dimensions management needs to
// record a usage row (provider / model / tokens / cost / groups).
var usageMetadataKeys = map[string]struct{}{
"llm.provider": {},
"llm.model": {},
"llm.resolved_provider_id": {},
"llm.input_tokens": {},
"llm.output_tokens": {},
"llm.total_tokens": {},
"cost.usd_total": {},
"llm.authorising_groups": {},
"llm.provider": {},
"llm.model": {},
"llm.resolved_provider_id": {},
"llm.input_tokens": {},
"llm.output_tokens": {},
"llm.total_tokens": {},
"llm.cached_input_tokens": {},
"llm.cache_creation_tokens": {},
"cost.usd_input": {},
"cost.usd_cached_input": {},
"cost.usd_cache_creation": {},
"cost.usd_output": {},
"cost.usd_total": {},
"cost.usd_cache": {},
"llm.authorising_groups": {},
}
// stripAgentNetworkEntryForUsage returns the entry reduced to what's needed to

View File

@@ -56,10 +56,12 @@ type bedrockResponse struct {
OutputTokens int64 `json:"output_tokens"`
CacheReadInputTokens int64 `json:"cache_read_input_tokens"`
CacheCreationInputTokens int64 `json:"cache_creation_input_tokens"`
// Converse — camelCase.
InputTokensCamel int64 `json:"inputTokens"`
OutputTokensCamel int64 `json:"outputTokens"`
TotalTokensCamel int64 `json:"totalTokens"`
// Converse — camelCase; cache buckets are additive to inputTokens (AWS names the write bucket cacheWriteInputTokens).
InputTokensCamel int64 `json:"inputTokens"`
OutputTokensCamel int64 `json:"outputTokens"`
TotalTokensCamel int64 `json:"totalTokens"`
CacheReadTokensCamel int64 `json:"cacheReadInputTokens"`
CacheWriteTokensCamel int64 `json:"cacheWriteInputTokens"`
} `json:"usage"`
}
@@ -83,16 +85,18 @@ func (BedrockParser) ParseResponse(status int, contentType string, body []byte)
}
inTok := firstNonZero(resp.Usage.InputTokens, resp.Usage.InputTokensCamel)
outTok := firstNonZero(resp.Usage.OutputTokens, resp.Usage.OutputTokensCamel)
cacheRead := firstNonZero(resp.Usage.CacheReadInputTokens, resp.Usage.CacheReadTokensCamel)
cacheWrite := firstNonZero(resp.Usage.CacheCreationInputTokens, resp.Usage.CacheWriteTokensCamel)
total := resp.Usage.TotalTokensCamel
if total == 0 {
total = inTok + outTok + resp.Usage.CacheReadInputTokens + resp.Usage.CacheCreationInputTokens
total = inTok + outTok + cacheRead + cacheWrite
}
return Usage{
InputTokens: inTok,
OutputTokens: outTok,
TotalTokens: total,
CachedInputTokens: resp.Usage.CacheReadInputTokens,
CacheCreationTokens: resp.Usage.CacheCreationInputTokens,
CachedInputTokens: cacheRead,
CacheCreationTokens: cacheWrite,
}, nil
}

View File

@@ -26,6 +26,18 @@ func TestBedrockParser_ParseResponse_Converse(t *testing.T) {
require.Equal(t, int64(14), u.TotalTokens, "converse uses provider total")
}
// Converse camelCase cache fields must land in the billed Usage buckets, same as the InvokeModel snake_case fields.
func TestBedrockParser_ParseResponse_ConverseCacheBuckets(t *testing.T) {
body := []byte(`{"usage":{"inputTokens":11,"outputTokens":3,"cacheReadInputTokens":7,"cacheWriteInputTokens":9}}`)
u, err := BedrockParser{}.ParseResponse(200, "application/json", body)
require.NoError(t, err)
require.Equal(t, int64(11), u.InputTokens, "converse input tokens")
require.Equal(t, int64(3), u.OutputTokens, "converse output tokens")
require.Equal(t, int64(7), u.CachedInputTokens, "converse cache-read tokens")
require.Equal(t, int64(9), u.CacheCreationTokens, "converse cache-write tokens")
require.Equal(t, int64(11+3+7+9), u.TotalTokens, "total backfill is additive when the provider omits totalTokens")
}
func TestBedrockParser_ParseResponse_StreamingUnsupported(t *testing.T) {
_, err := BedrockParser{}.ParseResponse(200, "application/vnd.amazon.eventstream", []byte("binary"))
require.ErrorIs(t, err, ErrStreamingUnsupported, "event-stream must route to the streaming accumulator")

View File

@@ -128,6 +128,46 @@ type Table struct {
// - 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 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:
//
// TotalUSD = InputUSD + CachedInputUSD + CacheCreationUSD + OutputUSD
// CacheUSD = CachedInputUSD + CacheCreationUSD
//
// InputUSD is always the cost of the *non-cached* input bucket, for both
// provider shapes: on OpenAI the cached subset is carved out of inTokens and
// billed as CachedInputUSD, so the two never double-count. Buckets a provider
// doesn't bill are zero, which keeps the identities above true everywhere.
type Costs struct {
InputUSD float64
CachedInputUSD float64
CacheCreationUSD float64
OutputUSD float64
TotalUSD float64
CacheUSD float64
}
// newCosts assembles a split from its per-bucket parts, deriving the two
// aggregates so TotalUSD and CacheUSD can never drift from the breakdown.
func newCosts(input, cachedInput, cacheCreation, output float64) Costs {
return Costs{
InputUSD: input,
CachedInputUSD: cachedInput,
CacheCreationUSD: cacheCreation,
OutputUSD: output,
TotalUSD: input + cachedInput + cacheCreation + output,
CacheUSD: cachedInput + cacheCreation,
}
}
// 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) {
// Clamp negatives to zero before any pricing math so a malformed
// upstream count can never produce a negative cost.
if inTokens < 0 {
@@ -143,15 +183,15 @@ func (t *Table) Cost(provider, model string, inTokens, outTokens, cachedInput, c
cacheCreation = 0
}
if t == nil {
return 0, false
return Costs{}, false
}
byModel, ok := t.entries[provider]
if !ok {
return 0, false
return Costs{}, false
}
entry, ok := byModel[model]
if !ok {
return 0, false
return Costs{}, false
}
output := (float64(outTokens) / 1000.0) * entry.OutputPer1K
switch provider {
@@ -168,7 +208,7 @@ func (t *Table) Cost(provider, model string, inTokens, outTokens, cachedInput, c
}
nonCached := float64(inTokens-clamped) / 1000.0 * entry.InputPer1K
cached := float64(clamped) / 1000.0 * cachedRate
return nonCached + cached + output, true
return newCosts(nonCached, cached, 0, output), true
case "anthropic", "bedrock":
// Bedrock-Anthropic returns the same additive cache buckets as
// first-party Anthropic; non-Anthropic Bedrock models simply report
@@ -184,10 +224,10 @@ func (t *Table) Cost(provider, model string, inTokens, outTokens, cachedInput, c
input := float64(inTokens) / 1000.0 * entry.InputPer1K
read := float64(cachedInput) / 1000.0 * readRate
create := float64(cacheCreation) / 1000.0 * createRate
return input + read + create + output, true
return newCosts(input, read, create, output), true
default:
input := float64(inTokens) / 1000.0 * entry.InputPer1K
return input + output, true
return newCosts(input, 0, 0, output), true
}
}

View File

@@ -0,0 +1,329 @@
package builtin_test
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"strconv"
"testing"
"github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream"
"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"
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin/cost_meter"
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin/llm_request_parser"
"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.
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)
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)
require.NoError(t, err, "build cost_meter")
t.Cleanup(func() { _ = costMW.Close() })
const jsonCT = "application/json"
const sseCT = "text/event-stream"
const awsCT = "application/vnd.amazon.eventstream"
cases := []struct {
name string
url string
reqBody []byte
respCT string
respBody []byte
wantProvider string
wantModel string
wantCost float64 // exact expected USD; ignored when wantSkip is set
wantCacheCost float64 // expected cost.usd_cache portion of wantCost
wantSkip string // expected cost.skipped reason, "" when priced
}{
{
// gpt-4o-mini $0.15/$0.60 per MTok: 1000×0.15/1M + 500×0.60/1M.
name: "openai chat completions",
url: "https://api.openai.com/v1/chat/completions",
reqBody: []byte(`{"model":"gpt-4o-mini","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"choices":[{"message":{"content":"pong"}}],"usage":{"prompt_tokens":1000,"completion_tokens":500,"total_tokens":1500}}`),
wantProvider: "openai",
wantModel: "gpt-4o-mini",
wantCost: 0.00045,
},
{
// OpenAI cached tokens are a SUBSET of prompt_tokens at a discount; gpt-4o $2.50/$10 per MTok, cached $1.25/M:
// 250×2.5/1M + 750×1.25/1M + 500×10/1M.
name: "openai cached subset discount",
url: "https://api.openai.com/v1/chat/completions",
reqBody: []byte(`{"model":"gpt-4o","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"usage":{"prompt_tokens":1000,"completion_tokens":500,"prompt_tokens_details":{"cached_tokens":750}}}`),
wantProvider: "openai",
wantModel: "gpt-4o",
wantCost: 0.0065625,
wantCacheCost: 0.0009375,
},
{
// OpenAI streaming: usage rides the final SSE frame.
name: "openai chat SSE stream",
url: "https://api.openai.com/v1/chat/completions",
reqBody: []byte(`{"model":"gpt-4o-mini","stream":true,"messages":[{"role":"user","content":"hi"}]}`),
respCT: sseCT,
respBody: sseBody(`{"choices":[{"delta":{"content":"po"}}]}`, `{"choices":[{"delta":{"content":"ng"}}]}`, `{"choices":[],"usage":{"prompt_tokens":1000,"completion_tokens":500}}`, "[DONE]"),
wantProvider: "openai",
wantModel: "gpt-4o-mini",
wantCost: 0.00045,
},
{
// Mistral speaks the OpenAI shape: mistral-large-latest $0.50/$1.50 per MTok.
name: "mistral via openai shape",
url: "https://api.mistral.ai/v1/chat/completions",
reqBody: []byte(`{"model":"mistral-large-latest","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"usage":{"prompt_tokens":1000,"completion_tokens":1000}}`),
wantProvider: "openai",
wantModel: "mistral-large-latest",
wantCost: 0.002,
},
{
// The field report, minus caching: Bedrock Sonnet 4.6 $3/$15 per MTok, 3×3/1M + 1514×15/1M = $0.022719.
// Also covers inference-profile normalization of the region-prefixed versioned id in the URL.
name: "bedrock invoke — reported scenario, no cache",
url: "https://bedrock-runtime.eu-central-1.amazonaws.com/model/global.anthropic.claude-sonnet-4-6-20260115-v1:0/invoke",
reqBody: []byte(`{"messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"content":[{"type":"text","text":"pong"}],"usage":{"input_tokens":3,"output_tokens":1514}}`),
wantProvider: "bedrock",
wantModel: "anthropic.claude-sonnet-4-6",
wantCost: 0.022719,
},
{
// The field report as observed: the FIRST call of a session also wrote a 30,528-token prompt cache at
// 1.25× input ($3.75/M): 0.022719 + 30528×3.75/1M = $0.137199 — the reported $0.1372.
name: "bedrock invoke — reported scenario with cache write",
url: "https://bedrock-runtime.eu-central-1.amazonaws.com/model/global.anthropic.claude-sonnet-4-6-20260115-v1:0/invoke",
reqBody: []byte(`{"messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"usage":{"input_tokens":3,"output_tokens":1514,"cache_creation_input_tokens":30528,"cache_read_input_tokens":0}}`),
wantProvider: "bedrock",
wantModel: "anthropic.claude-sonnet-4-6",
wantCost: 0.137199,
wantCacheCost: 0.11448,
},
{
// Same numbers over the InvokeModel event-stream: message_start carries input + cache, message_delta the output.
name: "bedrock invoke stream with cache write",
url: "https://bedrock-runtime.eu-central-1.amazonaws.com/model/global.anthropic.claude-sonnet-4-6-20260115-v1:0/invoke-with-response-stream",
reqBody: []byte(`{"messages":[{"role":"user","content":"hi"}]}`),
respCT: awsCT,
respBody: bedrockInvokeStream(t, `{"type":"message_start","message":{"usage":{"input_tokens":3,"output_tokens":1,"cache_creation_input_tokens":30528}}}`, `{"type":"content_block_delta","delta":{"type":"text_delta","text":"pong"}}`, `{"type":"message_delta","usage":{"output_tokens":1514}}`),
wantProvider: "bedrock",
wantModel: "anthropic.claude-sonnet-4-6",
wantCost: 0.137199,
wantCacheCost: 0.11448,
},
{
// Converse camelCase usage incl. cache buckets. Haiku 4.5 $1/$5 per MTok, read $0.10/M, write $1.25/M:
// 50×1/1M + 100×5/1M + 2000×0.1/1M + 1000×1.25/1M = $0.002.
name: "bedrock converse with cache buckets",
url: "https://bedrock-runtime.eu-central-1.amazonaws.com/model/eu.anthropic.claude-haiku-4-5-20251001-v1:0/converse",
reqBody: []byte(`{"messages":[{"role":"user","content":[{"text":"hi"}]}]}`),
respCT: jsonCT,
respBody: []byte(`{"output":{"message":{"content":[{"text":"pong"}]}},"usage":{"inputTokens":50,"outputTokens":100,"totalTokens":3150,"cacheReadInputTokens":2000,"cacheWriteInputTokens":1000}}`),
wantProvider: "bedrock",
wantModel: "anthropic.claude-haiku-4-5",
wantCost: 0.002,
wantCacheCost: 0.00145,
},
{
// Same numbers over converse-stream: usage rides the trailing metadata frame.
name: "bedrock converse stream with cache buckets",
url: "https://bedrock-runtime.eu-central-1.amazonaws.com/model/eu.anthropic.claude-haiku-4-5-20251001-v1:0/converse-stream",
reqBody: []byte(`{"messages":[{"role":"user","content":[{"text":"hi"}]}]}`),
respCT: awsCT,
respBody: bedrockConverseStream(t,
`{"delta":{"text":"pong"}}`,
`{"usage":{"inputTokens":50,"outputTokens":100,"totalTokens":3150,"cacheReadInputTokens":2000,"cacheWriteInputTokens":1000}}`,
),
wantProvider: "bedrock",
wantModel: "anthropic.claude-haiku-4-5",
wantCost: 0.002,
wantCacheCost: 0.00145,
},
{
// First-party Anthropic, additive cache buckets. Sonnet 4.6:
// 256×3/1M + 200×15/1M + 768×0.3/1M + 512×3.75/1M.
name: "anthropic messages with cache buckets",
url: "https://api.anthropic.com/v1/messages",
reqBody: []byte(`{"model":"claude-sonnet-4-6","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"content":[{"type":"text","text":"pong"}],"usage":{"input_tokens":256,"output_tokens":200,"cache_read_input_tokens":768,"cache_creation_input_tokens":512}}`),
wantProvider: "anthropic",
wantModel: "claude-sonnet-4-6",
wantCost: 0.0059184,
wantCacheCost: 0.0021504,
},
{
// Anthropic SSE: input from message_start, output from message_delta. Haiku 4.5: 1000×1/1M + 2000×5/1M.
name: "anthropic SSE stream",
url: "https://api.anthropic.com/v1/messages",
reqBody: []byte(`{"model":"claude-haiku-4-5","stream":true,"messages":[{"role":"user","content":"hi"}]}`),
respCT: sseCT,
respBody: sseBody(`{"type":"message_start","message":{"usage":{"input_tokens":1000,"output_tokens":2}}}`, `{"type":"content_block_delta","delta":{"type":"text_delta","text":"pong"}}`, `{"type":"message_delta","usage":{"output_tokens":2000}}`, `{"type":"message_stop"}`),
wantProvider: "anthropic",
wantModel: "claude-haiku-4-5",
wantCost: 0.011,
},
{
// Kimi's Anthropic-compatible endpoint: kimi-k3 $3/$15 per MTok under the anthropic table.
name: "kimi anthropic shape",
url: "https://api.moonshot.ai/anthropic/v1/messages",
reqBody: []byte(`{"model":"kimi-k3","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"usage":{"input_tokens":1000,"output_tokens":1000}}`),
wantProvider: "anthropic",
wantModel: "kimi-k3",
wantCost: 0.018,
},
{
// Vertex path-routed model with "@version" stripped; Anthropic-on-Vertex priced under the anthropic table.
name: "vertex anthropic path-routed",
url: "https://aiplatform.googleapis.com/v1/projects/p/locations/global/publishers/anthropic/models/claude-sonnet-4-6@20260115:rawPredict",
reqBody: []byte(`{"anthropic_version":"vertex-2023-10-16","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"usage":{"input_tokens":200,"output_tokens":100}}`),
wantProvider: "anthropic",
wantModel: "claude-sonnet-4-6",
wantCost: 0.0021,
},
{
// Gateway-prefixed model ids are not in the pricing table: the meter must SKIP, never guess a rate.
name: "gateway-prefixed model skips pricing",
url: "https://gateway.example.com/v1/chat/completions",
reqBody: []byte(`{"model":"openai/gpt-4o-mini","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"usage":{"prompt_tokens":1000,"completion_tokens":500}}`),
wantProvider: "openai",
wantModel: "openai/gpt-4o-mini",
wantSkip: "unknown_model",
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
in := &middleware.Input{
Method: "POST",
URL: tc.url,
Headers: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
Body: tc.reqBody,
}
reqOut, err := reqMW.Invoke(context.Background(), in)
require.NoError(t, err, "request parser")
in.Metadata = append(in.Metadata, reqOut.Metadata...)
require.Equal(t, tc.wantProvider, metaKV(in.Metadata, middleware.KeyLLMProvider), "detected provider")
require.Equal(t, tc.wantModel, metaKV(in.Metadata, middleware.KeyLLMModel), "detected (normalized) model")
in.Status = 200
in.RespHeaders = []middleware.KV{{Key: "Content-Type", Value: tc.respCT}}
in.RespBody = tc.respBody
respOut, err := respMW.Invoke(context.Background(), in)
require.NoError(t, err, "response parser")
in.Metadata = append(in.Metadata, respOut.Metadata...)
costOut, err := costMW.Invoke(context.Background(), in)
require.NoError(t, err, "cost meter")
if tc.wantSkip != "" {
assert.Equal(t, tc.wantSkip, metaKV(costOut.Metadata, middleware.KeyCostSkipped), "expected cost skip reason")
assert.Empty(t, metaKV(costOut.Metadata, middleware.KeyCostUSDTotal), "no cost may be emitted on skip")
return
}
raw := metaKV(costOut.Metadata, middleware.KeyCostUSDTotal)
require.NotEmpty(t, raw, "cost.usd_total must be emitted; skip=%q", metaKV(costOut.Metadata, middleware.KeyCostSkipped))
got, err := strconv.ParseFloat(raw, 64)
require.NoError(t, err, "cost must be a float")
// cost.usd_total is rendered with %.6f: allow half of the last printed digit on top of float error.
assert.InDelta(t, tc.wantCost, got, 5.1e-7, "USD cost for %s", tc.name)
rawCache := metaKV(costOut.Metadata, middleware.KeyCostUSDCache)
require.NotEmpty(t, rawCache, "cost.usd_cache must be emitted next to cost.usd_total")
gotCache, err := strconv.ParseFloat(rawCache, 64)
require.NoError(t, err, "cache cost must be a float")
assert.InDelta(t, tc.wantCacheCost, gotCache, 5.1e-7, "cache USD cost for %s", tc.name)
})
}
}
// metaKV returns the value for key in kvs, or "" when absent.
func metaKV(kvs []middleware.KV, key string) string {
for _, kv := range kvs {
if kv.Key == key {
return kv.Value
}
}
return ""
}
// sseBody renders data frames as a text/event-stream body.
func sseBody(frames ...string) []byte {
var b bytes.Buffer
for _, f := range frames {
b.WriteString("data: ")
b.WriteString(f)
b.WriteString("\n\n")
}
return b.Bytes()
}
// awsFrame encodes one AWS event-stream frame with the given :event-type.
func awsFrame(t *testing.T, eventType string, payload []byte) []byte {
t.Helper()
var buf bytes.Buffer
enc := eventstream.NewEncoder()
require.NoError(t, enc.Encode(&buf, eventstream.Message{
Headers: eventstream.Headers{{Name: ":event-type", Value: eventstream.StringValue(eventType)}},
Payload: payload,
}), "encode event-stream frame")
return buf.Bytes()
}
// bedrockInvokeStream builds an invoke-with-response-stream body: each "chunk" frame wraps a base64 Anthropic event.
func bedrockInvokeStream(t *testing.T, events ...string) []byte {
t.Helper()
var body bytes.Buffer
for _, ev := range events {
wrap, err := json.Marshal(map[string]string{"bytes": base64.StdEncoding.EncodeToString([]byte(ev))})
require.NoError(t, err)
body.Write(awsFrame(t, "chunk", wrap))
}
return body.Bytes()
}
// bedrockConverseStream builds a converse-stream body: contentBlockDelta frames plus a trailing metadata usage frame.
func bedrockConverseStream(t *testing.T, deltas ...string) []byte {
t.Helper()
var body bytes.Buffer
for i, ev := range deltas {
eventType := "contentBlockDelta"
if i == len(deltas)-1 {
eventType = "metadata"
}
body.Write(awsFrame(t, eventType, []byte(ev)))
}
return body.Bytes()
}

View File

@@ -32,7 +32,12 @@ const (
)
var metadataKeys = []string{
middleware.KeyCostUSDInput,
middleware.KeyCostUSDCachedInput,
middleware.KeyCostUSDCacheCreation,
middleware.KeyCostUSDOutput,
middleware.KeyCostUSDTotal,
middleware.KeyCostUSDCache,
middleware.KeyCostSkipped,
}
@@ -140,18 +145,38 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar
}
table := m.loader.Get()
cost, ok := table.Cost(provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens)
costs, ok := table.Costs(provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens)
if !ok {
out.Metadata = skip(skipUnknownModel)
return out, nil
}
// Per-bucket costs first: they're the base of the breakdown, and the two
// aggregates that follow are derived from exactly these four values.
out.Metadata = []middleware.KV{
{Key: middleware.KeyCostUSDTotal, Value: fmt.Sprintf("%.6f", cost)},
{Key: middleware.KeyCostUSDInput, Value: usd(costs.InputUSD)},
{Key: middleware.KeyCostUSDCachedInput, Value: usd(costs.CachedInputUSD)},
{Key: middleware.KeyCostUSDCacheCreation, Value: usd(costs.CacheCreationUSD)},
{Key: middleware.KeyCostUSDOutput, Value: usd(costs.OutputUSD)},
{Key: middleware.KeyCostUSDTotal, Value: usd(costs.TotalUSD)},
{Key: middleware.KeyCostUSDCache, Value: usd(costs.CacheUSD)},
}
return out, nil
}
// usd renders a cost as the fixed-precision string every cost.usd_* key
// carries, so the per-bucket values and the aggregates round identically.
//
// 9 decimals, not 6: these values are summed downstream — per request, per
// session, and per usage bucket — so the rounding step is applied once per
// bucket per row and then accumulated. At 6 decimals a single row loses up to
// 2e-6 across its four buckets (enough to break a 1e-6 reconciliation against
// published rates), and a bucket smaller than half a microdollar quantises to
// zero outright: 16 cache-read tokens on a cheap model is 1.6e-9, so summing
// 10k such rows reports 0.02 instead of 0.016. Nano-dollar precision keeps the
// per-row error ~1000x below the smallest realistic bucket.
func usd(v float64) string { return fmt.Sprintf("%.9f", v) }
// skip returns a single-entry metadata slice carrying the given skip
// reason under KeyCostSkipped.
func skip(reason string) []middleware.KV {

View File

@@ -67,7 +67,15 @@ func TestMiddleware_StaticSurface(t *testing.T) {
assert.NoError(t, mw.Close(), "Close on stateless middleware is a no-op")
keys := mw.MetadataKeys()
expected := []string{middleware.KeyCostUSDTotal, middleware.KeyCostSkipped}
expected := []string{
middleware.KeyCostUSDInput,
middleware.KeyCostUSDCachedInput,
middleware.KeyCostUSDCacheCreation,
middleware.KeyCostUSDOutput,
middleware.KeyCostUSDTotal,
middleware.KeyCostUSDCache,
middleware.KeyCostSkipped,
}
assert.Equal(t, expected, keys, "metadata key allowlist must match the spec")
}
@@ -105,7 +113,7 @@ func TestFactory_DefaultPricingPathLoadsFixture(t *testing.T) {
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")
assert.Equal(t, "0.000750000", value, "0.00015 + 0.0006 per 1k tokens, 9-decimal format")
}
func TestFactory_PricingPathOverride(t *testing.T) {
@@ -129,7 +137,7 @@ func TestFactory_PricingPathOverride(t *testing.T) {
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")
assert.Equal(t, "0.015000000", value, "2*0.0025 + 1*0.01 = 0.015 with 9-decimal format")
}
func TestInvoke_ComputesCostForKnownModel(t *testing.T) {
@@ -148,7 +156,7 @@ func TestInvoke_ComputesCostForKnownModel(t *testing.T) {
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")
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")
}
@@ -357,8 +365,25 @@ func TestInvoke_OpenAICachedSubsetDiscount(t *testing.T) {
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,
assert.Equal(t, "0.006562500", value,
"cached subset must be billed at the discount rate, non-cached at the full rate; never double-billed")
cache, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDCache)
require.True(t, ok, "cost.usd_cache must be emitted alongside cost.usd_total")
// 750 cached at 0.00125/1k = 0.0009375.
assert.Equal(t, "0.000937500", cache, "cache cost is the discounted cost of the cached subset")
// Per-bucket breakdown. On OpenAI the cached subset is carved out of the
// input bucket, so input covers only the 250 non-cached tokens — the two
// must never double-count the same 750 tokens.
assertBucket(t, out.Metadata, middleware.KeyCostUSDInput, "0.000625000",
"input bucket bills only the non-cached remainder")
assertBucket(t, out.Metadata, middleware.KeyCostUSDCachedInput, "0.000937500",
"cached-input bucket bills the discounted subset")
assertBucket(t, out.Metadata, middleware.KeyCostUSDCacheCreation, "0.000000000",
"OpenAI has no cache-write bucket")
assertBucket(t, out.Metadata, middleware.KeyCostUSDOutput, "0.005000000",
"output bucket bills 500 tokens at 0.01/1k")
}
// TestInvoke_AnthropicCacheBucketsAdditive proves the Anthropic
@@ -384,9 +409,33 @@ func TestInvoke_AnthropicCacheBucketsAdditive(t *testing.T) {
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,
// = 0.000768 + 0.0002304 + 0.00192 + 0.003 = 0.0059184.
assert.Equal(t, "0.005918400", value,
"each Anthropic input bucket must bill at its own rate — cache_read cheap, cache_creation expensive, regular input mid")
cache, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDCache)
require.True(t, ok, "cost.usd_cache must be emitted alongside cost.usd_total")
// 768 cache_read * 0.0003 + 512 cache_creation * 0.00375 = 0.0021504.
assert.Equal(t, "0.002150400", cache, "cache cost sums the read and creation buckets")
// Per-bucket breakdown: four separately-billed buckets, each at its own rate.
assertBucket(t, out.Metadata, middleware.KeyCostUSDInput, "0.000768000",
"input bucket bills 256 tokens at 0.003/1k")
assertBucket(t, out.Metadata, middleware.KeyCostUSDCachedInput, "0.000230400",
"cache-read bucket bills 768 tokens at the cheap 0.0003/1k")
assertBucket(t, out.Metadata, middleware.KeyCostUSDCacheCreation, "0.001920000",
"cache-write bucket bills 512 tokens at the expensive 0.00375/1k")
assertBucket(t, out.Metadata, middleware.KeyCostUSDOutput, "0.003000000",
"output bucket bills 200 tokens at 0.015/1k")
}
// assertBucket asserts one per-bucket cost key carries the expected
// 6-decimal value.
func assertBucket(t *testing.T, md []middleware.KV, key, want, msg string) {
t.Helper()
got, ok := metaValue(t, md, key)
require.Truef(t, ok, "%s must be emitted", key)
assert.Equal(t, want, got, msg)
}
// TestInvoke_CachedTokensAbsentFallsBackToBaseFormula covers the
@@ -411,7 +460,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.007500", value, "no cached metadata = same cost as before the feature landed")
assert.Equal(t, "0.007500000", value, "no cached metadata = same cost as before the feature landed")
}
// TestInvoke_UnparseableCachedTokensSkippedSilently proves the
@@ -435,7 +484,7 @@ func TestInvoke_UnparseableCachedTokensSkippedSilently(t *testing.T) {
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")
assert.Equal(t, "0.007500000", value, "same as the no-cached-metadata path")
}
// TestMiddleware_CloseCancelsReloader proves Close stops the per-instance

View File

@@ -69,15 +69,18 @@ func applyBedrockInvokeChunk(payload []byte, usage *llm.Usage, completion *strin
}
// converseStreamEvent captures the Converse stream frames carrying completion
// text (contentBlockDelta) and the final token usage (metadata).
// text (contentBlockDelta) and the final token usage (metadata). Cache buckets
// are additive to inputTokens (AWS write bucket: cacheWriteInputTokens).
type converseStreamEvent struct {
Delta *struct {
Text string `json:"text"`
} `json:"delta"`
Usage *struct {
InputTokens int64 `json:"inputTokens"`
OutputTokens int64 `json:"outputTokens"`
TotalTokens int64 `json:"totalTokens"`
InputTokens int64 `json:"inputTokens"`
OutputTokens int64 `json:"outputTokens"`
TotalTokens int64 `json:"totalTokens"`
CacheReadTokens int64 `json:"cacheReadInputTokens"`
CacheWriteTokens int64 `json:"cacheWriteInputTokens"`
} `json:"usage"`
}
@@ -105,6 +108,12 @@ func applyConverseStreamEvent(eventType string, payload []byte, usage *llm.Usage
if ev.Usage.TotalTokens > 0 {
usage.TotalTokens = ev.Usage.TotalTokens
}
if ev.Usage.CacheReadTokens > 0 {
usage.CachedInputTokens = ev.Usage.CacheReadTokens
}
if ev.Usage.CacheWriteTokens > 0 {
usage.CacheCreationTokens = ev.Usage.CacheWriteTokens
}
}
}
}

View File

@@ -66,6 +66,24 @@ func TestAccumulateBedrockStream_Converse(t *testing.T) {
require.Equal(t, "pong", completion, "converse text deltas concatenated")
}
// The converse-stream metadata frame's camelCase cache fields must reach the billed cache buckets.
func TestAccumulateBedrockStream_ConverseCacheBuckets(t *testing.T) {
var body bytes.Buffer
body.Write(bedrockFrame(t, "contentBlockDelta", mustJSON(t, map[string]any{"delta": map[string]any{"text": "pong"}})))
body.Write(bedrockFrame(t, "metadata", mustJSON(t, map[string]any{"usage": map[string]any{
"inputTokens": 11, "outputTokens": 3, "totalTokens": 30,
"cacheReadInputTokens": 7, "cacheWriteInputTokens": 9,
}})))
usage, completion := accumulateBedrockStream(body.Bytes())
require.Equal(t, int64(11), usage.InputTokens, "input tokens from metadata frame")
require.Equal(t, int64(3), usage.OutputTokens, "output tokens from metadata frame")
require.Equal(t, int64(7), usage.CachedInputTokens, "cache-read tokens from metadata frame")
require.Equal(t, int64(9), usage.CacheCreationTokens, "cache-write tokens from metadata frame")
require.Equal(t, int64(30), usage.TotalTokens, "provider-reported total wins")
require.Equal(t, "pong", completion)
}
func TestAccumulateBedrockStream_Truncated(t *testing.T) {
// A body cut mid-frame must not panic; partial usage is returned.
full := bedrockFrame(t, "metadata", mustJSON(t, map[string]any{"usage": map[string]any{"inputTokens": 11, "outputTokens": 3}}))

View File

@@ -75,8 +75,19 @@ const (
KeyLLMAttributionGroupID = "llm.attribution_group_id"
KeyLLMAttributionWindowS = "llm.attribution_window_seconds"
// Cost metering (emitted by cost_meter).
KeyCostUSDTotal = "cost.usd_total"
// Cost metering (emitted by cost_meter). The four per-bucket keys are the
// base of the breakdown — one per token bucket the provider bills
// separately — and the two aggregates below are derived from them:
// usd_total is their sum, usd_cache is cached_input + cache_creation.
KeyCostUSDInput = "cost.usd_input"
// KeyCostUSDCachedInput is the cost of the cache-read bucket (Anthropic cache_read; OpenAI's discounted cached subset of input).
KeyCostUSDCachedInput = "cost.usd_cached_input"
// KeyCostUSDCacheCreation is the cost of the cache-write bucket. Zero for providers without one.
KeyCostUSDCacheCreation = "cost.usd_cache_creation"
KeyCostUSDOutput = "cost.usd_output"
KeyCostUSDTotal = "cost.usd_total"
// KeyCostUSDCache is the portion of cost.usd_total billed for prompt-cache buckets (cache read/creation, or OpenAI's cached input subset).
KeyCostUSDCache = "cost.usd_cache"
KeyCostSkipped = "cost.skipped"
// Framework-emitted error markers. Use the mw.<id>.* prefix to

View File

@@ -5822,13 +5822,48 @@ components:
total_tokens:
type: integer
format: int64
description: Total tokens consumed.
description: Total tokens consumed, including prompt-cache tokens.
example: 1840
cached_input_tokens:
type: integer
format: int64
description: Input tokens read from the provider's prompt cache. Additive to input_tokens for Anthropic-shape providers; a subset of input_tokens for OpenAI.
example: 0
cache_creation_tokens:
type: integer
format: int64
description: Input tokens written to the provider's prompt cache. Zero for providers without a cache-write bucket.
example: 30528
cost_usd:
type: number
format: double
description: Estimated USD cost of the request.
example: 0.0231
input_cost_usd:
type: number
format: double
description: Cost of the non-cached input tokens. Base component of cost_usd.
example: 0.0048
cached_input_cost_usd:
type: number
format: double
description: Cost of the prompt-cache read tokens. Base component of cost_usd, and part of cache_cost_usd.
example: 0.0015
cache_creation_cost_usd:
type: number
format: double
description: Cost of the prompt-cache write tokens. Base component of cost_usd, and part of cache_cost_usd.
example: 0.1130
output_cost_usd:
type: number
format: double
description: Cost of the output tokens. Base component of cost_usd.
example: 0.0038
cache_cost_usd:
type: number
format: double
description: Portion of cost_usd billed for prompt-cache usage.
example: 0.1145
stream:
type: boolean
description: Whether the request was a streaming completion.
@@ -5852,7 +5887,14 @@ components:
- input_tokens
- output_tokens
- total_tokens
- cached_input_tokens
- cache_creation_tokens
- input_cost_usd
- cached_input_cost_usd
- cache_creation_cost_usd
- output_cost_usd
- cost_usd
- cache_cost_usd
AgentNetworkAccessLogsResponse:
type: object
properties:
@@ -5926,13 +5968,48 @@ components:
total_tokens:
type: integer
format: int64
description: Total tokens across the session.
description: Total tokens across the session, including prompt-cache tokens.
example: 12880
cached_input_tokens:
type: integer
format: int64
description: Total prompt-cache read tokens across the session.
example: 0
cache_creation_tokens:
type: integer
format: int64
description: Total prompt-cache write tokens across the session.
example: 30528
cost_usd:
type: number
format: double
description: Total estimated USD cost across the session.
example: 0.1617
input_cost_usd:
type: number
format: double
description: Total cost of non-cached input tokens across the session.
example: 0.0210
cached_input_cost_usd:
type: number
format: double
description: Total cost of prompt-cache read tokens across the session.
example: 0.0015
cache_creation_cost_usd:
type: number
format: double
description: Total cost of prompt-cache write tokens across the session.
example: 0.1130
output_cost_usd:
type: number
format: double
description: Total cost of output tokens across the session.
example: 0.0262
cache_cost_usd:
type: number
format: double
description: Portion of cost_usd billed for prompt-cache usage across the session.
example: 0.1145
providers:
type: array
items:
@@ -5959,7 +6036,14 @@ components:
- input_tokens
- output_tokens
- total_tokens
- cached_input_tokens
- cache_creation_tokens
- input_cost_usd
- cached_input_cost_usd
- cache_creation_cost_usd
- output_cost_usd
- cost_usd
- cache_cost_usd
- decision
- entries
AgentNetworkAccessLogSessionsResponse:
@@ -6013,19 +6097,61 @@ components:
total_tokens:
type: integer
format: int64
description: Total tokens in the bucket.
description: Total tokens in the bucket, including prompt-cache tokens.
example: 184000
cached_input_tokens:
type: integer
format: int64
description: Total prompt-cache read tokens in the bucket.
example: 20000
cache_creation_tokens:
type: integer
format: int64
description: Total prompt-cache write tokens in the bucket.
example: 45000
input_cost_usd:
type: number
format: double
description: Total cost of non-cached input tokens in the bucket.
example: 1.12
cached_input_cost_usd:
type: number
format: double
description: Total cost of prompt-cache read tokens in the bucket.
example: 0.06
cache_creation_cost_usd:
type: number
format: double
description: Total cost of prompt-cache write tokens in the bucket.
example: 0.36
output_cost_usd:
type: number
format: double
description: Total cost of output tokens in the bucket.
example: 0.77
cost_usd:
type: number
format: double
description: Total estimated USD spend in the bucket.
example: 2.31
cache_cost_usd:
type: number
format: double
description: Portion of cost_usd billed for prompt-cache usage in the bucket.
example: 0.42
required:
- period_start
- input_tokens
- output_tokens
- total_tokens
- cached_input_tokens
- cache_creation_tokens
- input_cost_usd
- cached_input_cost_usd
- cache_creation_cost_usd
- output_cost_usd
- cost_usd
- cache_cost_usd
AgentNetworkSettings:
type: object
description: Per-account Agent Network gateway settings. One row per account; cluster and subdomain are auto-assigned on first provider create and immutable thereafter.

View File

@@ -1732,6 +1732,21 @@ type AccountSettings struct {
// AgentNetworkAccessLog One per-request agent-network (LLM) access log entry with flattened, queryable LLM dimensions.
type AgentNetworkAccessLog struct {
// CacheCostUsd Portion of cost_usd billed for prompt-cache usage.
CacheCostUsd float64 `json:"cache_cost_usd"`
// CacheCreationCostUsd Cost of the prompt-cache write tokens. Base component of cost_usd, and part of cache_cost_usd.
CacheCreationCostUsd float64 `json:"cache_creation_cost_usd"`
// CacheCreationTokens Input tokens written to the provider's prompt cache. Zero for providers without a cache-write bucket.
CacheCreationTokens int64 `json:"cache_creation_tokens"`
// CachedInputCostUsd Cost of the prompt-cache read tokens. Base component of cost_usd, and part of cache_cost_usd.
CachedInputCostUsd float64 `json:"cached_input_cost_usd"`
// CachedInputTokens Input tokens read from the provider's prompt cache. Additive to input_tokens for Anthropic-shape providers; a subset of input_tokens for OpenAI.
CachedInputTokens int64 `json:"cached_input_tokens"`
// CostUsd Estimated USD cost of the request.
CostUsd float64 `json:"cost_usd"`
@@ -1753,6 +1768,9 @@ type AgentNetworkAccessLog struct {
// Id Unique identifier for the access log entry.
Id string `json:"id"`
// InputCostUsd Cost of the non-cached input tokens. Base component of cost_usd.
InputCostUsd float64 `json:"input_cost_usd"`
// InputTokens Input (prompt) tokens consumed.
InputTokens int64 `json:"input_tokens"`
@@ -1762,6 +1780,9 @@ type AgentNetworkAccessLog struct {
// Model Requested LLM model.
Model *string `json:"model,omitempty"`
// OutputCostUsd Cost of the output tokens. Base component of cost_usd.
OutputCostUsd float64 `json:"output_cost_usd"`
// OutputTokens Output (completion) tokens produced.
OutputTokens int64 `json:"output_tokens"`
@@ -1801,7 +1822,7 @@ type AgentNetworkAccessLog struct {
// Timestamp Timestamp when the request was made.
Timestamp time.Time `json:"timestamp"`
// TotalTokens Total tokens consumed.
// TotalTokens Total tokens consumed, including prompt-cache tokens.
TotalTokens int64 `json:"total_tokens"`
// UserId NetBird user id of the authenticated caller, if applicable.
@@ -1810,6 +1831,21 @@ type AgentNetworkAccessLog struct {
// AgentNetworkAccessLogSession A session-grouped view of agent-network access logs — all requests sharing a session id (or a single session-less request) folded into one summary plus its ordered entries.
type AgentNetworkAccessLogSession struct {
// CacheCostUsd Portion of cost_usd billed for prompt-cache usage across the session.
CacheCostUsd float64 `json:"cache_cost_usd"`
// CacheCreationCostUsd Total cost of prompt-cache write tokens across the session.
CacheCreationCostUsd float64 `json:"cache_creation_cost_usd"`
// CacheCreationTokens Total prompt-cache write tokens across the session.
CacheCreationTokens int64 `json:"cache_creation_tokens"`
// CachedInputCostUsd Total cost of prompt-cache read tokens across the session.
CachedInputCostUsd float64 `json:"cached_input_cost_usd"`
// CachedInputTokens Total prompt-cache read tokens across the session.
CachedInputTokens int64 `json:"cached_input_tokens"`
// CostUsd Total estimated USD cost across the session.
CostUsd float64 `json:"cost_usd"`
@@ -1825,12 +1861,18 @@ type AgentNetworkAccessLogSession struct {
// GroupIds Union of the authorising group ids across the session's entries.
GroupIds *[]string `json:"group_ids,omitempty"`
// InputCostUsd Total cost of non-cached input tokens across the session.
InputCostUsd float64 `json:"input_cost_usd"`
// InputTokens Total input (prompt) tokens across the session.
InputTokens int64 `json:"input_tokens"`
// Models Distinct models seen in the session.
Models *[]string `json:"models,omitempty"`
// OutputCostUsd Total cost of output tokens across the session.
OutputCostUsd float64 `json:"output_cost_usd"`
// OutputTokens Total output (completion) tokens across the session.
OutputTokens int64 `json:"output_tokens"`
@@ -1846,7 +1888,7 @@ type AgentNetworkAccessLogSession struct {
// StartedAt Timestamp of the session's earliest request.
StartedAt time.Time `json:"started_at"`
// TotalTokens Total tokens across the session.
// TotalTokens Total tokens across the session, including prompt-cache tokens.
TotalTokens int64 `json:"total_tokens"`
// UserId NetBird user id of the session's caller.
@@ -2347,19 +2389,40 @@ type AgentNetworkSettingsRequest struct {
// AgentNetworkUsageBucket One aggregated agent-network usage time bucket (UTC). The bucket width is set by the request's granularity.
type AgentNetworkUsageBucket struct {
// CacheCostUsd Portion of cost_usd billed for prompt-cache usage in the bucket.
CacheCostUsd float64 `json:"cache_cost_usd"`
// CacheCreationCostUsd Total cost of prompt-cache write tokens in the bucket.
CacheCreationCostUsd float64 `json:"cache_creation_cost_usd"`
// CacheCreationTokens Total prompt-cache write tokens in the bucket.
CacheCreationTokens int64 `json:"cache_creation_tokens"`
// CachedInputCostUsd Total cost of prompt-cache read tokens in the bucket.
CachedInputCostUsd float64 `json:"cached_input_cost_usd"`
// CachedInputTokens Total prompt-cache read tokens in the bucket.
CachedInputTokens int64 `json:"cached_input_tokens"`
// CostUsd Total estimated USD spend in the bucket.
CostUsd float64 `json:"cost_usd"`
// InputCostUsd Total cost of non-cached input tokens in the bucket.
InputCostUsd float64 `json:"input_cost_usd"`
// InputTokens Total input (prompt) tokens in the bucket.
InputTokens int64 `json:"input_tokens"`
// OutputCostUsd Total cost of output tokens in the bucket.
OutputCostUsd float64 `json:"output_cost_usd"`
// OutputTokens Total output (completion) tokens in the bucket.
OutputTokens int64 `json:"output_tokens"`
// PeriodStart Start of the bucket in YYYY-MM-DD (UTC) — the day, the week start (Monday), or the month start, depending on granularity.
PeriodStart string `json:"period_start"`
// TotalTokens Total tokens in the bucket.
// TotalTokens Total tokens in the bucket, including prompt-cache tokens.
TotalTokens int64 `json:"total_tokens"`
}