mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
f89ca64481
commit
0feca8641f
5 changed files with 138 additions and 38 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue