mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(proxy): skip other teams' deployments in provider-scoped credential fallback
This commit is contained in:
parent
384c4e6964
commit
3639e3619f
6 changed files with 173 additions and 38 deletions
|
|
@ -108,7 +108,7 @@
|
|||
"limit": 40525
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 20384
|
||||
"limit": 20383
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 32099
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import mimetypes
|
|||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, List, Literal, Mapping, Optional, Union
|
||||
from typing import TYPE_CHECKING, Iterator, List, Literal, Mapping, Optional, Union
|
||||
|
||||
from litellm.repositories.table_repositories import (
|
||||
ManagedFileRepository,
|
||||
|
|
@ -293,6 +293,54 @@ def get_credentials_for_model(
|
|||
return credentials
|
||||
|
||||
|
||||
def _deployment_provider_credentials(
|
||||
llm_router: "Router",
|
||||
custom_llm_provider: str,
|
||||
model_id: str,
|
||||
) -> "Mapping[str, object] | None":
|
||||
credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
|
||||
if credentials is not None and credentials.get("custom_llm_provider") == custom_llm_provider:
|
||||
return credentials
|
||||
return None
|
||||
|
||||
|
||||
def _team_byok_provider_credentials(
|
||||
llm_router: "Router",
|
||||
custom_llm_provider: str,
|
||||
team_id: str,
|
||||
) -> "Mapping[str, object] | None":
|
||||
for deployment in llm_router.model_list or []:
|
||||
model_info = deployment.get("model_info") or {}
|
||||
if model_info.get("team_id") != team_id:
|
||||
continue
|
||||
deployment_id = model_info.get("id")
|
||||
if deployment_id is None:
|
||||
continue
|
||||
credentials = _deployment_provider_credentials(llm_router, custom_llm_provider, deployment_id)
|
||||
if credentials is not None:
|
||||
return credentials
|
||||
return None
|
||||
|
||||
|
||||
def _authorized_deployment_ids(
|
||||
llm_router: "Router",
|
||||
model_name: str,
|
||||
team_id: "str | None",
|
||||
) -> Iterator[str]:
|
||||
name_matched = tuple(
|
||||
deployment for deployment in llm_router.model_list or [] if deployment.get("model_name") == model_name
|
||||
)
|
||||
candidate_deployments = name_matched or tuple(llm_router.pattern_router.route(model_name) or [])
|
||||
for deployment in candidate_deployments:
|
||||
model_info = deployment.get("model_info") or {}
|
||||
owner_team_id = model_info.get("team_id")
|
||||
if owner_team_id is not None and owner_team_id != team_id:
|
||||
continue
|
||||
deployment_id = model_info.get("id")
|
||||
if deployment_id is not None:
|
||||
yield deployment_id
|
||||
|
||||
|
||||
def get_team_provider_credentials(
|
||||
llm_router: Optional["Router"],
|
||||
team_models: List[str],
|
||||
|
|
@ -308,7 +356,11 @@ def get_team_provider_credentials(
|
|||
``model_info.team_id`` matches ``team_id``. This keeps team-scoped listings
|
||||
on the team's own provider account/key instead of a shared global one.
|
||||
2. Fallback: any deployment the team is granted access to for this provider,
|
||||
expanding wildcard routes and the all-proxy-models sentinel.
|
||||
expanding wildcard routes and the all-proxy-models sentinel. Candidate
|
||||
names resolve to concrete deployments (by name, then wildcard pattern),
|
||||
deployments owned by a different team are skipped, and credentials are
|
||||
always fetched by deployment id; a foreign team's deployment can never
|
||||
be selected just because it shares a model name with an authorized one.
|
||||
|
||||
Credential lookup is always scoped to the team's allowlist, so a team can
|
||||
never resolve a provider key for a deployment it isn't authorized to use.
|
||||
|
|
@ -318,24 +370,11 @@ def get_team_provider_credentials(
|
|||
if llm_router is None:
|
||||
return None
|
||||
|
||||
def _provider_credentials(model_id: str) -> Optional[dict]:
|
||||
credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
|
||||
if credentials is not None and credentials.get("custom_llm_provider") == custom_llm_provider:
|
||||
return credentials
|
||||
return None
|
||||
|
||||
# 1. Prefer the team's own BYOK deployment, matched by model_info.team_id.
|
||||
if team_id is not None:
|
||||
for deployment in llm_router.model_list or []:
|
||||
model_info = deployment.get("model_info") or {}
|
||||
if model_info.get("team_id") != team_id:
|
||||
continue
|
||||
deployment_id = model_info.get("id")
|
||||
if deployment_id is None:
|
||||
continue
|
||||
credentials = _provider_credentials(deployment_id)
|
||||
if credentials is not None:
|
||||
return credentials
|
||||
credentials = _team_byok_provider_credentials(llm_router, custom_llm_provider, team_id)
|
||||
if credentials is not None:
|
||||
return dict(credentials)
|
||||
|
||||
# 2. Fall back to deployments the team is allowed to access. The
|
||||
# all-proxy-models sentinel isn't expanded by get_complete_model_list, so
|
||||
|
|
@ -367,9 +406,10 @@ def get_team_provider_credentials(
|
|||
)
|
||||
)
|
||||
for model_name in models_to_try:
|
||||
credentials = _provider_credentials(model_name)
|
||||
if credentials is not None:
|
||||
return credentials
|
||||
for deployment_id in _authorized_deployment_ids(llm_router, model_name, team_id):
|
||||
credentials = _deployment_provider_credentials(llm_router, custom_llm_provider, deployment_id)
|
||||
if credentials is not None:
|
||||
return dict(credentials)
|
||||
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -363,6 +363,6 @@
|
|||
"limit": 105
|
||||
},
|
||||
"UP045": {
|
||||
"limit": 17824
|
||||
"limit": 17823
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -60,19 +60,23 @@ from fastapi import Response
|
|||
# produces wrong creds (or KeyError) and is impossible to hide.
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
AZURE_CREDS: Dict[str, str] = {
|
||||
"custom_llm_provider": "azure",
|
||||
"api_key": "sk-azure",
|
||||
"api_base": "https://azure.test",
|
||||
"model": "azure/gpt-4o-deployment",
|
||||
}
|
||||
VERTEX_CREDS: Dict[str, str] = {
|
||||
"custom_llm_provider": "vertex_ai",
|
||||
"api_key": "sk-vertex",
|
||||
"api_base": "https://vertex.test",
|
||||
"model": "vertex_ai/gemini-2.0",
|
||||
}
|
||||
CREDS: Dict[str, Dict[str, str]] = {
|
||||
"azure/gpt-4o": {
|
||||
"custom_llm_provider": "azure",
|
||||
"api_key": "sk-azure",
|
||||
"api_base": "https://azure.test",
|
||||
"model": "azure/gpt-4o-deployment",
|
||||
},
|
||||
"vertex-model": {
|
||||
"custom_llm_provider": "vertex_ai",
|
||||
"api_key": "sk-vertex",
|
||||
"api_base": "https://vertex.test",
|
||||
"model": "vertex_ai/gemini-2.0",
|
||||
},
|
||||
"azure/gpt-4o": AZURE_CREDS,
|
||||
"azure-dep-id": AZURE_CREDS,
|
||||
"vertex-model": VERTEX_CREDS,
|
||||
"vertex-dep-id": VERTEX_CREDS,
|
||||
}
|
||||
|
||||
# A real model-encoded file id: decodes to "azure/gpt-4o", strips to "file-original123".
|
||||
|
|
@ -152,9 +156,22 @@ def _creds_lookup(*, model_id: str) -> Dict[str, str]:
|
|||
|
||||
|
||||
def _configure_provider_scoped_lookup(router: MagicMock) -> None:
|
||||
router.model_list = []
|
||||
router.get_model_names = MagicMock(return_value=list(CREDS.keys()))
|
||||
router.model_list = [
|
||||
{
|
||||
"model_name": "azure/gpt-4o",
|
||||
"litellm_params": {"model": "azure/gpt-4o-deployment"},
|
||||
"model_info": {"id": "azure-dep-id"},
|
||||
},
|
||||
{
|
||||
"model_name": "vertex-model",
|
||||
"litellm_params": {"model": "vertex_ai/gemini-2.0"},
|
||||
"model_info": {"id": "vertex-dep-id"},
|
||||
},
|
||||
]
|
||||
router.get_model_names = MagicMock(return_value=["azure/gpt-4o", "vertex-model"])
|
||||
router.get_model_access_groups = MagicMock(return_value={})
|
||||
router.pattern_router = MagicMock()
|
||||
router.pattern_router.route = MagicMock(return_value=None)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -502,6 +519,10 @@ async def test_create__provider_only_merges_matching_deployment_credentials(harn
|
|||
"api_base": "https://vertex.test",
|
||||
"model": "vertex_ai/gemini-2.0",
|
||||
}
|
||||
assert [call.kwargs["model_id"] for call in harness.creds_resolver.call_args_list] == [
|
||||
"azure-dep-id",
|
||||
"vertex-dep-id",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -542,6 +563,44 @@ async def test_create__provider_only_prefers_team_own_deployment(harness):
|
|||
assert payload["api_base"] == "https://team.vertex.test"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__provider_only_skips_other_teams_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": "vertex-model",
|
||||
"litellm_params": {"model": "vertex_ai/gemini-2.0"},
|
||||
"model_info": {"id": "foreign-dep-id", "team_id": "team-other"},
|
||||
},
|
||||
{
|
||||
"model_name": "vertex-model",
|
||||
"litellm_params": {"model": "vertex_ai/gemini-2.0"},
|
||||
"model_info": {"id": "vertex-dep-id"},
|
||||
},
|
||||
]
|
||||
foreign_creds = {"custom_llm_provider": "vertex_ai", "api_key": "sk-foreign-team"}
|
||||
harness.creds_resolver.side_effect = lambda *, model_id: (
|
||||
dict(foreign_creds) if model_id in ("foreign-dep-id", "vertex-model") else dict(CREDS[model_id])
|
||||
)
|
||||
|
||||
await call_create(
|
||||
harness,
|
||||
user=UserAPIKeyAuth(api_key="sk-test", team_id="team-b", team_models=[]),
|
||||
)
|
||||
|
||||
payload = harness.acreate_kwargs()
|
||||
assert payload["api_key"] == "sk-vertex"
|
||||
assert all(call.kwargs["model_id"] != "foreign-dep-id" for call in harness.creds_resolver.call_args_list)
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
# 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.
|
||||
|
|
|
|||
|
|
@ -2742,3 +2742,39 @@ async def test_route_create_file_provider_only_falls_back_to_files_settings(
|
|||
|
||||
assert kwargs["custom_llm_provider"] == "openai"
|
||||
assert kwargs["api_key"] == "sk-files-settings"
|
||||
|
||||
|
||||
def test_provider_scoped_credentials_never_use_other_teams_deployment():
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
resolve_provider_scoped_credentials,
|
||||
)
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gemini-batch",
|
||||
"litellm_params": {
|
||||
"model": "vertex_ai/gemini-2.0-flash",
|
||||
"vertex_project": "other-team-project",
|
||||
},
|
||||
"model_info": {"id": "other-team-dep", "team_id": "team-other"},
|
||||
},
|
||||
{
|
||||
"model_name": "gemini-batch",
|
||||
"litellm_params": {
|
||||
"model": "vertex_ai/gemini-2.0-flash",
|
||||
"vertex_project": "shared-project",
|
||||
},
|
||||
"model_info": {"id": "shared-dep"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
credentials = resolve_provider_scoped_credentials(
|
||||
llm_router=router,
|
||||
custom_llm_provider="vertex_ai",
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", team_id="team-caller", team_models=[]),
|
||||
)
|
||||
|
||||
assert credentials is not None
|
||||
assert credentials["vertex_project"] == "shared-project"
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 23287
|
||||
"limit": 23286
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 27473
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue