fix(proxy): enforce key/team/org/project model grants on model-routed file and batch credentials

Files and batches routes take their model from a header, query param or a
model-encoded resource id, which the auth layer never sees, so any key could
name any deployment and act on that provider account with its server-side key.
Every caller-supplied model now goes through can_key_call_resolved_model before
deployment credentials are resolved, covering file create/retrieve/content/
delete/list, batch create/retrieve/list/cancel, and vector store files.
This commit is contained in:
mubashir1osmani 2026-09-10 17:28:27 -04:00
parent 1d911c766f
commit edd5727f3c
7 changed files with 322 additions and 77 deletions

View file

@ -34,9 +34,9 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
encode_batch_response_ids,
encode_file_id_with_model,
ensure_batch_response_managed_file_ids,
get_authorized_credentials_for_model,
get_batch_from_database,
get_batch_id_from_unified_batch_id,
get_credentials_for_model,
get_model_id_from_unified_batch_id,
get_models_from_unified_file_id,
get_original_file_id,
@ -218,9 +218,10 @@ async def create_batch(
# SCENARIO 1: File ID is encoded with model info
if model_from_file_id is not None and input_file_id:
credentials = get_credentials_for_model(
credentials = await get_authorized_credentials_for_model(
llm_router=llm_router,
model_id=model_from_file_id,
user_api_key_dict=user_api_key_dict,
operation_context="batch creation (file created with model)",
)
@ -310,9 +311,10 @@ async def create_batch(
# SCENARIO 2 & 3: Model from header/query OR custom_llm_provider fallback
if model_param:
# SCENARIO 2: Use model-based routing from header/query/body
credentials = get_credentials_for_model(
credentials = await get_authorized_credentials_for_model(
llm_router=llm_router,
model_id=model_param,
user_api_key_dict=user_api_key_dict,
operation_context="batch creation",
)
@ -540,9 +542,10 @@ async def retrieve_batch(
# Retrieve from provider (for non-terminal states or if DB lookup failed)
# SCENARIO 1: Batch ID is encoded with model info
if model_from_id is not None:
credentials: Final = get_credentials_for_model(
credentials: Final = await get_authorized_credentials_for_model(
llm_router=llm_router,
model_id=model_from_id,
user_api_key_dict=user_api_key_dict,
operation_context="batch retrieval (batch created with model)",
)
@ -764,9 +767,10 @@ async def list_batches(
data.get("model") or request.query_params.get("model") or request.headers.get("x-litellm-model")
):
# SCENARIO 2: Use model-based routing from header/query/body
credentials: Final = get_credentials_for_model(
credentials: Final = await get_authorized_credentials_for_model(
llm_router=llm_router,
model_id=model_param,
user_api_key_dict=user_api_key_dict,
operation_context="batch listing",
)
@ -952,9 +956,10 @@ async def cancel_batch(
# SCENARIO 1: Batch ID is encoded with model info
if model_from_id is not None:
credentials: Final = get_credentials_for_model(
credentials: Final = await get_authorized_credentials_for_model(
llm_router=llm_router,
model_id=model_from_id,
user_api_key_dict=user_api_key_dict,
operation_context="batch cancellation (batch created with model)",
)

View file

@ -351,6 +351,10 @@ def get_credentials_for_model(
"""
Retrieve API credentials for a model from the LLM Router.
Does not check whether the caller may use ``model_id``; use
``get_authorized_credentials_for_model`` for anything driven by a caller-supplied
model name (request body, header, query param, or a model-encoded resource id).
Args:
llm_router: LiteLLM Router instance
model_id: Model name or deployment ID
@ -381,6 +385,48 @@ def get_credentials_for_model(
return credentials
async def authorize_model_for_key(
model_id: str,
llm_router: Optional["Router"],
user_api_key_dict: "UserAPIKeyAuth",
) -> None:
"""
Enforce the caller's model grants on a model name the auth layer never saw.
The files and batches routes carry their model in a header, query param, or a
model-encoded resource id rather than the request body, so ``user_api_key_auth``
cannot check it. Run the same key, team (incl. team-member and access-group
fallbacks), org and project allowlist checks a chat request would get, so a
restricted key cannot borrow another deployment's server-side credentials.
Raises:
ProxyException (403): the caller is not allowed to use ``model_id``
"""
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
await can_key_call_resolved_model(
model=model_id,
llm_model_list=None,
valid_token=user_api_key_dict,
llm_router=llm_router,
)
async def get_authorized_credentials_for_model(
llm_router: Optional["Router"],
model_id: str,
user_api_key_dict: "UserAPIKeyAuth",
operation_context: str = "file operation",
) -> dict: # mutable-ok: same contract as get_credentials_for_model, callers merge it into request data
"""``get_credentials_for_model`` gated by ``authorize_model_for_key``."""
await authorize_model_for_key(model_id=model_id, llm_router=llm_router, user_api_key_dict=user_api_key_dict)
return get_credentials_for_model(
llm_router=llm_router,
model_id=model_id,
operation_context=operation_context,
)
def get_team_provider_credentials(
llm_router: Optional["Router"],
user_api_key_dict: "UserAPIKeyAuth",
@ -573,21 +619,27 @@ def prepare_data_with_credentials(
data["file_id"] = file_id
def handle_model_based_routing(
async def handle_model_based_routing(
file_id: str,
request, # FastAPI Request object
llm_router, # Router instance
data: dict,
user_api_key_dict: "UserAPIKeyAuth",
check_file_id_encoding: bool = True,
) -> tuple[bool, str | None, str | None, dict | None]:
"""
Orchestrate model-based credential routing for file operations.
The model name comes from the caller (embedded in the file id, or a header, query
param or body field), so it is authorized against the caller's key, team, org and
project grants before any deployment credentials are resolved.
Args:
file_id: File ID (may contain embedded model info)
request: FastAPI request object
llm_router: LiteLLM Router instance
data: Request data dictionary
user_api_key_dict: The authenticated caller
check_file_id_encoding: Whether to check for embedded model in file_id
Returns:
@ -599,6 +651,7 @@ def handle_model_based_routing(
Raises:
HTTPException: If router unavailable or model not found
ProxyException: If the caller is not allowed to use the model
"""
model_from_id, model_from_param = extract_model_from_sources(
file_id=file_id,
@ -608,9 +661,10 @@ def handle_model_based_routing(
# Priority 1: Model embedded in file_id
if check_file_id_encoding and model_from_id is not None:
credentials = get_credentials_for_model(
credentials = await get_authorized_credentials_for_model(
llm_router=llm_router,
model_id=model_from_id,
user_api_key_dict=user_api_key_dict,
operation_context=f"file operation (file created with model '{model_from_id}')",
)
original_file_id: Final = get_original_file_id(file_id)
@ -618,9 +672,10 @@ def handle_model_based_routing(
# Priority 2: Model from header/query/body
elif model_from_param is not None:
credentials = get_credentials_for_model(
credentials = await get_authorized_credentials_for_model(
llm_router=llm_router,
model_id=model_from_param,
user_api_key_dict=user_api_key_dict,
operation_context="file operation",
)
return True, model_from_param, None, credentials

View file

@ -68,7 +68,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
apply_team_provider_credentials,
encode_file_id_with_model,
extract_file_creation_params,
get_credentials_for_model,
get_authorized_credentials_for_model,
handle_model_based_routing,
prepare_data_with_credentials,
validate_file_list_limit,
@ -267,9 +267,10 @@ async def route_create_file(
# NEW: Handle model-based routing (no DB required)
if model is not None:
# Get credentials from model_list via router
credentials: Final = get_credentials_for_model(
credentials: Final = await get_authorized_credentials_for_model(
llm_router=llm_router,
model_id=model,
user_api_key_dict=user_api_key_dict,
operation_context="file upload",
)
@ -907,11 +908,12 @@ async def get_file_content(
model_used,
original_file_id,
credentials,
) = handle_model_based_routing(
) = await handle_model_based_routing(
file_id=file_id,
request=request,
llm_router=llm_router,
data=data,
user_api_key_dict=user_api_key_dict,
check_file_id_encoding=True,
)
@ -1122,11 +1124,12 @@ async def get_file(
model_used,
original_file_id,
credentials,
) = handle_model_based_routing(
) = await handle_model_based_routing(
file_id=file_id,
request=request,
llm_router=llm_router,
data=data,
user_api_key_dict=user_api_key_dict,
check_file_id_encoding=True,
)
@ -1330,11 +1333,12 @@ async def delete_file(
model_used,
original_file_id,
credentials,
) = handle_model_based_routing(
) = await handle_model_based_routing(
file_id=file_id,
request=request,
llm_router=llm_router,
data=data,
user_api_key_dict=user_api_key_dict,
check_file_id_encoding=True,
)
@ -1517,11 +1521,12 @@ async def list_files(
response: Any | None = None
# Check for model-based credential routing (no file_id encoding check for list)
should_route, model_used, _, credentials = handle_model_based_routing(
should_route, model_used, _, credentials = await handle_model_based_routing(
file_id="", # No file_id for list endpoint
request=request,
llm_router=llm_router,
data=data,
user_api_key_dict=user_api_key_dict,
check_file_id_encoding=False,
)
@ -1548,9 +1553,10 @@ async def list_files(
status_code=500,
detail="LLM Router not initialized. Ensure models added to proxy.",
)
credentials = get_credentials_for_model(
credentials = await get_authorized_credentials_for_model(
llm_router=llm_router,
model_id=target_model_names_list[0],
user_api_key_dict=user_api_key_dict,
operation_context="file list",
)
prepare_data_with_credentials(data=data, credentials=credentials)

View file

@ -144,11 +144,12 @@ async def _update_request_data_with_managed_file_id(
model_used,
original_file_id,
credentials,
) = handle_model_based_routing(
) = await handle_model_based_routing(
file_id=file_id,
request=request,
llm_router=llm_router,
data=data,
user_api_key_dict=user_api_key_dict,
check_file_id_encoding=True,
)
@ -273,11 +274,12 @@ async def _update_request_data_with_model_routing_hint(
_model_used,
_original_file_id,
credentials,
) = handle_model_based_routing(
) = await handle_model_based_routing(
file_id="",
request=request,
llm_router=llm_router,
data=data,
user_api_key_dict=user_api_key_dict,
check_file_id_encoding=False,
)

View file

@ -177,6 +177,8 @@ def harness():
logging.get_proxy_hook = MagicMock(return_value=None)
router = MagicMock(spec=Router)
router.model_group_alias = {}
router.get_model_access_groups = MagicMock(return_value={})
router.acreate_batch = AsyncMock(return_value=make_batch())
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
@ -1161,6 +1163,8 @@ def retrieve_harness():
logging.get_proxy_hook = MagicMock(return_value=None)
router = MagicMock(spec=Router)
router.model_group_alias = {}
router.get_model_access_groups = MagicMock(return_value={})
router.aretrieve_batch = AsyncMock(return_value=make_batch())
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
@ -1616,6 +1620,8 @@ def list_harness():
logging.get_proxy_hook = MagicMock(return_value=None)
router = MagicMock(spec=Router)
router.model_group_alias = {}
router.get_model_access_groups = MagicMock(return_value={})
router.alist_batches = AsyncMock(return_value=FakeListPage([]))
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
@ -2012,6 +2018,8 @@ def cancel_harness():
logging.get_proxy_hook = MagicMock(return_value=None)
router = MagicMock(spec=Router)
router.model_group_alias = {}
router.get_model_access_groups = MagicMock(return_value={})
router.acancel_batch = AsyncMock(return_value=make_batch())
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
@ -2733,8 +2741,6 @@ async def test_cancel__unified_batch_id_allowed_when_managed_files_required(canc
assert cancel_harness.router_acancel.call_count == 1
@pytest.mark.asyncio
async def test_retrieve__managed_batch_defers_cost_to_the_poller_when_it_is_running(retrieve_harness):
with patch.object(endpoints, "batch_cost_poller_is_active", MagicMock(return_value=True)):
@ -2762,3 +2768,51 @@ async def test_retrieve__raw_batch_id_is_untouched_by_the_poller_handoff(retriev
metadata = retrieve_harness.litellm_aretrieve.await_args.kwargs.get("litellm_metadata") or {}
assert metadata.get("batch_ignore_default_logging") is None
def _key_restricted_to(*models: str) -> UserAPIKeyAuth:
return UserAPIKeyAuth(api_key="sk-restricted", team_id="team-a", team_models=list(models), models=list(models))
@pytest.mark.asyncio
async def test_create__header_model_rejects_key_without_model_grant(harness):
"""A key not granted the model named in x-litellm-model must not receive that deployment's credentials."""
set_body(harness, {"input_file_id": "file-plain", "endpoint": "/v1/chat/completions", "completion_window": "24h"})
with pytest.raises(ProxyException) as exc_info:
await call_create(harness, user=_key_restricted_to("azure/gpt-4o"), headers={"x-litellm-model": "vertex-model"})
assert exc_info.value.code == "403"
harness.creds_resolver.assert_not_called()
harness.litellm_acreate.assert_not_called()
@pytest.mark.asyncio
async def test_create__header_model_allows_key_with_model_grant(harness):
set_body(harness, {"input_file_id": "file-plain", "endpoint": "/v1/chat/completions", "completion_window": "24h"})
await call_create(harness, user=_key_restricted_to("vertex-model"), headers={"x-litellm-model": "vertex-model"})
harness.creds_resolver.assert_called_once_with(model_id="vertex-model")
assert harness.acreate_kwargs()["custom_llm_provider"] == "vertex_ai"
@pytest.mark.asyncio
async def test_retrieve__model_encoded_id_rejects_key_without_model_grant(retrieve_harness):
"""The model embedded in a batch id is caller-controlled, so it is checked against the key's grants too."""
with pytest.raises(ProxyException) as exc_info:
await call_retrieve(retrieve_harness, AZURE_BATCH_ID, user=_key_restricted_to("vertex-model"))
assert exc_info.value.code == "403"
retrieve_harness.creds_resolver.assert_not_called()
retrieve_harness.litellm_aretrieve.assert_not_called()
@pytest.mark.asyncio
async def test_cancel__model_encoded_id_rejects_key_without_model_grant(cancel_harness):
with pytest.raises(ProxyException) as exc_info:
await call_cancel(cancel_harness, AZURE_BATCH_ID, user=_key_restricted_to("vertex-model"))
assert exc_info.value.code == "403"
cancel_harness.creds_resolver.assert_not_called()
cancel_harness.litellm_acancel.assert_not_called()

View file

@ -1877,7 +1877,7 @@ def test_get_file_content_streams_openai_direct_path(
monkeypatch.setattr(litellm, "afile_content", _mock_afile_content)
monkeypatch.setattr(
"litellm.proxy.openai_files_endpoints.files_endpoints.handle_model_based_routing",
lambda **kwargs: (False, None, None, None),
AsyncMock(return_value=(False, None, None, None)),
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
@ -1942,15 +1942,17 @@ def test_get_file_content_routed_provider_skips_streaming_when_resolved_provider
)
monkeypatch.setattr(
"litellm.proxy.openai_files_endpoints.files_endpoints.handle_model_based_routing",
lambda **kwargs: (
True,
"azure-gpt-3-5-turbo",
"file-original-123",
{
"custom_llm_provider": "azure",
"api_key": "azure-key",
"api_base": "https://azure.example.com",
},
AsyncMock(
return_value=(
True,
"azure-gpt-3-5-turbo",
"file-original-123",
{
"custom_llm_provider": "azure",
"api_key": "azure-key",
"api_base": "https://azure.example.com",
},
)
),
)
@ -2015,7 +2017,7 @@ def test_get_file_content_non_openai_provider_skips_streaming_handler(
)
monkeypatch.setattr(
"litellm.proxy.openai_files_endpoints.files_endpoints.handle_model_based_routing",
lambda **kwargs: (False, None, None, None),
AsyncMock(return_value=(False, None, None, None)),
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
@ -2463,14 +2465,16 @@ def test_list_files_model_routing_does_not_forward_custom_llm_provider_twice(
monkeypatch.setattr(litellm, "afile_list", _mock_afile_list)
monkeypatch.setattr(
"litellm.proxy.openai_files_endpoints.files_endpoints.handle_model_based_routing",
lambda **kwargs: (
True,
"azure-gpt-4o",
None,
{
"custom_llm_provider": "azure",
"api_key": "azure-key",
},
AsyncMock(
return_value=(
True,
"azure-gpt-4o",
None,
{
"custom_llm_provider": "azure",
"api_key": "azure-key",
},
)
),
)
@ -4878,3 +4882,145 @@ def test_get_file_model_routed_id_forwards_deployment_provider(mocker: MockerFix
assert captured_kwargs["api_key"] == "mistral-key"
assert captured_kwargs["file_id"] == "7a13fa8e-fcf8-42c5-aa61-c93c10e2c7df"
assert response.json()["id"] == encoded_id
def _mistral_plus_anthropic_router() -> Router:
return Router(
model_list=[
{
"model_name": "mistral-ocr",
"litellm_params": {"model": "mistral/mistral-ocr-latest", "api_key": "mistral-key"},
"model_info": {"id": "mistral-ocr-id"},
},
{
"model_name": "claude-opus-4-6",
"litellm_params": {"model": "anthropic/claude-opus-4-6", "api_key": "anthropic-key"},
"model_info": {"id": "claude-id"},
},
]
)
def _restricted_key(key_models: list[str]) -> UserAPIKeyAuth:
from litellm.proxy._types import LitellmUserRoles
return UserAPIKeyAuth(
api_key="test-key",
user_role=LitellmUserRoles.INTERNAL_USER,
user_id="test-user",
team_id="team-a",
team_models=["claude-opus-4-6", "mistral-ocr"],
models=key_models,
)
@pytest.mark.parametrize(
"http_method, path_suffix, litellm_fn",
[
("get", "", "afile_retrieve"),
("get", "/content", "afile_content"),
("delete", "", "afile_delete"),
],
)
def test_model_routed_file_ops_reject_key_without_model_grant(
mocker: MockerFixture, monkeypatch, http_method: str, path_suffix: str, litellm_fn: str
):
"""
Regression: a key whose allowlist does not include the deployment named in a
model-encoded file id must be refused before that deployment's server-side
credentials are resolved. Previously any key could name any deployment via the
id (or the x-litellm-model header) and act on that provider account's files.
"""
import litellm.proxy.proxy_server as ps
from litellm.proxy.openai_files_endpoints.common_utils import encode_file_id_with_model
router = _mistral_plus_anthropic_router()
proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router)
proxy_logging_obj.update_request_status = mocker.AsyncMock()
proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock()
upstream = mocker.AsyncMock(side_effect=AssertionError("provider must not be called"))
monkeypatch.setattr(litellm, litellm_fn, upstream)
app.dependency_overrides[ps.user_api_key_auth] = lambda: _restricted_key(["claude-opus-4-6"])
encoded_id = encode_file_id_with_model("7a13fa8e-fcf8-42c5-aa61-c93c10e2c7df", "mistral-ocr")
try:
response = getattr(client, http_method)(
f"/v1/files/{encoded_id}{path_suffix}", headers={"Authorization": "Bearer test-key"}
)
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert response.status_code == 403, response.text
assert "not allowed to access model" in response.text
upstream.assert_not_called()
def test_list_files_header_model_rejects_key_without_model_grant(mocker: MockerFixture, monkeypatch):
import litellm.proxy.proxy_server as ps
router = _mistral_plus_anthropic_router()
proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router)
proxy_logging_obj.update_request_status = mocker.AsyncMock()
proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock()
upstream = mocker.AsyncMock(side_effect=AssertionError("provider must not be called"))
monkeypatch.setattr(litellm, "afile_list", upstream)
app.dependency_overrides[ps.user_api_key_auth] = lambda: _restricted_key(["claude-opus-4-6"])
try:
response = client.get(
"/v1/files", headers={"Authorization": "Bearer test-key", "x-litellm-model": "mistral-ocr"}
)
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert response.status_code == 403, response.text
upstream.assert_not_called()
def test_model_routed_file_retrieve_allows_key_with_model_grant(mocker: MockerFixture, monkeypatch):
"""The grant check must not break the happy path: a key allowed the deployment still resolves its credentials."""
import litellm.proxy.proxy_server as ps
from litellm.proxy.openai_files_endpoints.common_utils import encode_file_id_with_model
router = _mistral_plus_anthropic_router()
proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router)
proxy_logging_obj.update_request_status = mocker.AsyncMock()
proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock()
captured_kwargs: dict = {}
async def _mock_afile_retrieve(**kwargs):
captured_kwargs.update(kwargs)
return OpenAIFileObject(
id="7a13fa8e-fcf8-42c5-aa61-c93c10e2c7df",
object="file",
bytes=2,
created_at=1234567890,
filename="batch.jsonl",
purpose="batch",
status="uploaded",
)
monkeypatch.setattr(litellm, "afile_retrieve", _mock_afile_retrieve)
app.dependency_overrides[ps.user_api_key_auth] = lambda: _restricted_key(["mistral-ocr"])
encoded_id = encode_file_id_with_model("7a13fa8e-fcf8-42c5-aa61-c93c10e2c7df", "mistral-ocr")
try:
response = client.get(f"/v1/files/{encoded_id}", headers={"Authorization": "Bearer test-key"})
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert response.status_code == 200, response.text
assert captured_kwargs["api_key"] == "mistral-key"
assert captured_kwargs["custom_llm_provider"] == "mistral"

View file

@ -58,10 +58,7 @@ def _make_batch_response(
def test_get_batch_id_from_unified_batch_id_handles_appended_fields():
decoded_id = (
"litellm_proxy;model_id:deployment-123;"
"llm_batch_id:batch_openai_123;llm_output_file_id:file-output"
)
decoded_id = "litellm_proxy;model_id:deployment-123;llm_batch_id:batch_openai_123;llm_output_file_id:file-output"
assert get_batch_id_from_unified_batch_id(decoded_id) == "batch_openai_123"
@ -107,12 +104,10 @@ async def test_create_batch_with_x_litellm_model_encodes_batch_id():
}
),
),
patch("litellm.proxy.batches_endpoints.endpoints.ProxyBaseLLMRequestProcessing") as mock_processor_cls,
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,
"litellm.proxy.batches_endpoints.endpoints.get_authorized_credentials_for_model",
new=AsyncMock(return_value=mock_credentials),
),
patch(
"litellm.proxy.batches_endpoints.endpoints.prepare_data_with_credentials",
@ -165,23 +160,15 @@ async def test_create_batch_with_x_litellm_model_encodes_batch_id():
)
# The batch_id should be encoded with model info
assert (
response.id != raw_batch_id
), f"Expected batch_id to be encoded, but got raw ID: {response.id}"
assert response.id.startswith(
"batch_"
), f"Encoded batch_id should keep batch_ prefix, got: {response.id}"
assert response.id != raw_batch_id, f"Expected batch_id to be encoded, but got raw ID: {response.id}"
assert response.id.startswith("batch_"), f"Encoded batch_id should keep batch_ prefix, got: {response.id}"
# Should be decodable back to the original
decoded_model = decode_model_from_file_id(response.id)
assert (
decoded_model == model_name
), f"Expected model '{model_name}' from decoded batch_id, got: {decoded_model}"
assert decoded_model == model_name, f"Expected model '{model_name}' from decoded batch_id, got: {decoded_model}"
original_id = get_original_file_id(response.id)
assert (
original_id == raw_batch_id
), f"Expected original ID '{raw_batch_id}', got: {original_id}"
assert original_id == raw_batch_id, f"Expected original ID '{raw_batch_id}', got: {original_id}"
assert mock_create_batch.call_args.kwargs["metadata"] == {"customer_id": "cust-123"}
@ -227,12 +214,10 @@ async def test_create_batch_with_x_litellm_model_encodes_output_and_error_file_i
}
),
),
patch("litellm.proxy.batches_endpoints.endpoints.ProxyBaseLLMRequestProcessing") as mock_processor_cls,
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,
"litellm.proxy.batches_endpoints.endpoints.get_authorized_credentials_for_model",
new=AsyncMock(return_value=mock_credentials),
),
patch(
"litellm.proxy.batches_endpoints.endpoints.prepare_data_with_credentials",
@ -316,9 +301,7 @@ async def test_create_batch_without_x_litellm_model_returns_raw_ids(monkeypatch)
}
),
),
patch(
"litellm.proxy.batches_endpoints.endpoints.ProxyBaseLLMRequestProcessing"
) as mock_processor_cls,
patch("litellm.proxy.batches_endpoints.endpoints.ProxyBaseLLMRequestProcessing") as mock_processor_cls,
patch(
"litellm.acreate_batch",
new=AsyncMock(return_value=mock_response),
@ -383,9 +366,7 @@ class TestBatchIdRoundTripWithRetrieve:
raw_batch_id = "batch_vllm_12345"
# What create_batch does:
encoded_id = encode_file_id_with_model(
file_id=raw_batch_id, model=model_name, id_type="batch"
)
encoded_id = encode_file_id_with_model(file_id=raw_batch_id, model=model_name, id_type="batch")
# What retrieve_batch does:
decoded_model = decode_model_from_file_id(encoded_id)
@ -410,9 +391,7 @@ class TestBatchIdRoundTripWithRetrieve:
]
for raw_id, model in test_cases:
encoded = encode_file_id_with_model(
file_id=raw_id, model=model, id_type="batch"
)
encoded = encode_file_id_with_model(file_id=raw_id, model=model, id_type="batch")
assert encoded.startswith("batch_")
assert decode_model_from_file_id(encoded) == model
assert get_original_file_id(encoded) == raw_id
@ -440,9 +419,7 @@ async def test_cancel_batch_with_unified_id_routes_with_decoded_model_and_batch_
mock_user_api_key_dict.team_metadata = {}
with (
patch(
"litellm.proxy.batches_endpoints.endpoints.ProxyBaseLLMRequestProcessing"
) as mock_processor_cls,
patch("litellm.proxy.batches_endpoints.endpoints.ProxyBaseLLMRequestProcessing") as mock_processor_cls,
patch(
"litellm.proxy.batches_endpoints.endpoints.update_batch_in_database",
new=AsyncMock(),