Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
54 changes: 37 additions & 17 deletions internal/routing/cost.go
Original file line number Diff line number Diff line change
Expand Up @@ -58,13 +58,21 @@ func EstimateWindowCost(features RequestFeatures, forecast TrafficForecast, pric

// InputInflation corrects for upstreams that inject additional tokens
// (system prompts, tool definitions) not visible to the router at decision
// time. The inflated InputTokens produces a realistic suffix size.
// time. These injected tokens are stable across requests in a session, so on
// cache-supporting upstreams they participate in the cache lifecycle just
// like the client-supplied prefix — they are billed as cache_read on hits
// and cache_creation on misses, not at the full input rate.
originalInputTokens := features.InputTokens
if cache.InputInflation > 1 {
inflated := int64(float64(features.InputTokens) * cache.InputInflation)
if inflated > features.InputTokens {
features.InputTokens = inflated
}
}
inflationTokens := float64(features.InputTokens - originalInputTokens)
if inflationTokens < 0 {
inflationTokens = 0
}

prefix := float64(features.ReusableInputTokens)
if !cache.Supported || (cache.MinTokens > 0 && int64(prefix) < cache.MinTokens) {
Expand All @@ -79,9 +87,13 @@ func EstimateWindowCost(features RequestFeatures, forecast TrafficForecast, pric
if coverage <= 0 || coverage > 1 || math.IsNaN(coverage) || math.IsInf(coverage, 0) {
coverage = 1
}
cachedPrefix := prefix * coverage
uncachedPrefix := prefix - cachedPrefix
suffix := float64(features.InputTokens) - prefix
// cacheable tokens = reusable prefix + upstream-injected inflation. Both are
// stable across the session and share the provider cache lifecycle.
cacheable := prefix + inflationTokens
cachedPortion := cacheable * coverage
uncachedPortion := cacheable - cachedPortion
// true suffix = the newly appended user turn, never cacheable.
suffix := float64(originalInputTokens) - prefix
if suffix < 0 {
suffix = 0
}
Expand Down Expand Up @@ -198,15 +210,17 @@ func EstimateWindowCost(features RequestFeatures, forecast TrafficForecast, pric
result.ExpectedCreates = misses
result.CacheLifetimes = lifetimes

result.CacheReadCost = hits * cachedPrefix * pricing.CacheReadPerToken * multiplier
result.CacheWriteCost = misses * cachedPrefix * pricing.CacheWritePerToken * multiplier
// suffix + uncached portion of prefix always pay full input price
result.CacheInputCost = n * (suffix + uncachedPrefix) * input * multiplier
result.CacheReadCost = hits * cachedPortion * pricing.CacheReadPerToken * multiplier
result.CacheWriteCost = misses * cachedPortion * pricing.CacheWritePerToken * multiplier
// true suffix (new user turn) + uncached portion of prefix+inflation always
// pay full input price. Inflation tokens ONLY appear at full input price via
// the uncached share of coverage; the cached share flows through read/write.
result.CacheInputCost = n * (suffix + uncachedPortion) * input * multiplier
if cache.CacheReadIncludesInput {
result.CacheReadCost += hits * cachedPrefix * input * multiplier
result.CacheReadCost += hits * cachedPortion * input * multiplier
}
if cache.CacheWriteIncludesInput {
result.CacheWriteCost += misses * cachedPrefix * input * multiplier
result.CacheWriteCost += misses * cachedPortion * input * multiplier
}
result.CacheTotal = result.CacheInputCost + result.CacheReadCost + result.CacheWriteCost + result.OutputCost
result.SelectedTotal = result.NoCacheTotal
Expand Down Expand Up @@ -263,12 +277,17 @@ func cacheLifetimes(cache CacheProfile, window time.Duration, requests float64,
// supplied prices, or the equation is not defined.
func breakEvenRequests(features RequestFeatures, forecast TrafficForecast, pricing Pricing, cache CacheProfile, multiplier float64) float64 {
features = features.Normalize()
originalInputTokens := features.InputTokens
if cache.InputInflation > 1 {
inflated := int64(float64(features.InputTokens) * cache.InputInflation)
if inflated > features.InputTokens {
features.InputTokens = inflated
}
}
inflationTokens := float64(features.InputTokens - originalInputTokens)
if inflationTokens < 0 {
inflationTokens = 0
}
prefix := float64(features.ReusableInputTokens)
if !cache.Supported || prefix == 0 || (cache.MinTokens > 0 && int64(prefix) < cache.MinTokens) {
return -1
Expand All @@ -277,11 +296,12 @@ func breakEvenRequests(features RequestFeatures, forecast TrafficForecast, prici
if coverage <= 0 || coverage > 1 || math.IsNaN(coverage) || math.IsInf(coverage, 0) {
coverage = 1
}
cachedPrefix := prefix * coverage
uncachedPrefix := prefix - cachedPrefix
cacheable := prefix + inflationTokens
cachedPortion := cacheable * coverage
uncachedPortion := cacheable - cachedPortion
input := pricing.InputPerToken * multiplier
out := forecast.OutputTokensPerReq * pricing.OutputPerToken * multiplier
suffix := float64(features.InputTokens) - prefix
suffix := float64(originalInputTokens) - prefix
if suffix < 0 {
suffix = 0
}
Expand All @@ -293,13 +313,13 @@ func breakEvenRequests(features RequestFeatures, forecast TrafficForecast, prici
if h < 0 || h > 1 || math.IsNaN(h) || math.IsInf(h, 0) {
h = 0
}
miss := (suffix+uncachedPrefix)*input + cachedPrefix*pricing.CacheWritePerToken*multiplier + out
hit := (suffix+uncachedPrefix)*input + cachedPrefix*pricing.CacheReadPerToken*multiplier + out
miss := (suffix+uncachedPortion)*input + cachedPortion*pricing.CacheWritePerToken*multiplier + out
hit := (suffix+uncachedPortion)*input + cachedPortion*pricing.CacheReadPerToken*multiplier + out
if cache.CacheWriteIncludesInput {
miss += cachedPrefix * input
miss += cachedPortion * input
}
if cache.CacheReadIncludesInput {
hit += cachedPrefix * input
hit += cachedPortion * input
}
// Expected subsequent request cost at the observed hit rate.
steady := h*hit + (1-h)*miss
Expand Down
83 changes: 83 additions & 0 deletions internal/routing/cost_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -112,3 +112,86 @@ func closeTo(t *testing.T, got, want float64) {
t.Fatalf("got %.12g, want %.12g", got, want)
}
}

// Upstream-injected inflation tokens (system prompts, tool schemas) are stable
// per session and share the provider cache lifecycle, so they must be priced
// as cache_read on hits — not as full input suffix. This regression guards
// against pre-fix behaviour where inflation flowed into "suffix" and was
// billed at pricing.InputPerToken.
func TestEstimateWindowCostInflationBilledAsCacheRead(t *testing.T) {
features := RequestFeatures{InputTokens: 100_000, ReusableInputTokens: 100_000, EstimatedOutputTokens: 100}
price := Pricing{
InputPerToken: 5e-6, OutputPerToken: 25e-6,
CacheReadPerToken: 5e-7, CacheWritePerToken: 6.25e-6,
InputKnown: true, OutputKnown: true, CacheReadKnown: true, CacheWriteKnown: true,
Multiplier: 1,
}
// InputInflation=1.8 means upstream really bills 180k input tokens.
// Under cache with hit_rate=1 and coverage=1, ALL 180k should flow through
// cache_read (0.5e-7), and the extra 80k must NOT show up at 5e-6 in suffix.
cache := CacheProfile{
Supported: true, TTL: 5 * time.Minute,
HitRate: 1, HitRateSource: HitRateObserved,
CoverageRatio: 1, InputInflation: 1.8,
}
cost := EstimateWindowCost(features, TrafficForecast{Requests: 100, Window: 5 * time.Minute}, price, cache, time.Time{}, 5*time.Minute)

if !cost.CacheEligible {
t.Fatalf("cache should be eligible: %+v", cost)
}
// Original InputTokens == ReusableInputTokens, so true suffix is 0. The
// cache branch must contain no full-price input; the 80k inflation goes
// through cache_read.
if cost.CacheInputCost > 0.001 {
t.Fatalf("inflation must not appear as full-price suffix; CacheInputCost=%.6f", cost.CacheInputCost)
}
// With one guaranteed miss per lifetime and 99 hits, cacheable=180k:
expectedRead := 99 * 180_000.0 * 5e-7
if math.Abs(cost.CacheReadCost-expectedRead) > 1e-6 {
t.Fatalf("CacheReadCost=%.6f want ~%.6f", cost.CacheReadCost, expectedRead)
}
expectedWrite := 1 * 180_000.0 * 6.25e-6
if math.Abs(cost.CacheWriteCost-expectedWrite) > 1e-6 {
t.Fatalf("CacheWriteCost=%.6f want ~%.6f", cost.CacheWriteCost, expectedWrite)
}
// no_cache branch unchanged: 100 * inflated 180k * 5e-6 = 90.
if math.Abs(cost.NoCacheTotal-90.25) > 0.01 {
t.Fatalf("NoCacheTotal=%.4f want ~90.25", cost.NoCacheTotal)
}
if !cost.CacheUsed || cost.CacheTotal >= cost.NoCacheTotal {
t.Fatalf("cache should win over no-cache for high-inflation upstream: %+v", cost)
}
}

// The true suffix (newly appended user turn) must still be billed at the full
// input rate. Only the *inflation* portion moves to cache pricing.
func TestEstimateWindowCostTrueSuffixStillFullPrice(t *testing.T) {
// User-visible payload: 100k prefix + 5k new user turn = 105k. Inflation 1.5
// adds 52.5k more tokens the upstream injects.
features := RequestFeatures{InputTokens: 105_000, ReusableInputTokens: 100_000, EstimatedOutputTokens: 50}
price := Pricing{
InputPerToken: 5e-6, OutputPerToken: 25e-6,
CacheReadPerToken: 5e-7, CacheWritePerToken: 6.25e-6,
InputKnown: true, OutputKnown: true, CacheReadKnown: true, CacheWriteKnown: true,
Multiplier: 1,
}
cache := CacheProfile{
Supported: true, TTL: 5 * time.Minute,
HitRate: 1, HitRateSource: HitRateObserved,
CoverageRatio: 1, InputInflation: 1.5,
}
cost := EstimateWindowCost(features, TrafficForecast{Requests: 10, Window: 5 * time.Minute}, price, cache, time.Time{}, 5*time.Minute)

// True suffix = 5_000 tokens per request at full price: 10 * 5000 * 5e-6 = 0.25.
if math.Abs(cost.CacheInputCost-0.25) > 1e-9 {
t.Fatalf("true suffix must be at full input price; got CacheInputCost=%.6f want 0.25", cost.CacheInputCost)
}
// Cacheable = 100_000 + 52_500 = 152_500.
// hits = 9, misses = 1 (one guaranteed miss for lifetime).
if math.Abs(cost.CacheReadCost-0.68625) > 1e-6 {
t.Fatalf("CacheReadCost=%.6f want 0.68625", cost.CacheReadCost)
}
if math.Abs(cost.CacheWriteCost-0.953125) > 1e-6 {
t.Fatalf("CacheWriteCost=%.6f want 0.953125", cost.CacheWriteCost)
}
}
Loading