mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(proxy): resolve deployment credentials for provider-only batch and file calls
This commit is contained in:
parent
7cd009caf7
commit
384c4e6964
5 changed files with 393 additions and 33 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue