mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(batches): resolve proxy model alias to deployment model before provider call
POST /v1/batches with a model-group alias (model-encoded input_file_id, or a
model header/query/body param) passed that alias straight to
litellm.acreate_batch. create_batch runs the model through get_llm_provider,
which cannot resolve a proxy alias, so the alias reached the provider transform
unchanged; the Bedrock batch transform forwards model as the batch modelId and
AWS rejected it ("The provided model identifier is invalid").
Resolve the alias to the deployment's real litellm_params.model at the proxy
batch-create endpoint and swap it onto the request before the provider call,
the same resolution the router's own _acreate_batch already does internally.
Response IDs stay encoded with the alias so retrieve/cancel route back. Adds
Router.get_deployment_model_for_alias on top of a shared
_resolve_unblocked_deployment helper extracted from the credential resolver.
This commit is contained in:
parent
63a9bb556c
commit
5a2b7ead30
4 changed files with 225 additions and 21 deletions
|
|
@ -5,7 +5,7 @@
|
|||
|
||||
######################################################################
|
||||
import asyncio
|
||||
from typing import Any, Dict, Optional, cast
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response
|
||||
|
||||
|
|
@ -40,9 +40,34 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
|||
from litellm.proxy.utils import handle_exception_on_proxy, is_known_model
|
||||
from litellm.types.llms.openai import LiteLLMBatchCreateRequest
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _swap_alias_for_deployment_model(
|
||||
create_batch_data: LiteLLMBatchCreateRequest,
|
||||
alias: str,
|
||||
llm_router: Optional["Router"],
|
||||
) -> None:
|
||||
"""
|
||||
Replace a proxy model-group alias on the batch request with the
|
||||
deployment's real provider model (in place).
|
||||
|
||||
``litellm.create_batch`` runs the model through ``get_llm_provider``, which
|
||||
cannot resolve a proxy alias, so a provider transform (e.g. Bedrock's, which
|
||||
forwards ``model`` as the batch ``modelId``) would otherwise receive the
|
||||
alias and the provider would reject it. Falls back to the alias when the
|
||||
router is unavailable or the alias resolves to nothing.
|
||||
"""
|
||||
if llm_router is None:
|
||||
return
|
||||
resolved_model = llm_router.get_deployment_model_for_alias(model_id=alias)
|
||||
if resolved_model is not None:
|
||||
create_batch_data["model"] = resolved_model
|
||||
|
||||
|
||||
@router.post(
|
||||
"/{provider}/v1/batches",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
|
|
@ -173,6 +198,11 @@ async def create_batch(
|
|||
data=_create_batch_data, # type: ignore
|
||||
credentials=credentials,
|
||||
)
|
||||
_swap_alias_for_deployment_model(
|
||||
create_batch_data=_create_batch_data,
|
||||
alias=model_from_file_id,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
# Create batch using model credentials
|
||||
response = await litellm.acreate_batch(
|
||||
|
|
@ -269,6 +299,11 @@ async def create_batch(
|
|||
data=_create_batch_data, # type: ignore
|
||||
credentials=credentials,
|
||||
)
|
||||
_swap_alias_for_deployment_model(
|
||||
create_batch_data=_create_batch_data,
|
||||
alias=model_param,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
# Create batch using model credentials
|
||||
response = await litellm.acreate_batch(
|
||||
|
|
|
|||
|
|
@ -9154,6 +9154,48 @@ class Router:
|
|||
raise Exception("Model Name invalid - {}".format(type(model)))
|
||||
return None
|
||||
|
||||
def _resolve_unblocked_deployment(self, model_id: str) -> Optional[Deployment]:
|
||||
"""
|
||||
Resolve a model id, model-group alias, or wildcard pattern to a single
|
||||
deployment, returning None when nothing matches or the match is paused
|
||||
via ``LiteLLM_ProxyModelTable.blocked``.
|
||||
"""
|
||||
deployment = self.get_deployment(model_id=model_id)
|
||||
|
||||
if deployment is None:
|
||||
deployment = self.get_deployment_by_model_group_name(
|
||||
model_group_name=model_id
|
||||
)
|
||||
|
||||
if deployment is None:
|
||||
potential_wildcard_models = self.pattern_router.route(model_id) or []
|
||||
if potential_wildcard_models:
|
||||
deployment_dict = 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
|
||||
return deployment
|
||||
|
||||
def get_deployment_model_for_alias(self, model_id: str) -> Optional[str]:
|
||||
"""
|
||||
Resolve a model-group alias (or deployment id / wildcard) to the
|
||||
deployment's underlying ``litellm_params.model``.
|
||||
|
||||
Callers that hand a model to provider SDKs (e.g. the proxy batch-create
|
||||
path) need the real provider model id, not the proxy alias:
|
||||
``get_llm_provider`` cannot resolve an alias, so passing it straight
|
||||
through reaches the provider as an invalid model identifier. Returns
|
||||
None when the alias resolves to nothing or to a paused deployment.
|
||||
"""
|
||||
deployment = self._resolve_unblocked_deployment(model_id=model_id)
|
||||
if deployment is None:
|
||||
return None
|
||||
return deployment.litellm_params.model
|
||||
|
||||
def get_deployment_credentials_with_provider(
|
||||
self, model_id: str
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
|
|
@ -9177,27 +9219,8 @@ class Router:
|
|||
credentials = router.get_deployment_credentials_with_provider("gpt-4o-litellm")
|
||||
# Returns: {"api_key": "sk-...", "custom_llm_provider": "openai", ...}
|
||||
"""
|
||||
# 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 = self._resolve_unblocked_deployment(model_id=model_id)
|
||||
if deployment is None:
|
||||
deployment = self.get_deployment_by_model_group_name(
|
||||
model_group_name=model_id
|
||||
)
|
||||
|
||||
# If still not found, check for wildcard pattern matches
|
||||
if deployment is None:
|
||||
potential_wildcard_models = self.pattern_router.route(model_id) or []
|
||||
if potential_wildcard_models:
|
||||
# Use the first matching wildcard deployment
|
||||
deployment_dict = 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
|
||||
|
|
|
|||
|
|
@ -416,6 +416,110 @@ class TestBatchIdRoundTripWithRetrieve:
|
|||
assert get_original_file_id(encoded) == raw_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_batch_swaps_alias_for_deployment_model_before_provider_call():
|
||||
"""
|
||||
SCENARIO 1 (model-encoded input_file_id): the proxy must hand
|
||||
litellm.acreate_batch the deployment's real provider model, not the proxy
|
||||
alias. get_llm_provider cannot resolve an alias, so passing it straight
|
||||
through reaches the Bedrock batch transform as an invalid modelId. The
|
||||
response IDs must still be encoded with the ALIAS so retrieve routes back.
|
||||
"""
|
||||
from litellm.proxy.batches_endpoints.endpoints import create_batch
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
encode_file_id_with_model,
|
||||
)
|
||||
|
||||
alias = "bedrock-batch-haiku"
|
||||
real_model = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
|
||||
raw_input_file_id = "s3://bucket/litellm-bedrock-files/in.jsonl"
|
||||
encoded_input_file_id = encode_file_id_with_model(
|
||||
file_id=raw_input_file_id, model=alias
|
||||
)
|
||||
raw_batch_id = "batch_bedrock_123"
|
||||
|
||||
mock_response = _make_batch_response(batch_id=raw_batch_id)
|
||||
mock_request = _make_mock_request(headers={})
|
||||
mock_fastapi_response = MagicMock()
|
||||
mock_user_api_key_dict = MagicMock()
|
||||
mock_user_api_key_dict.parent_otel_span = None
|
||||
mock_user_api_key_dict.user_id = "test_user"
|
||||
mock_user_api_key_dict.team_metadata = {}
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_deployment_model_for_alias = MagicMock(return_value=real_model)
|
||||
|
||||
mock_credentials = {
|
||||
"custom_llm_provider": "bedrock",
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
|
||||
request_body = {
|
||||
"input_file_id": encoded_input_file_id,
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
"model": alias,
|
||||
}
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.batches_endpoints.endpoints._read_request_body",
|
||||
new=AsyncMock(return_value=dict(request_body)),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.batches_endpoints.endpoints.ProxyBaseLLMRequestProcessing"
|
||||
) as mock_processor_cls,
|
||||
patch(
|
||||
"litellm.proxy.batches_endpoints.endpoints.get_credentials_for_model",
|
||||
return_value=mock_credentials,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.batches_endpoints.endpoints.prepare_data_with_credentials",
|
||||
),
|
||||
patch(
|
||||
"litellm.acreate_batch",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_create_batch,
|
||||
patch(
|
||||
"litellm.proxy.batches_endpoints.endpoints.is_known_model",
|
||||
return_value=False,
|
||||
),
|
||||
patch("litellm.proxy.proxy_server.general_settings", {}),
|
||||
patch("litellm.proxy.proxy_server.llm_router", mock_router),
|
||||
patch("litellm.proxy.proxy_server.proxy_config", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.version", "1.0.0"),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj",
|
||||
MagicMock(
|
||||
post_call_success_hook=AsyncMock(return_value=mock_response),
|
||||
update_request_status=AsyncMock(),
|
||||
),
|
||||
),
|
||||
):
|
||||
mock_create_batch.return_value = mock_response
|
||||
mock_processor = MagicMock()
|
||||
mock_processor.common_processing_pre_call_logic = AsyncMock(
|
||||
return_value=(dict(request_body), MagicMock())
|
||||
)
|
||||
mock_processor_cls.return_value = mock_processor
|
||||
|
||||
response = await create_batch(
|
||||
request=mock_request,
|
||||
fastapi_response=mock_fastapi_response,
|
||||
provider=None,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
mock_router.get_deployment_model_for_alias.assert_called_once_with(model_id=alias)
|
||||
create_kwargs = mock_create_batch.call_args.kwargs
|
||||
assert create_kwargs["model"] == real_model, (
|
||||
"Bedrock batch transform receives modelId from this 'model'; it must be the "
|
||||
f"deployment's real model, got {create_kwargs.get('model')!r}"
|
||||
)
|
||||
# The encoded response id must carry the ALIAS so retrieve routes back.
|
||||
assert decode_model_from_file_id(response.id) == alias
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel_batch_with_unified_id_routes_with_decoded_model_and_batch_id():
|
||||
from litellm.proxy.batches_endpoints.endpoints import cancel_batch
|
||||
|
|
|
|||
|
|
@ -5073,6 +5073,48 @@ def test_get_deployment_credentials_with_provider_returns_none_for_blocked_deplo
|
|||
assert router.get_deployment_credentials_with_provider(model_id="dep-1") is not None
|
||||
|
||||
|
||||
def test_get_deployment_model_for_alias_resolves_underlying_model():
|
||||
"""
|
||||
The proxy batch-create path resolves a model-group alias to its deployment
|
||||
so it can hand the provider the deployment's real model id, not the alias.
|
||||
get_llm_provider cannot resolve a proxy alias, so without this the Bedrock
|
||||
batch transform receives the alias as a modelId and AWS rejects it.
|
||||
"""
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "bedrock-batch-haiku",
|
||||
"litellm_params": {
|
||||
"model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
|
||||
"aws_region_name": "us-east-1",
|
||||
},
|
||||
"model_info": {"id": "bedrock-batch-dep-0"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
assert (
|
||||
router.get_deployment_model_for_alias(model_id="bedrock-batch-haiku")
|
||||
== "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
|
||||
)
|
||||
# Resolving by deployment id returns the same underlying model.
|
||||
assert (
|
||||
router.get_deployment_model_for_alias(model_id="bedrock-batch-dep-0")
|
||||
== "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
|
||||
)
|
||||
|
||||
|
||||
def test_get_deployment_model_for_alias_returns_none_for_unknown_model():
|
||||
router = _router_with_two_deployments([False, False])
|
||||
assert router.get_deployment_model_for_alias(model_id="does-not-exist") is None
|
||||
|
||||
|
||||
def test_get_deployment_model_for_alias_returns_none_for_blocked_deployment():
|
||||
router = _router_with_two_deployments([True, False])
|
||||
assert router.get_deployment_model_for_alias(model_id="dep-0") is None
|
||||
assert router.get_deployment_model_for_alias(model_id="dep-1") == "openai/gpt-4o-1"
|
||||
|
||||
|
||||
def test_is_deployment_blocked_static_helper_reflects_blocked_flag():
|
||||
"""
|
||||
Exercises Router._is_deployment_blocked so router_code_coverage.py (AST call graph)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue