diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 2b9236bad8f..3b7ed112932 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -126,7 +126,10 @@ def _calculate_key_rotation_time(rotation_interval: str) -> datetime: def _set_key_rotation_fields( - data: dict, auto_rotate: bool, rotation_interval: Optional[str], existing_key_alias: Optional[str] = None + data: dict, + auto_rotate: bool, + rotation_interval: Optional[str], + existing_key_alias: Optional[str] = None, ) -> None: """ Helper function to set rotation fields in key data if auto_rotate is enabled. @@ -948,7 +951,16 @@ async def _validate_key_models_against_effective_team_models( team_member_models=member_models, ) - # 3. Fallback: if effective models empty but team has models, use team.models + # 3. Cap effective models to team.models (same intersection the runtime + # auth check applies) so the key is never stamped with models the + # runtime would block. + if ( + team_table.models + and SpecialModelNames.all_proxy_models.value not in team_table.models + ): + effective_models = list(set(effective_models) & set(team_table.models)) + + # 4. Fallback: if effective models empty but team has models, use team.models if not effective_models: if team_table.models: effective_models = team_table.models @@ -960,7 +972,7 @@ async def _validate_key_models_against_effective_team_models( }, ) - # 4. Step 6b: If data.models is empty, default to effective models + # 5. If data.models is empty, default to effective models if not data.models: data.models = effective_models else: @@ -3156,7 +3168,10 @@ async def delete_verification_tokens( hashed_token = hash_token(cast(str, key)) user_api_key_cache.delete_cache(hashed_token) - return {"deleted_keys": deleted_tokens, "failed_tokens": failed_tokens}, _keys_being_deleted + return { + "deleted_keys": deleted_tokens, + "failed_tokens": failed_tokens, + }, _keys_being_deleted def _transform_verification_tokens_to_deleted_records( @@ -3259,7 +3274,7 @@ async def delete_key_aliases( ) -async def _rotate_master_key( # noqa: PLR0915 +async def _rotate_master_key( # noqa: PLR0915 prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, current_master_key: str, @@ -3474,6 +3489,8 @@ async def _insert_deprecated_key( "Failed to insert deprecated key for grace period: %s", deprecated_err, ) + + async def _execute_virtual_key_regeneration( *, prisma_client: PrismaClient, @@ -4037,8 +4054,7 @@ def _get_member_team_ids_from_objects( team.team_id for team in team_objects if any( - member.user_id is not None - and member.user_id == user_api_key_dict.user_id + member.user_id is not None and member.user_id == user_api_key_dict.user_id for member in team.members_with_roles ) ] @@ -4282,9 +4298,7 @@ async def key_aliases( where_sql = " AND ".join(where_parts) - count_sql = ( - f'SELECT COUNT(*) AS count FROM "LiteLLM_VerificationToken" WHERE {where_sql}' - ) + count_sql = f'SELECT COUNT(*) AS count FROM "LiteLLM_VerificationToken" WHERE {where_sql}' count_rows = await prisma_client.db.query_raw(count_sql, *query_params) total_count = int(count_rows[0]["count"]) if count_rows else 0 @@ -4299,7 +4313,9 @@ async def key_aliases( f" LIMIT ${limit_idx} OFFSET ${offset_idx}" ) alias_rows = await prisma_client.db.query_raw(aliases_sql, *aliases_params) - aliases: List[str] = [row["key_alias"] for row in alias_rows if row.get("key_alias")] + aliases: List[str] = [ + row["key_alias"] for row in alias_rows if row.get("key_alias") + ] total_pages = -(-total_count // size) if total_count > 0 else 0 verbose_proxy_logger.debug( diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 573232fd437..37a0b5ed01d 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -1744,25 +1744,23 @@ async def _process_team_members( updated_users: List[LiteLLM_UserTable] = [] updated_team_memberships: List[LiteLLM_TeamMembership] = [] + # Always validate member models ⊆ team.models regardless of feature flag. + # This prevents storing out-of-bounds data that could become effective + # if the flag is enabled later. if data.models is not None: - from litellm.proxy.management_endpoints.common_utils import ( - _is_team_model_overrides_enabled, - ) - - if _is_team_model_overrides_enabled(): - if ( - complete_team_data.models - and SpecialModelNames.all_proxy_models.value - not in complete_team_data.models - ): - invalid = set(data.models) - set(complete_team_data.models) - if invalid: - raise HTTPException( - status_code=400, - detail={ - "error": f"Models {list(invalid)} not in team's allowed models: {complete_team_data.models}" - }, - ) + if ( + complete_team_data.models + and SpecialModelNames.all_proxy_models.value + not in complete_team_data.models + ): + invalid = set(data.models) - set(complete_team_data.models) + if invalid: + raise HTTPException( + status_code=400, + detail={ + "error": f"Models {list(invalid)} not in team's allowed models: {complete_team_data.models}" + }, + ) default_team_budget_id = ( complete_team_data.metadata.get("team_member_budget_id") @@ -2426,26 +2424,20 @@ async def team_member_update( identified_budget_id = tm.budget_id break + # Always validate member models ⊆ team.models regardless of feature flag. if data.models is not None: - from litellm.proxy.management_endpoints.common_utils import ( - _is_team_model_overrides_enabled, - ) - - if _is_team_model_overrides_enabled(): - # Validate models are within team's allowed set (team.models) - if ( - existing_team_row.models - and SpecialModelNames.all_proxy_models.value - not in existing_team_row.models - ): - invalid = set(data.models) - set(existing_team_row.models) - if invalid: - raise HTTPException( - status_code=400, - detail={ - "error": f"Models {list(invalid)} not in team's allowed models: {existing_team_row.models}" - }, - ) + if ( + existing_team_row.models + and SpecialModelNames.all_proxy_models.value not in existing_team_row.models + ): + invalid = set(data.models) - set(existing_team_row.models) + if invalid: + raise HTTPException( + status_code=400, + detail={ + "error": f"Models {list(invalid)} not in team's allowed models: {existing_team_row.models}" + }, + ) ### upsert new budget async with prisma_client.db.tx() as tx: diff --git a/tests/test_litellm/proxy/auth/test_team_model_overrides.py b/tests/test_litellm/proxy/auth/test_team_model_overrides.py new file mode 100644 index 00000000000..b2741725805 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_team_model_overrides.py @@ -0,0 +1,400 @@ +""" +Unit tests for team-scoped model overrides. + +Tests cover: +- compute_effective_team_models (union logic) +- can_team_access_model with overrides (runtime enforcement) +- default_models ⊆ team.models validation +- member models ⊆ team.models validation +- _validate_key_models_against_effective_team_models (key creation) +""" + +import os +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import ( + LiteLLM_TeamTable, + UserAPIKeyAuth, +) +from litellm.proxy.auth.auth_checks import ( + can_team_access_model, + compute_effective_team_models, +) + + +# ── compute_effective_team_models ──────────────────────────────────────────── + + +class TestComputeEffectiveTeamModels: + def test_union_of_defaults_and_member(self): + result = compute_effective_team_models( + team_default_models=["gpt-4o"], + team_member_models=["claude-sonnet"], + ) + assert set(result) == {"gpt-4o", "claude-sonnet"} + + def test_deduplicates(self): + result = compute_effective_team_models( + team_default_models=["gpt-4o", "claude-sonnet"], + team_member_models=["claude-sonnet"], + ) + assert sorted(result) == sorted(["gpt-4o", "claude-sonnet"]) + + def test_none_defaults(self): + result = compute_effective_team_models( + team_default_models=None, + team_member_models=["claude-sonnet"], + ) + assert result == ["claude-sonnet"] + + def test_none_member(self): + result = compute_effective_team_models( + team_default_models=["gpt-4o"], + team_member_models=None, + ) + assert result == ["gpt-4o"] + + def test_both_none(self): + result = compute_effective_team_models( + team_default_models=None, + team_member_models=None, + ) + assert result == [] + + def test_both_empty(self): + result = compute_effective_team_models( + team_default_models=[], + team_member_models=[], + ) + assert result == [] + + +# ── can_team_access_model (runtime check) ──────────────────────────────────── + + +class TestCanTeamAccessModelOverrides: + @pytest.mark.asyncio + @patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"}) + async def test_allowed_model_in_defaults(self): + """Model in default_models should be allowed.""" + team_object = LiteLLM_TeamTable( + team_id="team-1", + models=["gpt-4o", "gpt-4o-mini", "claude-sonnet"], + ) + valid_token = UserAPIKeyAuth( + token="test-token", + team_id="team-1", + team_default_models=["gpt-4o"], + team_member_models=None, + ) + mock_router = MagicMock() + mock_router.get_model_group_info.return_value = None + + result = await can_team_access_model( + model="gpt-4o", + team_object=team_object, + llm_router=mock_router, + team_model_aliases=None, + valid_token=valid_token, + ) + assert result is True + + @pytest.mark.asyncio + @patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"}) + async def test_allowed_model_in_member_override(self): + """Model in member models should be allowed.""" + team_object = LiteLLM_TeamTable( + team_id="team-1", + models=["gpt-4o", "gpt-4o-mini", "claude-sonnet"], + ) + valid_token = UserAPIKeyAuth( + token="test-token", + team_id="team-1", + team_default_models=["gpt-4o"], + team_member_models=["claude-sonnet"], + ) + mock_router = MagicMock() + mock_router.get_model_group_info.return_value = None + + result = await can_team_access_model( + model="claude-sonnet", + team_object=team_object, + llm_router=mock_router, + team_model_aliases=None, + valid_token=valid_token, + ) + assert result is True + + @pytest.mark.asyncio + @patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"}) + async def test_blocked_model_not_in_effective(self): + """Model in team.models but NOT in effective models should be blocked.""" + + team_object = LiteLLM_TeamTable( + team_id="team-1", + models=["gpt-4o", "gpt-4o-mini", "claude-sonnet"], + ) + valid_token = UserAPIKeyAuth( + token="test-token", + team_id="team-1", + team_default_models=["gpt-4o"], + team_member_models=["claude-sonnet"], + ) + mock_router = MagicMock() + mock_router.get_model_group_info.return_value = None + + with pytest.raises(Exception): + await can_team_access_model( + model="gpt-4o-mini", + team_object=team_object, + llm_router=mock_router, + team_model_aliases=None, + valid_token=valid_token, + ) + + @pytest.mark.asyncio + @patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"}) + async def test_effective_models_intersected_with_team_models(self): + """Even if default_models has out-of-bounds model, runtime should block it.""" + team_object = LiteLLM_TeamTable( + team_id="team-1", + models=["gpt-4o"], # team only allows gpt-4o + ) + valid_token = UserAPIKeyAuth( + token="test-token", + team_id="team-1", + team_default_models=[ + "gpt-4o", + "claude-sonnet", + ], # claude-sonnet is out of bounds + team_member_models=None, + ) + mock_router = MagicMock() + mock_router.get_model_group_info.return_value = None + + # claude-sonnet should be blocked even though it's in default_models + with pytest.raises(Exception): + await can_team_access_model( + model="claude-sonnet", + team_object=team_object, + llm_router=mock_router, + team_model_aliases=None, + valid_token=valid_token, + ) + + @pytest.mark.asyncio + @patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "false"}) + async def test_flag_off_uses_team_models(self): + """When flag is off, should use team.models as before.""" + team_object = LiteLLM_TeamTable( + team_id="team-1", + models=["gpt-4o", "gpt-4o-mini"], + ) + valid_token = UserAPIKeyAuth( + token="test-token", + team_id="team-1", + team_default_models=["gpt-4o"], + team_member_models=None, + ) + mock_router = MagicMock() + mock_router.get_model_group_info.return_value = None + + # gpt-4o-mini should be allowed because flag is off, team.models is used + result = await can_team_access_model( + model="gpt-4o-mini", + team_object=team_object, + llm_router=mock_router, + team_model_aliases=None, + valid_token=valid_token, + ) + assert result is True + + @pytest.mark.asyncio + @patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"}) + async def test_no_overrides_configured_uses_team_models(self): + """Team with no default_models/member_models uses team.models unchanged.""" + team_object = LiteLLM_TeamTable( + team_id="team-1", + models=["gpt-4o", "gpt-4o-mini"], + ) + valid_token = UserAPIKeyAuth( + token="test-token", + team_id="team-1", + team_default_models=None, + team_member_models=None, + ) + mock_router = MagicMock() + mock_router.get_model_group_info.return_value = None + + result = await can_team_access_model( + model="gpt-4o-mini", + team_object=team_object, + llm_router=mock_router, + team_model_aliases=None, + valid_token=valid_token, + ) + assert result is True + + +# ── _validate_key_models_against_effective_team_models ─────────────────────── + + +class TestValidateKeyModelsAgainstEffective: + @pytest.mark.asyncio + @patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"}) + async def test_empty_data_models_gets_effective(self): + """Key with no models should inherit effective models.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_key_models_against_effective_team_models, + ) + + mock_prisma = MagicMock() + mock_membership = MagicMock() + mock_membership.models = ["claude-sonnet"] + mock_prisma.db.litellm_teammembership.find_unique = AsyncMock( + return_value=mock_membership + ) + + team_table = MagicMock() + team_table.default_models = ["gpt-4o"] + team_table.models = ["gpt-4o", "claude-sonnet", "gpt-4o-mini"] + + data = MagicMock() + data.models = [] + + await _validate_key_models_against_effective_team_models( + team_id="team-1", + user_id="user-1", + data=data, + team_table=team_table, + prisma_client=mock_prisma, + ) + + assert set(data.models) == {"gpt-4o", "claude-sonnet"} + + @pytest.mark.asyncio + @patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"}) + async def test_effective_models_capped_to_team_models(self): + """Key effective models should be intersected with team.models.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_key_models_against_effective_team_models, + ) + + mock_prisma = MagicMock() + mock_membership = MagicMock() + mock_membership.models = ["claude-sonnet"] # out-of-bounds + mock_prisma.db.litellm_teammembership.find_unique = AsyncMock( + return_value=mock_membership + ) + + team_table = MagicMock() + team_table.default_models = ["gpt-4o"] + team_table.models = ["gpt-4o"] # team only allows gpt-4o + + data = MagicMock() + data.models = [] + + await _validate_key_models_against_effective_team_models( + team_id="team-1", + user_id="user-1", + data=data, + team_table=team_table, + prisma_client=mock_prisma, + ) + + # claude-sonnet should be capped out + assert data.models == ["gpt-4o"] + + @pytest.mark.asyncio + @patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"}) + async def test_disallowed_model_in_key_raises(self): + """Key requesting model outside effective set should raise 403.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_key_models_against_effective_team_models, + ) + + mock_prisma = MagicMock() + mock_membership = MagicMock() + mock_membership.models = ["claude-sonnet"] + mock_prisma.db.litellm_teammembership.find_unique = AsyncMock( + return_value=mock_membership + ) + + team_table = MagicMock() + team_table.default_models = ["gpt-4o"] + team_table.models = ["gpt-4o", "claude-sonnet", "gpt-4o-mini"] + + data = MagicMock() + data.models = ["gpt-4o-mini"] # not in effective models + + with pytest.raises(HTTPException) as exc_info: + await _validate_key_models_against_effective_team_models( + team_id="team-1", + user_id="user-1", + data=data, + team_table=team_table, + prisma_client=mock_prisma, + ) + + assert exc_info.value.status_code == 403 + + @pytest.mark.asyncio + @patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"}) + async def test_no_overrides_skips_validation(self): + """Teams without default_models or member models skip override validation.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_key_models_against_effective_team_models, + ) + + mock_prisma = MagicMock() + mock_membership = MagicMock() + mock_membership.models = [] + mock_prisma.db.litellm_teammembership.find_unique = AsyncMock( + return_value=mock_membership + ) + + team_table = MagicMock() + team_table.default_models = [] + team_table.models = ["gpt-4o", "gpt-4o-mini"] + + data = MagicMock() + data.models = ["gpt-4o-mini"] + + # Should return without modifying data.models + await _validate_key_models_against_effective_team_models( + team_id="team-1", + user_id="user-1", + data=data, + team_table=team_table, + prisma_client=mock_prisma, + ) + + assert data.models == ["gpt-4o-mini"] + + @pytest.mark.asyncio + @patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "false"}) + async def test_flag_off_skips_entirely(self): + """When flag is off, validation is skipped entirely.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_key_models_against_effective_team_models, + ) + + mock_prisma = MagicMock() + team_table = MagicMock() + + data = MagicMock() + data.models = ["anything"] + + await _validate_key_models_against_effective_team_models( + team_id="team-1", + user_id="user-1", + data=data, + team_table=team_table, + prisma_client=mock_prisma, + ) + + # Should be unchanged — no validation happened + assert data.models == ["anything"]