From 4f6247b5c3c6ee5284bfcf59d3dba0d2c327548d Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Mon, 27 Jul 2026 04:42:41 +0900 Subject: [PATCH] [management, proxy] Add prompt-cache token and cost accounting to agent network usage (#6900) Co-authored-by: braginini --- .github/workflows/agent-network-e2e.yml | 9 + e2e/agentnetwork/chat_test.go | 234 ++++++++++++- e2e/harness/client.go | 10 +- e2e/harness/combined.go | 23 ++ .../modules/agentnetwork/accesslog_ingest.go | 97 +++--- .../accesslog_ingest_realstore_test.go | 49 ++- .../accesslog_sessions_realstore_test.go | 6 +- .../modules/agentnetwork/types/accesslog.go | 141 ++++++-- .../agentnetwork/types/accesslogfilter.go | 4 +- .../modules/agentnetwork/types/cost_test.go | 124 +++++++ .../modules/agentnetwork/types/usage.go | 26 +- .../agentnetwork/types/usageoverview.go | 51 ++- management/server/migration/migration.go | 78 +++++ management/server/migration/migration_test.go | 97 ++++++ .../server/store/sql_store_agentnetwork.go | 2 +- .../sql_store_agentnetwork_accesslog_test.go | 12 +- management/server/store/store.go | 8 + proxy/internal/accesslog/logger.go | 23 +- proxy/internal/llm/bedrock.go | 18 +- proxy/internal/llm/bedrock_test.go | 12 + proxy/internal/llm/pricing/pricing.go | 52 ++- .../builtin/cost_calculation_matrix_test.go | 329 ++++++++++++++++++ .../builtin/cost_meter/middleware.go | 29 +- .../builtin/cost_meter/middleware_test.go | 67 +++- .../llm_response_parser/streaming_bedrock.go | 17 +- .../streaming_bedrock_test.go | 18 + proxy/internal/middleware/keys.go | 15 +- shared/management/http/api/openapi.yml | 132 ++++++- shared/management/http/api/types.gen.go | 69 +++- 29 files changed, 1582 insertions(+), 170 deletions(-) create mode 100644 management/internals/modules/agentnetwork/types/cost_test.go create mode 100644 proxy/internal/middleware/builtin/cost_calculation_matrix_test.go diff --git a/.github/workflows/agent-network-e2e.yml b/.github/workflows/agent-network-e2e.yml index bf4868871..88b98293d 100644 --- a/.github/workflows/agent-network-e2e.yml +++ b/.github/workflows/agent-network-e2e.yml @@ -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 }} diff --git a/e2e/agentnetwork/chat_test.go b/e2e/agentnetwork/chat_test.go index 90c3766ec..65c4d813f 100644 --- a/e2e/agentnetwork/chat_test.go +++ b/e2e/agentnetwork/chat_test.go @@ -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) } diff --git a/e2e/harness/client.go b/e2e/harness/client.go index 2ffcf653d..f53d0ea64 100644 --- a/e2e/harness/client.go +++ b/e2e/harness/client.go @@ -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("", err) } defer r.Close() - b, _ := io.ReadAll(io.LimitReader(r, 256<<10)) + b, _ := io.ReadAll(io.LimitReader(r, 4<<20)) return string(b) } diff --git a/e2e/harness/combined.go b/e2e/harness/combined.go index a6f43a139..5723100ca 100644 --- a/e2e/harness/combined.go +++ b/e2e/harness/combined.go @@ -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) diff --git a/management/internals/modules/agentnetwork/accesslog_ingest.go b/management/internals/modules/agentnetwork/accesslog_ingest.go index 59e53efa2..ecc1780f3 100644 --- a/management/internals/modules/agentnetwork/accesslog_ingest.go +++ b/management/internals/modules/agentnetwork/accesslog_ingest.go @@ -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)) diff --git a/management/internals/modules/agentnetwork/accesslog_ingest_realstore_test.go b/management/internals/modules/agentnetwork/accesslog_ingest_realstore_test.go index 431ce680e..cd81cfbe4 100644 --- a/management/internals/modules/agentnetwork/accesslog_ingest_realstore_test.go +++ b/management/internals/modules/agentnetwork/accesslog_ingest_realstore_test.go @@ -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") diff --git a/management/internals/modules/agentnetwork/accesslog_sessions_realstore_test.go b/management/internals/modules/agentnetwork/accesslog_sessions_realstore_test.go index 7d53d7547..94518c2f7 100644 --- a/management/internals/modules/agentnetwork/accesslog_sessions_realstore_test.go +++ b/management/internals/modules/agentnetwork/accesslog_sessions_realstore_test.go @@ -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") diff --git a/management/internals/modules/agentnetwork/types/accesslog.go b/management/internals/modules/agentnetwork/types/accesslog.go index 92b8bc358..cde7be7de 100644 --- a/management/internals/modules/agentnetwork/types/accesslog.go +++ b/management/internals/modules/agentnetwork/types/accesslog.go @@ -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) diff --git a/management/internals/modules/agentnetwork/types/accesslogfilter.go b/management/internals/modules/agentnetwork/types/accesslogfilter.go index d571a87b6..d35516ffa 100644 --- a/management/internals/modules/agentnetwork/types/accesslogfilter.go +++ b/management/internals/modules/agentnetwork/types/accesslogfilter.go @@ -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(*)", diff --git a/management/internals/modules/agentnetwork/types/cost_test.go b/management/internals/modules/agentnetwork/types/cost_test.go new file mode 100644 index 000000000..9194ef4a3 --- /dev/null +++ b/management/internals/modules/agentnetwork/types/cost_test.go @@ -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") +} diff --git a/management/internals/modules/agentnetwork/types/usage.go b/management/internals/modules/agentnetwork/types/usage.go index dd01d4300..658fa8e59 100644 --- a/management/internals/modules/agentnetwork/types/usage.go +++ b/management/internals/modules/agentnetwork/types/usage.go @@ -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. diff --git a/management/internals/modules/agentnetwork/types/usageoverview.go b/management/internals/modules/agentnetwork/types/usageoverview.go index 658832bec..81e02c6b8 100644 --- a/management/internals/modules/agentnetwork/types/usageoverview.go +++ b/management/internals/modules/agentnetwork/types/usageoverview.go @@ -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)) diff --git a/management/server/migration/migration.go b/management/server/migration/migration.go index ae26a254e..6d8ed90cc 100644 --- a/management/server/migration/migration.go +++ b/management/server/migration/migration.go @@ -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 +} diff --git a/management/server/migration/migration_test.go b/management/server/migration/migration_test.go index cc97c2dff..b60a50ee1 100644 --- a/management/server/migration/migration_test.go +++ b/management/server/migration/migration_test.go @@ -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") +} diff --git a/management/server/store/sql_store_agentnetwork.go b/management/server/store/sql_store_agentnetwork.go index b0df0cd2a..b72dc735f 100644 --- a/management/server/store/sql_store_agentnetwork.go +++ b/management/server/store/sql_store_agentnetwork.go @@ -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) } diff --git a/management/server/store/sql_store_agentnetwork_accesslog_test.go b/management/server/store/sql_store_agentnetwork_accesslog_test.go index 793c82d79..8ba79a062 100644 --- a/management/server/store/sql_store_agentnetwork_accesslog_test.go +++ b/management/server/store/sql_store_agentnetwork_accesslog_test.go @@ -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 diff --git a/management/server/store/store.go b/management/server/store/store.go index b78dd9d0f..1beea72fd 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -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) + }, } } diff --git a/proxy/internal/accesslog/logger.go b/proxy/internal/accesslog/logger.go index d47c71ca4..a438f42d2 100644 --- a/proxy/internal/accesslog/logger.go +++ b/proxy/internal/accesslog/logger.go @@ -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 diff --git a/proxy/internal/llm/bedrock.go b/proxy/internal/llm/bedrock.go index f7802beb2..eb64167a1 100644 --- a/proxy/internal/llm/bedrock.go +++ b/proxy/internal/llm/bedrock.go @@ -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 } diff --git a/proxy/internal/llm/bedrock_test.go b/proxy/internal/llm/bedrock_test.go index ca6f092f3..e99ee55df 100644 --- a/proxy/internal/llm/bedrock_test.go +++ b/proxy/internal/llm/bedrock_test.go @@ -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") diff --git a/proxy/internal/llm/pricing/pricing.go b/proxy/internal/llm/pricing/pricing.go index 09afec5ff..b77000000 100644 --- a/proxy/internal/llm/pricing/pricing.go +++ b/proxy/internal/llm/pricing/pricing.go @@ -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 } } diff --git a/proxy/internal/middleware/builtin/cost_calculation_matrix_test.go b/proxy/internal/middleware/builtin/cost_calculation_matrix_test.go new file mode 100644 index 000000000..f479eb563 --- /dev/null +++ b/proxy/internal/middleware/builtin/cost_calculation_matrix_test.go @@ -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() +} diff --git a/proxy/internal/middleware/builtin/cost_meter/middleware.go b/proxy/internal/middleware/builtin/cost_meter/middleware.go index 4da620310..63da6d17b 100644 --- a/proxy/internal/middleware/builtin/cost_meter/middleware.go +++ b/proxy/internal/middleware/builtin/cost_meter/middleware.go @@ -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 { diff --git a/proxy/internal/middleware/builtin/cost_meter/middleware_test.go b/proxy/internal/middleware/builtin/cost_meter/middleware_test.go index d1c161cab..e5d431d77 100644 --- a/proxy/internal/middleware/builtin/cost_meter/middleware_test.go +++ b/proxy/internal/middleware/builtin/cost_meter/middleware_test.go @@ -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 diff --git a/proxy/internal/middleware/builtin/llm_response_parser/streaming_bedrock.go b/proxy/internal/middleware/builtin/llm_response_parser/streaming_bedrock.go index a82a9cdbc..30809b46e 100644 --- a/proxy/internal/middleware/builtin/llm_response_parser/streaming_bedrock.go +++ b/proxy/internal/middleware/builtin/llm_response_parser/streaming_bedrock.go @@ -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 + } } } } diff --git a/proxy/internal/middleware/builtin/llm_response_parser/streaming_bedrock_test.go b/proxy/internal/middleware/builtin/llm_response_parser/streaming_bedrock_test.go index f93505882..f66347591 100644 --- a/proxy/internal/middleware/builtin/llm_response_parser/streaming_bedrock_test.go +++ b/proxy/internal/middleware/builtin/llm_response_parser/streaming_bedrock_test.go @@ -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}})) diff --git a/proxy/internal/middleware/keys.go b/proxy/internal/middleware/keys.go index 9c584ad82..336bed19f 100644 --- a/proxy/internal/middleware/keys.go +++ b/proxy/internal/middleware/keys.go @@ -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..* prefix to diff --git a/shared/management/http/api/openapi.yml b/shared/management/http/api/openapi.yml index 47ca80a7c..e3d11227a 100644 --- a/shared/management/http/api/openapi.yml +++ b/shared/management/http/api/openapi.yml @@ -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. diff --git a/shared/management/http/api/types.gen.go b/shared/management/http/api/types.gen.go index a9e98cf84..a4de48a09 100644 --- a/shared/management/http/api/types.gen.go +++ b/shared/management/http/api/types.gen.go @@ -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"` }