mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
69def0545d
commit
726abc68a6
3 changed files with 62 additions and 12 deletions
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue