diff --git a/litellm/llms/vertex_ai/cost_calculator.py b/litellm/llms/vertex_ai/cost_calculator.py index 2e6e27c0cb7..0fca557e1f9 100644 --- a/litellm/llms/vertex_ai/cost_calculator.py +++ b/litellm/llms/vertex_ai/cost_calculator.py @@ -189,16 +189,19 @@ def _handle_128k_pricing( (prompt_tokens_details.cache_creation_tokens or 0) if prompt_tokens_details is not None else 0 ) text_tokens: Final = max(usage.prompt_tokens - cache_read_tokens - cache_creation_tokens, 0) + tier_tokens: Final = max(usage.prompt_tokens - cache_creation_tokens, 0) completion_tokens: Final = usage.completion_tokens input_rate: Final = ( input_cost_per_token_above_128k_tokens - if input_cost_per_token_above_128k_tokens is not None and _is_above_128k(tokens=text_tokens) + if input_cost_per_token_above_128k_tokens is not None and _is_above_128k(tokens=tier_tokens) else (model_info["input_cost_per_token"] or 0.0) ) cache_read_rate: Final = model_info.get("cache_read_input_token_cost") or input_rate - cache_creation_rate: Final = model_info.get("cache_creation_input_token_cost") or input_rate + cache_creation_rate: Final = model_info.get("cache_creation_input_token_cost") or ( + model_info["input_cost_per_token"] or 0.0 + ) prompt_cost = ( text_tokens * input_rate + cache_read_tokens * cache_read_rate + cache_creation_tokens * cache_creation_rate diff --git a/tests/test_litellm/llms/vertex_ai/test_cost_calculator.py b/tests/test_litellm/llms/vertex_ai/test_cost_calculator.py index 8e6e49f4ca2..7c558c3f549 100644 --- a/tests/test_litellm/llms/vertex_ai/test_cost_calculator.py +++ b/tests/test_litellm/llms/vertex_ai/test_cost_calculator.py @@ -8,17 +8,28 @@ from litellm.types.utils import PromptTokensDetailsWrapper, Usage @pytest.mark.parametrize( - ("text_tokens", "expected_prompt_cost"), + ("text_tokens", "cache_read_tokens", "cache_creation_tokens", "expected_prompt_cost"), [ - (140_000, 140_000 * 0.002 + 120_000 * 0.0005), - (120_000, 120_000 * 0.001 + 120_000 * 0.0005), + (140_000, 0, 120_000, 140_000 * 0.002 + 120_000 * 0.0005), + (120_000, 0, 120_000, 120_000 * 0.001 + 120_000 * 0.0005), + (100_000, 50_000, 0, 100_000 * 0.002 + 50_000 * 0.00025), + (100_000, 0, 120_000, 100_000 * 0.001 + 120_000 * 0.0005), + ], + ids=[ + "creation_tokens_do_not_change_the_tier_rate", + "creation_tokens_cannot_push_the_tier_threshold", + "cache_read_tokens_count_toward_the_tier", + "small_text_plus_creation_stays_below_the_tier", ], - ids=["creation_tokens_do_not_change_the_tier_rate", "creation_tokens_cannot_push_the_tier_threshold"], ) -def test_above_128k_pricing_splits_cache_creation_tokens_out_of_the_prompt( - monkeypatch: pytest.MonkeyPatch, text_tokens: int, expected_prompt_cost: float +def test_above_128k_pricing_splits_cache_tokens_out_of_the_prompt( + monkeypatch: pytest.MonkeyPatch, + text_tokens: int, + cache_read_tokens: int, + cache_creation_tokens: int, + expected_prompt_cost: float, ) -> None: - """Cache creation tokens bill at the cache-creation rate and never count toward the above-128k tier.""" + """Cache reads count toward the above-128k tier, cache creation bills at the cache-creation rate.""" model: Final = "vertex_ai/fake-above-128k-model" monkeypatch.setitem( litellm.model_cost, @@ -27,16 +38,20 @@ def test_above_128k_pricing_splits_cache_creation_tokens_out_of_the_prompt( "litellm_provider": "vertex_ai", "input_cost_per_token": 0.001, "input_cost_per_token_above_128k_tokens": 0.002, + "cache_read_input_token_cost": 0.00025, "cache_creation_input_token_cost": 0.0005, "output_cost_per_token": 0.003, }, ) usage: Final = Usage( - prompt_tokens=text_tokens + 120_000, + prompt_tokens=text_tokens + cache_read_tokens + cache_creation_tokens, completion_tokens=10, - total_tokens=text_tokens + 120_010, - prompt_tokens_details=PromptTokensDetailsWrapper(cache_creation_tokens=120_000), + total_tokens=text_tokens + cache_read_tokens + cache_creation_tokens + 10, + prompt_tokens_details=PromptTokensDetailsWrapper( + cached_tokens=cache_read_tokens or 0, + cache_creation_tokens=cache_creation_tokens or 0, + ), ) prompt_cost, completion_cost = cost_per_token( model=model,