diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index f1b4c6b5b17..7a132c01399 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -41,6 +41,7 @@ from litellm.proxy._types import ( CallTypes, LiteLLM_ManagedFileTable, LiteLLM_ManagedObjectTable, + ProxyException, UserAPIKeyAuth, ) from litellm.proxy.openai_files_endpoints.common_utils import ( @@ -423,13 +424,23 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): # This is because the encoded object ids stored in the managed objects table do not contain the provider information # To support provider filtering, we would need to store the provider information in the encoded object ids if provider: - raise Exception("Filtering by 'provider' is not supported when using managed batches.") + raise ProxyException( + message="Filtering by 'provider' is not supported when using managed batches.", + type="invalid_request_error", + param="provider", + code=400, + ) # Model name filtering is not supported for managed batches # This is because the encoded object ids stored in the managed objects table do not contain the model name # A hash of the model name + litellm_params for the model name is encoded as the model id. This is not sufficient to reliably map the target model names to the model ids. if target_model_names: - raise Exception("Filtering by 'target_model_names' is not supported when using managed batches.") + raise ProxyException( + message="Filtering by 'target_model_names' is not supported when using managed batches.", + type="invalid_request_error", + param="target_model_names", + code=400, + ) owner_filter = build_owner_filter(user_api_key_dict) if owner_filter is None: diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 623900322be..f8b4e65c270 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -5,6 +5,8 @@ ###################################################################### import asyncio +import os +from collections.abc import Mapping from typing import Any, Final, cast from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response @@ -48,6 +50,23 @@ from litellm.types.llms.openai import LiteLLMBatchCreateRequest router: Final = APIRouter() +def _raise_not_found_when_openai_fallback_unservable( + requested_provider: "str | None", + data: Mapping[str, object], + not_found_message: str, +) -> None: + if requested_provider is not None: + return + if data.get("api_key") or litellm.api_key or litellm.openai_key or os.getenv("OPENAI_API_KEY"): + return + raise ProxyException( + message=not_found_message, + type="invalid_request_error", + param=None, + code=404, + ) + + async def _resolve_managed_input_file_storage_url(input_file_id: str) -> "str | None": """Resolve a managed (unified) input_file_id to its backend storage_url. @@ -150,12 +169,12 @@ async def create_batch( router_model = data.get("model", None) is_router_model = is_known_model(model=router_model, llm_router=llm_router) - custom_llm_provider: Final = ( + requested_provider: Final = ( provider or data.pop("custom_llm_provider", None) or get_custom_llm_provider_from_request_headers(request=request) - or "openai" ) + custom_llm_provider: Final = requested_provider or "openai" _create_batch_data: Final = LiteLLMBatchCreateRequest(**data) # Apply team-level batch output expiry enforcement @@ -317,6 +336,11 @@ async def create_batch( user_api_key_dict=user_api_key_dict, custom_llm_provider=custom_llm_provider, ) + _raise_not_found_when_openai_fallback_unservable( + requested_provider=requested_provider, + data=cast(dict, _create_batch_data), # cast-ok: TypedDict is a dict at runtime + not_found_message=f"No such File object: {input_file_id}", + ) response = await litellm.acreate_batch( custom_llm_provider=custom_llm_provider, **_create_batch_data, @@ -566,18 +590,23 @@ async def retrieve_batch( # SCENARIO 3: Fallback to custom_llm_provider (uses env variables) else: - custom_llm_provider: Final = ( + requested_provider: Final = ( provider or get_custom_llm_provider_from_request_headers(request=request) or get_custom_llm_provider_from_request_query(request=request) - or "openai" ) + custom_llm_provider: Final = requested_provider or "openai" apply_team_provider_credentials( data=data, llm_router=llm_router, user_api_key_dict=user_api_key_dict, custom_llm_provider=custom_llm_provider, ) + _raise_not_found_when_openai_fallback_unservable( + requested_provider=requested_provider, + data=data, + not_found_message=f"No batch found with id '{batch_id}'.", + ) response = await litellm.aretrieve_batch( custom_llm_provider=custom_llm_provider, **data, @@ -970,13 +999,13 @@ async def cancel_batch( # SCENARIO 3: Fallback to custom_llm_provider (uses env variables) else: body_custom_llm_provider = data.pop("custom_llm_provider", None) - custom_llm_provider: Final = ( + requested_provider: Final = ( provider or body_custom_llm_provider or get_custom_llm_provider_from_request_headers(request=request) or get_custom_llm_provider_from_request_query(request=request) - or "openai" ) + custom_llm_provider: Final = requested_provider or "openai" # Extract batch_id from data to avoid "multiple values for keyword argument" error # data was cast from CancelBatchRequest which already contains batch_id data.pop("batch_id", None) @@ -986,6 +1015,11 @@ async def cancel_batch( user_api_key_dict=user_api_key_dict, custom_llm_provider=custom_llm_provider, ) + _raise_not_found_when_openai_fallback_unservable( + requested_provider=requested_provider, + data=data, + not_found_message=f"No batch found with id '{batch_id}'.", + ) _cancel_batch_data: Final = CancelBatchRequest(batch_id=batch_id, **data) response = await litellm.acancel_batch( custom_llm_provider=custom_llm_provider, diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index fde1feb80e2..22d5afd5d42 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -3147,3 +3147,43 @@ async def test_file_list_cursors_follow_the_owner_scoped_page(): assert response.first_id == "litellm_proxy:mine" assert response.last_id == "litellm_proxy:mine" assert response.has_more is False + + +@pytest.mark.asyncio +async def test_list_user_batches_provider_filter_rejected_with_400(): + from litellm.proxy._types import ProxyException, UserAPIKeyAuth + + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=MagicMock() + ) + + with pytest.raises(ProxyException) as exc: + await proxy_managed_files.list_user_batches( + user_api_key_dict=UserAPIKeyAuth(user_id="123"), + provider="openai", + ) + + assert exc.value.code == "400" + assert exc.value.type == "invalid_request_error" + assert exc.value.param == "provider" + assert exc.value.message == "Filtering by 'provider' is not supported when using managed batches." + + +@pytest.mark.asyncio +async def test_list_user_batches_target_model_names_filter_rejected_with_400(): + from litellm.proxy._types import ProxyException, UserAPIKeyAuth + + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=MagicMock() + ) + + with pytest.raises(ProxyException) as exc: + await proxy_managed_files.list_user_batches( + user_api_key_dict=UserAPIKeyAuth(user_id="123"), + target_model_names="gpt-4o", + ) + + assert exc.value.code == "400" + assert exc.value.type == "invalid_request_error" + assert exc.value.param == "target_model_names" + assert exc.value.message == "Filtering by 'target_model_names' is not supported when using managed batches." diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index b720ea16968..e7e754c0d8f 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -117,6 +117,21 @@ class FakeRequest: self.query_params = query or {} +@pytest.fixture +def openai_env_creds(monkeypatch): + """Deterministic env creds so the implicit-openai fallback forwards instead + of tripping the no-creds 404 gate, regardless of the host environment.""" + monkeypatch.setenv("OPENAI_API_KEY", "sk-env-openai") + + +@pytest.fixture +def no_openai_creds(monkeypatch): + """Neutralize every credential source the 404 gate checks.""" + monkeypatch.delenv("OPENAI_API_KEY", raising=False) + monkeypatch.setattr(litellm, "api_key", None) + monkeypatch.setattr(litellm, "openai_key", None) + + @dataclass class Harness: """Holds every mocked seam so a test can configure inputs and assert calls.""" @@ -409,7 +424,7 @@ async def test_create__body_model_beats_header_and_query(harness): @pytest.mark.asyncio -async def test_create__fallback_default_openai(harness): +async def test_create__fallback_default_openai(harness, openai_env_creds): set_body( harness, { @@ -427,6 +442,63 @@ async def test_create__fallback_default_openai(harness): assert harness.acreate_kwargs()["custom_llm_provider"] == "openai" +@pytest.mark.asyncio +async def test_create__fallback_no_creds_404(harness, no_openai_creds): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + with pytest.raises(ProxyException) as exc: + await call_create(harness) + + assert exc.value.code == "404" + assert exc.value.type == "invalid_request_error" + assert exc.value.param is None + assert exc.value.message == "No such File object: file-plain" + harness.litellm_acreate.assert_not_called() + harness.router_acreate.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__fallback_explicit_provider_bypasses_not_found_gate(harness, no_openai_creds): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness, provider="anthropic") + + assert harness.acreate_kwargs()["custom_llm_provider"] == "anthropic" + + +@pytest.mark.asyncio +async def test_create__fallback_env_key_alone_forwards(harness, monkeypatch): + monkeypatch.setattr(litellm, "api_key", None) + monkeypatch.setattr(litellm, "openai_key", None) + monkeypatch.setenv("OPENAI_API_KEY", "sk-env-openai") + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness) + + assert harness.acreate_kwargs()["custom_llm_provider"] == "openai" + + @pytest.mark.asyncio async def test_create__fallback_provider_path_param(harness): set_body( @@ -795,7 +867,7 @@ def _user_with_expiry(expiry: Any) -> UserAPIKeyAuth: @pytest.mark.asyncio -async def test_create__team_expiry_injected(harness): +async def test_create__team_expiry_injected(harness, openai_env_creds): set_body( harness, { @@ -814,7 +886,7 @@ async def test_create__team_expiry_injected(harness): @pytest.mark.asyncio -async def test_create__no_team_expiry_not_injected(harness): +async def test_create__no_team_expiry_not_injected(harness, openai_env_creds): set_body( harness, { @@ -862,7 +934,7 @@ async def test_create__team_expiry_malformed_500(harness, expiry): @pytest.mark.asyncio -async def test_create__uses_acreate_batch_route_type(harness): +async def test_create__uses_acreate_batch_route_type(harness, openai_env_creds): set_body( harness, { @@ -878,7 +950,7 @@ async def test_create__uses_acreate_batch_route_type(harness): @pytest.mark.asyncio -async def test_create__metadata_sanitized_before_forwarding(harness): +async def test_create__metadata_sanitized_before_forwarding(harness, openai_env_creds): set_body( harness, { @@ -896,7 +968,7 @@ async def test_create__metadata_sanitized_before_forwarding(harness): @pytest.mark.asyncio -async def test_create__exception_calls_failure_hook(harness): +async def test_create__exception_calls_failure_hook(harness, openai_env_creds): set_body( harness, { @@ -1204,7 +1276,7 @@ async def test_retrieve__loadbalancing_raw_id_routes_to_router(retrieve_harness) @pytest.mark.asyncio -async def test_retrieve__fallback_default_openai(retrieve_harness): +async def test_retrieve__fallback_default_openai(retrieve_harness, openai_env_creds): await call_retrieve(retrieve_harness, "batch-raw-xyz") assert retrieve_harness.litellm_aretrieve.call_count == 1 @@ -1217,6 +1289,40 @@ async def test_retrieve__fallback_default_openai(retrieve_harness): assert retrieve_harness.update_batch_in_db.call_count == 1 +@pytest.mark.asyncio +async def test_retrieve__fallback_no_creds_404(retrieve_harness, no_openai_creds): + with pytest.raises(ProxyException) as exc: + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert exc.value.code == "404" + assert exc.value.type == "invalid_request_error" + assert exc.value.param is None + assert exc.value.message == "No batch found with id 'batch-raw-xyz'." + retrieve_harness.litellm_aretrieve.assert_not_called() + retrieve_harness.router_aretrieve.assert_not_called() + + +@pytest.mark.asyncio +async def test_retrieve__fallback_explicit_provider_bypasses_not_found_gate(retrieve_harness, no_openai_creds): + await call_retrieve(retrieve_harness, "batch-raw-xyz", provider="anthropic") + + assert retrieve_harness.aretrieve_kwargs()["custom_llm_provider"] == "anthropic" + + +@pytest.mark.asyncio +async def test_retrieve__fallback_env_key_alone_forwards(retrieve_harness, monkeypatch): + monkeypatch.setattr(litellm, "api_key", None) + monkeypatch.setattr(litellm, "openai_key", None) + monkeypatch.setenv("OPENAI_API_KEY", "sk-env-openai") + + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert retrieve_harness.aretrieve_kwargs() == { + "custom_llm_provider": "openai", + "batch_id": "batch-raw-xyz", + } + + @pytest.mark.asyncio async def test_retrieve__fallback_provider_path_param(retrieve_harness): await call_retrieve(retrieve_harness, "batch-raw-xyz", provider="anthropic") @@ -1302,7 +1408,7 @@ async def test_retrieve__db_terminal_unified_resolves_file_ids(retrieve_harness) @pytest.mark.asyncio -async def test_retrieve__db_non_terminal_state_syncs_with_provider(retrieve_harness): +async def test_retrieve__db_non_terminal_state_syncs_with_provider(retrieve_harness, openai_env_creds): """A non-terminal DB row must NOT short-circuit; the endpoint syncs with the provider to refresh state.""" db_response = make_batch(id="batch-from-db", status="validating") @@ -1321,14 +1427,14 @@ async def test_retrieve__db_non_terminal_state_syncs_with_provider(retrieve_harn @pytest.mark.asyncio -async def test_retrieve__uses_aretrieve_batch_route_type(retrieve_harness): +async def test_retrieve__uses_aretrieve_batch_route_type(retrieve_harness, openai_env_creds): await call_retrieve(retrieve_harness, "batch-raw-xyz") assert retrieve_harness.pre_call.call_args.kwargs["route_type"] == "aretrieve_batch" @pytest.mark.asyncio -async def test_retrieve__exception_calls_failure_hook(retrieve_harness): +async def test_retrieve__exception_calls_failure_hook(retrieve_harness, openai_env_creds): retrieve_harness.litellm_aretrieve.side_effect = ValueError("provider boom") with pytest.raises(Exception): @@ -1609,7 +1715,9 @@ async def test_list__target_model_names_takes_first_only(list_harness): @pytest.mark.asyncio -async def test_list__fallback_default_openai(list_harness): +async def test_list__fallback_default_openai(list_harness, no_openai_creds): + """list stays ungated by the no-creds 404 guard: it answers about a + collection, not a specific id, so there is nothing to 404 about.""" await call_list(list_harness) assert list_harness.litellm_alist.call_count == 1 @@ -1960,7 +2068,7 @@ async def test_cancel__unified_no_router_500(cancel_harness): @pytest.mark.asyncio -async def test_cancel__fallback_default_openai(cancel_harness): +async def test_cancel__fallback_default_openai(cancel_harness, openai_env_creds): await call_cancel(cancel_harness, "batch-raw-xyz") assert cancel_harness.litellm_acancel.call_count == 1 @@ -1974,6 +2082,40 @@ async def test_cancel__fallback_default_openai(cancel_harness): assert cancel_harness.update_batch_in_db.call_count == 1 +@pytest.mark.asyncio +async def test_cancel__fallback_no_creds_404(cancel_harness, no_openai_creds): + with pytest.raises(ProxyException) as exc: + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert exc.value.code == "404" + assert exc.value.type == "invalid_request_error" + assert exc.value.param is None + assert exc.value.message == "No batch found with id 'batch-raw-xyz'." + cancel_harness.litellm_acancel.assert_not_called() + cancel_harness.router_acancel.assert_not_called() + + +@pytest.mark.asyncio +async def test_cancel__fallback_explicit_provider_bypasses_not_found_gate(cancel_harness, no_openai_creds): + await call_cancel(cancel_harness, "batch-raw-xyz", provider="anthropic") + + assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "anthropic" + + +@pytest.mark.asyncio +async def test_cancel__fallback_env_key_alone_forwards(cancel_harness, monkeypatch): + monkeypatch.setattr(litellm, "api_key", None) + monkeypatch.setattr(litellm, "openai_key", None) + monkeypatch.setenv("OPENAI_API_KEY", "sk-env-openai") + + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert cancel_harness.acancel_kwargs() == { + "custom_llm_provider": "openai", + "batch_id": "batch-raw-xyz", + } + + @pytest.mark.asyncio async def test_cancel__fallback_provider_path_param(cancel_harness): await call_cancel(cancel_harness, "batch-raw-xyz", provider="anthropic") @@ -2028,14 +2170,14 @@ async def test_cancel__fallback_provider_precedence_path_over_body(cancel_harnes @pytest.mark.asyncio -async def test_cancel__uses_acancel_batch_route_type(cancel_harness): +async def test_cancel__uses_acancel_batch_route_type(cancel_harness, openai_env_creds): await call_cancel(cancel_harness, "batch-raw-xyz") assert cancel_harness.pre_call.call_args.kwargs["route_type"] == "acancel_batch" @pytest.mark.asyncio -async def test_cancel__exception_calls_failure_hook(cancel_harness): +async def test_cancel__exception_calls_failure_hook(cancel_harness, openai_env_creds): cancel_harness.litellm_acancel.side_effect = ValueError("provider boom") with pytest.raises(Exception): @@ -2357,7 +2499,7 @@ async def test_create__model_encoded_input_file_id_rejected_when_managed_files_r @pytest.mark.asyncio -async def test_create__raw_input_file_id_allowed_when_managed_files_not_required(harness): +async def test_create__raw_input_file_id_allowed_when_managed_files_not_required(harness, openai_env_creds): set_body( harness, { @@ -2459,7 +2601,7 @@ async def test_retrieve__managed_batch_still_accounts_inline_without_a_poller(re @pytest.mark.asyncio -async def test_retrieve__raw_batch_id_is_untouched_by_the_poller_handoff(retrieve_harness): +async def test_retrieve__raw_batch_id_is_untouched_by_the_poller_handoff(retrieve_harness, openai_env_creds): with patch.object(endpoints, "batch_cost_poller_is_active", MagicMock(return_value=True)): await call_retrieve(retrieve_harness, "batch-raw-xyz") diff --git a/tests/test_litellm/proxy/test_batch_x_litellm_model_encoding.py b/tests/test_litellm/proxy/test_batch_x_litellm_model_encoding.py index 101dc48603a..3d1831bb4cd 100644 --- a/tests/test_litellm/proxy/test_batch_x_litellm_model_encoding.py +++ b/tests/test_litellm/proxy/test_batch_x_litellm_model_encoding.py @@ -287,13 +287,15 @@ async def test_create_batch_with_x_litellm_model_encodes_output_and_error_file_i @pytest.mark.asyncio -async def test_create_batch_without_x_litellm_model_returns_raw_ids(): +async def test_create_batch_without_x_litellm_model_returns_raw_ids(monkeypatch): """ Without x-litellm-model header, create_batch should NOT encode batch IDs (falls through to Scenario 3 / custom_llm_provider fallback). """ from litellm.proxy.batches_endpoints.endpoints import create_batch + monkeypatch.setenv("OPENAI_API_KEY", "sk-env-openai") + raw_batch_id = "batch_abc123" mock_response = _make_batch_response(batch_id=raw_batch_id) mock_request = _make_mock_request(headers={})