diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index a2dc5e1caf5..35c9b8e13ea 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -35,6 +35,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( prepare_data_with_credentials, resolve_input_file_id_to_unified, resolve_output_file_ids_to_unified, + resolve_provider_scoped_credentials, update_batch_in_database, ) from litellm.proxy.utils import handle_exception_on_proxy, is_known_model @@ -295,6 +296,16 @@ async def create_batch( verbose_proxy_logger.debug(f"Created batch using model: {model_param}") else: # SCENARIO 3: Fallback to custom_llm_provider (uses env variables) + provider_credentials = resolve_provider_scoped_credentials( + llm_router=llm_router, + custom_llm_provider=custom_llm_provider, + user_api_key_dict=user_api_key_dict, + ) + if provider_credentials is not None: + prepare_data_with_credentials( + data=cast(dict, _create_batch_data), # cast-ok: TypedDict is a plain dict at runtime + credentials=dict(provider_credentials), + ) response = await litellm.acreate_batch( custom_llm_provider=custom_llm_provider, **_create_batch_data, # type: ignore @@ -525,6 +536,13 @@ async def retrieve_batch( or get_custom_llm_provider_from_request_query(request=request) or "openai" ) + provider_credentials = resolve_provider_scoped_credentials( + llm_router=llm_router, + custom_llm_provider=custom_llm_provider, + user_api_key_dict=user_api_key_dict, + ) + if provider_credentials is not None: + prepare_data_with_credentials(data=data, credentials=dict(provider_credentials)) response = await litellm.aretrieve_batch( custom_llm_provider=custom_llm_provider, **data, # type: ignore @@ -718,6 +736,13 @@ async def list_batches( or get_custom_llm_provider_from_request_query(request=request) or "openai" ) + provider_credentials = resolve_provider_scoped_credentials( + llm_router=llm_router, + custom_llm_provider=custom_llm_provider, + user_api_key_dict=user_api_key_dict, + ) + if provider_credentials is not None: + prepare_data_with_credentials(data=data, credentials=dict(provider_credentials)) response = await litellm.alist_batches( custom_llm_provider=custom_llm_provider, # type: ignore after=after, @@ -908,6 +933,13 @@ async def cancel_batch( # 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) + provider_credentials = resolve_provider_scoped_credentials( + llm_router=llm_router, + custom_llm_provider=custom_llm_provider, + user_api_key_dict=user_api_key_dict, + ) + if provider_credentials is not None: + prepare_data_with_credentials(data=data, credentials=dict(provider_credentials)) _cancel_batch_data = CancelBatchRequest(batch_id=batch_id, **data) response = await litellm.acancel_batch( custom_llm_provider=custom_llm_provider, # type: ignore diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index efd1d6b6cee..b0fba78c045 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -3,7 +3,7 @@ import mimetypes import re from dataclasses import dataclass, field from types import MappingProxyType -from typing import TYPE_CHECKING, List, Literal, Optional, Union +from typing import TYPE_CHECKING, List, Literal, Mapping, Optional, Union from litellm.repositories.table_repositories import ( ManagedFileRepository, @@ -14,6 +14,7 @@ from litellm.types.utils import SpecialEnums if TYPE_CHECKING: from fastapi import Request + from litellm.proxy._types import UserAPIKeyAuth from litellm.router import Router @@ -373,6 +374,26 @@ def get_team_provider_credentials( return None +def resolve_provider_scoped_credentials( + llm_router: Optional["Router"], + custom_llm_provider: str, + user_api_key_dict: "UserAPIKeyAuth", +) -> "Mapping[str, object] | None": + """ + Resolve the caller's deployment credentials for ``custom_llm_provider`` on + provider-scoped batch/file operations that don't pin a model, so those + calls don't fall through to environment defaults (e.g. + ``google.auth.default()`` for Vertex AI). Returns None when no authorized + deployment matches the provider. + """ + return get_team_provider_credentials( + llm_router=llm_router, + team_models=user_api_key_dict.team_models or [], + custom_llm_provider=custom_llm_provider, + team_id=user_api_key_dict.team_id, + ) + + def prepare_data_with_credentials( data: dict, credentials: dict, diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 0c9aa667751..3acc811292b 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -46,9 +46,9 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( encode_file_id_with_model, extract_file_creation_params, get_credentials_for_model, - get_team_provider_credentials, handle_model_based_routing, prepare_data_with_credentials, + resolve_provider_scoped_credentials, validate_managed_files_requirement, ) from litellm.proxy.utils import ProxyLogging, is_known_model @@ -253,11 +253,22 @@ async def route_create_file( _create_file_request=_create_file_request, ) else: - # get configs for custom_llm_provider - llm_provider_config = get_files_provider_config(custom_llm_provider=custom_llm_provider) - if llm_provider_config is not None: - # add llm_provider_config to data - _create_file_request.update(llm_provider_config) + provider_credentials = resolve_provider_scoped_credentials( + llm_router=llm_router, + custom_llm_provider=custom_llm_provider, + user_api_key_dict=user_api_key_dict, + ) + if provider_credentials is not None: + prepare_data_with_credentials( + data=cast(dict, _create_file_request), # cast-ok: TypedDict is a plain dict at runtime + credentials=dict(provider_credentials), + ) + else: + # get configs for custom_llm_provider + llm_provider_config = get_files_provider_config(custom_llm_provider=custom_llm_provider) + if llm_provider_config is not None: + # add llm_provider_config to data + _create_file_request.update(llm_provider_config) _create_file_request.pop("custom_llm_provider", None) # type: ignore # for now use custom_llm_provider=="openai" -> this will change as LiteLLM adds more providers for acreate_batch response = await litellm.acreate_file(**_create_file_request, custom_llm_provider=custom_llm_provider) # type: ignore @@ -735,6 +746,15 @@ async def get_file_content( check_file_id_encoding=True, ) + if not should_route: + provider_credentials = resolve_provider_scoped_credentials( + llm_router=llm_router, + custom_llm_provider=custom_llm_provider, + user_api_key_dict=user_api_key_dict, + ) + if provider_credentials is not None: + prepare_data_with_credentials(data=data, credentials=dict(provider_credentials)) + from litellm.proxy.openai_files_endpoints.file_content_streaming_handler import ( FileContentStreamingHandler, ) @@ -983,6 +1003,13 @@ async def get_file( # Remove file_id from data to avoid "multiple values for keyword argument" error # data was initialized with {"file_id": file_id} data.pop("file_id", None) + provider_credentials = resolve_provider_scoped_credentials( + llm_router=llm_router, + custom_llm_provider=custom_llm_provider, + user_api_key_dict=user_api_key_dict, + ) + if provider_credentials is not None: + prepare_data_with_credentials(data=data, credentials=dict(provider_credentials)) response = await litellm.afile_retrieve( custom_llm_provider=custom_llm_provider, file_id=file_id, @@ -1183,6 +1210,13 @@ async def delete_file( ) else: data.pop("file_id", None) + provider_credentials = resolve_provider_scoped_credentials( + llm_router=llm_router, + custom_llm_provider=custom_llm_provider, + user_api_key_dict=user_api_key_dict, + ) + if provider_credentials is not None: + prepare_data_with_credentials(data=data, credentials=dict(provider_credentials)) response = await litellm.afile_delete( custom_llm_provider=custom_llm_provider, file_id=file_id, @@ -1354,14 +1388,13 @@ async def list_files( # No model/target_model_names pinned: resolve upstream credentials from # the team's deployment for this provider so the call is authenticated # against the team's own account (e.g. the team's openai deployment). - team_credentials = get_team_provider_credentials( + provider_credentials = resolve_provider_scoped_credentials( llm_router=llm_router, - team_models=user_api_key_dict.team_models or [], custom_llm_provider=custom_llm_provider, - team_id=user_api_key_dict.team_id, + user_api_key_dict=user_api_key_dict, ) - if team_credentials is not None: - prepare_data_with_credentials(data=data, credentials=team_credentials) + if provider_credentials is not None: + prepare_data_with_credentials(data=data, credentials=dict(provider_credentials)) response = await litellm.afile_list( custom_llm_provider=custom_llm_provider, diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 9573fddd435..426385a958b 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -151,6 +151,12 @@ def _creds_lookup(*, model_id: str) -> Dict[str, str]: return dict(CREDS[model_id]) +def _configure_provider_scoped_lookup(router: MagicMock) -> None: + router.model_list = [] + router.get_model_names = MagicMock(return_value=list(CREDS.keys())) + router.get_model_access_groups = MagicMock(return_value={}) + + @pytest.fixture def harness(): """Seam harness. Patches only true I/O boundaries; pure encode/decode/merge @@ -165,6 +171,7 @@ def harness(): router = MagicMock(spec=Router) router.acreate_batch = AsyncMock(return_value=make_batch()) router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup) + _configure_provider_scoped_lookup(router) read_body = AsyncMock(side_effect=lambda request: body_holder["body"]) pre_call = AsyncMock(side_effect=lambda **kw: (body_holder["body"], MagicMock())) @@ -403,8 +410,9 @@ async def test_create__body_model_beats_header_and_query(harness): # =========================================================================== # -# SCENARIO 3 - fallback to custom_llm_provider (env-var creds). MUST NOT touch -# the credential resolver. +# SCENARIO 3 - provider-only fallback. Deployment credentials for the provider +# are resolved from the router and merged; a provider with no configured +# deployment forwards the payload untouched (env-var creds). # =========================================================================== # @@ -423,8 +431,13 @@ async def test_create__fallback_default_openai(harness): assert harness.litellm_acreate.call_count == 1 harness.router_acreate.assert_not_called() - harness.creds_resolver.assert_not_called() # inverse-bug guard - assert harness.acreate_kwargs()["custom_llm_provider"] == "openai" + assert harness.acreate_kwargs() == { + "custom_llm_provider": "openai", + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "metadata": None, + } @pytest.mark.asyncio @@ -440,8 +453,9 @@ async def test_create__fallback_provider_path_param(harness): await call_create(harness, provider="anthropic") - harness.creds_resolver.assert_not_called() - assert harness.acreate_kwargs()["custom_llm_provider"] == "anthropic" + payload = harness.acreate_kwargs() + assert payload["custom_llm_provider"] == "anthropic" + assert "api_key" not in payload @pytest.mark.asyncio @@ -462,6 +476,72 @@ async def test_create__fallback_body_custom_llm_provider(harness): assert payload["custom_llm_provider"] == "bedrock" +@pytest.mark.asyncio +async def test_create__provider_only_merges_matching_deployment_credentials(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + harness.provider_from_headers.return_value = "vertex_ai" + + await call_create(harness) + + assert harness.litellm_acreate.call_count == 1 + harness.router_acreate.assert_not_called() + assert harness.acreate_kwargs() == { + "custom_llm_provider": "vertex_ai", + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "metadata": None, + "api_key": "sk-vertex", + "api_base": "https://vertex.test", + "model": "vertex_ai/gemini-2.0", + } + + +@pytest.mark.asyncio +async def test_create__provider_only_prefers_team_own_deployment(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + harness.provider_from_headers.return_value = "vertex_ai" + harness.router.model_list = [ + { + "model_name": "team-gemini", + "litellm_params": {"model": "vertex_ai/gemini-2.0"}, + "model_info": {"id": "team-dep-id", "team_id": "team-a"}, + } + ] + team_creds = { + "custom_llm_provider": "vertex_ai", + "api_key": "sk-team-vertex", + "api_base": "https://team.vertex.test", + } + harness.creds_resolver.side_effect = lambda *, model_id: ( + dict(team_creds) if model_id == "team-dep-id" else dict(CREDS[model_id]) + ) + + await call_create( + harness, + user=UserAPIKeyAuth(api_key="sk-test", team_id="team-a", team_models=[]), + ) + + harness.creds_resolver.assert_called_once_with(model_id="team-dep-id") + payload = harness.acreate_kwargs() + assert payload["api_key"] == "sk-team-vertex" + assert payload["api_base"] == "https://team.vertex.test" + + # =========================================================================== # # Unified file id routing (-> llm_router). Helpers mocked only here because a # real unified id is opaque base64; the routing contract is what we lock. @@ -964,6 +1044,7 @@ def retrieve_harness(): router = MagicMock(spec=Router) router.aretrieve_batch = AsyncMock(return_value=make_batch()) router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup) + _configure_provider_scoped_lookup(router) pre_call = AsyncMock(side_effect=lambda **kw: (data_holder["data"], MagicMock())) get_headers = MagicMock(return_value={}) @@ -1174,8 +1255,9 @@ async def test_retrieve__loadbalancing_raw_id_routes_to_router(retrieve_harness) # --------------------------------------------------------------------------- # -# SCENARIO 3 - fallback to custom_llm_provider (env-var creds). MUST NOT touch -# the credential resolver or the router. +# SCENARIO 3 - provider-only fallback. Deployment credentials for the provider +# are resolved from the router and merged; a provider with no configured +# deployment forwards the payload untouched (env-var creds). # --------------------------------------------------------------------------- # @@ -1185,7 +1267,6 @@ async def test_retrieve__fallback_default_openai(retrieve_harness): assert retrieve_harness.litellm_aretrieve.call_count == 1 retrieve_harness.router_aretrieve.assert_not_called() - retrieve_harness.creds_resolver.assert_not_called() # inverse-bug guard assert retrieve_harness.aretrieve_kwargs() == { "custom_llm_provider": "openai", "batch_id": "batch-raw-xyz", @@ -1197,8 +1278,9 @@ async def test_retrieve__fallback_default_openai(retrieve_harness): async def test_retrieve__fallback_provider_path_param(retrieve_harness): await call_retrieve(retrieve_harness, "batch-raw-xyz", provider="anthropic") - retrieve_harness.creds_resolver.assert_not_called() - assert retrieve_harness.aretrieve_kwargs()["custom_llm_provider"] == "anthropic" + payload = retrieve_harness.aretrieve_kwargs() + assert payload["custom_llm_provider"] == "anthropic" + assert "api_key" not in payload @pytest.mark.asyncio @@ -1219,6 +1301,25 @@ async def test_retrieve__fallback_provider_from_query(retrieve_harness): assert retrieve_harness.aretrieve_kwargs()["custom_llm_provider"] == "vertex_ai" +@pytest.mark.asyncio +async def test_retrieve__provider_only_merges_matching_deployment_credentials( + retrieve_harness, +): + retrieve_harness.provider_from_headers.return_value = "vertex_ai" + + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert retrieve_harness.litellm_aretrieve.call_count == 1 + retrieve_harness.router_aretrieve.assert_not_called() + assert retrieve_harness.aretrieve_kwargs() == { + "custom_llm_provider": "vertex_ai", + "batch_id": "batch-raw-xyz", + "api_key": "sk-vertex", + "api_base": "https://vertex.test", + "model": "vertex_ai/gemini-2.0", + } + + @pytest.mark.asyncio async def test_retrieve__fallback_provider_precedence_path_over_header( retrieve_harness, @@ -1382,6 +1483,7 @@ def list_harness(): router = MagicMock(spec=Router) router.alist_batches = AsyncMock(return_value=FakeListPage([])) router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup) + _configure_provider_scoped_lookup(router) read_body = AsyncMock(side_effect=lambda request: body_holder["body"]) pre_call = AsyncMock(side_effect=lambda **kw: (body_holder["body"], MagicMock())) @@ -1587,8 +1689,9 @@ async def test_list__target_model_names_takes_first_only(list_harness): # --------------------------------------------------------------------------- # -# Branch 4 - fallback to custom_llm_provider (env-var creds). MUST NOT touch -# the credential resolver or the router. +# Branch 4 - provider-only fallback. Deployment credentials for the provider +# are resolved from the router and merged; a provider with no configured +# deployment forwards the payload untouched (env-var creds). # --------------------------------------------------------------------------- # @@ -1598,7 +1701,6 @@ async def test_list__fallback_default_openai(list_harness): assert list_harness.litellm_alist.call_count == 1 list_harness.router_alist.assert_not_called() - list_harness.creds_resolver.assert_not_called() # inverse-bug guard assert list_harness.alist_kwargs() == { "custom_llm_provider": "openai", "after": None, @@ -1610,8 +1712,9 @@ async def test_list__fallback_default_openai(list_harness): async def test_list__fallback_provider_path_param(list_harness): await call_list(list_harness, provider="anthropic") - list_harness.creds_resolver.assert_not_called() - assert list_harness.alist_kwargs()["custom_llm_provider"] == "anthropic" + payload = list_harness.alist_kwargs() + assert payload["custom_llm_provider"] == "anthropic" + assert "api_key" not in payload @pytest.mark.asyncio @@ -1632,6 +1735,24 @@ async def test_list__fallback_provider_from_query(list_harness): assert list_harness.alist_kwargs()["custom_llm_provider"] == "vertex_ai" +@pytest.mark.asyncio +async def test_list__provider_only_merges_matching_deployment_credentials(list_harness): + list_harness.provider_from_headers.return_value = "vertex_ai" + + await call_list(list_harness, limit=5, after="cur") + + assert list_harness.litellm_alist.call_count == 1 + list_harness.router_alist.assert_not_called() + assert list_harness.alist_kwargs() == { + "custom_llm_provider": "vertex_ai", + "after": "cur", + "limit": 5, + "api_key": "sk-vertex", + "api_base": "https://vertex.test", + "model": "vertex_ai/gemini-2.0", + } + + @pytest.mark.asyncio async def test_list__fallback_after_and_limit_forwarded(list_harness): await call_list(list_harness, after="cursor-9", limit=42) @@ -1735,6 +1856,7 @@ def cancel_harness(): router = MagicMock(spec=Router) router.acancel_batch = AsyncMock(return_value=make_batch()) router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup) + _configure_provider_scoped_lookup(router) pre_call = AsyncMock(side_effect=lambda **kw: (data_holder["data"], MagicMock())) # add_litellm_data_to_request is a passthrough that returns the data it got. @@ -1927,8 +2049,9 @@ async def test_cancel__unified_no_router_500(cancel_harness): # --------------------------------------------------------------------------- # -# SCENARIO 3 - fallback to custom_llm_provider. Rebuilds a CancelBatchRequest -# and forwards only {custom_llm_provider, batch_id}. +# SCENARIO 3 - provider-only fallback. Rebuilds a CancelBatchRequest, then +# merges deployment credentials resolved for the provider; a provider with no +# configured deployment forwards only {custom_llm_provider, batch_id}. # --------------------------------------------------------------------------- # @@ -1938,7 +2061,6 @@ async def test_cancel__fallback_default_openai(cancel_harness): assert cancel_harness.litellm_acancel.call_count == 1 cancel_harness.router_acancel.assert_not_called() - cancel_harness.creds_resolver.assert_not_called() # inverse-bug guard # current behavior: enrichment keys dropped; only these two forwarded. assert cancel_harness.acancel_kwargs() == { "custom_llm_provider": "openai", @@ -1951,8 +2073,28 @@ async def test_cancel__fallback_default_openai(cancel_harness): async def test_cancel__fallback_provider_path_param(cancel_harness): await call_cancel(cancel_harness, "batch-raw-xyz", provider="anthropic") - cancel_harness.creds_resolver.assert_not_called() - assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "anthropic" + payload = cancel_harness.acancel_kwargs() + assert payload["custom_llm_provider"] == "anthropic" + assert "api_key" not in payload + + +@pytest.mark.asyncio +async def test_cancel__provider_only_merges_matching_deployment_credentials( + cancel_harness, +): + cancel_harness.provider_from_headers.return_value = "vertex_ai" + + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert cancel_harness.litellm_acancel.call_count == 1 + cancel_harness.router_acancel.assert_not_called() + assert cancel_harness.acancel_kwargs() == { + "custom_llm_provider": "vertex_ai", + "batch_id": "batch-raw-xyz", + "api_key": "sk-vertex", + "api_base": "https://vertex.test", + "model": "vertex_ai/gemini-2.0", + } @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 46ecb31e1c8..e1d66e9a478 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -2610,3 +2610,135 @@ def test_list_files_with_all_proxy_models_team_uses_openai_deployment( assert captured_kwargs.get("api_key") == "team-openai-key" assert captured_kwargs.get("custom_llm_provider") == "openai" proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def _openai_file_object() -> OpenAIFileObject: + return OpenAIFileObject( + id="file-abc", + object="file", + bytes=2, + created_at=1, + filename="data.jsonl", + purpose="batch", + status="uploaded", + ) + + +async def _run_route_create_file( + mocker: MockerFixture, + router: Router, + custom_llm_provider: str, +) -> dict: + from litellm.proxy.openai_files_endpoints import files_endpoints as fe + from litellm.types.llms.openai import CreateFileRequest + + acreate_file = mocker.patch.object( + litellm, + "acreate_file", + new=mocker.AsyncMock(return_value=_openai_file_object()), + ) + + await fe.route_create_file( + llm_router=router, + _create_file_request=CreateFileRequest( + file=("data.jsonl", b"{}", "application/json"), + purpose="batch", + ), + purpose="batch", + proxy_logging_obj=mocker.MagicMock(), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + target_model_names_list=[], + is_router_model=False, + router_model=None, + custom_llm_provider=custom_llm_provider, + ) + + assert acreate_file.call_count == 1 + return dict(acreate_file.call_args.kwargs) + + +@pytest.mark.asyncio +async def test_route_create_file_provider_only_uses_deployment_credentials( + mocker: MockerFixture, +): + router = Router( + model_list=[ + { + "model_name": "gemini-batch", + "litellm_params": { + "model": "vertex_ai/gemini-2.0-flash", + "vertex_project": "configured-project", + "vertex_location": "us-central1", + "vertex_credentials": "/etc/keys/vertex-sa.json", + }, + "model_info": {"id": "gemini-batch-id"}, + } + ] + ) + + kwargs = await _run_route_create_file(mocker, router, "vertex_ai") + + assert kwargs["custom_llm_provider"] == "vertex_ai" + assert kwargs["vertex_project"] == "configured-project" + assert kwargs["vertex_location"] == "us-central1" + assert kwargs["vertex_credentials"] == "/etc/keys/vertex-sa.json" + + +@pytest.mark.asyncio +async def test_route_create_file_provider_only_deployment_beats_files_settings( + mocker: MockerFixture, monkeypatch +): + from litellm.proxy.openai_files_endpoints import files_endpoints as fe + + monkeypatch.setattr( + fe, + "files_config", + [{"custom_llm_provider": "openai", "api_key": "sk-files-settings"}], + ) + router = Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": { + "model": "openai/gpt-4o", + "api_key": "sk-deployment", + }, + "model_info": {"id": "gpt-4o-id"}, + } + ] + ) + + kwargs = await _run_route_create_file(mocker, router, "openai") + + assert kwargs["custom_llm_provider"] == "openai" + assert kwargs["api_key"] == "sk-deployment" + + +@pytest.mark.asyncio +async def test_route_create_file_provider_only_falls_back_to_files_settings( + mocker: MockerFixture, monkeypatch +): + from litellm.proxy.openai_files_endpoints import files_endpoints as fe + + monkeypatch.setattr( + fe, + "files_config", + [{"custom_llm_provider": "openai", "api_key": "sk-files-settings"}], + ) + router = Router( + model_list=[ + { + "model_name": "gemini-batch", + "litellm_params": { + "model": "vertex_ai/gemini-2.0-flash", + "vertex_project": "configured-project", + }, + "model_info": {"id": "gemini-batch-id"}, + } + ] + ) + + kwargs = await _run_route_create_file(mocker, router, "openai") + + assert kwargs["custom_llm_provider"] == "openai" + assert kwargs["api_key"] == "sk-files-settings"