fix(batches): price model-encoded batch retrievals by their deployment

A batch retrieved by its model-encoded id takes the direct (non-router) path,
which resolved credentials without stamping the deployment's model_info, so a
completed batch on a deployment with its own per-page pricing was billed at
the published rate with an empty model_id on the spend row.

Extract the router's credential lookup into get_credential_deployment and
stamp the resolved deployment's model_info onto the retrieve call the way the
router does for routed calls.
This commit is contained in:
mateo-berri 2026-09-19 00:30:15 -07:00
parent f89ca64481
commit 0feca8641f
5 changed files with 138 additions and 38 deletions

View file

@ -32,6 +32,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
from litellm.proxy.openai_files_endpoints.common_utils import (
BATCH_CREATE_HIDDEN_PARAM,
_is_base64_encoded_unified_file_id,
add_deployment_model_info,
add_internal_model_credentials,
apply_team_provider_credentials,
authorize_model_for_key,
@ -580,6 +581,7 @@ async def retrieve_batch(
# so litellm.aretrieve_batch can load BedrockBatchesConfig. Without
# it the call falls into the legacy provider switch and 400s.
data["model"] = model_from_id
add_deployment_model_info(data=data, llm_router=llm_router, model_id=model_from_id)
# Retrieve batch using model credentials
response = await litellm.aretrieve_batch(

View file

@ -593,6 +593,25 @@ def add_internal_model_credentials(
data["_litellm_internal_model_credentials"] = MappingProxyType(dict(credentials))
def add_deployment_model_info(
data: dict,
llm_router: Optional["Router"],
model_id: str,
) -> None:
"""
Stamp the resolved deployment's `model_info` onto a direct (non-router) batch call
(in-place), the way the router does for routed calls, so the completed batch is
priced by its deployment id instead of the published model rate.
"""
deployment: Final = llm_router.get_credential_deployment(model_id=model_id) if llm_router is not None else None
if deployment is None:
return
data["litellm_metadata"] = {
**(data.get("litellm_metadata") or {}),
"model_info": deployment.model_info.model_dump(),
}
def prepare_data_with_credentials(
data: dict,
credentials: dict,

View file

@ -10313,6 +10313,55 @@ class Router:
return display_name
return None
def get_credential_deployment(self, model_id: str, team_id: str | None = None) -> Deployment | None:
"""
The deployment a passthrough endpoint (files, batches, etc.) resolves for a
model id or model name: by deployment id first, then by model_name, then by
the team's exact public model name, then by wildcard pattern (team wildcards
before global ones, so a global "openai/*" never shadows the team's own
entry). Name and wildcard lookups never resolve another team's deployment.
Returns None when nothing matches or the match is paused via
`LiteLLM_ProxyModelTable.blocked`, so callers cannot bypass an admin pause
by resolving the deployment directly.
"""
deployment: Final = (
self.get_deployment(model_id=model_id)
or self._get_model_group_deployment_usable_by_team(model_group_name=model_id, team_id=team_id)
or self._get_team_public_name_deployment(model_id=model_id, team_id=team_id)
or self._get_wildcard_deployment_usable_by_team(model_id=model_id, team_id=team_id)
)
if deployment is None or self._is_deployment_blocked(deployment):
return None
return deployment
def _get_team_public_name_deployment(self, model_id: str, team_id: str | None) -> Deployment | None:
if team_id is None:
return None
team_indices: Final = self.team_model_to_deployment_indices.get((team_id, model_id))
if not team_indices:
return None
team_model: Final = self.model_list[team_indices[0]]
return Deployment(**team_model) if isinstance(team_model, dict) else team_model
def _get_wildcard_deployment_usable_by_team(self, model_id: str, team_id: str | None) -> Deployment | None:
team_pattern_router: Final = self.team_pattern_routers.get(team_id) if team_id is not None else None
team_wildcard_models: Final = team_pattern_router.route(model_id) if team_pattern_router else None
global_wildcard_models: Final = tuple(
wildcard_model
for wildcard_model in (self.pattern_router.route(model_id) or ())
if self._deployment_usable_by_team(wildcard_model, team_id)
)
potential_wildcard_models: Final = team_wildcard_models or global_wildcard_models
if not potential_wildcard_models:
return None
wildcard_deployment: Final = potential_wildcard_models[0]
if isinstance(wildcard_deployment, dict):
return Deployment(**wildcard_deployment)
if isinstance(wildcard_deployment, Deployment):
return wildcard_deployment
return None
def get_deployment_credentials_with_provider(
self, model_id: str, team_id: str | None = None
) -> dict[str, Any] | None:
@ -10320,8 +10369,8 @@ class Router:
Get API credentials and provider info from a model name in model_list.
Useful for passthrough endpoints (files, batches, etc.) that need credentials.
This method tries to find a deployment by model_id first, and if not found,
it tries to find by model_group_name (model_name).
Resolves the deployment with `get_credential_deployment` (by deployment id,
then model_name, team public model name, and wildcard pattern).
Args:
model_id: Model ID or model name from model_list (e.g., "gpt-4o-litellm")
@ -10342,43 +10391,8 @@ class Router:
credentials = router.get_deployment_credentials_with_provider("gpt-4o-litellm")
# Returns: {"api_key": "sk-...", "custom_llm_provider": "openai", "model": "gpt-4o", ...}
"""
# Try to get deployment by model_id first
deployment = self.get_deployment(model_id=model_id)
# If not found, try by model_group_name
deployment: Final = self.get_credential_deployment(model_id=model_id, team_id=team_id)
if deployment is None:
deployment = self._get_model_group_deployment_usable_by_team(model_group_name=model_id, team_id=team_id)
# If not found, check team-scoped deployments whose team public model
# name exactly matches model_id (wildcard team names are matched via
# team_pattern_routers below).
if deployment is None and team_id is not None:
team_indices: Final = self.team_model_to_deployment_indices.get((team_id, model_id), [])
if team_indices:
team_model: Final = self.model_list[team_indices[0]]
deployment = Deployment(**team_model) if isinstance(team_model, dict) else team_model
# If still not found, check for wildcard pattern matches. Team wildcard
# matches take priority so a global pattern (e.g. "openai/*") doesn't
# shadow the team's own entry.
if deployment is None:
team_pattern_router: Final = self.team_pattern_routers.get(team_id) if team_id is not None else None
team_wildcard_models: Final = (team_pattern_router.route(model_id) or []) if team_pattern_router else []
global_wildcard_models: Final = [
wildcard_model
for wildcard_model in (self.pattern_router.route(model_id) or [])
if self._deployment_usable_by_team(wildcard_model, team_id)
]
potential_wildcard_models: Final = team_wildcard_models or global_wildcard_models
if potential_wildcard_models:
# Use the first matching wildcard deployment
deployment_dict: Final = potential_wildcard_models[0]
if isinstance(deployment_dict, dict):
deployment = Deployment(**deployment_dict)
elif isinstance(deployment_dict, Deployment):
deployment = deployment_dict
if deployment is None or self._is_deployment_blocked(deployment):
return None
# Get basic credentials

View file

@ -51,6 +51,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
from litellm.proxy.utils import ProxyLogging
from litellm.router import Router
from litellm.types.llms.openai import BatchJobStatus
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
from litellm.types.utils import CredentialItem, LiteLLMBatch
from fastapi import Request, Response
@ -1194,6 +1195,7 @@ def retrieve_harness():
router.model_list = []
router.aretrieve_batch = AsyncMock(return_value=make_batch())
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
router.get_credential_deployment = MagicMock(return_value=None)
pre_call = AsyncMock(side_effect=lambda **kw: (data_holder["data"], MagicMock()))
get_headers = MagicMock(return_value={})
@ -1312,6 +1314,30 @@ async def test_retrieve__model_encoded_id(retrieve_harness):
assert retrieve_harness.update_batch_in_db.call_args.kwargs["operation"] == "retrieve"
@pytest.mark.asyncio
async def test_retrieve__model_encoded_id__stamps_deployment_model_info_for_cost(retrieve_harness):
"""Regression: this path calls litellm.aretrieve_batch directly, so nothing stamped the
deployment's model_info the way the router does for routed calls. Cost tracking then never
saw the deployment id, and a completed batch on a deployment with its own per-page pricing
was billed at the published rate with an empty model_id on the spend row."""
retrieve_harness.router.get_credential_deployment.return_value = Deployment(
model_name="azure-gpt",
litellm_params=LiteLLM_Params(model="azure/gpt-4o"),
model_info=ModelInfo(id="dep-123"),
)
retrieve_harness.pre_call.side_effect = lambda **kw: (
{**retrieve_harness.data["data"], "litellm_metadata": {"user_api_key_alias": "qa-key"}},
MagicMock(),
)
await call_retrieve(retrieve_harness, AZURE_BATCH_ID)
retrieve_harness.router.get_credential_deployment.assert_called_once_with(model_id="azure/gpt-4o")
litellm_metadata = retrieve_harness.aretrieve_kwargs()["litellm_metadata"]
assert litellm_metadata["model_info"]["id"] == "dep-123"
assert litellm_metadata["user_api_key_alias"] == "qa-key"
@pytest.mark.asyncio
async def test_retrieve__model_encoded_id__forwards_decoded_model_not_deployment(
retrieve_harness,

View file

@ -5745,6 +5745,45 @@ async def test_router_unknown_model_error_message_renders_model_name_literally()
assert " " not in message # no padding run from an expanded format field
def test_get_credential_deployment_is_the_deployment_credentials_resolve_to():
"""Regression: a batch retrieved with credentials resolved by model name was priced
without its deployment id, so per-deployment pricing never applied. The deployment
behind the credentials must be reachable by name and by id, carrying its model_info."""
router = litellm.Router(
model_list=[
{
"model_name": "mistral-ocr",
"litellm_params": {"model": "mistral/mistral-ocr-latest", "api_key": "sk-ocr"},
"model_info": {"id": "ocr-dep", "ocr_cost_per_page_batches": 0.0123},
}
]
)
by_name = router.get_credential_deployment(model_id="mistral-ocr")
by_id = router.get_credential_deployment(model_id="ocr-dep")
assert by_name is not None and by_id is not None
assert by_name.model_info.id == by_id.model_info.id == "ocr-dep"
assert by_name.model_info.model_dump()["ocr_cost_per_page_batches"] == 0.0123
assert router.get_deployment_credentials_with_provider(model_id="mistral-ocr")["api_key"] == "sk-ocr"
assert router.get_credential_deployment(model_id="no-such-model") is None
def test_get_credential_deployment_skips_a_paused_deployment():
router = litellm.Router(
model_list=[
{
"model_name": "paused-ocr",
"litellm_params": {"model": "mistral/mistral-ocr-latest", "api_key": "sk-ocr"},
"model_info": {"id": "paused-dep", "blocked": True},
}
]
)
assert router.get_credential_deployment(model_id="paused-ocr") is None
assert router.get_credential_deployment(model_id="paused-dep") is None
def test_get_deployment_credentials_with_provider_aws_bedrock_runtime_endpoint():
"""
Test that get_deployment_credentials_with_provider correctly copies