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 0851fd5325d..5a2a080a495 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 @@ -35,6 +35,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.management_endpoints.key_management_endpoints import ( _check_org_key_limits, _check_project_key_limits, + _check_project_key_limits_on_mutation, _check_team_key_limits, _common_key_generation_helper, _enforce_upperbound_key_params, @@ -18301,265 +18302,192 @@ async def test_key_creator_cannot_detach_project_without_admin_access(): # --- Tests: a project may only be attached to keys of the team that owns it --- +_OWNED_PROJECT: Final = "proj-owned-1" -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 [], + +@pytest.mark.parametrize( + "project_team_id, key_team_id, expected_status, expected_in_detail", + [ + # A key on another team must not point at this project. + ("team-b", "team-a", 403, ["team-b", "team-a"]), + # The issue's step 5: no team at all still charges a team's project. + ("team-b", None, 403, ["no team"]), + # Accept controls. Without these the check is a wall, not a boundary. + ("team-a", "team-a", None, []), + # A project with no owning team has nobody to protect. + (None, "team-a", None, []), + (None, None, None, []), + ], +) +@pytest.mark.asyncio +async def test_a_project_is_only_for_keys_of_its_own_team( + project_team_id, key_team_id, expected_status, expected_in_detail +): + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, [], team_id=project_team_id ) - -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, - ): + async def check(): await _check_project_key_limits( - project_id=project.project_id, - data=data, + project_id=_OWNED_PROJECT, + data=GenerateKeyRequest(team_id=key_team_id), prisma_client=MagicMock(), - user_api_key_cache=MagicMock(), + user_api_key_cache=user_api_key_cache, key_team_id=key_team_id, ) + if expected_status is None: + await check() + # The project was read and accepted, rather than never reached. + assert await user_api_key_cache.async_get_cache(key=_project_cache_key(_OWNED_PROJECT)) + return -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, - ) + await check() + assert exc.value.status_code == expected_status + for fragment in expected_in_detail: + assert fragment in str(exc.value.detail) @pytest.mark.asyncio async def test_a_foreign_project_is_refused_as_foreign_not_as_a_model_problem(): + """A 403 reads as a tenancy problem; the 400 would read as a configuration one.""" + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, ["gpt-4o-mini"], team_id="team-b" + ) + 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"]), + await _check_project_key_limits( + project_id=_OWNED_PROJECT, + data=GenerateKeyRequest(team_id="team-a", models=["gpt-4o"]), + prisma_client=MagicMock(), + user_api_key_cache=user_api_key_cache, key_team_id="team-a", ) + assert exc.value.status_code == 403 +@pytest.mark.parametrize( + "existing_team_id, data_kwargs, expected_status", + [ + # Moves only the project: the key's stored team is what counts. + ("team-a", {"project_id": _OWNED_PROJECT}, 403), + # Moves only the team: the project stays attached, so it still counts. + ("team-b", {"team_id": "team-a"}, 403), + # Moves the key onto the project's own team. Accept control. + (None, {"team_id": "team-b"}, None), + # Touches neither, so there is nothing to re-check. + ("team-a", {"key_alias": "renamed"}, None), + ], +) @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(): +async def test_an_update_checks_the_project_the_mutation_leaves_attached( + existing_team_id, data_kwargs, expected_status +): existing: Final = LiteLLM_VerificationToken( - token="sk-hash", team_id="team-b", project_id="proj-owned-1" + token="sk-hash", team_id=existing_team_id, project_id=_OWNED_PROJECT ) - 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) + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") - -@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, - ): + async def check(): await _check_project_key_limits_on_mutation( - data=UpdateKeyRequest(key="sk-x", key_alias="renamed"), + data=UpdateKeyRequest(key="sk-x", **data_kwargs), existing_key_row=existing, prisma_client=MagicMock(), - user_api_key_cache=MagicMock(), + user_api_key_cache=user_api_key_cache, ) - assert lookup.await_count == 0 + + if expected_status is None: + await check() + assert await user_api_key_cache.async_get_cache(key=_project_cache_key(_OWNED_PROJECT)) + return + + with pytest.raises(HTTPException) as exc: + await check() + assert exc.value.status_code == expected_status @pytest.mark.asyncio -async def test_regenerate_may_not_attach_another_teams_project(): +async def test_an_update_that_touches_none_of_the_fields_does_not_read_the_project(): + existing: Final = LiteLLM_VerificationToken( + token="sk-hash", team_id="team-a", project_id=_OWNED_PROJECT + ) + prisma_client: Final = MagicMock() + + # The cache is empty, so any lookup would have to reach the database. + await _check_project_key_limits_on_mutation( + data=UpdateKeyRequest(key="sk-x", key_alias="renamed"), + existing_key_row=existing, + prisma_client=prisma_client, + user_api_key_cache=UserApiKeyCache(), + ) + + assert prisma_client.mock_calls == [] + + +@pytest.mark.parametrize( + "key_team_id, expected_status, expected_updates", + [ + # /key/regenerate wrote project_id and team_id with no project check. + ("team-a", 403, 0), + # Accept control: a key already on the project's team regenerates. + ("team-b", None, 1), + ], +) +@pytest.mark.asyncio +async def test_regenerate_checks_the_project_it_attaches( + key_team_id, expected_status, expected_updates +): 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") + existing_key: Final = LiteLLM_VerificationToken(token="abc123", team_id=key_team_id) mock_prisma_client: Final = _make_regenerate_mock_prisma() + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], 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="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: + async def regenerate(): + with ( + patch( # test-quality-ok: a fresh random token is a side effect of regeneration, not the project check under test; the file's other regenerate cells stub the same seam + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( # test-quality-ok: the deprecated-key row is a side effect of regeneration; stubbing it keeps the DB assertion below about the key update alone + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: the cache delete is a side effect of regeneration and would evict the seeded project this cell reads + "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"), + data=RegenerateKeyRequest(project_id=_OWNED_PROJECT), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, - user_api_key_cache=MagicMock(), + user_api_key_cache=user_api_key_cache, proxy_logging_obj=MagicMock(), ) - assert exc_info.value.status_code == 403 + if expected_status is None: + await regenerate() + else: + with pytest.raises(HTTPException) as exc: + await regenerate() + assert exc.value.status_code == expected_status + # 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, + assert ( + mock_prisma_client.db.litellm_verificationtoken.update.await_count == expected_updates ) - 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.""" @@ -18582,19 +18510,32 @@ def _make_generate_mock_prisma(): return mock_prisma_client -async def _generate_key_with_defaulted_team(monkeypatch, project_team_id): +@pytest.mark.parametrize( + "project_team_id, expected_status", + [ + # default_key_generate_params fills team_id after the request is parsed, + # so the key ends up on team-b and may use its projects. + ("team-b", None), + # ... and still may not use another team's. + ("team-c", 403), + ], +) +@pytest.mark.asyncio +async def test_a_team_supplied_by_defaults_is_what_the_project_is_checked_against( + monkeypatch, project_team_id, expected_status +): monkeypatch.setattr( "litellm.proxy.proxy_server.prisma_client", _make_generate_mock_prisma() ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", + await _cache_with_project(_OWNED_PROJECT, [], team_id=project_team_id), + ) 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), - ): + async def generate(): return await _common_key_generation_helper( - data=GenerateKeyRequest(project_id="proj-owned-1"), + data=GenerateKeyRequest(project_id=_OWNED_PROJECT), user_api_key_dict=UserAPIKeyAuth( user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234", user_id="1234" ), @@ -18602,16 +18543,11 @@ async def _generate_key_with_defaulted_team(monkeypatch, project_team_id): team_table=None, ) + if expected_status is None: + response = await generate() + assert response.team_id == "team-b" + return -@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 + await generate() + assert exc.value.status_code == expected_status