Merge pull request #8934 from BerriAI/litellm_fix_team_model_access_checks

JWT Auth Fix - [Bug]: JWT access with Groups not working when team is assigned All Proxy Models access
This commit is contained in:
Ishaan Jaff 2025-03-10 21:01:41 -07:00 • committed by GitHub
commit 2903caaf54
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 88 additions and 78 deletions

View file

@ -69,11 +69,13 @@ class TritonConfig(BaseConfig):
def get_complete_url(
self,
api_base: str,
api_base: Optional[str],
model: str,
optional_params: dict,
stream: Optional[bool] = None,
) -> str:
if api_base is None:
raise ValueError("api_base is required")
llm_type = self._get_triton_llm_type(api_base)
if llm_type == "generate" and stream:
return api_base + "_stream"

View file

@ -2054,6 +2054,7 @@ class ProxyErrorTypes(str, enum.Enum):
budget_exceeded = "budget_exceeded"
key_model_access_denied = "key_model_access_denied"
team_model_access_denied = "team_model_access_denied"
user_model_access_denied = "user_model_access_denied"
expired_key = "expired_key"
auth_error = "auth_error"
internal_server_error = "internal_server_error"
@ -2062,6 +2063,20 @@ class ProxyErrorTypes(str, enum.Enum):
validation_error = "bad_request_error"
cache_ping_error = "cache_ping_error"
@classmethod
def get_model_access_error_type_for_object(
cls, object_type: Literal["key", "user", "team"]
) -> "ProxyErrorTypes":
"""
Get the model access error type for object_type
"""
if object_type == "key":
return cls.key_model_access_denied
elif object_type == "team":
return cls.team_model_access_denied
elif object_type == "user":
return cls.user_model_access_denied
DB_CONNECTION_ERROR_TYPES = (httpx.ConnectError, httpx.ReadError, httpx.ReadTimeout)

View file

@ -98,12 +98,19 @@ async def common_checks(
)
# 2. If team can call model
_team_model_access_check(
team_object=team_object,
model=_model,
llm_router=llm_router,
team_model_aliases=valid_token.team_model_aliases if valid_token else None,
)
if _model and team_object:
if not await can_team_access_model(
model=_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=valid_token.team_model_aliases if valid_token else None,
):
raise ProxyException(
message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}",
type=ProxyErrorTypes.team_model_access_denied,
param="model",
code=status.HTTP_401_UNAUTHORIZED,
)
## 2.1 If user can call model (if personal key)
if team_object is None and user_object is not None:
@ -971,10 +978,18 @@ async def _can_object_call_model(
llm_router: Optional[Router],
models: List[str],
team_model_aliases: Optional[Dict[str, str]] = None,
object_type: Literal["user", "team", "key"] = "user",
) -> Literal[True]:
"""
Checks if token can call a given model
Args:
- model: str
- llm_router: Optional[Router]
- models: List[str]
- team_model_aliases: Optional[Dict[str, str]]
- object_type: Literal["user", "team", "key"]. We use the object type to raise the correct exception type
Returns:
- True: if token allowed to call model
@ -1018,10 +1033,15 @@ async def _can_object_call_model(
if (len(filtered_models) == 0 and len(models) == 0) or "*" in filtered_models:
all_model_access = True
if SpecialModelNames.all_proxy_models.value in filtered_models:
all_model_access = True
if model is not None and model not in filtered_models and all_model_access is False:
raise ProxyException(
message=f"API Key not allowed to access model. This token can only access models={models}. Tried to access {model}",
type=ProxyErrorTypes.key_model_access_denied,
message=f"{object_type} not allowed to access model. This {object_type} can only access models={models}. Tried to access {model}",
type=ProxyErrorTypes.get_model_access_error_type_for_object(
object_type=object_type
),
param="model",
code=status.HTTP_401_UNAUTHORIZED,
)
@ -1072,6 +1092,26 @@ async def can_key_call_model(
llm_router=llm_router,
models=valid_token.models,
team_model_aliases=valid_token.team_model_aliases,
object_type="key",
)
async def can_team_access_model(
model: str,
team_object: Optional[LiteLLM_TeamTable],
llm_router: Optional[Router],
team_model_aliases: Optional[Dict[str, str]] = None,
) -> Literal[True]:
"""
Returns True if the team can access a specific model.
"""
return await _can_object_call_model(
model=model,
llm_router=llm_router,
models=team_object.models if team_object else [],
team_model_aliases=team_model_aliases,
object_type="team",
)
@ -1096,6 +1136,7 @@ async def can_user_call_model(
model=model,
llm_router=llm_router,
models=user_object.models,
object_type="user",
)
@ -1248,53 +1289,6 @@ async def _team_max_budget_check(
)
def _team_model_access_check(
model: Optional[str],
team_object: Optional[LiteLLM_TeamTable],
llm_router: Optional[Router],
team_model_aliases: Optional[Dict[str, str]] = None,
):
"""
Access check for team models
Raises:
Exception if the team is not allowed to call the`model`
"""
if (
model is not None
and team_object is not None
and team_object.models is not None
and len(team_object.models) > 0
and model not in team_object.models
):
# this means the team has access to all models on the proxy
if "all-proxy-models" in team_object.models or "*" in team_object.models:
# this means the team has access to all models on the proxy
pass
# check if the team model is an access_group
elif (
model_in_access_group(
model=model, team_models=team_object.models, llm_router=llm_router
)
is True
):
pass
elif model and "*" in model:
pass
elif _model_in_team_aliases(model=model, team_model_aliases=team_model_aliases):
pass
elif _model_matches_any_wildcard_pattern_in_list(
model=model, allowed_model_list=team_object.models
):
pass
else:
raise ProxyException(
message=f"Team not allowed to access model. Team={team_object.team_id}, Model={model}. Allowed team models = {team_object.models}",
type=ProxyErrorTypes.team_model_access_denied,
param="model",
code=status.HTTP_401_UNAUTHORIZED,
)
def is_model_allowed_by_pattern(model: str, allowed_model_pattern: str) -> bool:
"""
Check if a model matches an allowed pattern.

View file

@ -33,6 +33,7 @@ from litellm.proxy._types import (
ScopeMapping,
Span,
)
from litellm.proxy.auth.auth_checks import can_team_access_model
from litellm.proxy.utils import PrismaClient, ProxyLogging
from .auth_checks import (
@ -723,8 +724,12 @@ class JWTAuthManager:
team_models = team_object.models
if isinstance(team_models, list) and (
not requested_model
or requested_model in team_models
or "*" in team_models
or can_team_access_model(
model=requested_model,
team_object=team_object,
llm_router=None,
team_model_aliases=None,
)
):
is_allowed = allowed_routes_check(
user_role=LitellmUserRoles.TEAM,

View file

@ -9,7 +9,7 @@ from typing import Any, Optional, List, Literal
async def generate_key(
session, models: Optional[List[str]] = None, team_id: Optional[str] = None
):
"""Helper function to generate a key with specific model access"""
"""Helper function to generate a key with specific model access controls"""
url = "http://0.0.0.0:4000/key/generate"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {}
@ -94,7 +94,7 @@ async def test_model_access_patterns(key_models, test_model, expect_success):
assert _error_body["type"] == "key_model_access_denied"
assert _error_body["param"] == "model"
assert _error_body["code"] == "401"
assert "API Key not allowed to access model" in _error_body["message"]
assert "key not allowed to access model" in _error_body["message"]
@pytest.mark.asyncio
@ -159,12 +159,6 @@ async def test_model_access_update():
"team_models, test_model, expect_success",
[
(["openai/*"], "anthropic/claude-2", False), # Non-matching model
(["gpt-4"], "gpt-4", True), # Exact model match
(["bedrock/*"], "bedrock/anthropic.claude-3", True), # Bedrock wildcard
(["bedrock/anthropic.*"], "bedrock/anthropic.claude-3", True), # Pattern match
(["bedrock/anthropic.*"], "bedrock/amazon.titan", False), # Pattern non-match
(None, "gpt-4", True), # No model restrictions
([], "gpt-4", True), # Empty model list
],
)
@pytest.mark.asyncio
@ -285,6 +279,6 @@ def _validate_model_access_exception(
assert _error_body["param"] == "model"
assert _error_body["code"] == "401"
if expected_type == "key_model_access_denied":
assert "API Key not allowed to access model" in _error_body["message"]
assert "key not allowed to access model" in _error_body["message"]
elif expected_type == "team_model_access_denied":
assert "Team not allowed to access model" in _error_body["message"]
assert "eam not allowed to access model" in _error_body["message"]

View file

@ -27,7 +27,7 @@ from litellm.proxy._types import (
)
from litellm.proxy.utils import PrismaClient
from litellm.proxy.auth.auth_checks import (
_team_model_access_check,
can_team_access_model,
_virtual_key_soft_budget_check,
)
from litellm.proxy.utils import ProxyLogging
@ -427,9 +427,9 @@ async def test_virtual_key_max_budget_check(
],
)
@pytest.mark.asyncio
async def test_team_model_access_check(model, team_models, expect_to_work):
async def test_can_team_access_model(model, team_models, expect_to_work):
"""
Test cases for _team_model_access_check:
Test cases for can_team_access_model:
1. Exact model match
2. all-proxy-models access
3. Wildcard (*) access
@ -438,16 +438,16 @@ async def test_team_model_access_check(model, team_models, expect_to_work):
6. Empty model list
7. None model list
"""
team_object = LiteLLM_TeamTable(
team_id="test-team",
models=team_models,
)
try:
_team_model_access_check(
team_object = LiteLLM_TeamTable(
team_id="test-team",
models=team_models,
)
result = await can_team_access_model(
model=model,
team_object=team_object,
llm_router=None,
team_model_aliases=None,
)
if not expect_to_work:
pytest.fail(

View file

@ -308,7 +308,7 @@ async def test_chat_completion():
model="gpt-4",
messages=[{"role": "user", "content": "Hello!"}],
)
assert "API Key not allowed to access model." in str(e)
assert "key not allowed to access model." in str(e)
@pytest.mark.asyncio