diff --git a/e2e/agentnetwork/chat_test.go b/e2e/agentnetwork/chat_test.go index d5582205f..4e49ec252 100644 --- a/e2e/agentnetwork/chat_test.go +++ b/e2e/agentnetwork/chat_test.go @@ -61,8 +61,30 @@ SELECT u.output_tokens, u.cached_input_tokens, u.cache_creation_tokens, - u.cost_usd, - u.cache_cost_usd, + 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 @@ -102,19 +124,31 @@ func verifyUsageRowsSQL(t *testing.T, srv *harness.Combined) { for rows.Next() { var provider, model string var inTok, outTok, readTok, writeTok int64 - var cost, cacheCost, wantTotal, wantCache float64 - require.NoError(t, rows.Scan(&provider, &model, &inTok, &outTok, &readTok, &writeTok, &cost, &cacheCost, &wantTotal, &wantCache), "scan usage row") - t.Logf("[sql] %s/%s: in=%d out=%d cache_read=%d cache_write=%d stored=$%.6f/$%.6f expected=$%.6f/$%.6f", - provider, model, inTok, outTok, readTok, writeTok, cost, cacheCost, wantTotal, wantCache) - assert.InDeltaf(t, wantTotal, cost, 1e-6, "stored cost_usd for %s/%s must match the published-rate recompute", provider, model) - assert.InDeltaf(t, wantCache, cacheCost, 1e-6, "stored cache_cost_usd for %s/%s must match the published-rate recompute", provider, model) + var inCost, cachedInCost, cacheCreateCost, outCost, cost, cacheCost float64 + var wantIn, wantCachedIn, wantCacheCreate, wantOut, wantTotal, wantCache float64 + require.NoError(t, rows.Scan(&provider, &model, &inTok, &outTok, &readTok, &writeTok, + &inCost, &cachedInCost, &cacheCreateCost, &outCost, &cost, &cacheCost, + &wantIn, &wantCachedIn, &wantCacheCreate, &wantOut, &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, wantIn, inCost, 1e-6, "stored input_cost_usd for %s/%s must match the published-rate recompute", provider, model) + assert.InDeltaf(t, wantCachedIn, cachedInCost, 1e-6, "stored cached_input_cost_usd for %s/%s must match the published-rate recompute", provider, model) + assert.InDeltaf(t, wantCacheCreate, cacheCreateCost, 1e-6, "stored cache_creation_cost_usd for %s/%s must match the published-rate recompute", provider, model) + assert.InDeltaf(t, wantOut, 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, cost_usd FROM agent_network_request_usage WHERE model LIKE '%/%'`).Rows() + 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() { @@ -156,20 +190,37 @@ func validateAccessLogCost(t *testing.T, pc providerCase, row api.AgentNetworkAc require.Positive(t, row.OutputTokens, "priced row must carry output tokens") require.Positive(t, row.TotalTokens, "priced row must carry total tokens") - var wantTotal, wantCache float64 + var wantInput, wantCachedInput, wantCacheCreation float64 if provider == "openai" { cached := min(row.CachedInputTokens, row.InputTokens) // cached is a subset of input - wantCache = float64(cached) / 1000 * rates.read - wantTotal = float64(row.InputTokens-cached)/1000*rates.in + wantCache + float64(row.OutputTokens)/1000*rates.out + 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. - wantCache = float64(row.CachedInputTokens)/1000*rates.read + float64(row.CacheCreationTokens)/1000*rates.write - wantTotal = float64(row.InputTokens)/1000*rates.in + wantCache + float64(row.OutputTokens)/1000*rates.out + 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 total=$%.6f cache=$%.6f from published rates", pc.name, wantTotal, wantCache) - assert.InDeltaf(t, wantTotal, row.CostUsd, 1e-6, "stored cost_usd for %s (%s)", pc.name, model) - assert.InDeltaf(t, wantCache, row.CacheCostUsd, 1e-6, "stored cache_cost_usd for %s (%s)", pc.name, model) + 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 diff --git a/management/internals/modules/agentnetwork/accesslog_ingest.go b/management/internals/modules/agentnetwork/accesslog_ingest.go index eaa84f79d..ecc1780f3 100644 --- a/management/internals/modules/agentnetwork/accesslog_ingest.go +++ b/management/internals/modules/agentnetwork/accesslog_ingest.go @@ -29,8 +29,10 @@ const ( 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 - metaKeyCostUSDTotal = "cost.usd_total" - metaKeyCostUSDCache = "cost.usd_cache" + 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" @@ -111,23 +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), - CachedInputTokens: parseMetaInt(meta, metaKeyCachedInputTokens), - CacheCreationTokens: parseMetaInt(meta, metaKeyCacheCreationTokens), - CostUSD: parseMetaFloat(meta, metaKeyCostUSDTotal), - CacheCostUSD: parseMetaFloat(meta, metaKeyCostUSDCache), - 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 @@ -146,21 +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, - CachedInputTokens: e.CachedInputTokens, - CacheCreationTokens: e.CacheCreationTokens, - CostUSD: e.CostUSD, - CacheCostUSD: e.CacheCostUSD, + 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 3265f26bd..cd81cfbe4 100644 --- a/management/internals/modules/agentnetwork/accesslog_ingest_realstore_test.go +++ b/management/internals/modules/agentnetwork/accesslog_ingest_realstore_test.go @@ -37,8 +37,10 @@ func newIngestTestEntry() *accesslogs.AccessLogEntry { metaKeyTotalTokens: "1174", metaKeyCachedInputTokens: "256", metaKeyCacheCreationTokens: "768", - metaKeyCostUSDTotal: "0.0123", - metaKeyCostUSDCache: "0.0029", + metaKeyCostUSDInput: "0.0071", + metaKeyCostUSDCachedInput: "0.0009", + metaKeyCostUSDCacheCreate: "0.0020", + metaKeyCostUSDOutput: "0.0023", metaKeyStream: "true", metaKeyRequestPrompt: "hello", metaKeyResponseCompletion: "world", @@ -70,8 +72,17 @@ func TestIngestAccessLog_RealStore_LogCollectionOff(t *testing.T) { assert.Equal(t, int64(50), usage[0].OutputTokens, "output tokens 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") - assert.InDelta(t, 0.0123, usage[0].CostUSD, 1e-9, "cost must round-trip from metadata") - assert.InDelta(t, 0.0029, usage[0].CacheCostUSD, 1e-9, "cache cost 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) @@ -104,7 +115,12 @@ func TestIngestAccessLog_RealStore_LogCollectionOn(t *testing.T) { 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 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 347086d15..cde7be7de 100644 --- a/management/internals/modules/agentnetwork/types/accesslog.go +++ b/management/internals/modules/agentnetwork/types/accesslog.go @@ -41,12 +41,24 @@ type AgentNetworkAccessLog struct { InputTokens int64 OutputTokens int64 TotalTokens int64 - // Prompt-cache buckets: read + write token counts and the portion of CostUSD they account for. + // Prompt-cache buckets: read + write token counts. CachedInputTokens int64 CacheCreationTokens int64 - CostUSD float64 - CacheCostUSD float64 - Stream bool + // 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. @@ -64,22 +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, - CachedInputTokens: a.CachedInputTokens, - CacheCreationTokens: a.CacheCreationTokens, - CostUsd: a.CostUSD, - CacheCostUsd: a.CacheCostUSD, - 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) @@ -119,23 +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 - CachedInputTokens int64 - CacheCreationTokens int64 - CostUSD float64 - CacheCostUSD 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 @@ -217,8 +264,10 @@ func (sess *AgentNetworkAccessLogSession) foldEntry(sk *sessionSeen, e *AgentNet sess.TotalTokens += e.TotalTokens sess.CachedInputTokens += e.CachedInputTokens sess.CacheCreationTokens += e.CacheCreationTokens - sess.CostUSD += e.CostUSD - sess.CacheCostUSD += e.CacheCostUSD + 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 } @@ -261,18 +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, - CachedInputTokens: sess.CachedInputTokens, - CacheCreationTokens: sess.CacheCreationTokens, - CostUsd: sess.CostUSD, - CacheCostUsd: sess.CacheCostUSD, - 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..51268c686 --- /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 wantIn, wantCachedIn, wantCacheCreate, wantOut float64 + for _, e := range entries { + wantIn += e.InputCostUSD + wantCachedIn += e.CachedInputCostUSD + wantCacheCreate += e.CacheCreationCostUSD + wantOut += e.OutputCostUSD + } + + assert.InDelta(t, wantIn, sess.InputCostUsd, 1e-12, "session input cost is the sum of its entries") + assert.InDelta(t, wantCachedIn, sess.CachedInputCostUsd, 1e-12, "session cache-read cost is the sum of its entries") + assert.InDelta(t, wantCacheCreate, sess.CacheCreationCostUsd, 1e-12, "session cache-write cost is the sum of its entries") + assert.InDelta(t, wantOut, sess.OutputCostUsd, 1e-12, "session output cost is the sum of its entries") + assert.InDelta(t, wantIn+wantCachedIn+wantCacheCreate+wantOut, 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 911bae375..658fa8e59 100644 --- a/management/internals/modules/agentnetwork/types/usage.go +++ b/management/internals/modules/agentnetwork/types/usage.go @@ -25,12 +25,19 @@ type AgentNetworkUsage struct { InputTokens int64 OutputTokens int64 TotalTokens int64 - // Prompt-cache buckets: read + write token counts and the portion of CostUSD they account for. + // Prompt-cache buckets: read + write token counts. CachedInputTokens int64 CacheCreationTokens int64 - CostUSD float64 - CacheCostUSD float64 - CreatedAt time.Time + // 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 @@ -38,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 f12f06022..81e02c6b8 100644 --- a/management/internals/modules/agentnetwork/types/usageoverview.go +++ b/management/internals/modules/agentnetwork/types/usageoverview.go @@ -33,27 +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 - CachedInputTokens int64 - CacheCreationTokens int64 - CostUSD float64 - CacheCostUSD 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, - CachedInputTokens: b.CachedInputTokens, - CacheCreationTokens: b.CacheCreationTokens, - CostUsd: b.CostUSD, - CacheCostUsd: b.CacheCostUSD, + 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(), } } @@ -92,8 +110,10 @@ func AggregateUsageByGranularity(rows []*AgentNetworkUsage, g UsageGranularity) b.TotalTokens += r.TotalTokens b.CachedInputTokens += r.CachedInputTokens b.CacheCreationTokens += r.CacheCreationTokens - b.CostUSD += r.CostUSD - b.CacheCostUSD += r.CacheCostUSD + 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..65f62798b 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,96 @@ 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{})) + require.NoError(t, db.Create(&agentNetworkTypes.AgentNetworkUsage{ + ID: "u1", AccountID: "acct-1", Model: "claude-sonnet-4-6", + 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 ed6058b70..a438f42d2 100644 --- a/proxy/internal/accesslog/logger.go +++ b/proxy/internal/accesslog/logger.go @@ -229,6 +229,10 @@ var usageMetadataKeys = map[string]struct{}{ "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": {}, diff --git a/proxy/internal/llm/pricing/pricing.go b/proxy/internal/llm/pricing/pricing.go index 5320245d5..b77000000 100644 --- a/proxy/internal/llm/pricing/pricing.go +++ b/proxy/internal/llm/pricing/pricing.go @@ -132,11 +132,37 @@ func (t *Table) Cost(provider, model string, inTokens, outTokens, cachedInput, c return c.TotalUSD, ok } -// Costs is a per-request cost split. CacheUSD is the portion of TotalUSD billed for -// prompt-cache buckets and is always <= TotalUSD. +// 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 { - TotalUSD float64 - CacheUSD float64 + 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 @@ -182,7 +208,7 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, } nonCached := float64(inTokens-clamped) / 1000.0 * entry.InputPer1K cached := float64(clamped) / 1000.0 * cachedRate - return Costs{TotalUSD: nonCached + cached + output, CacheUSD: cached}, 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 @@ -198,10 +224,10 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, input := float64(inTokens) / 1000.0 * entry.InputPer1K read := float64(cachedInput) / 1000.0 * readRate create := float64(cacheCreation) / 1000.0 * createRate - return Costs{TotalUSD: input + read + create + output, CacheUSD: read + create}, true + return newCosts(input, read, create, output), true default: input := float64(inTokens) / 1000.0 * entry.InputPer1K - return Costs{TotalUSD: input + output}, true + return newCosts(input, 0, 0, output), true } } diff --git a/proxy/internal/middleware/builtin/cost_meter/middleware.go b/proxy/internal/middleware/builtin/cost_meter/middleware.go index d0d715d9e..63da6d17b 100644 --- a/proxy/internal/middleware/builtin/cost_meter/middleware.go +++ b/proxy/internal/middleware/builtin/cost_meter/middleware.go @@ -32,6 +32,10 @@ const ( ) var metadataKeys = []string{ + middleware.KeyCostUSDInput, + middleware.KeyCostUSDCachedInput, + middleware.KeyCostUSDCacheCreation, + middleware.KeyCostUSDOutput, middleware.KeyCostUSDTotal, middleware.KeyCostUSDCache, middleware.KeyCostSkipped, @@ -147,13 +151,32 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar 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", costs.TotalUSD)}, - {Key: middleware.KeyCostUSDCache, Value: fmt.Sprintf("%.6f", costs.CacheUSD)}, + {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 37811be40..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.KeyCostUSDCache, 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,13 +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, rendered 0.000937 by %.6f on the binary float. - assert.Equal(t, "0.000937", cache, "cache cost is the discounted cost of the cached subset") + // 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 @@ -389,14 +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 → "0.002150". - assert.Equal(t, "0.002150", cache, "cache cost sums the read and creation buckets") + // 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 @@ -421,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 @@ -445,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/keys.go b/proxy/internal/middleware/keys.go index c031b0364..336bed19f 100644 --- a/proxy/internal/middleware/keys.go +++ b/proxy/internal/middleware/keys.go @@ -75,8 +75,17 @@ 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" diff --git a/shared/management/http/api/openapi.yml b/shared/management/http/api/openapi.yml index 830d163c2..e3d11227a 100644 --- a/shared/management/http/api/openapi.yml +++ b/shared/management/http/api/openapi.yml @@ -5839,6 +5839,26 @@ components: 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 @@ -5869,6 +5889,10 @@ components: - 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: @@ -5961,6 +5985,26 @@ components: 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 @@ -5994,6 +6038,10 @@ components: - 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 @@ -6061,6 +6109,26 @@ components: 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 @@ -6078,6 +6146,10 @@ components: - 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: diff --git a/shared/management/http/api/types.gen.go b/shared/management/http/api/types.gen.go index 38eb9b959..a4de48a09 100644 --- a/shared/management/http/api/types.gen.go +++ b/shared/management/http/api/types.gen.go @@ -1735,9 +1735,15 @@ 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"` @@ -1762,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"` @@ -1771,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"` @@ -1822,9 +1834,15 @@ 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"` @@ -1843,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"` @@ -2368,18 +2392,30 @@ 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"`