fix(cost_calc): apply xAI's inclusive 200k threshold to the token-type breakdown

Derive threshold inclusivity from the provider inside generic_cost_per_token and
get_token_type_cost_breakdown so the spend-log breakdown can never disagree with
the billed totals at exactly 200k prompt tokens
This commit is contained in:
mateo-berri 2026-08-10 10:51:52 -07:00
parent 69def0545d
commit 726abc68a6
3 changed files with 62 additions and 12 deletions

View file

@ -49,6 +49,12 @@ _SERVICE_TIER_TO_COST_KEY_SUFFIX: Final[Mapping[str, str]] = MappingProxyType(
}
)
_INCLUSIVE_THRESHOLD_PROVIDERS: Final = frozenset({"xai"})
def _uses_inclusive_token_thresholds(custom_llm_provider: str | None) -> bool:
return custom_llm_provider in _INCLUSIVE_THRESHOLD_PROVIDERS
def _get_token_detail_value(details: object, key: str) -> int | None:
if isinstance(details, dict):
@ -712,7 +718,6 @@ def generic_cost_per_token(
service_tier: str | None = None,
data_residency: str | None = None,
model_info: ModelInfo | None = None,
threshold_is_inclusive: bool = False,
) -> tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -724,8 +729,6 @@ def generic_cost_per_token(
- usage: LiteLLM Usage block, containing anthropic caching information
- data_residency: optional OpenAI data-residency region (e.g. "eu", "us"),
used to apply the per-model regional-processing uplift multiplier.
- threshold_is_inclusive: bill the above-threshold tier when the prompt is exactly
at the threshold, as xAI does.
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
@ -791,7 +794,7 @@ def generic_cost_per_token(
model_info=model_info,
usage=usage,
service_tier=service_tier,
threshold_is_inclusive=threshold_is_inclusive,
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
)
prompt_cost = _calculate_input_cost(
@ -924,7 +927,12 @@ def get_token_type_cost_breakdown(
cache_creation_cost_rate,
cache_creation_cost_above_1hr_rate,
cache_read_cost_rate,
) = _get_token_base_cost(model_info=model_info, usage=usage, service_tier=service_tier)
) = _get_token_base_cost(
model_info=model_info,
usage=usage,
service_tier=service_tier,
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
)
reasoning_tokens = (
_parse_completion_tokens_details(usage)["reasoning_tokens"]

View file

@ -47,13 +47,7 @@ def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
completion_tokens_details=None,
)
# xAI bills the higher tier once the prompt reaches 200k tokens, not strictly above it
prompt_cost, completion_cost = generic_cost_per_token(
model=model,
usage=modified_usage,
custom_llm_provider="xai",
threshold_is_inclusive=True,
)
prompt_cost, completion_cost = generic_cost_per_token(model=model, usage=modified_usage, custom_llm_provider="xai")
return prompt_cost, completion_cost

View file

@ -2143,6 +2143,54 @@ def test_token_type_cost_breakdown_matches_real_gemini_numbers():
assert breakdown.cache_creation_cost == 0.0
def test_token_type_cost_breakdown_xai_at_exactly_200k_uses_higher_tier_rates():
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
usage = Usage(
prompt_tokens=200_000,
completion_tokens=2_000,
total_tokens=202_000,
completion_tokens_details=CompletionTokensDetailsWrapper(
reasoning_tokens=1_500, text_tokens=500
),
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=50_000, text_tokens=150_000
),
)
breakdown = get_token_type_cost_breakdown(
model="grok-4.20-0309-reasoning", custom_llm_provider="xai", usage=usage
)
assert breakdown.reasoning_cost == pytest.approx(1_500 * 5e-06)
assert breakdown.cache_read_cost == pytest.approx(50_000 * 4e-07)
def test_token_type_cost_breakdown_xai_just_below_200k_uses_base_tier_rates():
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
usage = Usage(
prompt_tokens=199_999,
completion_tokens=2_000,
total_tokens=201_999,
completion_tokens_details=CompletionTokensDetailsWrapper(
reasoning_tokens=1_500, text_tokens=500
),
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=50_000, text_tokens=149_999
),
)
breakdown = get_token_type_cost_breakdown(
model="grok-4.20-0309-reasoning", custom_llm_provider="xai", usage=usage
)
assert breakdown.reasoning_cost == pytest.approx(1_500 * 2.5e-06)
assert breakdown.cache_read_cost == pytest.approx(50_000 * 2e-07)
def test_token_type_cost_breakdown_includes_cache_creation_from_top_level_usage():
"""
Bedrock/Anthropic report cache tokens as top-level usage fields; the Usage