diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 23dfe9e0c07..346b46d2052 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1066,6 +1066,19 @@ async def _common_key_generation_helper( # check if user set upperbound key/generate params on config.yaml _enforce_upperbound_key_params(data, fill_defaults=True) + # Checked after the defaults, because default_key_generate_params can supply + # team_id and the project's owner is checked against the key's final team. + if data.project_id is not None and prisma_client is not None: + from litellm.proxy.proxy_server import user_api_key_cache + + await _check_project_key_limits( + project_id=data.project_id, + data=data, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + key_team_id=data.team_id, + ) + # Delegated-authority ceiling (GHSA-q775-qw9r-2r4g): a non-admin caller # cannot grant a key a higher budget than their own authority. is_ui_session_team_key = user_api_key_dict.team_id == UI_SESSION_TOKEN_TEAM_ID and _requested_team_id is not None @@ -1553,11 +1566,8 @@ async def _check_project_key_limits( Validate that the key belongs to the project's team, and that its models and budget respect the project's limits. - - The project's owning team must be the key's team. A project is created - under exactly one team and its budget and models are that team's, so a - key on another team recorded under it charges a tenant that never granted - anything — and issuing one needs no proxy-admin rights, only the project - id (#41089). A project with no team belongs to nobody and is left alone. + - The project's owning team must be the key's team. A project with no team + has no owner to protect, so it is not restricted - Key models must be a subset of project models, except the all-team-models / all-proxy-models sentinels, which inherit a parent scope and are narrowed by the project at request time - Key max_budget must be <= project max_budget @@ -1574,7 +1584,6 @@ async def _check_project_key_limits( detail={"error": f"Project not found, project_id={project_id}"}, ) - # Validate the project's team owns the key if project_obj.team_id is not None and project_obj.team_id != key_team_id: raise HTTPException( status_code=403, @@ -1610,6 +1619,33 @@ async def _check_project_key_limits( ) +# Touching any of these can change the project a key is under, the team it is on, or what the project must allow. +_PROJECT_LIMIT_FIELDS: Final = frozenset({"project_id", "team_id", "models", "max_budget"}) + + +async def _check_project_key_limits_on_mutation( + data: UpdateKeyRequest | RegenerateKeyRequest, + existing_key_row: LiteLLM_VerificationToken, + prisma_client: PrismaClient, + user_api_key_cache: UserApiKeyCache, +) -> None: + """Run _check_project_key_limits against the key as the mutation leaves it.""" + if not data.model_fields_set & _PROJECT_LIMIT_FIELDS: + return + + project_id: Final = data.project_id if "project_id" in data.model_fields_set else existing_key_row.project_id + if project_id is None: + return + + await _check_project_key_limits( + project_id=project_id, + data=data, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + key_team_id=(data.team_id if "team_id" in data.model_fields_set else existing_key_row.team_id), + ) + + def check_org_key_model_specific_limits( keys: Sequence[LiteLLM_VerificationToken], org_table: LiteLLM_OrganizationTable, @@ -1933,16 +1969,6 @@ async def generate_key_fn( prisma_client=prisma_client, ) - # Validate key against project limits if project_id is set - if data.project_id is not None: - await _check_project_key_limits( - project_id=data.project_id, - data=data, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - key_team_id=data.team_id, - ) - return await _common_key_generation_helper( data=data, user_api_key_dict=user_api_key_dict, @@ -2879,23 +2905,12 @@ async def _validate_update_key_data( access_group_ids=data.access_group_ids, ) - # Validate key against project limits if project_id is being set - _project_id_to_check: Final = ( - data.project_id if "project_id" in data.model_fields_set else existing_key_row.project_id + await _check_project_key_limits_on_mutation( + data=data, + existing_key_row=existing_key_row, + prisma_client=checked_prisma_client, + user_api_key_cache=user_api_key_cache, ) - # Also when the project itself is being set or changed: that is exactly when - # the team that owns it has to be checked, and a request that moves only the - # project carries neither models nor max_budget (#41089). - if _project_id_to_check is not None and ( - "project_id" in data.model_fields_set or data.models is not None or data.max_budget is not None - ): - await _check_project_key_limits( - project_id=_project_id_to_check, - data=data, - prisma_client=checked_prisma_client, - user_api_key_cache=user_api_key_cache, - key_team_id=(data.team_id if "team_id" in data.model_fields_set else existing_key_row.team_id), - ) # When the caller asks to change the key's organization_id, require that # they are a member of (or a proxy admin over) the target organization. @@ -5154,6 +5169,12 @@ async def _execute_virtual_key_regeneration( if data is not None: # Enforce upperbound key params on regenerate (don't fill defaults) _enforce_upperbound_key_params(data, fill_defaults=False) + await _check_project_key_limits_on_mutation( + data=data, + existing_key_row=key_in_db, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) non_default_values = await prepare_key_update_data(data=data, existing_key_row=key_in_db) # Only validate key_alias format if it's actually being changed new_key_alias: Final = non_default_values.get("key_alias") diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 37e7bc95285..0851fd5325d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -18297,3 +18297,321 @@ async def test_key_creator_cannot_detach_project_without_admin_access(): ) assert exc.value.status_code == 403 assert "Only proxy admins, team admins, or org admins" in str(exc.value.detail) + + +# --- Tests: a project may only be attached to keys of the team that owns it --- + + +def _make_owned_project(team_id, models=None, project_id="proj-owned-1"): + return LiteLLM_ProjectTableCachedObj( + project_id=project_id, + team_id=team_id, + models=models or [], + ) + + +async def _check_project_limits_with(project, data, key_team_id): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _check_project_key_limits, + ) + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_project_object", + new_callable=AsyncMock, + return_value=project, + ): + await _check_project_key_limits( + project_id=project.project_id, + data=data, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + key_team_id=key_team_id, + ) + + +async def _check_project_limits_on_mutation_with(project, data, existing_key_row): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _check_project_key_limits_on_mutation, + ) + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_project_object", + new_callable=AsyncMock, + return_value=project, + ): + await _check_project_key_limits_on_mutation( + data=data, + existing_key_row=existing_key_row, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + ) + + +@pytest.mark.asyncio +async def test_a_key_may_not_point_at_another_teams_project(): + with pytest.raises(HTTPException) as exc: + await _check_project_limits_with( + _make_owned_project(team_id="team-b"), + GenerateKeyRequest(team_id="team-a"), + key_team_id="team-a", + ) + assert exc.value.status_code == 403 + detail = str(exc.value.detail) + assert "team-b" in detail and "team-a" in detail + + +@pytest.mark.asyncio +async def test_a_key_with_no_team_may_not_point_at_a_teams_project(): + with pytest.raises(HTTPException) as exc: + await _check_project_limits_with( + _make_owned_project(team_id="team-b"), + GenerateKeyRequest(), + key_team_id=None, + ) + assert exc.value.status_code == 403 + assert "no team" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_a_key_on_the_owning_team_is_accepted(): + await _check_project_limits_with( + _make_owned_project(team_id="team-a", models=["gpt-4o"]), + GenerateKeyRequest(team_id="team-a", models=["gpt-4o"]), + key_team_id="team-a", + ) + + +@pytest.mark.asyncio +async def test_a_project_with_no_team_is_not_restricted(): + await _check_project_limits_with( + _make_owned_project(team_id=None), + GenerateKeyRequest(team_id="team-a"), + key_team_id="team-a", + ) + await _check_project_limits_with( + _make_owned_project(team_id=None), + GenerateKeyRequest(), + key_team_id=None, + ) + + +@pytest.mark.asyncio +async def test_a_foreign_project_is_refused_as_foreign_not_as_a_model_problem(): + with pytest.raises(HTTPException) as exc: + await _check_project_limits_with( + _make_owned_project(team_id="team-b", models=["gpt-4o-mini"]), + GenerateKeyRequest(team_id="team-a", models=["gpt-4o"]), + key_team_id="team-a", + ) + assert exc.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_an_update_that_moves_only_the_project_uses_the_keys_stored_team(): + existing: Final = LiteLLM_VerificationToken(token="sk-hash", team_id="team-a") + with pytest.raises(HTTPException) as exc: + await _check_project_limits_on_mutation_with( + _make_owned_project(team_id="team-b"), + UpdateKeyRequest(key="sk-x", project_id="proj-owned-1"), + existing, + ) + assert exc.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_an_update_that_moves_only_the_team_still_checks_the_attached_project(): + existing: Final = LiteLLM_VerificationToken( + token="sk-hash", team_id="team-b", project_id="proj-owned-1" + ) + with pytest.raises(HTTPException) as exc: + await _check_project_limits_on_mutation_with( + _make_owned_project(team_id="team-b"), + UpdateKeyRequest(key="sk-x", team_id="team-a"), + existing, + ) + assert exc.value.status_code == 403 + assert "team-a" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_an_update_that_moves_the_key_to_the_projects_own_team_is_accepted(): + existing: Final = LiteLLM_VerificationToken( + token="sk-hash", team_id=None, project_id="proj-owned-1" + ) + await _check_project_limits_on_mutation_with( + _make_owned_project(team_id="team-b"), + UpdateKeyRequest(key="sk-x", team_id="team-b"), + existing, + ) + + +@pytest.mark.asyncio +async def test_an_update_that_touches_neither_does_not_look_the_project_up(): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _check_project_key_limits_on_mutation, + ) + + existing: Final = LiteLLM_VerificationToken( + token="sk-hash", team_id="team-a", project_id="proj-owned-1" + ) + lookup: Final = AsyncMock(return_value=_make_owned_project(team_id="team-b")) + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_project_object", + new=lookup, + ): + await _check_project_key_limits_on_mutation( + data=UpdateKeyRequest(key="sk-x", key_alias="renamed"), + existing_key_row=existing, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + ) + assert lookup.await_count == 0 + + +@pytest.mark.asyncio +async def test_regenerate_may_not_attach_another_teams_project(): + from litellm.proxy._types import RegenerateKeyRequest + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _execute_virtual_key_regeneration, + ) + + existing_key: Final = LiteLLM_VerificationToken(token="abc123", team_id="team-a") + mock_prisma_client: Final = _make_regenerate_mock_prisma() + + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_project_object", + new_callable=AsyncMock, + return_value=_make_owned_project(team_id="team-b"), + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + ): + with pytest.raises(HTTPException) as exc_info: + await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=existing_key, + hashed_api_key="abc123", + key="abc123", + data=RegenerateKeyRequest(project_id="proj-owned-1"), + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + assert exc_info.value.status_code == 403 + # A refused regenerate must not reach the DB update. + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 0 + + +@pytest.mark.asyncio +async def test_regenerate_may_attach_a_project_of_the_keys_own_team(): + from litellm.proxy._types import RegenerateKeyRequest + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _execute_virtual_key_regeneration, + ) + + existing_key: Final = LiteLLM_VerificationToken(token="abc123", team_id="team-b") + mock_prisma_client: Final = _make_regenerate_mock_prisma() + + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_project_object", + new_callable=AsyncMock, + return_value=_make_owned_project(team_id="team-b"), + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + ): + await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=existing_key, + hashed_api_key="abc123", + key="abc123", + data=RegenerateKeyRequest(project_id="proj-owned-1"), + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1 + + +def _make_generate_mock_prisma(): + """Mock prisma client shaped for _common_key_generation_helper.""" + mock_prisma_client = AsyncMock() + mock_prisma_client.insert_data = AsyncMock( + return_value=MagicMock( + token="hashed_token_123", litellm_budget_table=None, object_permission=None + ) + ) + mock_prisma_client.db = MagicMock() + mock_prisma_client.db.litellm_verificationtoken = MagicMock() + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0) + mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock( + return_value=MagicMock( + token="hashed_token_123", litellm_budget_table=None, object_permission=None + ) + ) + return mock_prisma_client + + +async def _generate_key_with_defaulted_team(monkeypatch, project_team_id): + monkeypatch.setattr( + "litellm.proxy.proxy_server.prisma_client", _make_generate_mock_prisma() + ) + monkeypatch.setattr(litellm, "default_key_generate_params", {"team_id": "team-b"}) + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_project_object", + new_callable=AsyncMock, + return_value=_make_owned_project(team_id=project_team_id), + ): + return await _common_key_generation_helper( + data=GenerateKeyRequest(project_id="proj-owned-1"), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234", user_id="1234" + ), + litellm_changed_by=None, + team_table=None, + ) + + +@pytest.mark.asyncio +async def test_a_team_supplied_by_defaults_is_what_the_project_is_checked_against(monkeypatch): + # default_key_generate_params fills team_id after the request is parsed, so a + # request with no team_id still ends up on team-b and may use its projects. + await _generate_key_with_defaulted_team(monkeypatch, project_team_id="team-b") + + +@pytest.mark.asyncio +async def test_a_team_supplied_by_defaults_does_not_open_another_teams_project(monkeypatch): + with pytest.raises(HTTPException) as exc: + await _generate_key_with_defaulted_team(monkeypatch, project_team_id="team-c") + assert exc.value.status_code == 403 diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_project_team_ownership.py b/tests/test_litellm/proxy/management_endpoints/test_key_project_team_ownership.py deleted file mode 100644 index 082c9a0e501..00000000000 --- a/tests/test_litellm/proxy/management_endpoints/test_key_project_team_ownership.py +++ /dev/null @@ -1,115 +0,0 @@ -"""A project may only be attached to keys of the team that owns it (#41089). - -A project is created under exactly one team, and its budget and models are that -team's. Nothing checked that the key's team matched, so an admin of team-a — a -member of no other team — could issue a key on their own team pointing at a -team-b project, and team-b would see the spend under a project they never -granted anything on. Only the project id was needed. - -The negative controls are the point: a project with no owning team is left -alone, and a key on the owning team still passes, or the check would be a wall -rather than a boundary. -""" - -from unittest.mock import AsyncMock, MagicMock - -import pytest -from fastapi import HTTPException - -from litellm.models.project import LiteLLM_ProjectTable -from litellm.proxy._types import GenerateKeyRequest, UpdateKeyRequest - - -def _project(team_id, models=None, project_id="proj-1"): - return LiteLLM_ProjectTable( - project_id=project_id, - team_id=team_id, - models=models or [], - ) - - -async def _check(project, data, key_team_id): - """Drive _check_project_key_limits with the project the store would return.""" - from litellm.proxy.management_endpoints import key_management_endpoints as kme - - original = kme.get_project_object - kme.get_project_object = AsyncMock(return_value=project) - try: - await kme._check_project_key_limits( - project_id=project.project_id, - data=data, - prisma_client=MagicMock(), - user_api_key_cache=MagicMock(), - key_team_id=key_team_id, - ) - finally: - kme.get_project_object = original - - -@pytest.mark.asyncio -async def test_a_key_may_not_point_at_another_teams_project(): - with pytest.raises(HTTPException) as exc: - await _check( - _project(team_id="team-b"), - GenerateKeyRequest(team_id="team-a"), - key_team_id="team-a", - ) - assert exc.value.status_code == 403 - detail = str(exc.value.detail) - assert "team-b" in detail and "team-a" in detail - - -@pytest.mark.asyncio -async def test_a_key_with_no_team_may_not_point_at_a_teams_project(): - # The issue's step 5: no team at all still charges a team's project. - with pytest.raises(HTTPException) as exc: - await _check( - _project(team_id="team-b"), - GenerateKeyRequest(), - key_team_id=None, - ) - assert exc.value.status_code == 403 - assert "no team" in str(exc.value.detail) - - -@pytest.mark.asyncio -async def test_a_key_on_the_owning_team_is_accepted(): - await _check( - _project(team_id="team-b"), - GenerateKeyRequest(team_id="team-b"), - key_team_id="team-b", - ) - - -@pytest.mark.asyncio -async def test_a_project_with_no_team_is_left_alone(): - # Nobody owns it, so there is no boundary to cross — rejecting here would - # break every project created outside a team. - await _check(_project(team_id=None), GenerateKeyRequest(team_id="team-a"), key_team_id="team-a") - await _check(_project(team_id=None), GenerateKeyRequest(), key_team_id=None) - - -@pytest.mark.asyncio -async def test_the_ownership_check_runs_before_the_model_check(): - # A foreign project must be refused as foreign, not as "model not allowed": - # the 400 would read as a configuration problem and hide the tenancy one. - with pytest.raises(HTTPException) as exc: - await _check( - _project(team_id="team-b", models=["gpt-4o-mini"]), - GenerateKeyRequest(team_id="team-a", models=["gpt-4o"]), - key_team_id="team-a", - ) - assert exc.value.status_code == 403 - - -@pytest.mark.asyncio -async def test_an_update_carries_the_keys_existing_team(): - # /key/update sends no team_id when only the project changes, so the check - # has to use the key's stored team rather than treating it as absent. - with pytest.raises(HTTPException) as exc: - await _check( - _project(team_id="team-b"), - UpdateKeyRequest(key="sk-x", project_id="proj-1"), - key_team_id="team-a", - ) - assert exc.value.status_code == 403