mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
c77b5bada5
commit
78ec018052
2 changed files with 38 additions and 3 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue