fix(cost): read deployment model_info from litellm_metadata too

Router OCR calls go through _ageneric_api_call_with_fallbacks, which stores the
deployment's model_info under litellm_metadata rather than metadata, so the
custom OCR pricing was still unreachable on that path. Check both keys, metadata
first, mirroring _get_base_model_from_litellm_call_metadata. Adds an end-to-end
test for the litellm_metadata shape.
This commit is contained in:
Mihidum Hettiyahandi 2026-08-24 09:23:39 +10:00
parent c77b5bada5
commit 78ec018052
2 changed files with 38 additions and 3 deletions

View file

@ -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,

View file

@ -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)