diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 8bb3a0ab1ee..b811dc0f6ee 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -48,6 +48,7 @@ async def _handle_completed_batch( custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"], model_name: str | None = None, litellm_params: dict | None = None, + model_info: ModelInfo | None = None, ) -> tuple[float, Usage, list[str]]: """Fetch a completed batch's output file and aggregate its cost, usage, and models in a single pass over the JSONL lines, so the parsed file content is @@ -58,6 +59,9 @@ async def _handle_completed_batch( custom_llm_provider: The LLM provider model_name: Optional model name litellm_params: Optional litellm parameters containing credentials (api_key, api_base, etc.) + model_info: Optional deployment-level model info with custom pricing, + threaded through so a deployment's configured rates win over the + global cost map. """ # A completed batch whose request lines all failed has no output file - the # results are written to a separate error_file_id and output_file_id is None. @@ -86,6 +90,7 @@ async def _handle_completed_batch( entries=_iter_batch_input_entries(file_content), custom_llm_provider=custom_llm_provider, model_name=model_name, + model_info=model_info, ) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 4a4a97b1d85..662ce38f746 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -108,6 +108,7 @@ from litellm.types.utils import ( LiteLLMBatch, LiteLLMLoggingBaseClass, LiteLLMRealtimeStreamLoggingObject, + ModelInfo, ModelResponse, ModelResponseStream, RawRequestTypedDict, @@ -579,6 +580,20 @@ class Logging(LiteLLMLoggingBaseClass): return model_id return None + def get_router_deployment_model_info(self) -> ModelInfo | None: + """Pricing the router registered under this deployment's model_info.id. + + Returns None when the deployment declares no pricing of its own, so the + caller falls back to the global cost map. + """ + model_id: Final = self.get_router_model_id() + if model_id is None: + return None + try: + return litellm.get_model_info(model=model_id) + except Exception: # noqa: BLE001 # get_model_info raises for any id with no registered pricing + return None + def update_environment_variables( self, litellm_params: dict, @@ -2600,7 +2615,9 @@ class Logging(LiteLLMLoggingBaseClass): ) = await _handle_completed_batch( batch=result, custom_llm_provider=self.custom_llm_provider, + model_name=self.model, litellm_params=self.litellm_params, + model_info=self.get_router_deployment_model_info(), ) result._hidden_params["response_cost"] = response_cost diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index 70d4ce2cebd..033febabd45 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -1320,3 +1320,98 @@ async def test_output_file_content_bedrock_reads_with_deployment_aws_credentials assert captured["aws_region_name"] == "us-west-2" assert captured["_litellm_internal_model_credentials"] is snapshot assert "model" not in captured + + +# =========================================================================== # +# _handle_completed_batch threads the deployment's model identity + pricing +# +# Regression: the retrieve path called _handle_completed_batch with neither +# model_name nor model_info. For bedrock that left cost_model falling back to +# the provider's own response model ("claude-sonnet-4-6"), which does not +# resolve under custom_llm_provider="bedrock", so cost silently became $0 while +# usage stayed correct. Dropping model_info separately discarded a deployment's +# configured rates, billing a zero-cost deployment at the public rate. +# =========================================================================== # + + +def _bedrock_row(model, input_tokens, output_tokens): + return { + "modelInput": {"messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]}, + "modelOutput": { + "model": model, + "id": "msg_1", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "stop_reason": "end_turn", + "usage": { + "input_tokens": input_tokens, + "output_tokens": output_tokens, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + }, + }, + "recordId": "r", + } + + +@pytest.mark.asyncio +async def test_handle_completed_bedrock_batch_prices_from_deployment_model(monkeypatch): + """A bedrock batch must price from the deployment model, not the response model.""" + rows = [_bedrock_row("claude-sonnet-4-6", 18, 10)] * 100 + + async def fake_fetch(batch, custom_llm_provider, litellm_params=None): + return _vertex_jsonl(rows) + + monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch) + + cost, usage, _ = await bu._handle_completed_batch( + _batch("of"), + custom_llm_provider="bedrock", + model_name="bedrock/global.anthropic.claude-sonnet-4-6", + ) + + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (1800, 1000, 2800) + # 3e-06 / 1.5e-05 on-demand, halved for batch. + assert cost == pytest.approx(1800 * 3e-06 / 2 + 1000 * 1.5e-05 / 2) + + # The response model alone cannot price a bedrock batch: this is the $0 bug. + zero_cost, zero_usage, _ = await bu._handle_completed_batch( + _batch("of"), + custom_llm_provider="bedrock", + model_name=None, + ) + assert zero_cost == 0.0 + assert zero_usage.total_tokens == 2800 + + +@pytest.mark.asyncio +async def test_handle_completed_batch_honors_deployment_pricing(monkeypatch): + """A deployment's configured rates must win over the global cost map.""" + rows = [_success_row(model="gemini-2.5-flash", usage=_usage(60, 75))] + + async def fake_fetch(batch, custom_llm_provider, litellm_params=None): + return _vertex_jsonl(rows) + + monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch) + + free_cost, _, _ = await bu._handle_completed_batch( + _batch("of"), + custom_llm_provider="vertex_ai", + model_name="vertex_ai/gemini-2.5-flash", + model_info={ + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "input_cost_per_token_batches": 0.0, + "output_cost_per_token_batches": 0.0, + }, + ) + assert free_cost == 0.0 + + billed_cost, _, _ = await bu._handle_completed_batch( + _batch("of"), + custom_llm_provider="vertex_ai", + model_name="vertex_ai/gemini-2.5-flash", + model_info=None, + ) + assert billed_cost > 0.0 diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 28a6c8dd18d..e9dc65e526d 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -1,3 +1,4 @@ +import contextlib import os import sys import asyncio @@ -340,6 +341,102 @@ class TestGetRouterModelId: assert obj.get_router_model_id() is None +class TestGetRouterDeploymentModelInfo: + """Pricing a deployment registered under its own model_info.id.""" + + def test_returns_registered_deployment_pricing(self, logging_obj): + deployment_id = "deploy-zero-cost-1" + litellm.model_cost[deployment_id] = { + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "input_cost_per_token_batches": 0.0, + "output_cost_per_token_batches": 0.0, + "litellm_provider": "vertex_ai", + "mode": "chat", + } + logging_obj.litellm_params = {"litellm_metadata": {"model_info": {"id": deployment_id}}} + try: + info = logging_obj.get_router_deployment_model_info() + assert info is not None + assert info["input_cost_per_token"] == 0.0 + assert info["output_cost_per_token_batches"] == 0.0 + finally: + litellm.model_cost.pop(deployment_id, None) + + def test_returns_none_for_unregistered_deployment(self, logging_obj): + logging_obj.litellm_params = {"litellm_metadata": {"model_info": {"id": "deploy-never-registered"}}} + assert logging_obj.get_router_deployment_model_info() is None + + def test_returns_none_without_a_deployment_id(self, logging_obj): + logging_obj.litellm_params = {"api_base": ""} + assert logging_obj.get_router_deployment_model_info() is None + + +class TestRetrieveBatchCostPassesModelIdentity: + """Regression: retrieving a batch priced it with no model identity at all. + + _handle_completed_batch was called without model_name or model_info, so a + bedrock batch fell back to the provider's own response model (unresolvable + under custom_llm_provider="bedrock") and silently cost $0, and a deployment's + configured rates were ignored entirely. + """ + + @pytest.mark.asyncio + async def test_forwards_deployment_model_and_pricing(self, monkeypatch): + from litellm.litellm_core_utils import litellm_logging as logging_module + from litellm.types.utils import LiteLLMBatch, Usage + + deployment_id = "deploy-batch-pricing-1" + litellm.model_cost[deployment_id] = { + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "bedrock", + "mode": "chat", + } + + captured: dict = {} + + async def fake_handle_completed_batch(**kwargs): + captured.update(kwargs) + return 1.25, Usage(prompt_tokens=1800, completion_tokens=1000, total_tokens=2800), ["m"] + + monkeypatch.setattr(logging_module, "_handle_completed_batch", fake_handle_completed_batch) + + obj = LitellmLogging( + model="bedrock/global.anthropic.claude-sonnet-4-6", + messages=[{"role": "user", "content": "Hey"}], + stream=False, + call_type="aretrieve_batch", + start_time=time.time(), + litellm_call_id="batch-call-1", + function_id="f", + ) + obj.custom_llm_provider = "bedrock" + obj.litellm_params = {"litellm_metadata": {"model_info": {"id": deployment_id}}} + + batch = LiteLLMBatch( + id="batch_abc", + completion_window="24h", + created_at=1, + endpoint="/v1/chat/completions", + input_file_id="file-in", + object="batch", + status="completed", + output_file_id="file-out", + ) + + try: + with contextlib.suppress(Exception): + await obj._async_success_handler_body(result=batch, start_time=None, end_time=None) + finally: + litellm.model_cost.pop(deployment_id, None) + + assert captured, "_handle_completed_batch was never called" + assert captured["model_name"] == "bedrock/global.anthropic.claude-sonnet-4-6" + assert captured["model_info"] is not None + assert captured["model_info"]["input_cost_per_token"] == 0.0 + + class TestAnthropicPassthroughCustomPricing: """Verify the Anthropic pass-through handler forwards custom pricing."""