From 51ef8a96f01f77e8e3b48e788ffc82c707f3d45e Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Mon, 23 Mar 2026 18:21:19 -0700 Subject: [PATCH] fix: resolve team id's on projects --- .pre-commit-config.yaml | 26 +-- litellm/proxy/auth/auth_checks.py | 29 ++- .../proxy/auth/test_auth_checks.py | 178 ++++++++++++++++-- 3 files changed, 202 insertions(+), 31 deletions(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 2bc361bc48f..8d601d6038d 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -7,19 +7,19 @@ repos: language: system types: [python] files: ^(litellm/|litellm_proxy_extras/|enterprise/) - - id: isort - name: isort - entry: isort - language: system - types: [python] - files: (litellm/|litellm_proxy_extras/|enterprise/).*\.py - exclude: ^litellm/__init__.py$ - - id: black - name: black - entry: poetry run black - language: system - types: [python] - files: (litellm/|litellm_proxy_extras/).*\.py + # - id: isort + # name: isort + # entry: isort + # language: system + # types: [python] + # files: (litellm/|litellm_proxy_extras/|enterprise/).*\.py + # exclude: ^litellm/__init__.py$ + # - id: black + # name: black + # entry: poetry run black + # language: system + # types: [python] + # files: (litellm/|litellm_proxy_extras/).*\.py - repo: https://github.com/pycqa/flake8 rev: 7.0.0 # The version of flake8 to use hooks: diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 815393467de..8facb49b6f1 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -188,6 +188,7 @@ async def _run_project_checks( skip_budget_checks: bool, valid_token: Optional[UserAPIKeyAuth], proxy_logging_obj: ProxyLogging, + team_object: Optional[LiteLLM_TeamTable] = None, ) -> None: """ Run all project-level checks: blocked, model access, budget, soft budget. @@ -204,11 +205,28 @@ async def _run_project_checks( # 2.2 If project can call model if _model and len(project_object.models) > 0: - can_project_access_model( - model=_model, - project_object=project_object, - llm_router=llm_router, - ) + resolved_models = list(project_object.models) + + if SpecialModelNames.all_proxy_models.value in resolved_models: + pass # all-proxy-models grants access to everything, skip the check + elif SpecialModelNames.all_team_models.value in resolved_models: + if team_object is not None and team_object.models: + resolved_models = list(team_object.models) + else: + resolved_models = [] + if len(resolved_models) > 0: + _can_object_call_model( + model=_model, + llm_router=llm_router, + models=resolved_models, + object_type="project", + ) + else: + can_project_access_model( + model=_model, + project_object=project_object, + llm_router=llm_router, + ) if not skip_budget_checks: # 3.0.2. If project is in budget @@ -464,6 +482,7 @@ async def common_checks( # noqa: PLR0915 skip_budget_checks=skip_budget_checks, valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, + team_object=team_object, ) # If this is a free model, skip all budget checks diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 69188fd200e..6b8225c5012 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -59,9 +59,9 @@ def reset_constants_module(): # Reload modules before test importlib.reload(constants) importlib.reload(auth_checks) - + yield - + # Reload modules after test to clean up importlib.reload(constants) importlib.reload(auth_checks) @@ -154,9 +154,9 @@ def test_experimental_ui_token_ignores_litellm_ui_session_duration( expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00")) now = get_utc_datetime() # Must be ~10 min, NOT 24h. If LITELLM_UI_SESSION_DURATION were incorrectly used, this would fail. - assert expires <= now + timedelta(minutes=11), ( - "Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION" - ) + assert expires <= now + timedelta( + minutes=11 + ), "Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION" def test_get_experimental_ui_login_jwt_auth_token_invalid( @@ -290,13 +290,15 @@ def test_get_cli_jwt_auth_token_custom_expiration( # Set custom expiration to 48 hours monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "48") - + # Reload the constants module to pick up the new env var importlib.reload(constants) # Also reload auth_checks to pick up the new constant value importlib.reload(auth_checks) - - token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values) + + token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token( + valid_sso_user_defined_values + ) # Decrypt and verify token contents decrypted_token = decrypt_value_helper( @@ -312,7 +314,6 @@ def test_get_cli_jwt_auth_token_custom_expiration( assert expires <= get_utc_datetime() + timedelta(hours=48, minutes=1) - @pytest.mark.asyncio async def test_default_internal_user_params_with_get_user_object(monkeypatch): """Test that default_internal_user_params is used when creating a new user via get_user_object""" @@ -433,7 +434,9 @@ async def test_get_user_object_upsert_includes_user_email(): mock_prisma_client.db.litellm_usertable.create.assert_called_once() creation_args = mock_prisma_client.db.litellm_usertable.create.call_args[1]["data"] - assert "user_email" in creation_args, "user_email should be included when upserting a new user" + assert ( + "user_email" in creation_args + ), "user_email should be included when upserting a new user" assert creation_args["user_email"] == "test@example.com" assert creation_args["user_id"] == "new_test_user" @@ -460,7 +463,9 @@ def test_log_budget_lookup_failure_skips_user_not_found(): @pytest.mark.asyncio -@patch("litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock) +@patch( + "litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock +) async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeypatch): """ Test that _get_team_db_check correctly calls the `new_team` function @@ -494,8 +499,12 @@ async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeyp @pytest.mark.asyncio -@patch("litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock) -async def test_get_team_db_check_does_not_call_new_team_if_exists(mock_new_team, monkeypatch): +@patch( + "litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock +) +async def test_get_team_db_check_does_not_call_new_team_if_exists( + mock_new_team, monkeypatch +): """ Test that _get_team_db_check does NOT call the `new_team` function if the team already exists in the database. @@ -1629,3 +1638,146 @@ async def test_custom_auth_common_checks_opt_in(): parent_otel_span=None, ) mock_common.assert_called_once() + + +# --------------------------------------------------------------------------- +# Tests for _run_project_checks: special model name resolution +# --------------------------------------------------------------------------- + + +def _make_project(models: list) -> "LiteLLM_ProjectTableCachedObj": + from litellm.proxy._types import LiteLLM_ProjectTableCachedObj + + return LiteLLM_ProjectTableCachedObj( + project_id="proj-1", + created_by="test", + updated_by="test", + models=models, + ) + + +def _make_team(models: list) -> LiteLLM_TeamTable: + return LiteLLM_TeamTable( + team_id="team-1", + models=models, + ) + + +@pytest.mark.asyncio +async def test_run_project_checks_all_team_models_allowed(): + """Project with all-team-models should allow a model that the team allows.""" + from litellm.proxy.auth.auth_checks import _run_project_checks + + project = _make_project(["all-team-models"]) + team = _make_team(["openai/gpt-5", "openai/gpt-4o"]) + proxy_logging = MagicMock() + + await _run_project_checks( + project_object=project, + _model="openai/gpt-5", + llm_router=None, + skip_budget_checks=True, + valid_token=None, + proxy_logging_obj=proxy_logging, + team_object=team, + ) + + +@pytest.mark.asyncio +async def test_run_project_checks_all_team_models_denied(): + """Project with all-team-models should deny a model the team does not allow.""" + from litellm.proxy.auth.auth_checks import _run_project_checks + + project = _make_project(["all-team-models"]) + team = _make_team(["openai/gpt-4o"]) + proxy_logging = MagicMock() + + with pytest.raises(ProxyException) as exc_info: + await _run_project_checks( + project_object=project, + _model="openai/gpt-5", + llm_router=None, + skip_budget_checks=True, + valid_token=None, + proxy_logging_obj=proxy_logging, + team_object=team, + ) + assert "project" in exc_info.value.message + + +@pytest.mark.asyncio +async def test_run_project_checks_all_team_models_no_team(): + """Project with all-team-models but no team object should skip the model check.""" + from litellm.proxy.auth.auth_checks import _run_project_checks + + project = _make_project(["all-team-models"]) + proxy_logging = MagicMock() + + await _run_project_checks( + project_object=project, + _model="openai/gpt-5", + llm_router=None, + skip_budget_checks=True, + valid_token=None, + proxy_logging_obj=proxy_logging, + team_object=None, + ) + + +@pytest.mark.asyncio +async def test_run_project_checks_all_proxy_models(): + """Project with all-proxy-models should allow any model.""" + from litellm.proxy.auth.auth_checks import _run_project_checks + + project = _make_project(["all-proxy-models"]) + proxy_logging = MagicMock() + + await _run_project_checks( + project_object=project, + _model="anything/any-model", + llm_router=None, + skip_budget_checks=True, + valid_token=None, + proxy_logging_obj=proxy_logging, + team_object=None, + ) + + +@pytest.mark.asyncio +async def test_run_project_checks_explicit_models_allowed(): + """Project with explicit model list should allow a listed model.""" + from litellm.proxy.auth.auth_checks import _run_project_checks + + project = _make_project(["openai/gpt-5"]) + proxy_logging = MagicMock() + + await _run_project_checks( + project_object=project, + _model="openai/gpt-5", + llm_router=None, + skip_budget_checks=True, + valid_token=None, + proxy_logging_obj=proxy_logging, + team_object=None, + ) + + +@pytest.mark.asyncio +async def test_run_project_checks_explicit_models_denied(): + """Project with explicit model list should deny an unlisted model.""" + from litellm.proxy.auth.auth_checks import _run_project_checks + + project = _make_project(["openai/gpt-4o"]) + proxy_logging = MagicMock() + + with pytest.raises(ProxyException) as exc_info: + await _run_project_checks( + project_object=project, + _model="openai/gpt-5", + llm_router=None, + skip_budget_checks=True, + valid_token=None, + proxy_logging_obj=proxy_logging, + team_object=None, + ) + assert "project" in exc_info.value.message