mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
commit
2903caaf54
7 changed files with 88 additions and 78 deletions
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue