fix:_get_response_model

This commit is contained in:
Ishaan Jaffer 2026-01-13 14:11:30 -08:00
parent 1d8c8ad750
commit 14d98430c5

View file

@ -587,6 +587,24 @@ def _model_contains_known_llm_provider(model: str) -> bool:
return _provider_prefix in LlmProvidersSet
def _get_response_model(completion_response: Any) -> Optional[str]:
"""
Extract the model name from a completion response object.
Used as a fallback for cost calculation when the input model name
doesn't exist in model_cost (e.g., Azure Model Router).
"""
if completion_response is None:
return None
if isinstance(completion_response, BaseModel):
return getattr(completion_response, "model", None)
elif isinstance(completion_response, dict):
return completion_response.get("model", None)
return None
def _get_usage_object(
completion_response: Any,
) -> Optional[Usage]:
@ -933,9 +951,8 @@ def completion_cost( # noqa: PLR0915
router_model_id=router_model_id,
)
potential_model_names = [selected_model]
if model is not None:
potential_model_names.append(model)
potential_model_names = [selected_model, _get_response_model(completion_response)]
for idx, model in enumerate(potential_model_names):
try: