fix(pricing): support Gemini context-tiered rates in pricing snapshot (#93469)

The pricing snapshot could only express flat per-million rates, so
gemini-3.1-pro sessions with prompts over 200k tokens under-counted
input 2x ($2 vs $4/M) and output 1.5x ($12 vs $18/M).

- Add optional tier fields to PricingEntry: tier_threshold_tokens,
  input/output/cache_read_cost_per_million_above (None = flat, falls
  back to base rate per-field).
- estimate_usage_cost selects the above-threshold rates for the WHOLE
  request once usage.prompt_tokens (input + cache read + cache write)
  exceeds the threshold, matching Google's billing semantics.
- Populate gemini-3.1-pro (4.00/18.00/0.40 above 200k; alias
  gemini-3.1-pro-preview inherits) and gemini-2.5-pro (2.50/15.00
  above 200k).
- Flat entries are untouched: no threshold means no behavior change.

Reported and tier-field shape designed by @tornike14 (#93469).

Tests: below/at threshold unchanged, above-threshold tiered whole-request
pricing, cache-read tier rate and base-rate fallback, preview alias,
flat entries unaffected.
This commit is contained in:
Teknium
2026-08-24 00:30:03 -07:00
parent a87d314e44
commit c9e2a46df6
2 changed files with 145 additions and 9 deletions

View File

@@ -127,6 +127,16 @@ class PricingEntry:
source_url: Optional[str] = None
pricing_version: Optional[str] = None
fetched_at: Optional[datetime] = None
# Context-tiered pricing (e.g. Gemini Pro models charge higher rates once
# the prompt exceeds 200k tokens). When ``tier_threshold_tokens`` is set
# and ``usage.prompt_tokens`` (input + cache read + cache write) exceeds
# it, the ``*_above`` rates replace the base rates for the WHOLE request —
# that matches Google's billing semantics (not marginal/bracketed rates).
# Any ``*_above`` field left as None falls back to its base rate.
tier_threshold_tokens: Optional[int] = None
input_cost_per_million_above: Optional[Decimal] = None
output_cost_per_million_above: Optional[Decimal] = None
cache_read_cost_per_million_above: Optional[Decimal] = None
@dataclass(frozen=True)
@@ -594,6 +604,10 @@ _OFFICIAL_DOCS_PRICING: Dict[tuple[str, str], PricingEntry] = {
input_cost_per_million=Decimal("2.00"),
output_cost_per_million=Decimal("12.00"),
cache_read_cost_per_million=Decimal("0.20"),
tier_threshold_tokens=200_000,
input_cost_per_million_above=Decimal("4.00"),
output_cost_per_million_above=Decimal("18.00"),
cache_read_cost_per_million_above=Decimal("0.40"),
source="official_docs_snapshot",
source_url="https://ai.google.dev/pricing",
pricing_version="google-pricing-2026-07-07",
@@ -638,6 +652,9 @@ _OFFICIAL_DOCS_PRICING: Dict[tuple[str, str], PricingEntry] = {
input_cost_per_million=Decimal("1.25"),
output_cost_per_million=Decimal("10.00"),
cache_read_cost_per_million=Decimal("0.125"),
tier_threshold_tokens=200_000,
input_cost_per_million_above=Decimal("2.50"),
output_cost_per_million_above=Decimal("15.00"),
source="official_docs_snapshot",
source_url="https://ai.google.dev/pricing",
pricing_version="google-pricing-2026-07-07",
@@ -1454,12 +1471,30 @@ def estimate_usage_cost(
notes: list[str] = []
amount = _ZERO
if usage.input_tokens and entry.input_cost_per_million is None:
# Whole-request context-tier selection (e.g. Gemini Pro >200k prompts):
# once the prompt (input + cache read + cache write) exceeds the entry's
# threshold, the above-threshold rates apply to the entire request. Any
# tier rate left as None falls back to the base rate.
input_rate = entry.input_cost_per_million
output_rate = entry.output_cost_per_million
cache_read_rate = entry.cache_read_cost_per_million
if (
entry.tier_threshold_tokens is not None
and usage.prompt_tokens > entry.tier_threshold_tokens
):
if entry.input_cost_per_million_above is not None:
input_rate = entry.input_cost_per_million_above
if entry.output_cost_per_million_above is not None:
output_rate = entry.output_cost_per_million_above
if entry.cache_read_cost_per_million_above is not None:
cache_read_rate = entry.cache_read_cost_per_million_above
if usage.input_tokens and input_rate is None:
return CostResult(amount_usd=None, status="unknown", source=entry.source, label="n/a")
if usage.output_tokens and entry.output_cost_per_million is None:
if usage.output_tokens and output_rate is None:
return CostResult(amount_usd=None, status="unknown", source=entry.source, label="n/a")
if usage.cache_read_tokens:
if entry.cache_read_cost_per_million is None:
if cache_read_rate is None:
return CostResult(
amount_usd=None,
status="unknown",
@@ -1477,12 +1512,12 @@ def estimate_usage_cost(
notes=("cache-write pricing unavailable for route",),
)
if entry.input_cost_per_million is not None:
amount += Decimal(usage.input_tokens) * entry.input_cost_per_million / _ONE_MILLION
if entry.output_cost_per_million is not None:
amount += Decimal(usage.output_tokens) * entry.output_cost_per_million / _ONE_MILLION
if entry.cache_read_cost_per_million is not None:
amount += Decimal(usage.cache_read_tokens) * entry.cache_read_cost_per_million / _ONE_MILLION
if input_rate is not None:
amount += Decimal(usage.input_tokens) * input_rate / _ONE_MILLION
if output_rate is not None:
amount += Decimal(usage.output_tokens) * output_rate / _ONE_MILLION
if cache_read_rate is not None:
amount += Decimal(usage.cache_read_tokens) * cache_read_rate / _ONE_MILLION
if entry.cache_write_cost_per_million is not None:
amount += Decimal(usage.cache_write_tokens) * entry.cache_write_cost_per_million / _ONE_MILLION
if entry.request_cost is not None and usage.request_count:

View File

@@ -766,3 +766,104 @@ def test_normalize_usage_nested_details_win_over_qwen_flat_top_level():
assert normalized.cache_read_tokens == 900
assert normalized.input_tokens == 1100
# ── Context-tiered pricing (Gemini Pro >200k prompts, #93469) ─────────────
def test_gemini_31_pro_below_tier_threshold_uses_base_rates():
"""Prompts at or below 200k tokens bill at the base rates — the tier
fields must not change any below-threshold estimate."""
result = estimate_usage_cost(
"gemini-3.1-pro",
CanonicalUsage(input_tokens=100_000, output_tokens=10_000),
provider="google",
)
# 100k * $2/M + 10k * $12/M
assert result.amount_usd == Decimal("0.32")
at_threshold = estimate_usage_cost(
"gemini-3.1-pro",
CanonicalUsage(input_tokens=200_000, output_tokens=10_000),
provider="google",
)
# Exactly 200k is still the lower tier (Google bills "> 200k" higher).
# 200k * $2/M + 10k * $12/M
assert at_threshold.amount_usd == Decimal("0.52")
def test_gemini_31_pro_above_tier_threshold_uses_tiered_rates_whole_request():
"""Once the prompt exceeds 200k tokens the >200k rates ($4 input /
$18 output per million) apply to the ENTIRE request, not just the
marginal tokens — matching Google's billing semantics (#93469).
Before the fix this request priced at 250k*$2/M + 10k*$12/M = $0.62,
under-counting input 2x and output 1.5x."""
result = estimate_usage_cost(
"gemini-3.1-pro",
CanonicalUsage(input_tokens=250_000, output_tokens=10_000),
provider="google",
)
# 250k * $4/M + 10k * $18/M
assert result.amount_usd == Decimal("1.18")
assert result.status == "estimated"
def test_gemini_31_pro_cache_read_tokens_count_toward_tier_and_tier_rate():
"""prompt_tokens (input + cache read + cache write) drives tier selection,
and cache reads above the threshold bill at the $0.40/M tier rate."""
result = estimate_usage_cost(
"gemini-3.1-pro",
CanonicalUsage(input_tokens=150_000, cache_read_tokens=100_000),
provider="google",
)
# prompt = 250k > 200k → 150k * $4/M + 100k * $0.40/M
assert result.amount_usd == Decimal("0.64")
def test_gemini_31_pro_preview_alias_shares_tiered_pricing():
"""The provider-emitted preview id aliases the canonical row, so it must
pick up the tier fields too."""
result = estimate_usage_cost(
"gemini-3.1-pro-preview",
CanonicalUsage(input_tokens=250_000, output_tokens=10_000),
provider="google",
)
assert result.amount_usd == Decimal("1.18")
def test_gemini_25_pro_tiered_rates_with_cache_read_fallback():
"""gemini-2.5-pro tiers at the same 200k threshold ($2.50 input / $15
output above). Its snapshot has no tiered cache-read rate, so cache reads
fall back to the base $0.125/M even above the threshold."""
result = estimate_usage_cost(
"gemini-2.5-pro",
CanonicalUsage(input_tokens=250_000, output_tokens=10_000),
provider="google",
)
# 250k * $2.50/M + 10k * $15/M
assert result.amount_usd == Decimal("0.775")
with_cache = estimate_usage_cost(
"gemini-2.5-pro",
CanonicalUsage(input_tokens=150_000, cache_read_tokens=100_000),
provider="google",
)
# prompt = 250k > 200k → 150k * $2.50/M + 100k * $0.125/M (base fallback)
assert with_cache.amount_usd == Decimal("0.3875")
def test_flat_entries_unaffected_by_tier_machinery():
"""Entries without tier fields keep pricing every token at the flat rate
no matter how large the prompt is."""
entry = get_pricing_entry("gemini-3.1-flash-lite", provider="google")
assert entry is not None
assert entry.tier_threshold_tokens is None
result = estimate_usage_cost(
"gemini-3.1-flash-lite",
CanonicalUsage(input_tokens=250_000, output_tokens=10_000),
provider="google",
)
# 250k * $0.25/M + 10k * $1.50/M
assert result.amount_usd == Decimal("0.0775")