diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index cb268513e00..182c176121d 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -1593,13 +1593,18 @@ def completion_cost( # (e.g. OCR per-page pricing). Same extraction as the video path # above, minus its `or {}` default: truthiness on the value adds # no mutable-collection construction (LIT002) and reads the same. + # Checked under both keys: router calls that go through + # `_ageneric_api_call_with_fallbacks` (OCR included) store the + # deployment's model_info under `litellm_metadata`, not `metadata`. _custom_model_info: ModelInfo | None = None if custom_pricing and litellm_logging_obj is not None: _cm_litellm_params = getattr(litellm_logging_obj, "litellm_params", None) if _cm_litellm_params is not None: - _cm_metadata = _cm_litellm_params.get("metadata") - if _cm_metadata: - _custom_model_info = _cm_metadata.get("model_info", None) + for _cm_metadata_key in ("metadata", "litellm_metadata"): + _cm_metadata = _cm_litellm_params.get(_cm_metadata_key) + if _cm_metadata and _cm_metadata.get("model_info") is not None: + _custom_model_info = _cm_metadata.get("model_info") + break ( prompt_tokens_cost_usd_dollar, diff --git a/tests/test_litellm/test_ocr_custom_pricing.py b/tests/test_litellm/test_ocr_custom_pricing.py index 7cd4fc1dec1..4dffc43c23f 100644 --- a/tests/test_litellm/test_ocr_custom_pricing.py +++ b/tests/test_litellm/test_ocr_custom_pricing.py @@ -111,3 +111,33 @@ def test_ocr_custom_pricing_end_to_end_through_completion_cost() -> None: litellm_logging_obj=logging_obj, ) assert cost == pytest.approx(CUSTOM_COST_PER_PAGE * 3) + + +def test_ocr_custom_pricing_end_to_end_via_litellm_metadata() -> None: + """Router OCR calls go through `_ageneric_api_call_with_fallbacks`, which + stores the deployment's model_info under `litellm_metadata` rather than + `metadata`. The extraction must read that key too.""" + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging + + logging_obj = LiteLLMLogging( + model=UNMAPPED_MODEL, + messages=[], + stream=False, + call_type="ocr", + start_time=None, + litellm_call_id="test-ocr-custom-pricing-litellm-metadata", + function_id="1234", + ) + logging_obj.litellm_params = { + "litellm_metadata": {"model_info": {"ocr_cost_per_page": CUSTOM_COST_PER_PAGE}}, + } + + cost = completion_cost( + completion_response=_ocr_response(UNMAPPED_MODEL, pages_processed=3), + model=UNMAPPED_MODEL, + custom_llm_provider="azure_ai", + call_type="ocr", + custom_pricing=True, + litellm_logging_obj=logging_obj, + ) + assert cost == pytest.approx(CUSTOM_COST_PER_PAGE * 3)