From eeb13fffbd66cabaae5397197aad46cfa647c254 Mon Sep 17 00:00:00 2001 From: L4XB Date: Mon, 14 Sep 2026 18:20:51 +0200 Subject: [PATCH 01/20] fix(key mgmt): a project may only be attached to keys of its own team MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A project is created under exactly one team, and its budget and models are validated against that team's. Nothing checked that a key's team matched, so `POST /key/generate` with `team_id: team-a` and a `project_id` owned by `team-b` answered 200 and stored exactly that — a key belonging to team-a, recorded under team-b's project and validated against its models and budget. The same held with no team at all. It needs no proxy-admin rights: an admin of team-a who is a member of no other team can issue keys under another tenant's project with nothing but the project id, and that tenant sees the spend without having granted anything. `_check_project_key_limits` now takes the key's team and refuses a project owned by a different one, before the model and budget checks so the refusal reads as what it is. A project with no owning team is left alone — nobody owns it, so there is no boundary to cross. `/key/update` passes the key's stored team when the request does not carry one, and now runs the check whenever the project itself is set or changed, which a request that moves only the project previously skipped. Two existing cells needed the project's team passed explicitly: they measure the model allowlist, not tenancy, and would otherwise have been asserting model behaviour on a request the new gate refuses. A third builds an unowned project for the same reason. Fixes #41089 --- .../key_management_endpoints.py | 29 ++++- .../test_key_management_endpoints.py | 13 +- .../test_key_project_team_ownership.py | 115 ++++++++++++++++++ 3 files changed, 152 insertions(+), 5 deletions(-) create mode 100644 tests/test_litellm/proxy/management_endpoints/test_key_project_team_ownership.py diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 95ccb7bbe0b..23dfe9e0c07 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1547,10 +1547,17 @@ async def _check_project_key_limits( data: GenerateKeyRequest | UpdateKeyRequest, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, + key_team_id: str | None = None, ) -> None: """ - Validate that key's models and budget respect its project's 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. - 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 @@ -1567,6 +1574,17 @@ 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, + detail={ + "error": f"Project {project_id} belongs to team {project_obj.team_id}, " + f"but the key belongs to {key_team_id if key_team_id is not None else 'no team'}. " + "A project can only be attached to keys of the team that owns it." + }, + ) + # Validate key models are a subset of project models if data.models and len(project_obj.models) > 0: for m in data.models: @@ -1922,6 +1940,7 @@ async def generate_key_fn( 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( @@ -2864,12 +2883,18 @@ async def _validate_update_key_data( _project_id_to_check: Final = ( data.project_id if "project_id" in data.model_fields_set else existing_key_row.project_id ) - if _project_id_to_check is not None and (data.models is not None or data.max_budget is not None): + # 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 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 2ac52da57df..37e7bc95285 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 @@ -17849,11 +17849,13 @@ async def test_regenerate_key_repoints_live_membership_not_the_key_row_it_read( ) == ["attached-model"] -async def _cache_with_project(project_id: str, project_models: list[str]) -> UserApiKeyCache: +async def _cache_with_project( + project_id: str, project_models: list[str], team_id: str | None = "team-lit-5823" +) -> UserApiKeyCache: user_api_key_cache = UserApiKeyCache() await user_api_key_cache.async_set_cache( key=_project_cache_key(project_id), - value=LiteLLM_ProjectTableCachedObj(project_id=project_id, team_id="team-lit-5823", models=project_models), + value=LiteLLM_ProjectTableCachedObj(project_id=project_id, team_id=team_id, models=project_models), model_type=LiteLLM_ProjectTableCachedObj, ) return user_api_key_cache @@ -17871,6 +17873,7 @@ async def test_check_project_key_limits_accepts_inherited_model_sentinels(reques data=request_cls(key="sk-lit-5823", models=[sentinel]), prisma_client=MagicMock(), user_api_key_cache=user_api_key_cache, + key_team_id="team-lit-5823", ) @@ -17889,6 +17892,7 @@ async def test_check_project_key_limits_still_rejects_real_model_outside_project data=request_cls(key="sk-lit-5823", models=key_models), prisma_client=MagicMock(), user_api_key_cache=user_api_key_cache, + key_team_id="team-lit-5823", ) assert exc_info.value.status_code == 400 @@ -18257,7 +18261,10 @@ async def test_project_detachment_preserves_omission_and_other_key_fields(): @pytest.mark.asyncio async def test_project_detachment_uses_effective_project_for_validation(project_id: str | None): existing: Final = LiteLLM_VerificationToken(token="project-detach-token", project_id="project-orbit") - cache: Final = await _cache_with_project("project-orbit", ["model-orbit"]) + # An unowned project: this cell is about WHICH project the validation uses, + # so the ownership gate (#41089) must not be what it measures. Giving the + # key a team instead would pull the whole team lookup into a MagicMock db. + cache: Final = await _cache_with_project("project-orbit", ["model-orbit"], team_id=None) data: Final = UpdateKeyRequest(key=existing.token, project_id=project_id, models=["model-other"]) if project_id is None: await _validate_update_key_data( 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 new file mode 100644 index 00000000000..082c9a0e501 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_key_project_team_ownership.py @@ -0,0 +1,115 @@ +"""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 From 4440d4e6c754d96edf5f6942ce0cbd06d5cf1a9e Mon Sep 17 00:00:00 2001 From: L4XB Date: Mon, 14 Sep 2026 18:59:21 +0200 Subject: [PATCH 02/20] fix(key mgmt): check project ownership on every key mutation path /key/update ran the project check only when project_id, models or max_budget were supplied, so a request that changed team_id alone left a foreign project attached. /key/regenerate never ran it at all. Both now go through one helper that validates the key as the mutation leaves it. On /key/generate the check moves after default_key_generate_params is applied, because that can supply team_id; it was rejecting a valid key whose team came from the defaults. Tests move into the mapped test file per CLAUDE.md. --- .../key_management_endpoints.py | 85 +++-- .../test_key_management_endpoints.py | 318 ++++++++++++++++++ .../test_key_project_team_ownership.py | 115 ------- 3 files changed, 371 insertions(+), 147 deletions(-) delete mode 100644 tests/test_litellm/proxy/management_endpoints/test_key_project_team_ownership.py 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 From 3be888dec74883fcf6780585c5daa3e2d1ac9edb Mon Sep 17 00:00:00 2001 From: L4XB Date: Mon, 14 Sep 2026 19:58:35 +0200 Subject: [PATCH 03/20] test(key mgmt): read the project from the cache instead of patching the lookup The test-quality gate rejected the new cells: 12 TQ008 for patching `get_project_object`, an SDK internal, and 4 TQ001 for accept controls that could only fail by raising. The cells now seed `UserApiKeyCache` with the project, which is the idiom the neighbouring project cells already use and removes the patching. The accept controls are parametrised together with the rejecting ones, so each test function carries a real assertion. The three regenerate seams that stay stubbed each carry a reason. --- .../test_key_management_endpoints.py | 388 ++++++++---------- 1 file changed, 162 insertions(+), 226 deletions(-) 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 From 1f1dd295edce3180868e3bb1321acdf699185c35 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 09:26:15 +0000 Subject: [PATCH 04/20] fix(proxy): prevent cross-team project keys Co-authored-by: L4XB Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/project_endpoints.py | 18 + .../key_management_endpoints.py | 111 +++-- .../management/test_project_lifecycle.py | 210 ++++++++- .../test_project_endpoints_prisma.py | 53 +++ .../test_key_management_endpoints.py | 438 +++++++++++------- 5 files changed, 617 insertions(+), 213 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index d134c39c91b..652845f6268 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -767,6 +767,24 @@ async def update_project( detail={"error": "Cannot reassign project to a team you are not an admin of"}, ) + if data.team_id is not None and data.team_id != existing_project.team_id: + mismatched_key_count: Final = await _verification_token_table(prisma_client).count( + where={ + "project_id": data.project_id, + "OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}], + } + ) + if mismatched_key_count > 0: + raise HTTPException( + status_code=400, + detail={ + "error": ( + f"Project {data.project_id} has {mismatched_key_count} key(s) that do not belong to " + f"team {data.team_id}. Detach or delete them before moving the project." + ) + }, + ) + # Validate project limits against team limits if target_team_obj is not None: _check_team_project_limits( diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index e7fc9d761e0..60d0f292548 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1263,17 +1263,12 @@ 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( + await _check_key_project_team( 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, + prisma_client=prisma_client, + user_api_key_cache=proxy_server.user_api_key_cache, ) # Delegated-authority ceiling (GHSA-q775-qw9r-2r4g): a non-admin caller @@ -1772,14 +1767,10 @@ async def _check_project_key_limits( data: GenerateKeyRequest | UpdateKeyRequest, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, - key_team_id: str | None = None, ) -> None: """ - Validate that the key belongs to the project's team, and that its models - and budget respect the project's limits. + Validate that key's models and budget respect its project's limits. - - 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 @@ -1796,16 +1787,6 @@ async def _check_project_key_limits( detail={"error": f"Project not found, project_id={project_id}"}, ) - if project_obj.team_id is not None and project_obj.team_id != key_team_id: - raise HTTPException( - status_code=403, - detail={ - "error": f"Project {project_id} belongs to team {project_obj.team_id}, " - f"but the key belongs to {key_team_id if key_team_id is not None else 'no team'}. " - "A project can only be attached to keys of the team that owns it." - }, - ) - # Validate key models are a subset of project models if data.models and len(project_obj.models) > 0: for m in data.models: @@ -1831,30 +1812,61 @@ 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_key_project_team( + project_id: str, + key_team_id: str | None, + prisma_client: PrismaClient, + user_api_key_cache: UserApiKeyCache, +) -> None: + project_obj: Final = await get_project_object( + project_id=project_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + + if project_obj is None: + raise HTTPException( + status_code=404, + detail={"error": f"Project not found, project_id={project_id}"}, + ) + + if project_obj.team_id is None or project_obj.team_id == key_team_id: + return + + raise HTTPException( + status_code=400, + detail={ + "error": ( + f"Project {project_id} belongs to team {project_obj.team_id}, but the key belongs to " + f"{key_team_id if key_team_id is not None else 'no team'}. " + "A key can only be attached to a project owned by its own team." + ) + }, + ) -async def _check_project_key_limits_on_mutation( +async def _check_key_project_team_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: + fields_set: Final = data.model_fields_set + team_changed: Final = "team_id" in fields_set and data.team_id != existing_key_row.team_id + project_changed: Final = "project_id" in fields_set and data.project_id != existing_key_row.project_id + if not team_changed and not project_changed: return - project_id: Final = data.project_id if "project_id" in data.model_fields_set else existing_key_row.project_id + project_id: Final = data.project_id if "project_id" in fields_set else existing_key_row.project_id if project_id is None: return - await _check_project_key_limits( + team_id: Final = data.team_id if "team_id" in fields_set else existing_key_row.team_id + await _check_key_project_team( project_id=project_id, - data=data, + key_team_id=team_id, 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), ) @@ -2183,6 +2195,15 @@ 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, + ) + return await _common_key_generation_helper( data=data, user_api_key_dict=user_api_key_dict, @@ -3310,7 +3331,19 @@ async def _validate_update_key_data( access_group_ids=data.access_group_ids, ) - await _check_project_key_limits_on_mutation( + # 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 + ) + if _project_id_to_check is not None and (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, + ) + + await _check_key_project_team_on_mutation( data=data, existing_key_row=existing_key_row, prisma_client=checked_prisma_client, @@ -5609,6 +5642,12 @@ async def _execute_virtual_key_regeneration( ) if data is not None: + await _check_key_project_team_on_mutation( + data=data, + existing_key_row=key_in_db, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) _existing_key_metadata: Final = getattr(key_in_db, "metadata", None) enforce_output_token_estimates_are_admin_only( data=data, @@ -5643,12 +5682,6 @@ async def _execute_virtual_key_regeneration( await _enforce_custom_key_update_policy(hook=_custom_key_update_hook(proxy_server), data=update_request) # 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, prisma_client=prisma_client, llm_router=llm_router ) diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index 29a14b37ab9..867bb604162 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -1,9 +1,12 @@ +from collections.abc import Iterator from hashlib import sha256 from typing import Final +import httpx import pytest -from integration._support.client import Gateway, object_value, string_value -from integration._support.database import read_rows +from integration._support.client import JSON_OBJECT, Gateway, gateway_from_environment, object_value, string_value +from integration._support.database import read_rows, write_rows +from integration._support.process import owned_proxy from pydantic import JsonValue @@ -17,6 +20,27 @@ def _project_rows(project_id: str) -> list[dict[str, JsonValue]]: ) +def _key_rows(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT token, key_alias, team_id, project_id FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + +def _discard_unexpected_key(candidate: Gateway, response: httpx.Response) -> None: + if response.status_code != 200: + return + body: Final = JSON_OBJECT.validate_json(response.content) + candidate.post("/key/delete", {"keys": [string_value(body["key"])]}) + + +@pytest.fixture(scope="module") +def ownership_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + with gateway_from_environment() as gateway: + with owned_proxy(gateway, tmp_path_factory.mktemp("project-team-ownership"), {}, workers=2) as candidate: + yield candidate + + @pytest.mark.covers("mgmt.project.new.real_route_persists") def test_project_new_persists_real_state(gateway: Gateway) -> None: with gateway.scenario() as scenario: @@ -113,3 +137,185 @@ def test_project_delete_with_attached_key_refuses_and_preserves_state(gateway: G 'FROM "LiteLLM_VerificationToken" WHERE token = %s', (digest,), ) == key_before + + +def test_key_generate_rejects_foreign_team_project(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + cross_team: Final = ownership_gateway.request( + "POST", + "/key/generate", + {"team_id": team_a, "project_id": project_b, "models": [model]}, + ) + _discard_unexpected_key(ownership_gateway, cross_team) + assert cross_team.status_code == 400, cross_team.text + assert read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND team_id = %s', + (project_b, team_a), + ) == [] + + +def test_key_generate_rejects_missing_team_for_owned_project(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_b: Final = scenario.team(models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + unbound: Final = ownership_gateway.request( + "POST", "/key/generate", {"project_id": project_b, "models": [model]} + ) + _discard_unexpected_key(ownership_gateway, unbound) + assert unbound.status_code == 400, unbound.text + assert read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND team_id IS NULL', + (project_b,), + ) == [] + + +def test_key_generate_same_team_project_key_can_chat(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team: Final = scenario.team(models=[model]) + project: Final = scenario.project(team, models=[model]) + generated: Final = ownership_gateway.request( + "POST", "/key/generate", {"team_id": team, "project_id": project, "models": [model]} + ) + assert generated.status_code == 200, generated.text + key: Final = string_value(JSON_OBJECT.validate_json(generated.content)["key"]) + scenario.cleanups.callback(scenario.delete_key, key) + generated_rows: Final = _key_rows(key) + assert len(generated_rows) == 1 + assert generated_rows[0]["team_id"] == team + assert generated_rows[0]["project_id"] == project + chat: Final = ownership_gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "control"}]}, + key=key, + ) + assert chat.status_code == 200, chat.text + assert object_value(JSON_OBJECT.validate_json(chat.content)["usage"])["total_tokens"] == 40 + + +def test_key_update_rejects_team_change_and_allows_unchanged_values(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model]) + aliased: Final = ownership_gateway.request("POST", "/key/update", {"key": key, "key_alias": "updated"}) + assert aliased.status_code == 200, aliased.text + after_alias: Final = _key_rows(key) + assert len(after_alias) == 1 + assert after_alias[0]["key_alias"] == "updated" + assert after_alias[0]["team_id"] == team_a + assert after_alias[0]["project_id"] == project_a + unchanged_team: Final = ownership_gateway.request("POST", "/key/update", {"key": key, "team_id": team_a}) + assert unchanged_team.status_code == 200, unchanged_team.text + after_unchanged_team: Final = _key_rows(key) + assert after_unchanged_team == after_alias + before_reassignment: Final = _key_rows(key) + assert len(before_reassignment) == 1 + reassigned: Final = ownership_gateway.request("POST", "/key/update", {"key": key, "team_id": team_b}) + assert reassigned.status_code == 400, reassigned.text + assert _key_rows(key) == before_reassignment + + +def test_key_update_can_detach_project_and_change_team(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model]) + detached: Final = ownership_gateway.request( + "POST", "/key/update", {"key": key, "project_id": None, "team_id": team_b} + ) + assert detached.status_code == 200, detached.text + after_detach: Final = _key_rows(key) + assert len(after_detach) == 1 + assert after_detach[0]["project_id"] is None + assert after_detach[0]["team_id"] == team_b + assert after_detach[0]["key_alias"] is None + + +def test_key_regenerate_rejects_foreign_project_without_changing_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + generated: Final = ownership_gateway.request( + "POST", "/key/generate", {"team_id": team_a, "models": [model]} + ) + assert generated.status_code == 200, generated.text + key: Final = string_value(JSON_OBJECT.validate_json(generated.content)["key"]) + before: Final = _key_rows(key) + assert len(before) == 1 + response: Final = ownership_gateway.request( + "POST", f"/key/{key}/regenerate", {"project_id": project_b} + ) + _discard_unexpected_key(ownership_gateway, response) + if response.status_code != 200: + scenario.cleanups.callback(scenario.delete_key, key) + assert response.status_code == 400, response.text + assert _key_rows(key) == before + + +def test_project_update_rejects_moving_project_with_attached_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + attached_project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=attached_project, models=[model]) + project_before: Final = _project_rows(attached_project) + key_before: Final = _key_rows(key) + moved_with_key: Final = ownership_gateway.request( + "POST", "/project/update", {"project_id": attached_project, "team_id": team_b} + ) + assert moved_with_key.status_code == 400, moved_with_key.text + assert _project_rows(attached_project) == project_before + assert _key_rows(key) == key_before + + +def test_project_update_allows_moving_project_without_keys(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + unattached_project: Final = scenario.project(team_a, models=[model]) + moved_without_key: Final = ownership_gateway.request( + "POST", "/project/update", {"project_id": unattached_project, "team_id": team_b} + ) + assert moved_without_key.status_code == 200, moved_without_key.text + moved_project: Final = _project_rows(unattached_project) + assert len(moved_project) == 1 + assert moved_project[0]["team_id"] == team_b + assert read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (unattached_project,), + ) == [] + + +def test_project_update_rejects_moving_project_with_teamless_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project, models=[model]) + write_rows( + 'UPDATE "LiteLLM_VerificationToken" SET team_id = NULL WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + project_before: Final = _project_rows(project) + key_before: Final = _key_rows(key) + assert key_before[0]["team_id"] is None + moved: Final = ownership_gateway.request("POST", "/project/update", {"project_id": project, "team_id": team_b}) + assert moved.status_code == 400, moved.text + assert _project_rows(project) == project_before + assert _key_rows(key) == key_before diff --git a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py index 36878fa698c..83ad594c458 100644 --- a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py +++ b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py @@ -1,5 +1,7 @@ import os import traceback +from collections.abc import Mapping +from typing import Final from litellm._uuid import uuid from unittest import mock @@ -35,6 +37,7 @@ verbose_proxy_logger.setLevel(level=logging.DEBUG) from litellm.caching.caching import DualCache from litellm.proxy._types import ( + LiteLLM_TeamTable, NewProjectRequest, UpdateProjectRequest, DeleteProjectRequest, @@ -1256,6 +1259,56 @@ def _written_project_data(mock_prisma: mock.MagicMock) -> dict: return mock_prisma.db.litellm_projecttable.update.await_args.kwargs["data"] +@pytest.mark.asyncio +async def test_update_project_rejects_move_when_attached_teamless_key_exists( + monkeypatch: pytest.MonkeyPatch, +) -> None: + project_id: Final = "project-teamless-key" + destination_team_id: Final = "team-b" + mock_prisma: Final = _project_update_mocks(monkeypatch, {}) + mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = "team-a" + mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock( + return_value=LiteLLM_TeamTable(team_id=destination_team_id) + ) + + async def count_teamless_keys(*, where: Mapping[str, object]) -> int: + conditions: Final = where.get("OR") + return int(isinstance(conditions, list) and {"team_id": None} in conditions) + + mock_prisma.db.litellm_verificationtoken.count = mock.AsyncMock(side_effect=count_teamless_keys) + + with pytest.raises(ProxyException) as error: + await _run_project_update(project_id, team_id=destination_team_id) + + expected_detail: Final = { + "error": ( + f"Project {project_id} has 1 key(s) that do not belong to team {destination_team_id}. " + "Detach or delete them before moving the project." + ) + } + assert error.value.code == "400" + assert expected_detail["error"] in error.value.message + mock_prisma.db.litellm_verificationtoken.count.assert_awaited_once() + mock_prisma.db.litellm_projecttable.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_update_project_allows_move_when_no_attached_keys_exist(monkeypatch: pytest.MonkeyPatch) -> None: + project_id: Final = "project-without-keys" + destination_team_id: Final = "team-b" + mock_prisma: Final = _project_update_mocks(monkeypatch, {}) + mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = "team-a" + mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock( + return_value=LiteLLM_TeamTable(team_id=destination_team_id) + ) + mock_prisma.db.litellm_verificationtoken.count = mock.AsyncMock(return_value=0) + + await _run_project_update(project_id, team_id=destination_team_id) + + mock_prisma.db.litellm_verificationtoken.count.assert_awaited_once() + mock_prisma.db.litellm_projecttable.update.assert_awaited_once() + + @pytest.mark.asyncio async def test_update_project_clears_model_itpm_limit_sent_as_an_empty_map(monkeypatch): """ diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 2e7b595e317..b4284d3d7a9 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -45,9 +45,10 @@ from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, project_cache_key from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy.management_endpoints.key_management_endpoints import ( + _check_key_project_team, + _check_key_project_team_on_mutation, _check_org_key_limits, _check_project_key_limits, - _check_project_key_limits_on_mutation, _check_team_key_limits, _common_key_generation_helper, _effective_key_after_update, @@ -72,6 +73,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( delete_verification_tokens, generate_key_fn, generate_key_helper_fn, + generate_service_account_key_fn, key_aliases, key_generation_check, list_keys, @@ -20079,7 +20081,6 @@ async def test_check_project_key_limits_accepts_inherited_model_sentinels(reques data=request_cls(key="sk-lit-5823", models=[sentinel]), prisma_client=MagicMock(), user_api_key_cache=user_api_key_cache, - key_team_id="team-lit-5823", ) @@ -20098,7 +20099,6 @@ async def test_check_project_key_limits_still_rejects_real_model_outside_project data=request_cls(key="sk-lit-5823", models=key_models), prisma_client=MagicMock(), user_api_key_cache=user_api_key_cache, - key_team_id="team-lit-5823", ) assert exc_info.value.status_code == 400 @@ -20514,125 +20514,274 @@ async def test_key_creator_cannot_detach_project_without_admin_access(): 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 --- - _OWNED_PROJECT: Final = "proj-owned-1" +_OWNERSHIP_PROJECT_TEAM: Final = "ownership-project-team" +_OWNERSHIP_KEY_TEAM: Final = "ownership-key-team" +_OWNERSHIP_DESTINATION_TEAM: Final = "ownership-destination-team" -@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, []), - ], -) +def _make_generate_mock_prisma() -> AsyncMock: + mock_prisma_client: Final = 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) + ) + mock_prisma_client.db.litellm_teamtable = MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) + return mock_prisma_client + + +def _configure_key_endpoints( + monkeypatch: pytest.MonkeyPatch, + user_api_key_cache: UserApiKeyCache, +) -> AsyncMock: + mock_prisma_client: Final = _make_generate_mock_prisma() + mock_prisma_client.writer_db = mock_prisma_client.db + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache) + return mock_prisma_client + + +@pytest.mark.parametrize("key_team_id", ["team-a", 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 -): +async def test_key_generation_rejects_foreign_project_team( + monkeypatch: pytest.MonkeyPatch, + key_team_id: str | None, +) -> None: + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") + _configure_key_endpoints(monkeypatch, user_api_key_cache) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=GenerateKeyRequest(project_id=_OWNED_PROJECT, team_id=key_team_id), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + team_table=None, + ) + + expected_detail: Final = { + "error": ( + f"Project {_OWNED_PROJECT} belongs to team team-b, but the key belongs to " + f"{key_team_id if key_team_id is not None else 'no team'}. " + "A key can only be attached to a project owned by its own team." + ) + } + assert error.value.status_code == 400 + assert error.value.detail == expected_detail + + +@pytest.mark.parametrize("project_team_id", ["team-a", None]) +@pytest.mark.asyncio +async def test_key_generation_accepts_same_team_and_unowned_projects( + monkeypatch: pytest.MonkeyPatch, + project_team_id: str | None, +) -> None: user_api_key_cache: Final = await _cache_with_project( _OWNED_PROJECT, [], team_id=project_team_id ) + _configure_key_endpoints(monkeypatch, user_api_key_cache) + monkeypatch.setattr(litellm, "default_key_generate_params", None) - async def check(): - await _check_project_key_limits( - project_id=_OWNED_PROJECT, - data=GenerateKeyRequest(team_id=key_team_id), - prisma_client=MagicMock(), - user_api_key_cache=user_api_key_cache, - key_team_id=key_team_id, - ) + response: Final = await _common_key_generation_helper( + data=GenerateKeyRequest(project_id=_OWNED_PROJECT, team_id="team-a"), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + team_table=None, + ) - 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)) + assert response.team_id == "team-a" + + +@pytest.mark.parametrize(("project_team_id", "expected_status"), [("team-b", None), ("team-c", 400)]) +@pytest.mark.asyncio +async def test_key_generation_uses_default_team_for_project_ownership( + monkeypatch: pytest.MonkeyPatch, + project_team_id: str, + expected_status: int | None, +) -> None: + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=project_team_id) + _configure_key_endpoints(monkeypatch, user_api_key_cache) + monkeypatch.setattr(litellm, "default_key_generate_params", {"team_id": "team-b"}) + data: Final = GenerateKeyRequest(project_id=_OWNED_PROJECT) + + if expected_status is not None: + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=data, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + team_table=None, + ) + assert error.value.status_code == expected_status + assert "belongs to team team-c, but the key belongs to team-b" in str(error.value.detail) return - with pytest.raises(HTTPException) as exc: - await check() - assert exc.value.status_code == expected_status - for fragment in expected_in_detail: - assert fragment in str(exc.value.detail) + response: Final = await _common_key_generation_helper( + data=data, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + team_table=None, + ) + assert response.team_id == "team-b" @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" +async def test_service_account_generation_rejects_foreign_project_team(monkeypatch: pytest.MonkeyPatch) -> None: + team_id: Final = "service-account-team" + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=LiteLLM_TeamTable(team_id=team_id) ) - with pytest.raises(HTTPException) as exc: - 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", + with pytest.raises(HTTPException) as error: + await generate_service_account_key_fn( + data=GenerateKeyRequest(team_id=team_id, project_id=_OWNED_PROJECT), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, ) - assert exc.value.status_code == 403 + assert error.value.status_code == 400 + assert "belongs to team team-b, but the key belongs to service-account-team" in str( + error.value.detail + ) + + +@pytest.mark.asyncio +async def test_key_generation_default_budget_does_not_reject_project_budget(monkeypatch: pytest.MonkeyPatch) -> None: + user_api_key_cache: Final = UserApiKeyCache() + project: Final = LiteLLM_ProjectTableCachedObj( + project_id=_OWNED_PROJECT, + team_id="team-a", + models=[], + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ) + await user_api_key_cache.async_set_cache( + key=project_cache_key(_OWNED_PROJECT), + value=project, + model_type=LiteLLM_ProjectTableCachedObj, + ) + _configure_key_endpoints(monkeypatch, user_api_key_cache) + monkeypatch.setattr(litellm, "default_key_generate_params", {"team_id": "team-a", "max_budget": 10.0}) + + response: Final = await generate_key_fn( + data=GenerateKeyRequest(project_id=_OWNED_PROJECT), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + ) + + assert response.team_id == "team-a" @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), - ], + "request_fields", + [{"key_alias": "renamed"}, {"team_id": _OWNERSHIP_KEY_TEAM}], ) @pytest.mark.asyncio -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=existing_team_id, project_id=_OWNED_PROJECT +async def test_key_update_allows_legacy_project_mismatch_when_team_is_unchanged( + monkeypatch: pytest.MonkeyPatch, + request_fields: dict[str, str], +) -> None: + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM) + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_KEY_TEAM) + ) + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-key", + team_id=_OWNERSHIP_KEY_TEAM, + project_id=_OWNED_PROJECT, ) - user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") - async def check(): - await _check_project_key_limits_on_mutation( - data=UpdateKeyRequest(key="sk-x", **data_kwargs), - existing_key_row=existing, - prisma_client=MagicMock(), + await _validate_update_key_data( + data=UpdateKeyRequest(key="sk-key", **request_fields), + existing_key_row=existing_key_row, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + llm_router=None, + premium_user=True, + prisma_client=mock_prisma_client, + user_api_key_cache=user_api_key_cache, + ) + + +@pytest.mark.asyncio +async def test_key_update_rejects_team_change_for_project_bound_key(monkeypatch: pytest.MonkeyPatch) -> None: + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM + ) + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM) + ) + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-key", + team_id=_OWNERSHIP_KEY_TEAM, + project_id=_OWNED_PROJECT, + ) + + with pytest.raises(HTTPException) as error: + await _validate_update_key_data( + data=UpdateKeyRequest(key="sk-key", team_id=_OWNERSHIP_DESTINATION_TEAM), + existing_key_row=existing_key_row, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + llm_router=None, + premium_user=True, + prisma_client=mock_prisma_client, user_api_key_cache=user_api_key_cache, ) - if expected_status is None: - await check() - assert await user_api_key_cache.async_get_cache(key=project_cache_key(_OWNED_PROJECT)) - return + assert error.value.status_code == 400 + expected_detail: Final = ( + f"Project {_OWNED_PROJECT} belongs to team {_OWNERSHIP_PROJECT_TEAM}, " + f"but the key belongs to {_OWNERSHIP_DESTINATION_TEAM}" + ) + assert expected_detail in str(error.value.detail) - with pytest.raises(HTTPException) as exc: - await check() - assert exc.value.status_code == expected_status + +@pytest.mark.parametrize( + "request_fields", + [{"key_alias": "renamed"}, {"team_id": "team-a"}], +) +@pytest.mark.asyncio +async def test_key_team_ownership_mutation_allows_legacy_mismatch_without_changes( + request_fields: dict[str, str], +) -> None: + prisma_client: Final = MagicMock() + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-key", + team_id="team-a", + project_id=_OWNED_PROJECT, + ) + + await _check_key_project_team_on_mutation( + data=UpdateKeyRequest(key="sk-key", **request_fields), + existing_key_row=existing_key_row, + prisma_client=prisma_client, + user_api_key_cache=UserApiKeyCache(), + ) + + assert prisma_client.mock_calls == [] @pytest.mark.asyncio -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 - ) +async def test_key_team_ownership_mutation_allows_detach_with_team_change() -> None: prisma_client: Final = MagicMock() + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-key", + team_id="team-a", + project_id=_OWNED_PROJECT, + ) - # 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, + await _check_key_project_team_on_mutation( + data=UpdateKeyRequest(key="sk-key", project_id=None, team_id="team-b"), + existing_key_row=existing_key_row, prisma_client=prisma_client, user_api_key_cache=UserApiKeyCache(), ) @@ -20641,39 +20790,31 @@ async def test_an_update_that_touches_none_of_the_fields_does_not_read_the_proje @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), - ], + ("key_team_id", "expected_status", "expected_updates"), + [("team-a", 400, 0), ("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, - ) - +async def test_regenerate_checks_project_team_ownership( + key_team_id: str, + expected_status: int | None, + expected_updates: int, +) -> None: 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") - async def regenerate(): + async def regenerate() -> None: 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 + patch( "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 + patch( "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 + patch( "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", new_callable=AsyncMock, ), @@ -20690,81 +20831,34 @@ async def test_regenerate_checks_the_project_it_attaches( proxy_logging_obj=MagicMock(), ) - if expected_status is None: - await regenerate() - else: - with pytest.raises(HTTPException) as exc: + if expected_status is not None: + with pytest.raises(HTTPException) as error: await regenerate() - assert exc.value.status_code == expected_status + assert error.value.status_code == expected_status + assert "belongs to team team-b, but the key belongs to team-a" in str(error.value.detail) + else: + await regenerate() - # A refused regenerate must not reach the DB update. - assert ( - mock_prisma_client.db.litellm_verificationtoken.update.await_count == expected_updates - ) + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == expected_updates -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 - - -@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 -): +async def test_key_project_team_validation_uses_project_missing_404(monkeypatch: pytest.MonkeyPatch) -> None: + project_lookup: Final = AsyncMock(return_value=None) monkeypatch.setattr( - "litellm.proxy.proxy_server.prisma_client", _make_generate_mock_prisma() + "litellm.proxy.management_endpoints.key_management_endpoints.get_project_object", + project_lookup, ) - 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"}) - async def generate(): - return await _common_key_generation_helper( - data=GenerateKeyRequest(project_id=_OWNED_PROJECT), - user_api_key_dict=UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234", user_id="1234" - ), - litellm_changed_by=None, - team_table=None, + with pytest.raises(HTTPException) as error: + await _check_key_project_team( + project_id=_OWNED_PROJECT, + key_team_id="team-a", + prisma_client=MagicMock(), + user_api_key_cache=UserApiKeyCache(), ) - if expected_status is None: - response = await generate() - assert response.team_id == "team-b" - return - - with pytest.raises(HTTPException) as exc: - await generate() - assert exc.value.status_code == expected_status + assert error.value.status_code == 404 @pytest.mark.asyncio async def test_bulk_update_team_keys_runs_custom_key_policy_per_key(monkeypatch): From 49c6c070fbb4c304263cc6df800fe1159d2504b7 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 09:52:09 +0000 Subject: [PATCH 05/20] fix(proxy): validate bulk key team changes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../key_management_endpoints.py | 40 +++++++-- .../management/test_project_lifecycle.py | 45 +++++++++- .../test_key_management_endpoints.py | 89 +++++++++++++++++++ 3 files changed, 163 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 60d0f292548..0f9965f5629 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2761,6 +2761,23 @@ def _resolve_token_to_update(data: UpdateKeyRequest, existing_key_row: LiteLLM_V return existing_key_row.token +async def _check_single_key_update_team_permissions( + user_api_key_dict: UserAPIKeyAuth, + prisma_client: PrismaClient | None, + existing_key_row: LiteLLM_VerificationToken, + user_api_key_cache: UserApiKeyCache, +) -> None: + if prisma_client is None: + return + await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( + user_api_key_dict=user_api_key_dict, + route=KeyManagementRoutes.KEY_UPDATE, + prisma_client=prisma_client, + existing_key_row=existing_key_row, + user_api_key_cache=user_api_key_cache, + ) + + async def _process_single_key_update( update_key_request: UpdateKeyRequest, user_api_key_dict: UserAPIKeyAuth, @@ -2824,15 +2841,12 @@ async def _process_single_key_update( entity="key", ) - # Check team member permissions - if prisma_client is not None: - await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( - user_api_key_dict=user_api_key_dict, - route=KeyManagementRoutes.KEY_UPDATE, - prisma_client=prisma_client, - existing_key_row=existing_key_row, - user_api_key_cache=user_api_key_cache, - ) + await _check_single_key_update_team_permissions( + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + existing_key_row=existing_key_row, + user_api_key_cache=user_api_key_cache, + ) # Custom key update hook if user_custom_key_update is not None: @@ -2886,6 +2900,14 @@ async def _process_single_key_update( llm_router=llm_router, ) + if prisma_client is not None: + await _check_key_project_team_on_mutation( + data=update_key_request, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + key_request: Final = await _with_validated_object_permission( update_key_request=update_key_request, team_obj=team_obj, diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index 867bb604162..b13ed839fc1 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -253,18 +253,59 @@ def test_key_regenerate_rejects_foreign_project_without_changing_key(ownership_g ) assert generated.status_code == 200, generated.text key: Final = string_value(JSON_OBJECT.validate_json(generated.content)["key"]) + scenario.cleanups.callback(scenario.delete_key, key) before: Final = _key_rows(key) assert len(before) == 1 response: Final = ownership_gateway.request( "POST", f"/key/{key}/regenerate", {"project_id": project_b} ) _discard_unexpected_key(ownership_gateway, response) - if response.status_code != 200: - scenario.cleanups.callback(scenario.delete_key, key) assert response.status_code == 400, response.text assert _key_rows(key) == before +def test_key_bulk_update_rejects_foreign_team_project_and_preserves_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model]) + before: Final = _key_rows(key) + assert len(before) == 1 + assert before[0]["team_id"] == team_a + assert before[0]["project_id"] == project_a + + response: Final = ownership_gateway.request( + "POST", + "/key/bulk_update", + { + "keys": [ + {"key": key, "team_id": team_b}, + {"key": key, "max_budget": 10, "tags": ["bulk-update"]}, + ] + }, + ) + assert response.status_code == 200, response.text + body: Final = JSON_OBJECT.validate_json(response.content) + failed_updates: Final = body["failed_updates"] + successful_updates: Final = body["successful_updates"] + assert isinstance(failed_updates, list), response.text + assert isinstance(successful_updates, list), response.text + assert len(failed_updates) == 1 + assert len(successful_updates) == 1 + failed_update: Final = object_value(failed_updates[0]) + successful_update: Final = object_value(successful_updates[0]) + assert string_value(failed_update["key"]) == key + assert f"Project {project_a} belongs to team {team_a}" in string_value(failed_update["failed_reason"]) + assert string_value(successful_update["key"]) == key + + after: Final = _key_rows(key) + assert len(after) == 1 + assert after[0]["team_id"] == before[0]["team_id"] + assert after[0]["project_id"] == before[0]["project_id"] + + def test_project_update_rejects_moving_project_with_attached_key(ownership_gateway: Gateway) -> None: with ownership_gateway.scenario() as scenario: model: Final = scenario.model() diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index b4284d3d7a9..df4f159f156 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -20745,6 +20745,95 @@ async def test_key_update_rejects_team_change_for_project_bound_key(monkeypatch: assert expected_detail in str(error.value.detail) +@pytest.mark.asyncio +async def test_bulk_key_update_rejects_project_team_change_and_allows_other_fields( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy.management_endpoints.key_management_endpoints import bulk_update_keys + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateKeyRequest, + BulkUpdateKeyRequestItem, + ) + + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM + ) + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-bulk-key", + user_id=None, + models=[], + team_id=_OWNERSHIP_PROJECT_TEAM, + project_id=_OWNED_PROJECT, + ) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing_key_row) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM) + ) + mock_prisma_client.get_data = AsyncMock(return_value=existing_key_row) + updated_key: Final = MagicMock() + updated_key.model_dump.return_value = { + "max_budget": 10.0, + "tags": ["bulk-update"], + "team_id": _OWNERSHIP_PROJECT_TEAM, + "project_id": _OWNED_PROJECT, + } + mock_prisma_client.update_data = AsyncMock(return_value={"data": updated_key}) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + + request: Final = BulkUpdateKeyRequest( + keys=[ + BulkUpdateKeyRequestItem(key="sk-bulk-key", team_id=_OWNERSHIP_DESTINATION_TEAM), + BulkUpdateKeyRequestItem(key="sk-bulk-key", max_budget=10.0, tags=["bulk-update"]), + ] + ) + user_api_key_dict: Final = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin", + ) + + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + new_callable=AsyncMock, + return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM), + ), + patch("litellm.proxy.management_endpoints.common_utils._premium_user_check"), + patch("litellm.proxy.utils._premium_user_check"), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook", + new_callable=AsyncMock, + ), + ): + response: Final = await bulk_update_keys( + data=request, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + assert response.total_requested == 2 + assert [(item.key, item.failed_reason) for item in response.failed_updates] == [ + ( + "sk-bulk-key", + ( + f"Project {_OWNED_PROJECT} belongs to team {_OWNERSHIP_PROJECT_TEAM}, but the key belongs to " + f"{_OWNERSHIP_DESTINATION_TEAM}. A key can only be attached to a project owned by its own team." + ), + ) + ] + assert len(response.successful_updates) == 1 + assert response.successful_updates[0].key == "sk-bulk-key" + assert response.successful_updates[0].key_info["max_budget"] == 10.0 + assert response.successful_updates[0].key_info["tags"] == ["bulk-update"] + mock_prisma_client.update_data.assert_awaited_once() + + @pytest.mark.parametrize( "request_fields", [{"key_alias": "renamed"}, {"team_id": "team-a"}], From 1049ac4bffa87847d3ba8fc1b9fd41fc2affd1e9 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 10:30:18 +0000 Subject: [PATCH 06/20] test(proxy): use behavioral assertions in project ownership tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_key_management_endpoints.py | 46 +++++++++++-------- 1 file changed, 26 insertions(+), 20 deletions(-) diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index df4f159f156..0337a2bb290 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -16,6 +16,7 @@ from fastapi import HTTPException import inspect +from litellm.models.project import LiteLLM_ProjectTable from litellm.proxy._types import ( GenerateKeyRequest, KeyManagementRoutes, @@ -20474,11 +20475,10 @@ async def test_project_detachment_preserves_omission_and_other_key_fields(): @pytest.mark.parametrize("project_id", [None, "project-orbit", "project-other", ""]) @pytest.mark.asyncio -async def test_project_detachment_uses_effective_project_for_validation(project_id: str | None): +async def test_project_detachment_uses_effective_project_for_validation_on_unowned_project( + project_id: str | None, +): existing: Final = LiteLLM_VerificationToken(token="project-detach-token", project_id="project-orbit") - # An unowned project: this cell is about WHICH project the validation uses, - # so the ownership gate (#41089) must not be what it measures. Giving the - # key a team instead would pull the whole team lookup into a MagicMock db. cache: Final = await _cache_with_project("project-orbit", ["model-orbit"], team_id=None) data: Final = UpdateKeyRequest(key=existing.token, project_id=project_id, models=["model-other"]) if project_id is None: @@ -20807,8 +20807,8 @@ async def test_bulk_key_update_rejects_project_team_change_and_allows_other_fiel new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook", - new_callable=AsyncMock, + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks", + SimpleNamespace(async_key_updated_hook=AsyncMock()), ), ): response: Final = await bulk_update_keys( @@ -20836,16 +20836,22 @@ async def test_bulk_key_update_rejects_project_team_change_and_allows_other_fiel @pytest.mark.parametrize( "request_fields", - [{"key_alias": "renamed"}, {"team_id": "team-a"}], + [{"key_alias": "renamed"}, {"team_id": _OWNERSHIP_KEY_TEAM}], ) @pytest.mark.asyncio async def test_key_team_ownership_mutation_allows_legacy_mismatch_without_changes( request_fields: dict[str, str], ) -> None: prisma_client: Final = MagicMock() + prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable( + project_id=_OWNED_PROJECT, + team_id=_OWNERSHIP_PROJECT_TEAM, + ) + ) existing_key_row: Final = LiteLLM_VerificationToken( token="hashed-key", - team_id="team-a", + team_id=_OWNERSHIP_KEY_TEAM, project_id=_OWNED_PROJECT, ) @@ -20856,15 +20862,19 @@ async def test_key_team_ownership_mutation_allows_legacy_mismatch_without_change user_api_key_cache=UserApiKeyCache(), ) - assert prisma_client.mock_calls == [] - @pytest.mark.asyncio async def test_key_team_ownership_mutation_allows_detach_with_team_change() -> None: prisma_client: Final = MagicMock() + prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable( + project_id=_OWNED_PROJECT, + team_id=_OWNERSHIP_PROJECT_TEAM, + ) + ) existing_key_row: Final = LiteLLM_VerificationToken( token="hashed-key", - team_id="team-a", + team_id=_OWNERSHIP_KEY_TEAM, project_id=_OWNED_PROJECT, ) @@ -20875,8 +20885,6 @@ async def test_key_team_ownership_mutation_allows_detach_with_team_change() -> N user_api_key_cache=UserApiKeyCache(), ) - assert prisma_client.mock_calls == [] - @pytest.mark.parametrize( ("key_team_id", "expected_status", "expected_updates"), @@ -20932,23 +20940,21 @@ async def test_regenerate_checks_project_team_ownership( @pytest.mark.asyncio -async def test_key_project_team_validation_uses_project_missing_404(monkeypatch: pytest.MonkeyPatch) -> None: - project_lookup: Final = AsyncMock(return_value=None) - monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints.get_project_object", - project_lookup, - ) +async def test_key_project_team_validation_uses_project_missing_404() -> None: + prisma_client: Final = MagicMock() + prisma_client.db.litellm_projecttable.find_unique = AsyncMock(return_value=None) with pytest.raises(HTTPException) as error: await _check_key_project_team( project_id=_OWNED_PROJECT, key_team_id="team-a", - prisma_client=MagicMock(), + prisma_client=prisma_client, user_api_key_cache=UserApiKeyCache(), ) assert error.value.status_code == 404 + @pytest.mark.asyncio async def test_bulk_update_team_keys_runs_custom_key_policy_per_key(monkeypatch): from litellm.types.proxy.management_endpoints.key_management_endpoints import ( From 7e55e90847801ed141c21268864d87c2eb8cce93 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 17:15:07 +0000 Subject: [PATCH 07/20] fix(proxy): read project owner from the database for key ownership checks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../key_management_endpoints.py | 14 +--- .../test_key_management_endpoints.py | 67 ++++++++++++++++++- 2 files changed, 66 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 0f9965f5629..88cda0c526e 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -139,6 +139,7 @@ from litellm.repositories.config_repository import ConfigParam, ConfigRepository from litellm.repositories.credentials_repository import CredentialsRepository from litellm.repositories.model_repository import ModelRepository from litellm.repositories.prisma_protocols import TableActions +from litellm.repositories.project_repository import ProjectRepository from litellm.repositories.table_repositories import ( DeletedVerificationTokenRepository, DeprecatedVerificationTokenRepository, @@ -1268,7 +1269,6 @@ async def _common_key_generation_helper( project_id=data.project_id, key_team_id=data.team_id, prisma_client=prisma_client, - user_api_key_cache=proxy_server.user_api_key_cache, ) # Delegated-authority ceiling (GHSA-q775-qw9r-2r4g): a non-admin caller @@ -1816,13 +1816,8 @@ async def _check_key_project_team( project_id: str, key_team_id: str | None, prisma_client: PrismaClient, - user_api_key_cache: UserApiKeyCache, ) -> None: - project_obj: Final = await get_project_object( - project_id=project_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - ) + project_obj: Final = await ProjectRepository(prisma_client).find_by_id(project_id) if project_obj is None: raise HTTPException( @@ -1849,7 +1844,6 @@ async def _check_key_project_team_on_mutation( data: UpdateKeyRequest | RegenerateKeyRequest, existing_key_row: LiteLLM_VerificationToken, prisma_client: PrismaClient, - user_api_key_cache: UserApiKeyCache, ) -> None: fields_set: Final = data.model_fields_set team_changed: Final = "team_id" in fields_set and data.team_id != existing_key_row.team_id @@ -1866,7 +1860,6 @@ async def _check_key_project_team_on_mutation( project_id=project_id, key_team_id=team_id, prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, ) @@ -2905,7 +2898,6 @@ async def _process_single_key_update( data=update_key_request, existing_key_row=existing_key_row, prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, ) key_request: Final = await _with_validated_object_permission( @@ -3369,7 +3361,6 @@ async def _validate_update_key_data( data=data, existing_key_row=existing_key_row, prisma_client=checked_prisma_client, - user_api_key_cache=user_api_key_cache, ) # When the caller asks to change the key's organization_id, require that @@ -5668,7 +5659,6 @@ async def _execute_virtual_key_regeneration( data=data, existing_key_row=key_in_db, prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, ) _existing_key_metadata: Final = getattr(key_in_db, "metadata", None) enforce_output_token_estimates_are_admin_only( diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 0337a2bb290..42b5652144a 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -20544,6 +20544,18 @@ def _configure_key_endpoints( ) -> AsyncMock: mock_prisma_client: Final = _make_generate_mock_prisma() mock_prisma_client.writer_db = mock_prisma_client.db + project_obj: Final = user_api_key_cache.get_cache( + key=project_cache_key(_OWNED_PROJECT), + model_type=LiteLLM_ProjectTableCachedObj, + ) + mock_prisma_client.db.litellm_projecttable = MagicMock() + mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + return_value=( + LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=project_obj.team_id) + if project_obj is not None + else None + ) + ) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache) return mock_prisma_client @@ -20578,6 +20590,54 @@ async def test_key_generation_rejects_foreign_project_team( assert error.value.detail == expected_detail +@pytest.mark.parametrize( + ("key_team_id", "expected_status"), + [(_OWNERSHIP_PROJECT_TEAM, 200), (_OWNERSHIP_KEY_TEAM, 400)], +) +@pytest.mark.asyncio +async def test_key_generation_uses_database_project_team_when_cache_is_stale( + monkeypatch: pytest.MonkeyPatch, + key_team_id: str, + expected_status: int, +) -> None: + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_KEY_TEAM) + prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + prisma_client.db.litellm_projecttable = MagicMock() + prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) + ) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + data: Final = GenerateKeyRequest(project_id=_OWNED_PROJECT, team_id=key_team_id) + user_api_key_dict: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin") + + if expected_status == 200: + response: Final = await _common_key_generation_helper( + data=data, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + team_table=None, + ) + assert response.team_id == _OWNERSHIP_PROJECT_TEAM + return + + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=data, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + team_table=None, + ) + + expected_detail: Final = { + "error": ( + f"Project {_OWNED_PROJECT} belongs to team {_OWNERSHIP_PROJECT_TEAM}, but the key belongs to " + f"{_OWNERSHIP_KEY_TEAM}. A key can only be attached to a project owned by its own team." + ) + } + assert error.value.status_code == 400 + assert error.value.detail == expected_detail + + @pytest.mark.parametrize("project_team_id", ["team-a", None]) @pytest.mark.asyncio async def test_key_generation_accepts_same_team_and_unowned_projects( @@ -20859,7 +20919,6 @@ async def test_key_team_ownership_mutation_allows_legacy_mismatch_without_change data=UpdateKeyRequest(key="sk-key", **request_fields), existing_key_row=existing_key_row, prisma_client=prisma_client, - user_api_key_cache=UserApiKeyCache(), ) @@ -20882,7 +20941,6 @@ async def test_key_team_ownership_mutation_allows_detach_with_team_change() -> N data=UpdateKeyRequest(key="sk-key", project_id=None, team_id="team-b"), existing_key_row=existing_key_row, prisma_client=prisma_client, - user_api_key_cache=UserApiKeyCache(), ) @@ -20898,6 +20956,10 @@ async def test_regenerate_checks_project_team_ownership( ) -> None: existing_key: Final = LiteLLM_VerificationToken(token="abc123", team_id=key_team_id) mock_prisma_client: Final = _make_regenerate_mock_prisma() + mock_prisma_client.db.litellm_projecttable = MagicMock() + mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id="team-b") + ) user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") async def regenerate() -> None: @@ -20949,7 +21011,6 @@ async def test_key_project_team_validation_uses_project_missing_404() -> None: project_id=_OWNED_PROJECT, key_team_id="team-a", prisma_client=prisma_client, - user_api_key_cache=UserApiKeyCache(), ) assert error.value.status_code == 404 From 4964d55e6bd3fbf8d4d9f585b1bff6bc1e31f69f Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 18:19:01 +0000 Subject: [PATCH 08/20] fix(proxy): read project ownership from the primary database under the lookup deadline Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/project_endpoints.py | 14 ++-- .../key_management_endpoints.py | 15 ++++- .../test_project_endpoints_prisma.py | 8 ++- .../test_key_management_endpoints.py | 65 +++++++++++++++---- 4 files changed, 81 insertions(+), 21 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index 652845f6268..db36500935c 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -22,6 +22,7 @@ from litellm._uuid import uuid from litellm.proxy._types import * from litellm.proxy.auth.auth_checks import delete_cached_project_object from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.db.db_lookup_gate import bounded_db_lookup from litellm.proxy.management.teams.access import is_team_admin from litellm.proxy.management_endpoints.common_utils import _set_object_metadata_field from litellm.proxy.management_endpoints.team_admin_field_permissions import team_admin_may_manage_projects @@ -768,11 +769,14 @@ async def update_project( ) if data.team_id is not None and data.team_id != existing_project.team_id: - mismatched_key_count: Final = await _verification_token_table(prisma_client).count( - where={ - "project_id": data.project_id, - "OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}], - } + mismatched_key_count: Final = await bounded_db_lookup( + prisma_client.writer_db.litellm_verificationtoken.count( + where={ + "project_id": data.project_id, + "OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}], + } + ), + name="project_key_ownership", ) if mismatched_key_count > 0: raise HTTPException( diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 88cda0c526e..4b678a48eeb 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -42,6 +42,7 @@ from litellm.constants import ( from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.models.credentials import CredentialItem +from litellm.models.project import LiteLLM_ProjectTable from litellm.proxy._experimental.mcp_server.db import ( rotate_mcp_server_credentials_master_key, rotate_mcp_user_credentials_master_key, @@ -83,6 +84,7 @@ from litellm.proxy.common_utils.config_sync_pubsub import ( from litellm.proxy.common_utils.rbac_utils import check_org_admin_can_generate_keys from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.db.db_lookup_gate import bounded_db_lookup from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks from litellm.proxy.hooks.model_max_budget_limiter import build_model_max_budget_usage from litellm.proxy.management.teams.access import TEAM_ADMIN_ONLY, TEAM_OR_ORG_ADMIN, is_team_admin @@ -133,13 +135,12 @@ from litellm.proxy.utils import ( handle_exception_on_proxy, is_valid_api_key, ) -from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.base_repository import BaseRepository, record_to_dict from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.config_repository import ConfigParam, ConfigRepository from litellm.repositories.credentials_repository import CredentialsRepository from litellm.repositories.model_repository import ModelRepository from litellm.repositories.prisma_protocols import TableActions -from litellm.repositories.project_repository import ProjectRepository from litellm.repositories.table_repositories import ( DeletedVerificationTokenRepository, DeprecatedVerificationTokenRepository, @@ -1817,7 +1818,15 @@ async def _check_key_project_team( key_team_id: str | None, prisma_client: PrismaClient, ) -> None: - project_obj: Final = await ProjectRepository(prisma_client).find_by_id(project_id) + project_record: Final = await bounded_db_lookup( + prisma_client.writer_db.litellm_projecttable.find_unique(where={"project_id": project_id}), + name="project", + ) + project_obj: Final = ( + LiteLLM_ProjectTable.model_validate(record_to_dict(project_record)) + if project_record is not None + else None + ) if project_obj is None: raise HTTPException( diff --git a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py index 83ad594c458..b5046244c31 100644 --- a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py +++ b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py @@ -1236,6 +1236,7 @@ def _project_update_mocks(monkeypatch, stored_metadata: dict) -> mock.MagicMock: mock_prisma.jsonify_object = lambda data: data mock_prisma.db.litellm_projecttable.find_unique = mock.AsyncMock(return_value=existing_row) mock_prisma.db.litellm_projecttable.update = mock.AsyncMock(return_value=mock.MagicMock()) + mock_prisma.writer_db = mock_prisma.db monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", True) monkeypatch.setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma) @@ -1270,12 +1271,14 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists( mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock( return_value=LiteLLM_TeamTable(team_id=destination_team_id) ) + mock_prisma.db.litellm_verificationtoken.count = mock.AsyncMock(return_value=0) + mock_prisma.writer_db = mock.MagicMock() async def count_teamless_keys(*, where: Mapping[str, object]) -> int: conditions: Final = where.get("OR") return int(isinstance(conditions, list) and {"team_id": None} in conditions) - mock_prisma.db.litellm_verificationtoken.count = mock.AsyncMock(side_effect=count_teamless_keys) + mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(side_effect=count_teamless_keys) with pytest.raises(ProxyException) as error: await _run_project_update(project_id, team_id=destination_team_id) @@ -1288,7 +1291,8 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists( } assert error.value.code == "400" assert expected_detail["error"] in error.value.message - mock_prisma.db.litellm_verificationtoken.count.assert_awaited_once() + mock_prisma.writer_db.litellm_verificationtoken.count.assert_awaited_once() + mock_prisma.db.litellm_verificationtoken.count.assert_not_awaited() mock_prisma.db.litellm_projecttable.update.assert_not_awaited() diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 42b5652144a..8a3e31f1644 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -1,3 +1,4 @@ +import asyncio from collections.abc import Mapping from contextlib import ExitStack from typing import Final @@ -44,6 +45,7 @@ from litellm.proxy.auth.auth_checks import ( ) from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, project_cache_key +from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy.management_endpoints.key_management_endpoints import ( _check_key_project_team, @@ -20543,18 +20545,24 @@ def _configure_key_endpoints( user_api_key_cache: UserApiKeyCache, ) -> AsyncMock: mock_prisma_client: Final = _make_generate_mock_prisma() - mock_prisma_client.writer_db = mock_prisma_client.db project_obj: Final = user_api_key_cache.get_cache( key=project_cache_key(_OWNED_PROJECT), model_type=LiteLLM_ProjectTableCachedObj, ) + project_row: Final = ( + LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=project_obj.team_id) + if project_obj is not None + else None + ) mock_prisma_client.db.litellm_projecttable = MagicMock() mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock( - return_value=( - LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=project_obj.team_id) - if project_obj is not None - else None - ) + return_value=project_row + ) + mock_prisma_client.writer_db = MagicMock() + mock_prisma_client.writer_db.litellm_teamtable = mock_prisma_client.db.litellm_teamtable + mock_prisma_client.writer_db.litellm_projecttable = MagicMock() + mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( + return_value=project_row ) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache) @@ -20604,6 +20612,10 @@ async def test_key_generation_uses_database_project_team_when_cache_is_stale( prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) prisma_client.db.litellm_projecttable = MagicMock() prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_KEY_TEAM) + ) + prisma_client.writer_db.litellm_projecttable = MagicMock() + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) ) monkeypatch.setattr(litellm, "default_key_generate_params", None) @@ -20638,6 +20650,36 @@ async def test_key_generation_uses_database_project_team_when_cache_is_stale( assert error.value.detail == expected_detail +@pytest.mark.asyncio +async def test_key_generation_fails_when_writer_project_lookup_exceeds_deadline( + monkeypatch: pytest.MonkeyPatch, +) -> None: + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM) + prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + stalled_lookup: Final = asyncio.Event() + monkeypatch.setattr("litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.05) + + async def never_returns_project(*, where: Mapping[str, object]) -> None: + assert where == {"project_id": _OWNED_PROJECT} + await stalled_lookup.wait() + + prisma_client.writer_db.litellm_projecttable.find_unique = never_returns_project + monkeypatch.setattr(litellm, "default_key_generate_params", None) + + with pytest.raises(DBLookupDeadlineExceeded) as error: + await asyncio.wait_for( + _common_key_generation_helper( + data=GenerateKeyRequest(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + team_table=None, + ), + timeout=1.0, + ) + + assert error.value.lookup == "project" + + @pytest.mark.parametrize("project_team_id", ["team-a", None]) @pytest.mark.asyncio async def test_key_generation_accepts_same_team_and_unowned_projects( @@ -20903,7 +20945,7 @@ async def test_key_team_ownership_mutation_allows_legacy_mismatch_without_change request_fields: dict[str, str], ) -> None: prisma_client: Final = MagicMock() - prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( return_value=LiteLLM_ProjectTable( project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM, @@ -20925,7 +20967,7 @@ async def test_key_team_ownership_mutation_allows_legacy_mismatch_without_change @pytest.mark.asyncio async def test_key_team_ownership_mutation_allows_detach_with_team_change() -> None: prisma_client: Final = MagicMock() - prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( return_value=LiteLLM_ProjectTable( project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM, @@ -20956,8 +20998,9 @@ async def test_regenerate_checks_project_team_ownership( ) -> None: existing_key: Final = LiteLLM_VerificationToken(token="abc123", team_id=key_team_id) mock_prisma_client: Final = _make_regenerate_mock_prisma() - mock_prisma_client.db.litellm_projecttable = MagicMock() - mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + mock_prisma_client.writer_db = MagicMock() + mock_prisma_client.writer_db.litellm_projecttable = MagicMock() + mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id="team-b") ) user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") @@ -21004,7 +21047,7 @@ async def test_regenerate_checks_project_team_ownership( @pytest.mark.asyncio async def test_key_project_team_validation_uses_project_missing_404() -> None: prisma_client: Final = MagicMock() - prisma_client.db.litellm_projecttable.find_unique = AsyncMock(return_value=None) + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock(return_value=None) with pytest.raises(HTTPException) as error: await _check_key_project_team( From cd442a1999d5912424902059a1a8ed07c49ca72f Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 18:38:40 +0000 Subject: [PATCH 09/20] fix(proxy): check project moves against the primary database team Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/project_endpoints.py | 47 ++++++++++++------- .../key_management_endpoints.py | 4 +- .../test_project_endpoints_prisma.py | 47 ++++++++++++++++++- 3 files changed, 75 insertions(+), 23 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index db36500935c..edd865dbc01 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -30,6 +30,7 @@ from litellm.proxy.management_helpers.utils import ( management_endpoint_wrapper, ) from litellm.proxy.utils import PrismaClient, handle_exception_on_proxy +from litellm.repositories.base_repository import record_to_dict from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.object_permission_repository import ObjectPermissionRepository from litellm.repositories.prisma_protocols import TableActions @@ -768,26 +769,36 @@ async def update_project( detail={"error": "Cannot reassign project to a team you are not an admin of"}, ) - if data.team_id is not None and data.team_id != existing_project.team_id: - mismatched_key_count: Final = await bounded_db_lookup( - prisma_client.writer_db.litellm_verificationtoken.count( - where={ - "project_id": data.project_id, - "OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}], - } - ), - name="project_key_ownership", + if data.team_id is not None: + current_project_record: Final = await bounded_db_lookup( + prisma_client.writer_db.litellm_projecttable.find_unique(where={"project_id": data.project_id}), + name="project", ) - if mismatched_key_count > 0: - raise HTTPException( - status_code=400, - detail={ - "error": ( - f"Project {data.project_id} has {mismatched_key_count} key(s) that do not belong to " - f"team {data.team_id}. Detach or delete them before moving the project." - ) - }, + current_project: Final = ( + LiteLLM_ProjectTable.model_validate(record_to_dict(current_project_record)) + if current_project_record is not None + else None + ) + if current_project is not None and data.team_id != current_project.team_id: + mismatched_key_count: Final = await bounded_db_lookup( + prisma_client.writer_db.litellm_verificationtoken.count( + where={ + "project_id": data.project_id, + "OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}], + } + ), + name="project_key_ownership", ) + if mismatched_key_count > 0: + raise HTTPException( + status_code=400, + detail={ + "error": ( + f"Project {data.project_id} has {mismatched_key_count} key(s) that do not belong to " + f"team {data.team_id}. Detach or delete them before moving the project." + ) + }, + ) # Validate project limits against team limits if target_team_obj is not None: diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 4b678a48eeb..09ac97243cb 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1823,9 +1823,7 @@ async def _check_key_project_team( name="project", ) project_obj: Final = ( - LiteLLM_ProjectTable.model_validate(record_to_dict(project_record)) - if project_record is not None - else None + LiteLLM_ProjectTable.model_validate(record_to_dict(project_record)) if project_record is not None else None ) if project_obj is None: diff --git a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py index b5046244c31..d6ef9df3f2c 100644 --- a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py +++ b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py @@ -37,6 +37,7 @@ verbose_proxy_logger.setLevel(level=logging.DEBUG) from litellm.caching.caching import DualCache from litellm.proxy._types import ( + LiteLLM_ProjectTable, LiteLLM_TeamTable, NewProjectRequest, UpdateProjectRequest, @@ -1236,7 +1237,11 @@ def _project_update_mocks(monkeypatch, stored_metadata: dict) -> mock.MagicMock: mock_prisma.jsonify_object = lambda data: data mock_prisma.db.litellm_projecttable.find_unique = mock.AsyncMock(return_value=existing_row) mock_prisma.db.litellm_projecttable.update = mock.AsyncMock(return_value=mock.MagicMock()) - mock_prisma.writer_db = mock_prisma.db + mock_prisma.writer_db = mock.MagicMock() + mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock( + return_value={"project_id": "project-update-test", "team_id": None} + ) + mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(return_value=0) monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", True) monkeypatch.setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma) @@ -1273,6 +1278,9 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists( ) mock_prisma.db.litellm_verificationtoken.count = mock.AsyncMock(return_value=0) mock_prisma.writer_db = mock.MagicMock() + mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=project_id, team_id="team-a") + ) async def count_teamless_keys(*, where: Mapping[str, object]) -> int: conditions: Final = where.get("OR") @@ -1296,6 +1304,37 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists( mock_prisma.db.litellm_projecttable.update.assert_not_awaited() +@pytest.mark.asyncio +async def test_update_project_rejects_move_when_writer_team_differs_from_stale_reader( + monkeypatch: pytest.MonkeyPatch, +) -> None: + project_id: Final = "project-replica-lag" + destination_team_id: Final = "team-a" + mock_prisma: Final = _project_update_mocks(monkeypatch, {}) + mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = destination_team_id + mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock( + return_value=LiteLLM_TeamTable(team_id=destination_team_id) + ) + mock_prisma.db.litellm_verificationtoken.count = mock.AsyncMock(return_value=0) + mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=project_id, team_id="team-b") + ) + mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(return_value=1) + + with pytest.raises(ProxyException) as error: + await _run_project_update(project_id, team_id=destination_team_id) + + assert error.value.code == "400" + assert ( + f"Project {project_id} has 1 key(s) that do not belong to team {destination_team_id}. " + "Detach or delete them before moving the project." + ) in error.value.message + mock_prisma.writer_db.litellm_projecttable.find_unique.assert_awaited_once() + mock_prisma.writer_db.litellm_verificationtoken.count.assert_awaited_once() + mock_prisma.db.litellm_verificationtoken.count.assert_not_awaited() + mock_prisma.db.litellm_projecttable.update.assert_not_awaited() + + @pytest.mark.asyncio async def test_update_project_allows_move_when_no_attached_keys_exist(monkeypatch: pytest.MonkeyPatch) -> None: project_id: Final = "project-without-keys" @@ -1306,10 +1345,14 @@ async def test_update_project_allows_move_when_no_attached_keys_exist(monkeypatc return_value=LiteLLM_TeamTable(team_id=destination_team_id) ) mock_prisma.db.litellm_verificationtoken.count = mock.AsyncMock(return_value=0) + mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=project_id, team_id="team-a") + ) await _run_project_update(project_id, team_id=destination_team_id) - mock_prisma.db.litellm_verificationtoken.count.assert_awaited_once() + mock_prisma.writer_db.litellm_verificationtoken.count.assert_awaited_once() + mock_prisma.db.litellm_verificationtoken.count.assert_not_awaited() mock_prisma.db.litellm_projecttable.update.assert_awaited_once() From a105deab2b3dd7778908c67e25d81b7572a9d963 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 21:19:07 +0000 Subject: [PATCH 10/20] fix(proxy): run ownership checks after existing validation without the lookup deadline Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/project_endpoints.py | 59 +++-- .../key_management_endpoints.py | 56 +++-- .../management/test_project_lifecycle.py | 149 +++++++++++++ .../test_project_endpoints_prisma.py | 73 +++++++ .../test_key_management_endpoints.py | 201 +++++++++++++++--- 5 files changed, 449 insertions(+), 89 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index edd865dbc01..930cb7c73ae 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -22,7 +22,6 @@ from litellm._uuid import uuid from litellm.proxy._types import * from litellm.proxy.auth.auth_checks import delete_cached_project_object from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.db.db_lookup_gate import bounded_db_lookup from litellm.proxy.management.teams.access import is_team_admin from litellm.proxy.management_endpoints.common_utils import _set_object_metadata_field from litellm.proxy.management_endpoints.team_admin_field_permissions import team_admin_may_manage_projects @@ -769,37 +768,6 @@ async def update_project( detail={"error": "Cannot reassign project to a team you are not an admin of"}, ) - if data.team_id is not None: - current_project_record: Final = await bounded_db_lookup( - prisma_client.writer_db.litellm_projecttable.find_unique(where={"project_id": data.project_id}), - name="project", - ) - current_project: Final = ( - LiteLLM_ProjectTable.model_validate(record_to_dict(current_project_record)) - if current_project_record is not None - else None - ) - if current_project is not None and data.team_id != current_project.team_id: - mismatched_key_count: Final = await bounded_db_lookup( - prisma_client.writer_db.litellm_verificationtoken.count( - where={ - "project_id": data.project_id, - "OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}], - } - ), - name="project_key_ownership", - ) - if mismatched_key_count > 0: - raise HTTPException( - status_code=400, - detail={ - "error": ( - f"Project {data.project_id} has {mismatched_key_count} key(s) that do not belong to " - f"team {data.team_id}. Detach or delete them before moving the project." - ) - }, - ) - # Validate project limits against team limits if target_team_obj is not None: _check_team_project_limits( @@ -824,6 +792,33 @@ async def update_project( **({"max_budget": None} if "max_budget" in data.model_fields_set and data.max_budget is None else {}), } + if data.team_id is not None: + current_project_record: Final = await prisma_client.writer_db.litellm_projecttable.find_unique( + where={"project_id": data.project_id} + ) + current_project: Final = ( + LiteLLM_ProjectTable.model_validate(record_to_dict(current_project_record)) + if current_project_record is not None + else None + ) + if current_project is not None and data.team_id != current_project.team_id: + mismatched_key_count: Final = await prisma_client.writer_db.litellm_verificationtoken.count( + where={ + "project_id": data.project_id, + "OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}], + } + ) + if mismatched_key_count > 0: + raise HTTPException( + status_code=400, + detail={ + "error": ( + f"Project {data.project_id} has {mismatched_key_count} key(s) that do not belong to " + f"team {data.team_id}. Detach or delete them before moving the project." + ) + }, + ) + if budget_updates and existing_project.budget_id: # Update existing budget await _budget_table(prisma_client).update( diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 09ac97243cb..579d6848152 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -84,7 +84,6 @@ from litellm.proxy.common_utils.config_sync_pubsub import ( from litellm.proxy.common_utils.rbac_utils import check_org_admin_can_generate_keys from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache -from litellm.proxy.db.db_lookup_gate import bounded_db_lookup from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks from litellm.proxy.hooks.model_max_budget_limiter import build_model_max_budget_usage from litellm.proxy.management.teams.access import TEAM_ADMIN_ONLY, TEAM_OR_ORG_ADMIN, is_team_admin @@ -1265,13 +1264,6 @@ 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) - if data.project_id is not None and prisma_client is not None: - await _check_key_project_team( - project_id=data.project_id, - key_team_id=data.team_id, - prisma_client=prisma_client, - ) - # Delegated-authority ceiling (GHSA-q775-qw9r-2r4g): a non-admin caller # cannot grant a key a higher budget than their own authority. # UI session personal keys are capped by user_max_budget when it is available. @@ -1356,6 +1348,13 @@ async def _common_key_generation_helper( ), ) + if data.project_id is not None and prisma_client is not None: + await _check_key_project_team( + project_id=data.project_id, + key_team_id=data.team_id, + prisma_client=prisma_client, + ) + # TODO: @ishaan-jaff: Migrate all budget tracking to use LiteLLM_BudgetTable _budget_id = data.budget_id if prisma_client is not None and data.soft_budget is not None: @@ -1818,9 +1817,8 @@ async def _check_key_project_team( key_team_id: str | None, prisma_client: PrismaClient, ) -> None: - project_record: Final = await bounded_db_lookup( - prisma_client.writer_db.litellm_projecttable.find_unique(where={"project_id": project_id}), - name="project", + project_record: Final = await prisma_client.writer_db.litellm_projecttable.find_unique( + where={"project_id": project_id} ) project_obj: Final = ( LiteLLM_ProjectTable.model_validate(record_to_dict(project_record)) if project_record is not None else None @@ -2900,13 +2898,6 @@ async def _process_single_key_update( llm_router=llm_router, ) - if prisma_client is not None: - await _check_key_project_team_on_mutation( - data=update_key_request, - existing_key_row=existing_key_row, - prisma_client=prisma_client, - ) - key_request: Final = await _with_validated_object_permission( update_key_request=update_key_request, team_obj=team_obj, @@ -2938,6 +2929,12 @@ async def _process_single_key_update( detail={"error": "Database not connected"}, ) + await _check_key_project_team_on_mutation( + data=key_request, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + ) + update_values: Final = await _handle_update_object_permission( data_json=non_default_values, existing_key_row=existing_key_row, @@ -3364,12 +3361,6 @@ async def _validate_update_key_data( user_api_key_cache=user_api_key_cache, ) - await _check_key_project_team_on_mutation( - data=data, - existing_key_row=existing_key_row, - prisma_client=checked_prisma_client, - ) - # 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. # Without this gate, any caller could assign their key to an arbitrary @@ -3617,6 +3608,12 @@ async def update_key_fn( if prisma_client is None: raise Exception("Not connected to DB!") + await _check_key_project_team_on_mutation( + data=data, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + ) + update_values: Final = await _handle_update_object_permission( data_json=non_default_values, existing_key_row=existing_key_row, @@ -5662,11 +5659,6 @@ async def _execute_virtual_key_regeneration( ) if data is not None: - await _check_key_project_team_on_mutation( - data=data, - existing_key_row=key_in_db, - prisma_client=prisma_client, - ) _existing_key_metadata: Final = getattr(key_in_db, "metadata", None) enforce_output_token_estimates_are_admin_only( data=data, @@ -5718,6 +5710,12 @@ async def _execute_virtual_key_regeneration( request=data if data is not None else RegenerateKeyRequest(), ), ) + if data is not None: + await _check_key_project_team_on_mutation( + data=data, + existing_key_row=key_in_db, + prisma_client=prisma_client, + ) update_values: Final = await _handle_update_object_permission( data_json=non_default_values, existing_key_row=key_in_db, diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index b13ed839fc1..6982fbd581b 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -1,6 +1,7 @@ from collections.abc import Iterator from hashlib import sha256 from typing import Final +from uuid import uuid4 import httpx import pytest @@ -9,6 +10,9 @@ from integration._support.database import read_rows, write_rows from integration._support.process import owned_proxy from pydantic import JsonValue +from litellm.models.user import LiteLLM_UserTable +from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken + def _project_rows(project_id: str) -> list[dict[str, JsonValue]]: return read_rows( @@ -27,6 +31,22 @@ def _key_rows(key: str) -> list[dict[str, JsonValue]]: ) +def _cli_session_token(user_id: str, team_id: str, *, max_budget: float | None = None) -> str: + user: Final = LiteLLM_UserTable( + user_id=user_id, + user_role="internal_user", + teams=[team_id], + models=[], + max_budget=max_budget, + ) + return ExperimentalUIJWTToken.get_cli_jwt_auth_token( + user_info=user, + team_id=team_id, + team_alias="ownership-team", + max_budget=max_budget, + ) + + def _discard_unexpected_key(candidate: Gateway, response: httpx.Response) -> None: if response.status_code != 200: return @@ -264,6 +284,135 @@ def test_key_regenerate_rejects_foreign_project_without_changing_key(ownership_g assert _key_rows(key) == before +def test_key_generation_rejects_missing_project_without_writing_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team: Final = scenario.team(models=[model]) + missing_project_id: Final = f"missing-{uuid4()}" + response: Final = ownership_gateway.request( + "POST", + "/key/service-account/generate", + {"team_id": team, "project_id": missing_project_id, "models": [model]}, + ) + + assert response.status_code == 404, response.text + assert "Project not found" in response.text + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (missing_project_id,), + ) + == [] + ) + + +def test_key_regenerate_routes_reject_missing_project_without_changing_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team: Final = scenario.team(models=[model]) + key: Final = scenario.key(team_id=team, models=[model]) + missing_project_id: Final = f"missing-{uuid4()}" + before: Final = _key_rows(key) + responses: Final = ( + ownership_gateway.request( + "POST", + "/key/regenerate", + {"key": key, "project_id": missing_project_id}, + ), + ownership_gateway.request( + "POST", + f"/key/{key}/regenerate", + {"project_id": missing_project_id}, + ), + ) + + for response in responses: + assert response.status_code == 404, response.text + assert "Project not found" in response.text + assert _key_rows(key) == before + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (missing_project_id,), + ) + == [] + ) + + chat: Final = ownership_gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "rejected regeneration preserves key"}]}, + key=key, + ) + assert chat.status_code == 200, chat.text + + +def test_key_generate_budget_ceiling_precedes_project_ownership(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + caller_team: Final = scenario.team(models=[model]) + owner_team: Final = scenario.team(models=[model]) + project: Final = scenario.project(owner_team, models=[model]) + caller_id: Final = scenario.member(caller_team, role="admin") + caller_token: Final = _cli_session_token(caller_id, caller_team, max_budget=1) + response: Final = ownership_gateway.request( + "POST", + "/key/generate", + { + "team_id": caller_team, + "project_id": project, + "max_budget": 5, + "models": [model], + }, + key=caller_token, + ) + + assert response.status_code == 400, response.text + assert "max_budget (5.0) cannot exceed the caller's own max_budget (1.0)" in response.text + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (project,), + ) + == [] + ) + + +def test_key_regenerate_output_estimate_admin_error_precedes_project_ownership( + ownership_gateway: Gateway, +) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + project_team: Final = scenario.team(models=[model]) + destination_team: Final = scenario.team(models=[model]) + project: Final = scenario.project(project_team, models=[model]) + caller_id: Final = scenario.member(project_team, role="admin") + key: Final = scenario.key(team_id=project_team, user_id=caller_id, models=[model]) + caller_token: Final = _cli_session_token(caller_id, project_team) + before: Final = _key_rows(key) + response: Final = ownership_gateway.request( + "POST", + f"/key/{key}/regenerate", + { + "team_id": destination_team, + "project_id": project, + "metadata": {"default_estimated_output_tokens": 1}, + }, + key=caller_token, + ) + + assert response.status_code == 403, response.text + assert "Only proxy admins can set" in response.text + assert _key_rows(key) == before + chat: Final = ownership_gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "rejected regeneration keeps key valid"}]}, + key=key, + ) + assert chat.status_code == 200, chat.text + + def test_key_bulk_update_rejects_foreign_team_project_and_preserves_key(ownership_gateway: Gateway) -> None: with ownership_gateway.scenario() as scenario: model: Final = scenario.model() diff --git a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py index d6ef9df3f2c..e075d2c8aca 100644 --- a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py +++ b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py @@ -1,3 +1,4 @@ +import asyncio import os import traceback from collections.abc import Mapping @@ -8,6 +9,9 @@ from unittest import mock from dotenv import load_dotenv from fastapi import HTTPException, Request +from litellm.constants import PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS +from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker + load_dotenv() import time @@ -1335,6 +1339,75 @@ async def test_update_project_rejects_move_when_writer_team_differs_from_stale_r mock_prisma.db.litellm_projecttable.update.assert_not_awaited() +@pytest.mark.asyncio +async def test_update_project_slow_writer_ownership_reads_do_not_stall_db_tracker( + monkeypatch: pytest.MonkeyPatch, +) -> None: + project_id: Final = "project-slow-writer" + source_team_id: Final = "team-source" + destination_team_id: Final = "team-destination" + mock_prisma: Final = _project_update_mocks(monkeypatch, {}) + mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = source_team_id + mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock( + return_value=LiteLLM_TeamTable(team_id=destination_team_id, models=[]) + ) + + async def slow_project_lookup(*, where: Mapping[str, object]) -> LiteLLM_ProjectTable: + assert where == {"project_id": project_id} + await asyncio.sleep(0.2) + return LiteLLM_ProjectTable(project_id=project_id, team_id=source_team_id) + + async def slow_key_count(*, where: Mapping[str, object]) -> int: + assert where == { + "project_id": project_id, + "OR": [{"team_id": {"not": destination_team_id}}, {"team_id": None}], + } + await asyncio.sleep(0.2) + return 0 + + mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock(side_effect=slow_project_lookup) + mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(side_effect=slow_key_count) + monkeypatch.setattr("litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.05) + + db_lookup_stall_tracker.clear() + try: + await _run_project_update(project_id, team_id=destination_team_id) + assert mock_prisma.writer_db.litellm_projecttable.find_unique.await_count == 1 + assert mock_prisma.writer_db.litellm_verificationtoken.count.await_count == 1 + mock_prisma.db.litellm_projecttable.update.assert_awaited_once() + assert db_lookup_stall_tracker.stalled_within(PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS) is False + finally: + db_lookup_stall_tracker.clear() + + +@pytest.mark.asyncio +async def test_update_project_team_limit_error_precedes_mismatched_key_guard( + monkeypatch: pytest.MonkeyPatch, +) -> None: + project_id: Final = "project-team-limit-precedence" + source_team_id: Final = "team-source" + destination_team_id: Final = "team-destination" + mock_prisma: Final = _project_update_mocks(monkeypatch, {}) + mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = source_team_id + mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock( + return_value=LiteLLM_TeamTable( + team_id=destination_team_id, + models=["allowed-model"], + ) + ) + mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=project_id, team_id=source_team_id) + ) + mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(return_value=1) + + with pytest.raises(ProxyException, match="not in team's allowed models") as error: + await _run_project_update(project_id, team_id=destination_team_id, models=["disallowed-model"]) + + assert error.value.code == "400" + mock_prisma.writer_db.litellm_verificationtoken.count.assert_not_awaited() + mock_prisma.db.litellm_projecttable.update.assert_not_awaited() + + @pytest.mark.asyncio async def test_update_project_allows_move_when_no_attached_keys_exist(monkeypatch: pytest.MonkeyPatch) -> None: project_id: Final = "project-without-keys" diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 8a3e31f1644..df285ab5a69 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -45,7 +45,8 @@ from litellm.proxy.auth.auth_checks import ( ) from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, project_cache_key -from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded +from litellm.constants import PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS +from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy.management_endpoints.key_management_endpoints import ( _check_key_project_team, @@ -82,6 +83,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( list_keys, prepare_key_update_data, reset_key_spend_fn, + update_key_fn, validate_key_list_check, validate_key_team_change, ) @@ -2963,7 +2965,10 @@ async def test_update_key_nonexistent_key_returns_404(monkeypatch): def _setup_update_key_mocks(monkeypatch, mock_prisma_client): monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) - monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + proxy_logging_obj: Final = MagicMock() + proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock() + proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj) monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) monkeypatch.setattr("litellm.store_audit_logs", False) @@ -18797,7 +18802,9 @@ def _estimate_key_row(token: str, metadata: dict): return existing_key -def _wire_update_key_fn(monkeypatch, existing_key): +def _wire_update_key_fn( + monkeypatch: pytest.MonkeyPatch, existing_key: LiteLLM_VerificationToken | MagicMock +) -> AsyncMock: mock_prisma_client = AsyncMock() updated_key = MagicMock() updated_key.token = existing_key.token @@ -18826,6 +18833,7 @@ def _wire_update_key_fn(monkeypatch, existing_key): "litellm.proxy.management_endpoints.key_management_endpoints._enforce_unique_key_alias", _noop, ) + return mock_prisma_client @pytest.mark.asyncio @@ -20598,6 +20606,37 @@ async def test_key_generation_rejects_foreign_project_team( assert error.value.detail == expected_detail +@pytest.mark.asyncio +async def test_key_generation_budget_ceiling_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM) + _configure_key_endpoints(monkeypatch, user_api_key_cache) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=GenerateKeyRequest( + project_id=_OWNED_PROJECT, + team_id=_OWNERSHIP_KEY_TEAM, + max_budget=10, + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="litellm_login_session", + user_id="session-user", + team_id=_OWNERSHIP_KEY_TEAM, + is_session_token=True, + max_budget=5, + ), + litellm_changed_by=None, + team_table=LiteLLM_TeamTableCachedObj(team_id=_OWNERSHIP_KEY_TEAM), + ) + + assert error.value.status_code == 400 + assert error.value.detail == {"error": "max_budget (10.0) cannot exceed the caller's own max_budget (5.0)."} + + @pytest.mark.parametrize( ("key_team_id", "expected_status"), [(_OWNERSHIP_PROJECT_TEAM, 200), (_OWNERSHIP_KEY_TEAM, 400)], @@ -20651,23 +20690,24 @@ async def test_key_generation_uses_database_project_team_when_cache_is_stale( @pytest.mark.asyncio -async def test_key_generation_fails_when_writer_project_lookup_exceeds_deadline( +async def test_key_generation_slow_writer_project_lookup_does_not_stall_db_tracker( monkeypatch: pytest.MonkeyPatch, ) -> None: user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM) prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) - stalled_lookup: Final = asyncio.Event() monkeypatch.setattr("litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.05) - async def never_returns_project(*, where: Mapping[str, object]) -> None: + async def slow_project_lookup(*, where: Mapping[str, object]) -> LiteLLM_ProjectTable: assert where == {"project_id": _OWNED_PROJECT} - await stalled_lookup.wait() + await asyncio.sleep(0.2) + return LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) - prisma_client.writer_db.litellm_projecttable.find_unique = never_returns_project + prisma_client.writer_db.litellm_projecttable.find_unique = slow_project_lookup monkeypatch.setattr(litellm, "default_key_generate_params", None) - with pytest.raises(DBLookupDeadlineExceeded) as error: - await asyncio.wait_for( + db_lookup_stall_tracker.clear() + try: + response: Final = await asyncio.wait_for( _common_key_generation_helper( data=GenerateKeyRequest(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM), user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), @@ -20676,8 +20716,10 @@ async def test_key_generation_fails_when_writer_project_lookup_exceeds_deadline( ), timeout=1.0, ) - - assert error.value.lookup == "project" + assert response.team_id == _OWNERSHIP_PROJECT_TEAM + assert db_lookup_stall_tracker.stalled_within(PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS) is False + finally: + db_lookup_stall_tracker.clear() @pytest.mark.parametrize("project_team_id", ["team-a", None]) @@ -20815,36 +20857,98 @@ async def test_key_update_allows_legacy_project_mismatch_when_team_is_unchanged( @pytest.mark.asyncio async def test_key_update_rejects_team_change_for_project_bound_key(monkeypatch: pytest.MonkeyPatch) -> None: - user_api_key_cache: Final = await _cache_with_project( - _OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM - ) - mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM) - ) existing_key_row: Final = LiteLLM_VerificationToken( token="hashed-key", team_id=_OWNERSHIP_KEY_TEAM, project_id=_OWNED_PROJECT, ) + mock_prisma_client: Final = _wire_update_key_fn(monkeypatch, existing_key_row) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + AsyncMock(return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM, team_members=[])), + ) + mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock()) + mock_request: Final = MagicMock() + mock_request.query_params = {} - with pytest.raises(HTTPException) as error: - await _validate_update_key_data( + with pytest.raises((HTTPException, ProxyException)) as error: + await update_key_fn( + request=mock_request, data=UpdateKeyRequest(key="sk-key", team_id=_OWNERSHIP_DESTINATION_TEAM), - existing_key_row=existing_key_row, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), - llm_router=None, - premium_user=True, - prisma_client=mock_prisma_client, - user_api_key_cache=user_api_key_cache, + litellm_changed_by=None, ) - assert error.value.status_code == 400 + assert str(getattr(error.value, "status_code", None) or getattr(error.value, "code", None)) == "400" expected_detail: Final = ( f"Project {_OWNED_PROJECT} belongs to team {_OWNERSHIP_PROJECT_TEAM}, " f"but the key belongs to {_OWNERSHIP_DESTINATION_TEAM}" ) - assert expected_detail in str(error.value.detail) + assert expected_detail in str(getattr(error.value, "detail", None) or getattr(error.value, "message", None)) + mock_prisma_client.update_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_key_update_organization_membership_error_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-key", + user_id="key-owner", + created_by="key-owner", + team_id=_OWNERSHIP_PROJECT_TEAM, + project_id=_OWNED_PROJECT, + ) + mock_prisma_client: Final = _wire_update_key_fn(monkeypatch, existing_key_row) + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache) + mock_prisma_client.db.litellm_usertable = MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=LiteLLM_UserTable(user_id="key-owner", organization_memberships=[]) + ) + mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) + ) + team_object: Final = LiteLLM_TeamTableCachedObj( + team_id=_OWNERSHIP_KEY_TEAM, + members_with_roles=[Member(user_id="key-owner", role="admin")], + team_member_permissions=["/key/update"], + ) + team_lookup: Final = AsyncMock(return_value=team_object) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + team_lookup, + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.team_member_permission_checks.get_team_object", + team_lookup, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock()) + mock_request: Final = MagicMock() + mock_request.query_params = {} + + with pytest.raises(ProxyException) as error: + await update_key_fn( + request=mock_request, + data=UpdateKeyRequest( + key="sk-key", + team_id=_OWNERSHIP_KEY_TEAM, + organization_id="org-not-member", + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-user", + user_id="key-owner", + team_id=_OWNERSHIP_PROJECT_TEAM, + ), + litellm_changed_by=None, + ) + + assert error.value.code == "403" + assert error.value.message == "Caller is not a member of organization_id=org-not-member" @pytest.mark.asyncio @@ -21044,6 +21148,47 @@ async def test_regenerate_checks_project_team_ownership( assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == expected_updates +@pytest.mark.asyncio +async def test_regenerate_output_estimate_admin_error_precedes_project_ownership() -> None: + existing_key: Final = LiteLLM_VerificationToken( + token="abc123", + user_id="key-owner", + team_id=_OWNERSHIP_PROJECT_TEAM, + project_id=_OWNED_PROJECT, + metadata={}, + ) + mock_prisma_client: Final = _make_regenerate_mock_prisma() + mock_prisma_client.writer_db = MagicMock() + mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) + ) + + with pytest.raises(HTTPException) as error: + 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=_OWNED_PROJECT, + team_id=_OWNERSHIP_KEY_TEAM, + metadata={"default_estimated_output_tokens": 1}, + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-user", + user_id="key-owner", + ), + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + assert error.value.status_code == 403 + assert "Only proxy admins can set" in str(error.value.detail) + mock_prisma_client.db.litellm_verificationtoken.update.assert_not_awaited() + + @pytest.mark.asyncio async def test_key_project_team_validation_uses_project_missing_404() -> None: prisma_client: Final = MagicMock() From aba2597203b3327b40f32c25ba02d1103ee334b4 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 21:49:05 +0000 Subject: [PATCH 11/20] fix(proxy): type ownership writer reads Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/project_endpoints.py | 13 +++++++++++-- .../key_management_endpoints.py | 16 ++++++++++++---- 2 files changed, 23 insertions(+), 6 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index 930cb7c73ae..c90917d5341 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -12,7 +12,7 @@ Endpoints for /project operations import json from collections.abc import Mapping, Sequence -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, cast from fastapi import APIRouter, Depends, HTTPException, Request from pydantic import TypeAdapter @@ -56,6 +56,15 @@ def _project_table(prisma_client: PrismaClient) -> TableActions["prisma_models.L return ProjectRepository(prisma_client).table +def _writer_project_table( + prisma_client: PrismaClient, +) -> TableActions["prisma_models.LiteLLM_ProjectTable"]: + return cast( # cast-ok: writer_db exposes generated Prisma tables through a dynamic wrapper + "TableActions[prisma_models.LiteLLM_ProjectTable]", + prisma_client.writer_db.litellm_projecttable, + ) + + def _verification_token_table( prisma_client: PrismaClient, ) -> TableActions["prisma_models.LiteLLM_VerificationToken"]: @@ -793,7 +802,7 @@ async def update_project( } if data.team_id is not None: - current_project_record: Final = await prisma_client.writer_db.litellm_projecttable.find_unique( + current_project_record: Final = await _writer_project_table(prisma_client).find_unique( where={"project_id": data.project_id} ) current_project: Final = ( diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 579d6848152..75c79fd77e1 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -262,6 +262,15 @@ def _prisma_table( ) +def _writer_project_table( + prisma_client: PrismaClient, +) -> "TableActions[prisma_models.LiteLLM_ProjectTable]": + return cast( # cast-ok: writer_db exposes generated Prisma tables through a dynamic wrapper + "TableActions[prisma_models.LiteLLM_ProjectTable]", + prisma_client.writer_db.litellm_projecttable, + ) + + def _deleted_verification_token_table( prisma_client: PrismaClient, ) -> "TableActions[prisma_models.LiteLLM_DeletedVerificationToken]": @@ -1332,7 +1341,8 @@ async def _common_key_generation_helper( apply_enterprise_key_management_params, ) - data = apply_enterprise_key_management_params(data, team_table) + enterprise_data: Final[object] = apply_enterprise_key_management_params(data, team_table) + data = GenerateKeyRequest.model_validate(enterprise_data) except Exception as e: verbose_proxy_logger.debug( "litellm.proxy.proxy_server.generate_key_fn(): Enterprise key management params not applied - %s", e @@ -1817,9 +1827,7 @@ async def _check_key_project_team( key_team_id: str | None, prisma_client: PrismaClient, ) -> None: - project_record: Final = await prisma_client.writer_db.litellm_projecttable.find_unique( - where={"project_id": project_id} - ) + project_record: Final = await _writer_project_table(prisma_client).find_unique(where={"project_id": project_id}) project_obj: Final = ( LiteLLM_ProjectTable.model_validate(record_to_dict(project_record)) if project_record is not None else None ) From b5cfe6665425f65eeae6793710490017ece29ffb Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 22:52:42 +0000 Subject: [PATCH 12/20] test(proxy): make ownership tests independent of runner salt and clock Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management/test_project_lifecycle.py | 30 +++++++++++++++---- .../test_project_endpoints_prisma.py | 8 +++-- .../test_key_management_endpoints.py | 5 ++-- 3 files changed, 32 insertions(+), 11 deletions(-) diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index 6982fbd581b..212229e3535 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -13,6 +13,8 @@ from pydantic import JsonValue from litellm.models.user import LiteLLM_UserTable from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken +_OWNED_PROXY_SALT_KEY: Final = "sk-integration-salt" + def _project_rows(project_id: str) -> list[dict[str, JsonValue]]: return read_rows( @@ -31,7 +33,14 @@ def _key_rows(key: str) -> list[dict[str, JsonValue]]: ) -def _cli_session_token(user_id: str, team_id: str, *, max_budget: float | None = None) -> str: +def _cli_session_token( + user_id: str, + team_id: str, + *, + monkeypatch: pytest.MonkeyPatch, + max_budget: float | None = None, +) -> str: + monkeypatch.setenv("LITELLM_SALT_KEY", _OWNED_PROXY_SALT_KEY) user: Final = LiteLLM_UserTable( user_id=user_id, user_role="internal_user", @@ -57,7 +66,12 @@ def _discard_unexpected_key(candidate: Gateway, response: httpx.Response) -> Non @pytest.fixture(scope="module") def ownership_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: with gateway_from_environment() as gateway: - with owned_proxy(gateway, tmp_path_factory.mktemp("project-team-ownership"), {}, workers=2) as candidate: + with owned_proxy( + gateway, + tmp_path_factory.mktemp("project-team-ownership"), + {"LITELLM_SALT_KEY": _OWNED_PROXY_SALT_KEY}, + workers=2, + ) as candidate: yield candidate @@ -347,14 +361,18 @@ def test_key_regenerate_routes_reject_missing_project_without_changing_key(owner assert chat.status_code == 200, chat.text -def test_key_generate_budget_ceiling_precedes_project_ownership(ownership_gateway: Gateway) -> None: +def test_key_generate_budget_ceiling_precedes_project_ownership( + ownership_gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: with ownership_gateway.scenario() as scenario: model: Final = scenario.model() caller_team: Final = scenario.team(models=[model]) owner_team: Final = scenario.team(models=[model]) project: Final = scenario.project(owner_team, models=[model]) caller_id: Final = scenario.member(caller_team, role="admin") - caller_token: Final = _cli_session_token(caller_id, caller_team, max_budget=1) + caller_token: Final = _cli_session_token( + caller_id, caller_team, monkeypatch=monkeypatch, max_budget=1 + ) response: Final = ownership_gateway.request( "POST", "/key/generate", @@ -379,7 +397,7 @@ def test_key_generate_budget_ceiling_precedes_project_ownership(ownership_gatewa def test_key_regenerate_output_estimate_admin_error_precedes_project_ownership( - ownership_gateway: Gateway, + ownership_gateway: Gateway, monkeypatch: pytest.MonkeyPatch ) -> None: with ownership_gateway.scenario() as scenario: model: Final = scenario.model() @@ -388,7 +406,7 @@ def test_key_regenerate_output_estimate_admin_error_precedes_project_ownership( project: Final = scenario.project(project_team, models=[model]) caller_id: Final = scenario.member(project_team, role="admin") key: Final = scenario.key(team_id=project_team, user_id=caller_id, models=[model]) - caller_token: Final = _cli_session_token(caller_id, project_team) + caller_token: Final = _cli_session_token(caller_id, project_team, monkeypatch=monkeypatch) before: Final = _key_rows(key) response: Final = ownership_gateway.request( "POST", diff --git a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py index e075d2c8aca..98bb2969201 100644 --- a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py +++ b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py @@ -1354,7 +1354,8 @@ async def test_update_project_slow_writer_ownership_reads_do_not_stall_db_tracke async def slow_project_lookup(*, where: Mapping[str, object]) -> LiteLLM_ProjectTable: assert where == {"project_id": project_id} - await asyncio.sleep(0.2) + await asyncio.sleep(0) + await asyncio.sleep(0) return LiteLLM_ProjectTable(project_id=project_id, team_id=source_team_id) async def slow_key_count(*, where: Mapping[str, object]) -> int: @@ -1362,12 +1363,13 @@ async def test_update_project_slow_writer_ownership_reads_do_not_stall_db_tracke "project_id": project_id, "OR": [{"team_id": {"not": destination_team_id}}, {"team_id": None}], } - await asyncio.sleep(0.2) + await asyncio.sleep(0) + await asyncio.sleep(0) return 0 mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock(side_effect=slow_project_lookup) mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(side_effect=slow_key_count) - monkeypatch.setattr("litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.05) + monkeypatch.setattr("litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.0) db_lookup_stall_tracker.clear() try: diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index df285ab5a69..37d004b251b 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -20695,11 +20695,12 @@ async def test_key_generation_slow_writer_project_lookup_does_not_stall_db_track ) -> None: user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM) prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) - monkeypatch.setattr("litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.05) + monkeypatch.setattr("litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.0) async def slow_project_lookup(*, where: Mapping[str, object]) -> LiteLLM_ProjectTable: assert where == {"project_id": _OWNED_PROJECT} - await asyncio.sleep(0.2) + await asyncio.sleep(0) + await asyncio.sleep(0) return LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) prisma_client.writer_db.litellm_projecttable.find_unique = slow_project_lookup From f4c161136a6279039143afd7df6d51a5e6b274d9 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 00:28:00 +0000 Subject: [PATCH 13/20] fix(proxy): check key project ownership after existing validation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/project_endpoints.py | 54 ++++----- .../key_management_endpoints.py | 87 ++++++++------ .../management/test_project_lifecycle.py | 106 ++++++++++++++++-- .../test_key_management_endpoints.py | 84 ++++++++++++-- 4 files changed, 254 insertions(+), 77 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index c90917d5341..796b4acdc46 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -801,33 +801,6 @@ async def update_project( **({"max_budget": None} if "max_budget" in data.model_fields_set and data.max_budget is None else {}), } - if data.team_id is not None: - current_project_record: Final = await _writer_project_table(prisma_client).find_unique( - where={"project_id": data.project_id} - ) - current_project: Final = ( - LiteLLM_ProjectTable.model_validate(record_to_dict(current_project_record)) - if current_project_record is not None - else None - ) - if current_project is not None and data.team_id != current_project.team_id: - mismatched_key_count: Final = await prisma_client.writer_db.litellm_verificationtoken.count( - where={ - "project_id": data.project_id, - "OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}], - } - ) - if mismatched_key_count > 0: - raise HTTPException( - status_code=400, - detail={ - "error": ( - f"Project {data.project_id} has {mismatched_key_count} key(s) that do not belong to " - f"team {data.team_id}. Detach or delete them before moving the project." - ) - }, - ) - if budget_updates and existing_project.budget_id: # Update existing budget await _budget_table(prisma_client).update( @@ -870,6 +843,33 @@ async def update_project( # Remove budget fields (following organization_endpoints.py pattern) update_data = _remove_budget_fields_from_project_data(update_data) + if data.team_id is not None: + current_project_record: Final = await _writer_project_table(prisma_client).find_unique( + where={"project_id": data.project_id} + ) + current_project: Final = ( + LiteLLM_ProjectTable.model_validate(record_to_dict(current_project_record)) + if current_project_record is not None + else None + ) + if current_project is not None and data.team_id != current_project.team_id: + mismatched_key_count: Final = await prisma_client.writer_db.litellm_verificationtoken.count( + where={ + "project_id": data.project_id, + "OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}], + } + ) + if mismatched_key_count > 0: + raise HTTPException( + status_code=400, + detail={ + "error": ( + f"Project {data.project_id} has {mismatched_key_count} key(s) that do not belong to " + f"team {data.team_id}. Detach or delete them before moving the project." + ) + }, + ) + # Update project updated_project: Final = await _project_table(prisma_client).update( where={"project_id": data.project_id}, diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 75c79fd77e1..df429b5acbc 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1358,13 +1358,6 @@ async def _common_key_generation_helper( ), ) - if data.project_id is not None and prisma_client is not None: - await _check_key_project_team( - project_id=data.project_id, - key_team_id=data.team_id, - prisma_client=prisma_client, - ) - # TODO: @ishaan-jaff: Migrate all budget tracking to use LiteLLM_BudgetTable _budget_id = data.budget_id if prisma_client is not None and data.soft_budget is not None: @@ -1560,6 +1553,13 @@ async def _common_key_generation_helper( prisma_client=prisma_client, ) + if data.project_id is not None and prisma_client is not None: + await _check_key_project_team( + project_id=data.project_id, + key_team_id=data.team_id, + prisma_client=prisma_client, + ) + response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key", llm_router=llm_router) response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response @@ -1828,15 +1828,10 @@ async def _check_key_project_team( prisma_client: PrismaClient, ) -> None: project_record: Final = await _writer_project_table(prisma_client).find_unique(where={"project_id": project_id}) - project_obj: Final = ( - LiteLLM_ProjectTable.model_validate(record_to_dict(project_record)) if project_record is not None else None - ) + if project_record is None: + return - if project_obj is None: - raise HTTPException( - status_code=404, - detail={"error": f"Project not found, project_id={project_id}"}, - ) + project_obj: Final = LiteLLM_ProjectTable.model_validate(record_to_dict(project_record)) if project_obj.team_id is None or project_obj.team_id == key_team_id: return @@ -2514,6 +2509,11 @@ async def _update_key_row_with_soft_budget( existing_key_row=existing_key_row, changed_by=changed_by, ) + await _check_key_project_team_on_mutation( + data=data, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + ) include_object_permission: Final[prisma.types.LiteLLM_VerificationTokenInclude] = {"object_permission": True} updated_row: Final = await tx.litellm_verificationtoken.update( where=key_where, @@ -2529,6 +2529,25 @@ async def _update_key_row_with_soft_budget( return result +async def _update_key_row_with_project_team_check( + prisma_client: PrismaClient, + key: str, + data: UpdateKeyRequest, + update_values: Mapping[str, object], + existing_key_row: LiteLLM_VerificationToken, +) -> _KeyUpdateResult | None: + key_update_data: Final = MappingProxyType({**update_values, "token": key}) + await _check_key_project_team_on_mutation( + data=data, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + ) + response: Final = await prisma_client.update_data(token=key, data=key_update_data) + if response is None: + return None + return cast("_KeyUpdateResult", response) # cast-ok: key update_data returns token and data + + async def prepare_key_update_data( data: UpdateKeyRequest | RegenerateKeyRequest, existing_key_row: LiteLLM_VerificationToken, @@ -2937,18 +2956,17 @@ async def _process_single_key_update( detail={"error": "Database not connected"}, ) - await _check_key_project_team_on_mutation( - data=key_request, - existing_key_row=existing_key_row, - prisma_client=prisma_client, - ) - update_values: Final = await _handle_update_object_permission( data_json=non_default_values, existing_key_row=existing_key_row, prisma_client=prisma_client, ) _data: Final = {**update_values, "token": key_request.key} + await _check_key_project_team_on_mutation( + data=key_request, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + ) response: Final[Mapping[str, object] | None] = cast( # cast-ok: every update_data branch returns a str-keyed dict "Mapping[str, object] | None", await prisma_client.update_data(token=key_request.key, data=_data), @@ -3616,12 +3634,6 @@ async def update_key_fn( if prisma_client is None: raise Exception("Not connected to DB!") - await _check_key_project_team_on_mutation( - data=data, - existing_key_row=existing_key_row, - prisma_client=prisma_client, - ) - update_values: Final = await _handle_update_object_permission( data_json=non_default_values, existing_key_row=existing_key_row, @@ -3638,7 +3650,13 @@ async def update_key_fn( changed_by=changed_by, ) if "soft_budget" in data.model_fields_set - else await prisma_client.update_data(token=key, data=MappingProxyType({**update_values, "token": key})) + else await _update_key_row_with_project_team_check( + prisma_client=prisma_client, + key=key, + data=data, + update_values=update_values, + existing_key_row=existing_key_row, + ) ) # Delete - key from cache, since it's been updated! @@ -5718,12 +5736,6 @@ async def _execute_virtual_key_regeneration( request=data if data is not None else RegenerateKeyRequest(), ), ) - if data is not None: - await _check_key_project_team_on_mutation( - data=data, - existing_key_row=key_in_db, - prisma_client=prisma_client, - ) update_values: Final = await _handle_update_object_permission( data_json=non_default_values, existing_key_row=key_in_db, @@ -5754,6 +5766,13 @@ async def _execute_virtual_key_regeneration( grace_period=data.grace_period if data else None, ) + if data is not None: + await _check_key_project_team_on_mutation( + data=data, + existing_key_row=key_in_db, + prisma_client=prisma_client, + ) + updated_token: Final[LiteLLM_VerificationToken | None] = await _prisma_table( VerificationTokenRepository(prisma_client) ).update( diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index 212229e3535..84cfebc72a4 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -35,7 +35,7 @@ def _key_rows(key: str) -> list[dict[str, JsonValue]]: def _cli_session_token( user_id: str, - team_id: str, + team_id: str | None, *, monkeypatch: pytest.MonkeyPatch, max_budget: float | None = None, @@ -44,14 +44,14 @@ def _cli_session_token( user: Final = LiteLLM_UserTable( user_id=user_id, user_role="internal_user", - teams=[team_id], + teams=[team_id] if team_id is not None else [], models=[], max_budget=max_budget, ) return ExperimentalUIJWTToken.get_cli_jwt_auth_token( user_info=user, team_id=team_id, - team_alias="ownership-team", + team_alias="ownership-team" if team_id is not None else None, max_budget=max_budget, ) @@ -303,14 +303,28 @@ def test_key_generation_rejects_missing_project_without_writing_key(ownership_ga model: Final = scenario.model() team: Final = scenario.team(models=[model]) missing_project_id: Final = f"missing-{uuid4()}" - response: Final = ownership_gateway.request( + key_generation: Final = ownership_gateway.request( + "POST", + "/key/generate", + {"team_id": team, "project_id": missing_project_id, "models": [model]}, + ) + + assert key_generation.status_code == 404, key_generation.text + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (missing_project_id,), + ) + == [] + ) + + service_account_generation: Final = ownership_gateway.request( "POST", "/key/service-account/generate", {"team_id": team, "project_id": missing_project_id, "models": [model]}, ) - assert response.status_code == 404, response.text - assert "Project not found" in response.text + assert service_account_generation.status_code == 500, service_account_generation.text assert ( read_rows( 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', @@ -341,8 +355,7 @@ def test_key_regenerate_routes_reject_missing_project_without_changing_key(owner ) for response in responses: - assert response.status_code == 404, response.text - assert "Project not found" in response.text + assert response.status_code == 500, response.text assert _key_rows(key) == before assert ( read_rows( @@ -361,6 +374,83 @@ def test_key_regenerate_routes_reject_missing_project_without_changing_key(owner assert chat.status_code == 200, chat.text +def test_key_generate_nonmember_organization_error_precedes_project_ownership( + ownership_gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + organization_id: Final = scenario.organization() + owner_team: Final = scenario.team(models=[model]) + project: Final = scenario.project(owner_team, models=[model]) + caller_id: Final = scenario.user(user_role="internal_user") + caller_token: Final = _cli_session_token(caller_id, None, monkeypatch=monkeypatch) + response: Final = ownership_gateway.request( + "POST", + "/key/generate", + { + "project_id": project, + "organization_id": organization_id, + "models": [model], + }, + key=caller_token, + ) + + assert response.status_code == 403, response.text + assert f"Caller is not a member of organization_id={organization_id}" in response.text + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (project,), + ) + == [] + ) + + +def test_key_generate_duplicate_alias_error_precedes_project_ownership( + ownership_gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + caller_team: Final = scenario.team(models=[model]) + owner_team: Final = scenario.team(models=[model]) + project: Final = scenario.project(owner_team, models=[model]) + key_alias: Final = f"duplicate-{uuid4()}" + existing_key: Final = scenario.key(team_id=caller_team, key_alias=key_alias, models=[model]) + caller_id: Final = scenario.member(caller_team, role="admin") + caller_token: Final = _cli_session_token(caller_id, caller_team, monkeypatch=monkeypatch) + existing_alias_rows: Final = read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE key_alias = %s', + (key_alias,), + ) + response: Final = ownership_gateway.request( + "POST", + "/key/generate", + { + "team_id": caller_team, + "project_id": project, + "key_alias": key_alias, + "models": [model], + }, + key=caller_token, + ) + + assert response.status_code == 400, response.text + assert f"Key with alias '{key_alias}' already exists" in response.text + assert len(existing_alias_rows) == 1 + assert read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE key_alias = %s', + (key_alias,), + ) == existing_alias_rows + assert len(_key_rows(existing_key)) == 1 + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (project,), + ) + == [] + ) + + def test_key_generate_budget_ceiling_precedes_project_ownership( ownership_gateway: Gateway, monkeypatch: pytest.MonkeyPatch ) -> None: diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 37d004b251b..8127d095755 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -20637,6 +20637,73 @@ async def test_key_generation_budget_ceiling_precedes_project_ownership( assert error.value.detail == {"error": "max_budget (10.0) cannot exceed the caller's own max_budget (5.0)."} +@pytest.mark.asyncio +async def test_key_generation_organization_membership_error_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM + ) + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + mock_prisma_client.db.litellm_usertable = MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=LiteLLM_UserTable(user_id="key-owner", organization_memberships=[]) + ) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=GenerateKeyRequest( + project_id=_OWNED_PROJECT, + organization_id="org-not-member", + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-user", + user_id="key-owner", + is_session_token=True, + ), + litellm_changed_by=None, + team_table=None, + ) + + assert error.value.status_code == 403 + assert error.value.detail == "Caller is not a member of organization_id=org-not-member" + + +@pytest.mark.asyncio +async def test_key_generation_duplicate_alias_error_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + key_alias: Final = "duplicate-alias" + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM + ) + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock( + return_value=LiteLLM_VerificationToken(token="existing-token", key_alias=key_alias) + ) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + + with pytest.raises(ProxyException) as error: + await _common_key_generation_helper( + data=GenerateKeyRequest( + project_id=_OWNED_PROJECT, + team_id=_OWNERSHIP_KEY_TEAM, + key_alias=key_alias, + ), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + team_table=None, + ) + + assert error.value.code == "400" + assert error.value.message == ( + f"Key with alias '{key_alias}' already exists. Unique key aliases across all keys are required." + ) + mock_prisma_client.insert_data.assert_not_awaited() + + @pytest.mark.parametrize( ("key_team_id", "expected_status"), [(_OWNERSHIP_PROJECT_TEAM, 200), (_OWNERSHIP_KEY_TEAM, 400)], @@ -21191,18 +21258,19 @@ async def test_regenerate_output_estimate_admin_error_precedes_project_ownership @pytest.mark.asyncio -async def test_key_project_team_validation_uses_project_missing_404() -> None: +async def test_key_project_team_validation_allows_missing_project() -> None: prisma_client: Final = MagicMock() prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock(return_value=None) - with pytest.raises(HTTPException) as error: - await _check_key_project_team( - project_id=_OWNED_PROJECT, - key_team_id="team-a", - prisma_client=prisma_client, - ) + await _check_key_project_team( + project_id=_OWNED_PROJECT, + key_team_id="team-a", + prisma_client=prisma_client, + ) - assert error.value.status_code == 404 + prisma_client.writer_db.litellm_projecttable.find_unique.assert_awaited_once_with( + where={"project_id": _OWNED_PROJECT} + ) @pytest.mark.asyncio From 1d24f7fdbc7d92977485b9521ef08d8880952caf Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 00:43:46 +0000 Subject: [PATCH 14/20] fix(proxy): run key project ownership right before the key row write Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../key_management_endpoints.py | 28 +++++------ .../management/test_project_lifecycle.py | 22 ++++++++- .../test_key_management_endpoints.py | 48 +++++++++++++++++-- 3 files changed, 78 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index df429b5acbc..c68316871f4 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1553,13 +1553,6 @@ async def _common_key_generation_helper( prisma_client=prisma_client, ) - if data.project_id is not None and prisma_client is not None: - await _check_key_project_team( - project_id=data.project_id, - key_team_id=data.team_id, - prisma_client=prisma_client, - ) - response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key", llm_router=llm_router) response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response @@ -4922,6 +4915,13 @@ async def generate_key_helper_fn( # the LiteLLM_VerificationToken table will increase in size if we don't do this check return user_data + if project_id is not None: + await _check_key_project_team( + project_id=project_id, + key_team_id=team_id, + prisma_client=prisma_client, + ) + ## CREATE KEY verbose_proxy_logger.debug( "prisma_client: Creating Key= %s", @@ -5742,6 +5742,13 @@ async def _execute_virtual_key_regeneration( prisma_client=prisma_client, ) update_data.update(update_values) + if data is not None: + await _check_key_project_team_on_mutation( + data=data, + existing_key_row=key_in_db, + prisma_client=prisma_client, + ) + jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data) # Snapshot before the token update: the FK cascade rewrites mapping rows to the new hash, @@ -5766,13 +5773,6 @@ async def _execute_virtual_key_regeneration( grace_period=data.grace_period if data else None, ) - if data is not None: - await _check_key_project_team_on_mutation( - data=data, - existing_key_row=key_in_db, - prisma_client=prisma_client, - ) - updated_token: Final[LiteLLM_VerificationToken | None] = await _prisma_table( VerificationTokenRepository(prisma_client) ).update( diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index 84cfebc72a4..e86a989271a 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -33,6 +33,20 @@ def _key_rows(key: str) -> list[dict[str, JsonValue]]: ) +def _deleted_key_rows(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT token FROM "LiteLLM_DeletedVerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + +def _deprecated_key_rows(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT token FROM "LiteLLM_DeprecatedVerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + def _cli_session_token( user_id: str, team_id: str | None, @@ -289,13 +303,19 @@ def test_key_regenerate_rejects_foreign_project_without_changing_key(ownership_g key: Final = string_value(JSON_OBJECT.validate_json(generated.content)["key"]) scenario.cleanups.callback(scenario.delete_key, key) before: Final = _key_rows(key) + before_deleted: Final = _deleted_key_rows(key) + before_deprecated: Final = _deprecated_key_rows(key) assert len(before) == 1 + assert before_deleted == [] + assert before_deprecated == [] response: Final = ownership_gateway.request( - "POST", f"/key/{key}/regenerate", {"project_id": project_b} + "POST", f"/key/{key}/regenerate", {"project_id": project_b, "grace_period": "1h"} ) _discard_unexpected_key(ownership_gateway, response) assert response.status_code == 400, response.text assert _key_rows(key) == before + assert _deleted_key_rows(key) == before_deleted + assert _deprecated_key_rows(key) == before_deprecated def test_key_generation_rejects_missing_project_without_writing_key(ownership_gateway: Gateway) -> None: diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 8127d095755..7e5b4987953 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -20671,6 +20671,37 @@ async def test_key_generation_organization_membership_error_precedes_project_own assert error.value.detail == "Caller is not a member of organization_id=org-not-member" +@pytest.mark.asyncio +async def test_key_generation_premium_permission_error_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM + ) + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False) + + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=GenerateKeyRequest( + project_id=_OWNED_PROJECT, + team_id=_OWNERSHIP_KEY_TEAM, + permissions={"get_spend_routes": True}, + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + ), + litellm_changed_by=None, + team_table=None, + ) + + assert error.value.status_code == 500 + assert error.value.detail == {"error": "Internal Server Error."} + mock_prisma_client.insert_data.assert_not_awaited() + + @pytest.mark.asyncio async def test_key_generation_duplicate_alias_error_precedes_project_ownership( monkeypatch: pytest.MonkeyPatch, @@ -21175,6 +21206,12 @@ async def test_regenerate_checks_project_team_ownership( mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id="team-b") ) + deleted_history_table: Final = MagicMock() + deleted_history_table.create_many = AsyncMock() + mock_prisma_client.db.litellm_deletedverificationtoken = deleted_history_table + deprecated_key_table: Final = MagicMock() + deprecated_key_table.upsert = AsyncMock() + mock_prisma_client.db.litellm_deprecatedverificationtoken = deprecated_key_table user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") async def regenerate() -> None: @@ -21184,10 +21221,6 @@ async def test_regenerate_checks_project_team_ownership( 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, @@ -21198,7 +21231,10 @@ async def test_regenerate_checks_project_team_ownership( key_in_db=existing_key, hashed_api_key="abc123", key="abc123", - data=RegenerateKeyRequest(project_id=_OWNED_PROJECT), + data=RegenerateKeyRequest( + project_id=_OWNED_PROJECT, + grace_period="1h" if expected_status is not None else None, + ), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, user_api_key_cache=user_api_key_cache, @@ -21210,6 +21246,8 @@ async def test_regenerate_checks_project_team_ownership( await regenerate() assert error.value.status_code == expected_status assert "belongs to team team-b, but the key belongs to team-a" in str(error.value.detail) + deleted_history_table.create_many.assert_not_awaited() + deprecated_key_table.upsert.assert_not_awaited() else: await regenerate() From a27bb1088e630cd295a65ee44d6514c08c156818 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 01:58:34 +0000 Subject: [PATCH 15/20] fix(proxy): reject key project ownership before permission and budget writes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/project_endpoints.py | 86 ++-- .../key_management_endpoints.py | 113 ++++-- .../management/test_project_lifecycle.py | 102 ++++- .../test_project_endpoints_prisma.py | 19 +- .../test_key_management_endpoints.py | 384 +++++++++++------- 5 files changed, 462 insertions(+), 242 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index 796b4acdc46..5a754141063 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -790,8 +790,16 @@ async def update_project( data, existing_project, _router_access_group_names(llm_router) ) + object_permission_data: Final = ( + data.object_permission.model_dump(exclude_none=True) if data.object_permission is not None else None + ) + object_permission_payload: Final = ( + _OBJECT_PERMISSION_PAYLOAD.validate_python(object_permission_data) if object_permission_data else None + ) + # Prepare update data update_data = _jsonified(prisma_client, data.model_dump(exclude_none=True, exclude={"project_id"})) + update_data.pop("object_permission", None) update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name # Handle budget updates @@ -801,48 +809,6 @@ async def update_project( **({"max_budget": None} if "max_budget" in data.model_fields_set and data.max_budget is None else {}), } - if budget_updates and existing_project.budget_id: - # Update existing budget - await _budget_table(prisma_client).update( - where={"budget_id": existing_project.budget_id}, - data={ - **budget_updates, - "updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name, - }, - ) - # Remove budget fields from project update - for field in budget_updates.keys(): - update_data.pop(field, None) - - # Handle object permissions - if "object_permission" in update_data: - object_permission_data = update_data.pop("object_permission") - if object_permission_data: - object_permission_payload: Final = _OBJECT_PERMISSION_PAYLOAD.validate_python(object_permission_data) - if existing_project.object_permission_id: - # Update existing permission - await _object_permission_table(prisma_client).update( - where={"object_permission_id": existing_project.object_permission_id}, - data=object_permission_payload, - ) - else: - # Create new permission - created_permission: Final = await _object_permission_table(prisma_client).create( - data=object_permission_payload, - ) - update_data["object_permission_id"] = created_permission.object_permission_id - - # Handle metadata fields - for field in LiteLLM_ManagementEndpoint_MetadataFields: - if field in update_data: - existing_metadata = update_data.get("metadata") - metadata_dict: dict[str, object] = existing_metadata if isinstance(existing_metadata, dict) else {} - metadata_dict[field] = update_data.pop(field) - update_data["metadata"] = metadata_dict - - # Remove budget fields (following organization_endpoints.py pattern) - update_data = _remove_budget_fields_from_project_data(update_data) - if data.team_id is not None: current_project_record: Final = await _writer_project_table(prisma_client).find_unique( where={"project_id": data.project_id} @@ -870,6 +836,42 @@ async def update_project( }, ) + if budget_updates and existing_project.budget_id: + # Update existing budget + await _budget_table(prisma_client).update( + where={"budget_id": existing_project.budget_id}, + data={ + **budget_updates, + "updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name, + }, + ) + # Remove budget fields from project update + for field in budget_updates.keys(): + update_data.pop(field, None) + + if object_permission_payload is not None: + if existing_project.object_permission_id: + await _object_permission_table(prisma_client).update( + where={"object_permission_id": existing_project.object_permission_id}, + data=object_permission_payload, + ) + else: + created_permission: Final = await _object_permission_table(prisma_client).create( + data=object_permission_payload, + ) + update_data["object_permission_id"] = created_permission.object_permission_id + + # Handle metadata fields + for field in LiteLLM_ManagementEndpoint_MetadataFields: + if field in update_data: + existing_metadata = update_data.get("metadata") + metadata_dict: dict[str, object] = existing_metadata if isinstance(existing_metadata, dict) else {} + metadata_dict[field] = update_data.pop(field) + update_data["metadata"] = metadata_dict + + # Remove budget fields (following organization_endpoints.py pattern) + update_data = _remove_budget_fields_from_project_data(update_data) + # Update project updated_project: Final = await _project_table(prisma_client).update( where={"project_id": data.project_id}, diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index c68316871f4..3b5b7264d1a 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -113,10 +113,11 @@ from litellm.proxy.management_helpers.access_group_key_sync import ( ) from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at from litellm.proxy.management_helpers.object_permission_utils import ( + ObjectPermissionUpsert, _set_object_permission, attach_object_permission_to_dict, - handle_update_object_permission_common, invalidate_cached_object_permissions, + prepare_object_permission_upsert, validate_key_mcp_servers_against_team, validate_key_search_tools_against_team, validate_key_vector_stores_against_team, @@ -139,6 +140,7 @@ from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.config_repository import ConfigParam, ConfigRepository from litellm.repositories.credentials_repository import CredentialsRepository from litellm.repositories.model_repository import ModelRepository +from litellm.repositories.object_permission_repository import ObjectPermissionRepository from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ( DeletedVerificationTokenRepository, @@ -2502,11 +2504,6 @@ async def _update_key_row_with_soft_budget( existing_key_row=existing_key_row, changed_by=changed_by, ) - await _check_key_project_team_on_mutation( - data=data, - existing_key_row=existing_key_row, - prisma_client=prisma_client, - ) include_object_permission: Final[prisma.types.LiteLLM_VerificationTokenInclude] = {"object_permission": True} updated_row: Final = await tx.litellm_verificationtoken.update( where=key_where, @@ -2522,19 +2519,12 @@ async def _update_key_row_with_soft_budget( return result -async def _update_key_row_with_project_team_check( +async def _update_key_row( prisma_client: PrismaClient, key: str, - data: UpdateKeyRequest, update_values: Mapping[str, object], - existing_key_row: LiteLLM_VerificationToken, ) -> _KeyUpdateResult | None: key_update_data: Final = MappingProxyType({**update_values, "token": key}) - await _check_key_project_team_on_mutation( - data=data, - existing_key_row=existing_key_row, - prisma_client=prisma_client, - ) response: Final = await prisma_client.update_data(token=key, data=key_update_data) if response is None: return None @@ -2647,27 +2637,42 @@ async def prepare_key_update_data( return non_default_values -async def _handle_update_object_permission( - data_json: dict, +async def _prepare_key_update_object_permission( + object_permission_data: object, existing_key_row: LiteLLM_VerificationToken, prisma_client: PrismaClient, -) -> dict: - """Persist the requested object permission row and swap it for its id, only after the key policy allowed the write.""" - if "object_permission" not in data_json: - return data_json +) -> ObjectPermissionUpsert | None: + if object_permission_data is None: + return None - object_permission_id: Final = await handle_update_object_permission_common( - data_json=data_json, + parsed_object_permission: Final[object] = ( + json.loads(object_permission_data) if isinstance(object_permission_data, str) else object_permission_data + ) + permission_data: Final[dict[str, object]] = ( + TypeAdapter(dict[str, object]).validate_python(parsed_object_permission) + if isinstance(parsed_object_permission, dict) + else {} + ) + return await prepare_object_permission_upsert( + new_object_permission=permission_data, existing_object_permission_id=existing_key_row.object_permission_id, prisma_client=prisma_client, ) - # Add the object_permission_id to data_json if one was created/updated - if object_permission_id is not None: - data_json["object_permission_id"] = object_permission_id - verbose_proxy_logger.debug("updated object_permission_id: %s", object_permission_id) - return data_json +async def _write_prepared_key_update_object_permission( + data_json: Mapping[str, object], + upsert: ObjectPermissionUpsert | None, + prisma_client: PrismaClient, +) -> Mapping[str, object]: + if upsert is None: + return data_json + + await ObjectPermissionRepository(prisma_client).table.upsert( + where={"object_permission_id": upsert.object_permission_id}, + data={"create": upsert.record, "update": upsert.record}, + ) + return MappingProxyType({**data_json, "object_permission_id": upsert.object_permission_id}) def is_different_team(data: UpdateKeyRequest, existing_key_row: LiteLLM_VerificationToken) -> bool: @@ -2949,17 +2954,25 @@ async def _process_single_key_update( detail={"error": "Database not connected"}, ) - update_values: Final = await _handle_update_object_permission( - data_json=non_default_values, + object_permission_upsert: Final = await _prepare_key_update_object_permission( + object_permission_data=non_default_values.get("object_permission"), existing_key_row=existing_key_row, prisma_client=prisma_client, ) - _data: Final = {**update_values, "token": key_request.key} + key_update_values: Final = MappingProxyType( + {field: value for field, value in non_default_values.items() if field != "object_permission"} + ) await _check_key_project_team_on_mutation( data=key_request, existing_key_row=existing_key_row, prisma_client=prisma_client, ) + update_values: Final = await _write_prepared_key_update_object_permission( + data_json=key_update_values, + upsert=object_permission_upsert, + prisma_client=prisma_client, + ) + _data: Final = {**update_values, "token": key_request.key} response: Final[Mapping[str, object] | None] = cast( # cast-ok: every update_data branch returns a str-keyed dict "Mapping[str, object] | None", await prisma_client.update_data(token=key_request.key, data=_data), @@ -3627,11 +3640,24 @@ async def update_key_fn( if prisma_client is None: raise Exception("Not connected to DB!") - update_values: Final = await _handle_update_object_permission( - data_json=non_default_values, + object_permission_upsert: Final = await _prepare_key_update_object_permission( + object_permission_data=non_default_values.get("object_permission"), existing_key_row=existing_key_row, prisma_client=prisma_client, ) + key_update_values: Final = MappingProxyType( + {field: value for field, value in non_default_values.items() if field != "object_permission"} + ) + await _check_key_project_team_on_mutation( + data=data, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + ) + update_values: Final = await _write_prepared_key_update_object_permission( + data_json=key_update_values, + upsert=object_permission_upsert, + prisma_client=prisma_client, + ) changed_by: Final = user_api_key_dict.user_id or litellm_proxy_admin_name response: Final = ( await _update_key_row_with_soft_budget( @@ -3643,12 +3669,10 @@ async def update_key_fn( changed_by=changed_by, ) if "soft_budget" in data.model_fields_set - else await _update_key_row_with_project_team_check( + else await _update_key_row( prisma_client=prisma_client, key=key, - data=data, update_values=update_values, - existing_key_row=existing_key_row, ) ) @@ -3657,7 +3681,7 @@ async def update_key_fn( await invalidate_cached_object_permissions( object_permission_ids=( existing_key_row.object_permission_id, - non_default_values.get("object_permission_id"), + update_values.get("object_permission_id"), ), user_api_key_cache=user_api_key_cache, ) @@ -5710,7 +5734,6 @@ async def _execute_virtual_key_regeneration( new_token: Final = await get_new_token(data=data) new_token_hash: Final = hash_token(new_token) new_token_key_name: Final = abbreviate_api_key(api_key=new_token) - update_data = {"token": new_token_hash, "key_name": new_token_key_name} non_default_values = {} if data is not None: @@ -5736,18 +5759,28 @@ async def _execute_virtual_key_regeneration( request=data if data is not None else RegenerateKeyRequest(), ), ) - update_values: Final = await _handle_update_object_permission( - data_json=non_default_values, + object_permission_upsert: Final = await _prepare_key_update_object_permission( + object_permission_data=non_default_values.get("object_permission"), existing_key_row=key_in_db, prisma_client=prisma_client, ) - update_data.update(update_values) + key_update_values: Final = MappingProxyType( + {field: value for field, value in non_default_values.items() if field != "object_permission"} + ) if data is not None: await _check_key_project_team_on_mutation( data=data, existing_key_row=key_in_db, prisma_client=prisma_client, ) + update_values: Final = await _write_prepared_key_update_object_permission( + data_json=key_update_values, + upsert=object_permission_upsert, + prisma_client=prisma_client, + ) + update_data: Final = MappingProxyType( + {"token": new_token_hash, "key_name": new_token_key_name, **update_values} + ) jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data) diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index e86a989271a..25b8da7875d 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -33,6 +33,49 @@ def _key_rows(key: str) -> list[dict[str, JsonValue]]: ) +def _key_permission_and_budget_ids(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT object_permission_id, budget_id FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + +def _object_permission_rows(permission_id: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT * FROM "LiteLLM_ObjectPermissionTable" WHERE object_permission_id = %s', + (permission_id,), + ) + + +def _budget_rows(budget_id: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT to_jsonb(b) AS row FROM "LiteLLM_BudgetTable" AS b WHERE budget_id = %s', + (budget_id,), + ) + + +def _clear_key_object_permission(key: str, permission_id: str) -> None: + write_rows( + 'UPDATE "LiteLLM_VerificationToken" SET object_permission_id = NULL WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + write_rows( + 'DELETE FROM "LiteLLM_ObjectPermissionTable" WHERE object_permission_id = %s', + (permission_id,), + ) + + +def _clear_project_object_permission(project_id: str, permission_id: str) -> None: + write_rows( + 'UPDATE "LiteLLM_ProjectTable" SET object_permission_id = NULL WHERE project_id = %s', + (project_id,), + ) + write_rows( + 'DELETE FROM "LiteLLM_ObjectPermissionTable" WHERE object_permission_id = %s', + (permission_id,), + ) + + def _deleted_key_rows(key: str) -> list[dict[str, JsonValue]]: return read_rows( 'SELECT token FROM "LiteLLM_DeletedVerificationToken" WHERE token = %s', @@ -253,7 +296,22 @@ def test_key_update_rejects_team_change_and_allows_unchanged_values(ownership_ga team_a: Final = scenario.team(models=[model]) team_b: Final = scenario.team(models=[model]) project_a: Final = scenario.project(team_a, models=[model]) - key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model], soft_budget=3.0) + permission_seed: Final = ownership_gateway.request( + "POST", + "/key/update", + {"key": key, "object_permission": {"vector_stores": ["existing-store"]}}, + ) + assert permission_seed.status_code == 200, permission_seed.text + key_permission_state: Final = _key_permission_and_budget_ids(key) + assert len(key_permission_state) == 1 + permission_id: Final = string_value(key_permission_state[0]["object_permission_id"]) + budget_id: Final = string_value(key_permission_state[0]["budget_id"]) + scenario.cleanups.callback(_clear_key_object_permission, key, permission_id) + permission_before: Final = _object_permission_rows(permission_id) + budget_before: Final = _budget_rows(budget_id) + assert len(permission_before) == 1 + assert len(budget_before) == 1 aliased: Final = ownership_gateway.request("POST", "/key/update", {"key": key, "key_alias": "updated"}) assert aliased.status_code == 200, aliased.text after_alias: Final = _key_rows(key) @@ -267,9 +325,20 @@ def test_key_update_rejects_team_change_and_allows_unchanged_values(ownership_ga assert after_unchanged_team == after_alias before_reassignment: Final = _key_rows(key) assert len(before_reassignment) == 1 - reassigned: Final = ownership_gateway.request("POST", "/key/update", {"key": key, "team_id": team_b}) + reassigned: Final = ownership_gateway.request( + "POST", + "/key/update", + { + "key": key, + "team_id": team_b, + "soft_budget": 5.0, + "object_permission": {"vector_stores": ["replacement-store"]}, + }, + ) assert reassigned.status_code == 400, reassigned.text assert _key_rows(key) == before_reassignment + assert _object_permission_rows(permission_id) == permission_before + assert _budget_rows(budget_id) == budget_before def test_key_update_can_detach_project_and_change_team(ownership_gateway: Gateway) -> None: @@ -588,16 +657,41 @@ def test_project_update_rejects_moving_project_with_attached_key(ownership_gatew model: Final = scenario.model() team_a: Final = scenario.team(models=[model]) team_b: Final = scenario.team(models=[model]) - attached_project: Final = scenario.project(team_a, models=[model]) + budget: Final = scenario.budget(max_budget=3) + attached_project: Final = scenario.project(team_a, budget_id=budget, models=[model]) key: Final = scenario.key(team_id=team_a, project_id=attached_project, models=[model]) + permission_seed: Final = ownership_gateway.request( + "POST", + "/project/update", + {"project_id": attached_project, "object_permission": {"vector_stores": ["existing-store"]}}, + ) + assert permission_seed.status_code == 200, permission_seed.text + project_permission: Final = read_rows( + 'SELECT object_permission_id FROM "LiteLLM_ProjectTable" WHERE project_id = %s', + (attached_project,), + ) + assert len(project_permission) == 1 + permission_id: Final = string_value(project_permission[0]["object_permission_id"]) + scenario.cleanups.callback(_clear_project_object_permission, attached_project, permission_id) project_before: Final = _project_rows(attached_project) key_before: Final = _key_rows(key) + permission_before: Final = _object_permission_rows(permission_id) + budget_before: Final = _budget_rows(budget) moved_with_key: Final = ownership_gateway.request( - "POST", "/project/update", {"project_id": attached_project, "team_id": team_b} + "POST", + "/project/update", + { + "project_id": attached_project, + "team_id": team_b, + "max_budget": 11, + "object_permission": {"vector_stores": ["replacement-store"]}, + }, ) assert moved_with_key.status_code == 400, moved_with_key.text assert _project_rows(attached_project) == project_before assert _key_rows(key) == key_before + assert _object_permission_rows(permission_id) == permission_before + assert _budget_rows(budget) == budget_before def test_project_update_allows_moving_project_without_keys(ownership_gateway: Gateway) -> None: diff --git a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py index 98bb2969201..cd185179ee9 100644 --- a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py +++ b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py @@ -1277,6 +1277,15 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists( destination_team_id: Final = "team-b" mock_prisma: Final = _project_update_mocks(monkeypatch, {}) mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = "team-a" + mock_prisma.db.litellm_projecttable.find_unique.return_value.budget_id = "budget-project" + mock_prisma.db.litellm_projecttable.find_unique.return_value.object_permission_id = "permission-project" + budget_table: Final = mock.MagicMock() + budget_table.update = mock.AsyncMock() + permission_table: Final = mock.MagicMock() + permission_table.update = mock.AsyncMock() + permission_table.create = mock.AsyncMock() + mock_prisma.db.litellm_budgettable = budget_table + mock_prisma.db.litellm_objectpermissiontable = permission_table mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock( return_value=LiteLLM_TeamTable(team_id=destination_team_id) ) @@ -1293,7 +1302,12 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists( mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(side_effect=count_teamless_keys) with pytest.raises(ProxyException) as error: - await _run_project_update(project_id, team_id=destination_team_id) + await _run_project_update( + project_id, + team_id=destination_team_id, + max_budget=50, + object_permission={"vector_stores": ["replacement-store"]}, + ) expected_detail: Final = { "error": ( @@ -1305,6 +1319,9 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists( assert expected_detail["error"] in error.value.message mock_prisma.writer_db.litellm_verificationtoken.count.assert_awaited_once() mock_prisma.db.litellm_verificationtoken.count.assert_not_awaited() + budget_table.update.assert_not_awaited() + permission_table.update.assert_not_awaited() + permission_table.create.assert_not_awaited() mock_prisma.db.litellm_projecttable.update.assert_not_awaited() diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 7e5b4987953..80c51d20f1e 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -1076,208 +1076,146 @@ async def test_key_generation_with_mcp_tool_permissions(monkeypatch): @pytest.mark.asyncio -async def test_key_update_object_permissions_existing_permission(): - """ - Test updating object permissions when a key already has an existing object_permission_id. - - This test verifies that when updating vector stores for a key that already has an - object_permission_id, the existing LiteLLM_ObjectPermissionTable record is updated - with the new permissions and the object_permission_id remains the same. - """ +async def test_key_update_prepares_existing_object_permission_without_writing(): from unittest.mock import AsyncMock, MagicMock - import pytest - - from litellm.proxy._types import ( - LiteLLM_ObjectPermissionBase, - LiteLLM_VerificationToken, - ) + from litellm.proxy._types import LiteLLM_ObjectPermissionBase, LiteLLM_VerificationToken from litellm.proxy.management_endpoints.key_management_endpoints import ( - _handle_update_object_permission, + _prepare_key_update_object_permission, + _write_prepared_key_update_object_permission, ) mock_prisma_client = AsyncMock() - - # Mock existing key with object_permission_id existing_key_row = LiteLLM_VerificationToken( token="test_token_hash", object_permission_id="existing_perm_id_123", user_id="user123", team_id=None, ) - - # Mock existing object permission record - existing_object_permission = MagicMock() - existing_object_permission.model_dump.return_value = { + existing_permission = MagicMock() + existing_permission.model_dump.return_value = { "object_permission_id": "existing_perm_id_123", - "vector_stores": ["old_store_1", "old_store_2"], + "vector_stores": ["old_store"], } + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_permission) + permission_upsert = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=existing_object_permission + object_permission_data = LiteLLM_ObjectPermissionBase(vector_stores=["new_store"]).model_dump( + exclude_unset=True, exclude_none=True ) - - # Mock upsert operation - updated_permission = MagicMock() - updated_permission.object_permission_id = "existing_perm_id_123" - mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( - return_value=updated_permission - ) - - # Test data with new object permission - data_json = { - "object_permission": LiteLLM_ObjectPermissionBase( - vector_stores=["new_store_1", "new_store_2", "new_store_3"] - ).model_dump(exclude_unset=True, exclude_none=True), - "user_id": "user123", - } - - # Call the function - result = await _handle_update_object_permission( - data_json=data_json, + upsert = await _prepare_key_update_object_permission( + object_permission_data=object_permission_data, existing_key_row=existing_key_row, prisma_client=mock_prisma_client, ) - # Verify the object_permission was removed from data_json and object_permission_id was set - assert "object_permission" not in result - assert result["object_permission_id"] == "existing_perm_id_123" - - # Verify database operations were called correctly - mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( + assert upsert is not None + assert upsert.object_permission_id == "existing_perm_id_123" + assert upsert.record["vector_stores"] == ["new_store"] + permission_upsert.assert_not_awaited() + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_awaited_once_with( where={"object_permission_id": "existing_perm_id_123"} ) - mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() + + result = await _write_prepared_key_update_object_permission( + data_json={"user_id": "user123"}, + upsert=upsert, + prisma_client=mock_prisma_client, + ) + + assert result == {"user_id": "user123", "object_permission_id": "existing_perm_id_123"} + permission_upsert.assert_awaited_once_with( + where={"object_permission_id": "existing_perm_id_123"}, + data={"create": upsert.record, "update": upsert.record}, + ) @pytest.mark.asyncio -async def test_key_update_object_permissions_no_existing_permission(): - """ - Test creating object permissions when a key has no existing object_permission_id. +async def test_key_update_prepares_json_object_permission_and_upserts_new_row(): + import json + from unittest.mock import AsyncMock - This test verifies that when updating object permissions for a key that has - object_permission_id set to None, a new entry is created in the - LiteLLM_ObjectPermissionTable and the key is updated with the new object_permission_id. - """ - from unittest.mock import AsyncMock, MagicMock - - import pytest - - from litellm.proxy._types import ( - LiteLLM_ObjectPermissionBase, - LiteLLM_VerificationToken, - ) + from litellm.proxy._types import LiteLLM_VerificationToken from litellm.proxy.management_endpoints.key_management_endpoints import ( - _handle_update_object_permission, + _prepare_key_update_object_permission, + _write_prepared_key_update_object_permission, ) mock_prisma_client = AsyncMock() - - existing_key_row_no_perm = LiteLLM_VerificationToken( + existing_key_row = LiteLLM_VerificationToken( token="test_token_hash_2", object_permission_id=None, user_id="user456", team_id=None, ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) + permission_upsert = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert - # Mock find_unique to return None (no existing permission) - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=None - ) - - # Mock upsert to create new record - new_permission = MagicMock() - new_permission.object_permission_id = "new_perm_id_456" - mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( - return_value=new_permission - ) - - data_json = { - "object_permission": LiteLLM_ObjectPermissionBase( - vector_stores=["brand_new_store"] - ).model_dump(exclude_unset=True, exclude_none=True), - "user_id": "user456", - } - - result = await _handle_update_object_permission( - data_json=data_json, - existing_key_row=existing_key_row_no_perm, + upsert = await _prepare_key_update_object_permission( + object_permission_data=json.dumps({"vector_stores": ["brand_new_store"]}), + existing_key_row=existing_key_row, prisma_client=mock_prisma_client, ) - # Verify new object_permission_id was set - assert "object_permission" not in result - assert result["object_permission_id"] == "new_perm_id_456" - # Verify upsert was called to create new record - mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() + assert upsert is not None + assert upsert.record["vector_stores"] == ["brand_new_store"] + permission_upsert.assert_not_awaited() + result = await _write_prepared_key_update_object_permission( + data_json={}, + upsert=upsert, + prisma_client=mock_prisma_client, + ) + + assert result["object_permission_id"] == upsert.object_permission_id + permission_upsert.assert_awaited_once_with( + where={"object_permission_id": upsert.object_permission_id}, + data={"create": upsert.record, "update": upsert.record}, + ) @pytest.mark.asyncio -async def test_key_update_object_permissions_missing_permission_record(): - """ - Test creating object permissions when existing object_permission_id record is not found. +async def test_key_update_recreates_missing_object_permission_with_existing_id(): + from unittest.mock import AsyncMock - This test verifies that when updating object permissions for a key that has an - object_permission_id but the corresponding record cannot be found in the database, - a new entry is created in the LiteLLM_ObjectPermissionTable with the new permissions. - """ - from unittest.mock import AsyncMock, MagicMock - - import pytest - - from litellm.proxy._types import ( - LiteLLM_ObjectPermissionBase, - LiteLLM_VerificationToken, - ) + from litellm.proxy._types import LiteLLM_ObjectPermissionBase, LiteLLM_VerificationToken from litellm.proxy.management_endpoints.key_management_endpoints import ( - _handle_update_object_permission, + _prepare_key_update_object_permission, + _write_prepared_key_update_object_permission, ) mock_prisma_client = AsyncMock() - - existing_key_row_missing_perm = LiteLLM_VerificationToken( + existing_key_row = LiteLLM_VerificationToken( token="test_token_hash_3", object_permission_id="missing_perm_id_789", user_id="user789", team_id=None, ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) + permission_upsert = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert - # Mock find_unique to return None (permission record not found) - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=None - ) - - # Mock upsert to create new record - new_permission = MagicMock() - new_permission.object_permission_id = "recreated_perm_id_789" - mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( - return_value=new_permission - ) - - data_json = { - "object_permission": LiteLLM_ObjectPermissionBase( - vector_stores=["recreated_store"] - ).model_dump(exclude_unset=True, exclude_none=True), - "user_id": "user789", - } - - result = await _handle_update_object_permission( - data_json=data_json, - existing_key_row=existing_key_row_missing_perm, + upsert = await _prepare_key_update_object_permission( + object_permission_data=LiteLLM_ObjectPermissionBase(vector_stores=["recreated_store"]).model_dump( + exclude_unset=True, exclude_none=True + ), + existing_key_row=existing_key_row, prisma_client=mock_prisma_client, ) - # Verify new object_permission_id was set - assert "object_permission" not in result - assert result["object_permission_id"] == "recreated_perm_id_789" - - # Verify find_unique was called with the missing permission ID - mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( - where={"object_permission_id": "missing_perm_id_789"} + assert upsert is not None + assert upsert.object_permission_id == "missing_perm_id_789" + permission_upsert.assert_not_awaited() + await _write_prepared_key_update_object_permission( + data_json={}, + upsert=upsert, + prisma_client=mock_prisma_client, + ) + permission_upsert.assert_awaited_once_with( + where={"object_permission_id": "missing_perm_id_789"}, + data={"create": upsert.record, "update": upsert.record}, ) - - # Verify upsert was called to create new record - mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() @pytest.mark.asyncio @@ -7291,7 +7229,7 @@ async def test_bulk_update_keys_object_permission_is_granted_not_dropped(monkeyp upserted = prisma.db.litellm_objectpermissiontable.upsert.call_args.kwargs["data"]["create"] assert upserted["vector_stores"] == ["vs-1"] written = _written_key_row(prisma) - assert written["object_permission_id"] == "objperm-bulk" + assert written["object_permission_id"] == upserted["object_permission_id"] assert not {"max_budget", "team_id", "budget_id"} & written.keys() @@ -13170,6 +13108,7 @@ async def test_execute_virtual_key_regeneration_hides_the_untouched_modal_expiry _POLICY_DENIAL_MESSAGE = "key duration must be 7d or less" _POLICY_HASHED_TOKEN = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b" _POLICY_GENERATED_KEY = {"key": "sk-test-key", "expires": None, "user_id": "test-user", "team_id": None} +_OBJECT_PERMISSION_ID_AFTER_POLICY = "perm-after-policy" def _seven_day_policy(received: list[CustomKeyPolicyRequest]): @@ -13311,7 +13250,11 @@ async def test_regenerate_without_changes_still_runs_custom_key_policy(data): def _policy_existing_team_key() -> LiteLLM_VerificationToken: return LiteLLM_VerificationToken( - token=_POLICY_HASHED_TOKEN, user_id="test-user", team_id="team-a", max_budget=200.0 + token=_POLICY_HASHED_TOKEN, + user_id="test-user", + team_id="team-a", + max_budget=200.0, + object_permission_id=_OBJECT_PERMISSION_ID_AFTER_POLICY, ) @@ -13457,9 +13400,6 @@ async def test_process_single_key_update_rejects_when_custom_key_policy_denies() assert [policy_request.operation for policy_request in received] == ["update"] -_OBJECT_PERMISSION_ID_AFTER_POLICY = "perm-after-policy" - - def _record_object_permission_writes(mock_prisma_client: AsyncMock, events: list[str]) -> None: mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) @@ -13579,9 +13519,12 @@ async def test_regenerate_writes_the_object_permission_row_only_after_the_policy events: list[str] = [] _record_object_permission_writes(mock_prisma_client, events) data = RegenerateKeyRequest(max_budget=50.0, object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["vs-1"])) + existing_key: Final = _make_regenerate_existing_key().model_copy( + update={"object_permission_id": _OBJECT_PERMISSION_ID_AFTER_POLICY} + ) with _regenerate_policy_mocks(_recording_policy(events, allowed=True), AsyncMock(), AsyncMock()): - await _regenerate_under_policy(mock_prisma_client, _make_regenerate_existing_key(), data) + await _regenerate_under_policy(mock_prisma_client, existing_key, data) _assert_permission_row_written_after_policy( events, mock_prisma_client.db.litellm_verificationtoken.update.await_args.kwargs["data"] @@ -20960,11 +20903,25 @@ async def test_key_update_rejects_team_change_for_project_bound_key(monkeypatch: token="hashed-key", team_id=_OWNERSHIP_KEY_TEAM, project_id=_OWNED_PROJECT, + object_permission_id="permission-update", + budget_id="budget-update", ) mock_prisma_client: Final = _wire_update_key_fn(monkeypatch, existing_key_row) + existing_permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="permission-update", + vector_stores=["existing-store"], + ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_permission) + permission_upsert: Final = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert + permission_before: Final = existing_permission.model_dump() monkeypatch.setattr( "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", - AsyncMock(return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM, team_members=[])), + AsyncMock( + return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM).model_copy( + update={"team_members": []} + ) + ), ) mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) @@ -20976,7 +20933,12 @@ async def test_key_update_rejects_team_change_for_project_bound_key(monkeypatch: with pytest.raises((HTTPException, ProxyException)) as error: await update_key_fn( request=mock_request, - data=UpdateKeyRequest(key="sk-key", team_id=_OWNERSHIP_DESTINATION_TEAM), + data=UpdateKeyRequest( + key="sk-key", + team_id=_OWNERSHIP_DESTINATION_TEAM, + soft_budget=5.0, + object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["replacement-store"]), + ), user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), litellm_changed_by=None, ) @@ -20987,6 +20949,80 @@ async def test_key_update_rejects_team_change_for_project_bound_key(monkeypatch: f"but the key belongs to {_OWNERSHIP_DESTINATION_TEAM}" ) assert expected_detail in str(getattr(error.value, "detail", None) or getattr(error.value, "message", None)) + assert existing_permission.model_dump() == permission_before + permission_upsert.assert_not_awaited() + mock_prisma_client.tx.assert_not_called() + mock_prisma_client.update_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_key_update_ambiguous_permission_error_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-key", + team_id=_OWNERSHIP_KEY_TEAM, + project_id=_OWNED_PROJECT, + object_permission_id="permission-ambiguous", + ) + mock_prisma_client: Final = _wire_update_key_fn(monkeypatch, existing_key_row) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=LiteLLM_ObjectPermissionTable( + object_permission_id="permission-ambiguous", + mcp_tool_permissions={}, + ) + ) + permission_upsert: Final = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert + mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock( + return_value=[ + MagicMock(server_id="wiki-a-id", alias="wiki", server_name="wiki-a"), + MagicMock(server_id="wiki-b-id", alias="wiki", server_name="wiki-b"), + ] + ) + mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + AsyncMock( + return_value=LiteLLM_TeamTable( + team_id=_OWNERSHIP_DESTINATION_TEAM, + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="permission-team", + mcp_servers=["wiki-a-id", "wiki-b-id"], + ), + ) + ), + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock()) + + with pytest.raises((HTTPException, ProxyException)) as error: + await update_key_fn( + request=MagicMock(query_params={}), + data=UpdateKeyRequest( + key="sk-key", + team_id=_OWNERSHIP_DESTINATION_TEAM, + object_permission=LiteLLM_ObjectPermissionBase( + mcp_tool_permissions={"wiki": ["read_wiki_structure"]} + ), + ), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + ) + + assert str(getattr(error.value, "status_code", None) or getattr(error.value, "code", None)) == "400" + expected_detail: Final = { + "error": ( + "Ambiguous mcp_tool_permissions key: 'wiki' matches MCP servers ['wiki-a-id', 'wiki-b-id']. " + "Key tool permissions by server_id when servers share a name or alias." + ) + } + assert getattr(error.value, "detail", None) == expected_detail or getattr(error.value, "message", None) == str( + expected_detail + ) + permission_upsert.assert_not_awaited() + mock_prisma_client.writer_db.litellm_projecttable.find_unique.assert_not_awaited() mock_prisma_client.update_data.assert_not_awaited() @@ -21070,11 +21106,23 @@ async def test_bulk_key_update_rejects_project_team_change_and_allows_other_fiel models=[], team_id=_OWNERSHIP_PROJECT_TEAM, project_id=_OWNED_PROJECT, + object_permission_id="permission-bulk", + budget_id="budget-bulk", ) mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing_key_row) mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM) ) + existing_permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="permission-bulk", + vector_stores=["existing-store"], + ) + permission_before: Final = existing_permission.model_dump() + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_permission) + permission_upsert: Final = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert + budget_update: Final = AsyncMock() + mock_prisma_client.db.litellm_budgettable.update = budget_update mock_prisma_client.get_data = AsyncMock(return_value=existing_key_row) updated_key: Final = MagicMock() updated_key.model_dump.return_value = { @@ -21089,7 +21137,14 @@ async def test_bulk_key_update_rejects_project_team_change_and_allows_other_fiel request: Final = BulkUpdateKeyRequest( keys=[ - BulkUpdateKeyRequestItem(key="sk-bulk-key", team_id=_OWNERSHIP_DESTINATION_TEAM), + BulkUpdateKeyRequestItem.model_validate( + { + "key": "sk-bulk-key", + "team_id": _OWNERSHIP_DESTINATION_TEAM, + "soft_budget": 5.0, + "object_permission": LiteLLM_ObjectPermissionBase(vector_stores=["replacement-store"]), + } + ), BulkUpdateKeyRequestItem(key="sk-bulk-key", max_budget=10.0, tags=["bulk-update"]), ] ) @@ -21136,6 +21191,10 @@ async def test_bulk_key_update_rejects_project_team_change_and_allows_other_fiel assert response.successful_updates[0].key == "sk-bulk-key" assert response.successful_updates[0].key_info["max_budget"] == 10.0 assert response.successful_updates[0].key_info["tags"] == ["bulk-update"] + assert existing_permission.model_dump() == permission_before + permission_upsert.assert_not_awaited() + budget_update.assert_not_awaited() + mock_prisma_client.tx.assert_not_called() mock_prisma_client.update_data.assert_awaited_once() @@ -21199,7 +21258,11 @@ async def test_regenerate_checks_project_team_ownership( expected_status: int | None, expected_updates: int, ) -> None: - existing_key: Final = LiteLLM_VerificationToken(token="abc123", team_id=key_team_id) + existing_key: Final = LiteLLM_VerificationToken( + token="abc123", + team_id=key_team_id, + object_permission_id="permission-regenerate", + ) mock_prisma_client: Final = _make_regenerate_mock_prisma() mock_prisma_client.writer_db = MagicMock() mock_prisma_client.writer_db.litellm_projecttable = MagicMock() @@ -21212,6 +21275,14 @@ async def test_regenerate_checks_project_team_ownership( deprecated_key_table: Final = MagicMock() deprecated_key_table.upsert = AsyncMock() mock_prisma_client.db.litellm_deprecatedverificationtoken = deprecated_key_table + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=LiteLLM_ObjectPermissionTable( + object_permission_id="permission-regenerate", + vector_stores=["existing-store"], + ) + ) + permission_upsert: Final = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") async def regenerate() -> None: @@ -21234,6 +21305,7 @@ async def test_regenerate_checks_project_team_ownership( data=RegenerateKeyRequest( project_id=_OWNED_PROJECT, grace_period="1h" if expected_status is not None else None, + object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["replacement-store"]), ), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, @@ -21248,8 +21320,10 @@ async def test_regenerate_checks_project_team_ownership( assert "belongs to team team-b, but the key belongs to team-a" in str(error.value.detail) deleted_history_table.create_many.assert_not_awaited() deprecated_key_table.upsert.assert_not_awaited() + permission_upsert.assert_not_awaited() else: await regenerate() + permission_upsert.assert_awaited_once() assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == expected_updates From 0b5ea60bfd2a01545f1a05a9df6de150761672a8 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 02:10:22 +0000 Subject: [PATCH 16/20] fix(proxy): keep project object permission validation on the stored payload Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/project_endpoints.py | 13 +++---- .../management/test_project_lifecycle.py | 27 ------------- .../test_project_endpoints_prisma.py | 38 +++++++++++++++---- 3 files changed, 35 insertions(+), 43 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index 5a754141063..ccdf5057b1c 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -790,16 +790,8 @@ async def update_project( data, existing_project, _router_access_group_names(llm_router) ) - object_permission_data: Final = ( - data.object_permission.model_dump(exclude_none=True) if data.object_permission is not None else None - ) - object_permission_payload: Final = ( - _OBJECT_PERMISSION_PAYLOAD.validate_python(object_permission_data) if object_permission_data else None - ) - # Prepare update data update_data = _jsonified(prisma_client, data.model_dump(exclude_none=True, exclude={"project_id"})) - update_data.pop("object_permission", None) update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name # Handle budget updates @@ -809,6 +801,11 @@ async def update_project( **({"max_budget": None} if "max_budget" in data.model_fields_set and data.max_budget is None else {}), } + object_permission_data: Final = update_data.pop("object_permission", None) + object_permission_payload: Final = ( + _OBJECT_PERMISSION_PAYLOAD.validate_python(object_permission_data) if object_permission_data else None + ) + if data.team_id is not None: current_project_record: Final = await _writer_project_table(prisma_client).find_unique( where={"project_id": data.project_id} diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index 25b8da7875d..bf5dac0da89 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -65,17 +65,6 @@ def _clear_key_object_permission(key: str, permission_id: str) -> None: ) -def _clear_project_object_permission(project_id: str, permission_id: str) -> None: - write_rows( - 'UPDATE "LiteLLM_ProjectTable" SET object_permission_id = NULL WHERE project_id = %s', - (project_id,), - ) - write_rows( - 'DELETE FROM "LiteLLM_ObjectPermissionTable" WHERE object_permission_id = %s', - (permission_id,), - ) - - def _deleted_key_rows(key: str) -> list[dict[str, JsonValue]]: return read_rows( 'SELECT token FROM "LiteLLM_DeletedVerificationToken" WHERE token = %s', @@ -660,22 +649,8 @@ def test_project_update_rejects_moving_project_with_attached_key(ownership_gatew budget: Final = scenario.budget(max_budget=3) attached_project: Final = scenario.project(team_a, budget_id=budget, models=[model]) key: Final = scenario.key(team_id=team_a, project_id=attached_project, models=[model]) - permission_seed: Final = ownership_gateway.request( - "POST", - "/project/update", - {"project_id": attached_project, "object_permission": {"vector_stores": ["existing-store"]}}, - ) - assert permission_seed.status_code == 200, permission_seed.text - project_permission: Final = read_rows( - 'SELECT object_permission_id FROM "LiteLLM_ProjectTable" WHERE project_id = %s', - (attached_project,), - ) - assert len(project_permission) == 1 - permission_id: Final = string_value(project_permission[0]["object_permission_id"]) - scenario.cleanups.callback(_clear_project_object_permission, attached_project, permission_id) project_before: Final = _project_rows(attached_project) key_before: Final = _key_rows(key) - permission_before: Final = _object_permission_rows(permission_id) budget_before: Final = _budget_rows(budget) moved_with_key: Final = ownership_gateway.request( "POST", @@ -684,13 +659,11 @@ def test_project_update_rejects_moving_project_with_attached_key(ownership_gatew "project_id": attached_project, "team_id": team_b, "max_budget": 11, - "object_permission": {"vector_stores": ["replacement-store"]}, }, ) assert moved_with_key.status_code == 400, moved_with_key.text assert _project_rows(attached_project) == project_before assert _key_rows(key) == key_before - assert _object_permission_rows(permission_id) == permission_before assert _budget_rows(budget) == budget_before diff --git a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py index cd185179ee9..ad2ce582c1f 100644 --- a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py +++ b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py @@ -1269,6 +1269,36 @@ def _written_project_data(mock_prisma: mock.MagicMock) -> dict: return mock_prisma.db.litellm_projecttable.update.await_args.kwargs["data"] +@pytest.mark.asyncio +async def test_update_project_object_permission_validation_precedes_budget_write( + monkeypatch: pytest.MonkeyPatch, +) -> None: + project_id: Final = "project-object-permission-validation" + mock_prisma: Final = _project_update_mocks(monkeypatch, {}) + mock_prisma.db.litellm_projecttable.find_unique.return_value.budget_id = "budget-project" + budget_table: Final = mock.MagicMock() + budget_table.update = mock.AsyncMock() + mock_prisma.db.litellm_budgettable = budget_table + + def jsonify_object_permission_as_string(payload: dict[str, object]) -> dict[str, object]: + return {**payload, "object_permission": '{"vector_stores": ["replacement-store"]}'} + + mock_prisma.jsonify_object = jsonify_object_permission_as_string + + with pytest.raises(ProxyException) as error: + await _run_project_update( + project_id, + max_budget=50, + object_permission={"vector_stores": ["replacement-store"]}, + ) + + assert error.value.code == "500" + assert "Input should be a valid dictionary" in error.value.message + assert "input_type=str" in error.value.message + budget_table.update.assert_not_awaited() + mock_prisma.db.litellm_projecttable.update.assert_not_awaited() + + @pytest.mark.asyncio async def test_update_project_rejects_move_when_attached_teamless_key_exists( monkeypatch: pytest.MonkeyPatch, @@ -1278,14 +1308,9 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists( mock_prisma: Final = _project_update_mocks(monkeypatch, {}) mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = "team-a" mock_prisma.db.litellm_projecttable.find_unique.return_value.budget_id = "budget-project" - mock_prisma.db.litellm_projecttable.find_unique.return_value.object_permission_id = "permission-project" budget_table: Final = mock.MagicMock() budget_table.update = mock.AsyncMock() - permission_table: Final = mock.MagicMock() - permission_table.update = mock.AsyncMock() - permission_table.create = mock.AsyncMock() mock_prisma.db.litellm_budgettable = budget_table - mock_prisma.db.litellm_objectpermissiontable = permission_table mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock( return_value=LiteLLM_TeamTable(team_id=destination_team_id) ) @@ -1306,7 +1331,6 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists( project_id, team_id=destination_team_id, max_budget=50, - object_permission={"vector_stores": ["replacement-store"]}, ) expected_detail: Final = { @@ -1320,8 +1344,6 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists( mock_prisma.writer_db.litellm_verificationtoken.count.assert_awaited_once() mock_prisma.db.litellm_verificationtoken.count.assert_not_awaited() budget_table.update.assert_not_awaited() - permission_table.update.assert_not_awaited() - permission_table.create.assert_not_awaited() mock_prisma.db.litellm_projecttable.update.assert_not_awaited() From 1e9d10b531561f6d8703845271f80614868597eb Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 02:37:47 +0000 Subject: [PATCH 17/20] style(proxy): format key regeneration update payload Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/management_endpoints/key_management_endpoints.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 3b5b7264d1a..008f7f5692e 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -5778,9 +5778,7 @@ async def _execute_virtual_key_regeneration( upsert=object_permission_upsert, prisma_client=prisma_client, ) - update_data: Final = MappingProxyType( - {"token": new_token_hash, "key_name": new_token_key_name, **update_values} - ) + update_data: Final = MappingProxyType({"token": new_token_hash, "key_name": new_token_key_name, **update_values}) jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data) From 8d0f36bc59b4e33f7832bfc7fef7f8dd5f0253bb Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 04:43:05 +0000 Subject: [PATCH 18/20] fix(proxy): invalidate the written permission id on bulk update and regenerate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../key_management_endpoints.py | 4 +- .../test_key_management_endpoints.py | 71 ++++++++++++++++--- 2 files changed, 65 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 008f7f5692e..1e10834752a 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2982,7 +2982,7 @@ async def _process_single_key_update( await invalidate_cached_object_permissions( object_permission_ids=( existing_key_row.object_permission_id, - non_default_values.get("object_permission_id"), + update_values.get("object_permission_id"), ), user_api_key_cache=user_api_key_cache, ) @@ -5817,7 +5817,7 @@ async def _execute_virtual_key_regeneration( await invalidate_cached_object_permissions( object_permission_ids=( key_in_db.object_permission_id, - non_default_values.get("object_permission_id"), + update_values.get("object_permission_id"), ), user_api_key_cache=user_api_key_cache, ) diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 80c51d20f1e..e62be2499be 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -5,6 +5,7 @@ from typing import Final from types import SimpleNamespace import json from datetime import datetime, timedelta, timezone +from uuid import UUID import litellm import pytest @@ -7146,7 +7147,10 @@ _BULK_UPDATE_TEAM: Final = LiteLLM_TeamTableCachedObj(team_id="team-1") async def _run_bulk_update_on_one_key( - monkeypatch, item_payload: Mapping[str, object], team: LiteLLM_TeamTableCachedObj = _BULK_UPDATE_TEAM + monkeypatch: pytest.MonkeyPatch, + item_payload: Mapping[str, object], + team: LiteLLM_TeamTableCachedObj = _BULK_UPDATE_TEAM, + user_api_key_cache: UserApiKeyCache | None = None, ) -> tuple[BulkUpdateKeyResponse, AsyncMock]: from litellm.proxy.management_endpoints.key_management_endpoints import bulk_update_keys @@ -7162,6 +7166,8 @@ async def _run_bulk_update_on_one_key( ) mock_prisma_client.update_data = AsyncMock(return_value={"data": {"token": _BULK_UPDATE_TOKEN}}) _setup_update_key_mocks(monkeypatch, mock_prisma_client) + if user_api_key_cache is not None: + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache) monkeypatch.setattr( "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", AsyncMock(return_value=team) ) @@ -7233,6 +7239,47 @@ async def test_bulk_update_keys_object_permission_is_granted_not_dropped(monkeyp assert not {"max_budget", "team_id", "budget_id"} & written.keys() +@pytest.mark.asyncio +async def test_bulk_update_keys_invalidates_cache_for_new_object_permission(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy.common_utils.user_api_key_cache import object_permission_cache_key + + permission_id = str(UUID("00000000-0000-0000-0000-000000000001")) + permission_cache_key = object_permission_cache_key(permission_id) + user_api_key_cache = UserApiKeyCache() + user_api_key_cache.set_cache( + key=permission_cache_key, + value=LiteLLM_ObjectPermissionTable(object_permission_id=permission_id, vector_stores=["stale"]), + model_type=LiteLLM_ObjectPermissionTable, + ) + assert ( + user_api_key_cache.get_cache( + key=permission_cache_key, + model_type=LiteLLM_ObjectPermissionTable, + ) + is not None + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.object_permission_utils.uuid.uuid4", + lambda: UUID(permission_id), + ) + + response, prisma = await _run_bulk_update_on_one_key( + monkeypatch, + {"object_permission": {"vector_stores": ["vs-1"]}}, + user_api_key_cache=user_api_key_cache, + ) + + assert response.failed_updates == [] + assert _written_key_row(prisma)["object_permission_id"] == permission_id + assert ( + user_api_key_cache.get_cache( + key=permission_cache_key, + model_type=LiteLLM_ObjectPermissionTable, + ) + is None + ) + + @pytest.mark.asyncio async def test_bulk_update_keys_object_permission_outside_the_team_allowlist_is_refused(monkeypatch): """A bulk item's object_permission is checked against the key's team exactly as /key/update @@ -21634,16 +21681,16 @@ async def test_key_update_invalidates_cached_object_permission(monkeypatch): @pytest.mark.asyncio -async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch): +async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch: pytest.MonkeyPatch): """Regression: regenerating a key with new permissions must not keep serving the old grants.""" from litellm.proxy._types import LiteLLM_ObjectPermissionBase, RegenerateKeyRequest from litellm.proxy.auth.auth_checks import get_object_permission - from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key from litellm.proxy.management_endpoints.key_management_endpoints import ( _execute_virtual_key_regeneration, ) - permission_id = "objperm-regenerate" + permission_id = str(UUID("00000000-0000-0000-0000-000000000002")) grants = {"served": ["tool_a"]} def _row(**kwargs): @@ -21660,9 +21707,12 @@ async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch return_value=MagicMock(object_permission_id=permission_id) ) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr( + "litellm.proxy.management_helpers.object_permission_utils.uuid.uuid4", + lambda: UUID(permission_id), + ) existing_key = _make_regenerate_existing_key() - existing_key.object_permission_id = permission_id user_api_key_cache = UserApiKeyCache() assert ( await get_object_permission( @@ -21696,9 +21746,7 @@ async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch hashed_api_key="abc123", key="abc123", data=RegenerateKeyRequest( - object_permission=LiteLLM_ObjectPermissionBase( - mcp_tool_permissions={"server-1": ["tool_a", "tool_b"]} - ) + object_permission=LiteLLM_ObjectPermissionBase(mcp_tool_permissions={"server-1": ["tool_a", "tool_b"]}) ), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, @@ -21706,6 +21754,13 @@ async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch proxy_logging_obj=AsyncMock(), ) + assert ( + user_api_key_cache.get_cache( + key=object_permission_cache_key(permission_id), + model_type=LiteLLM_ObjectPermissionTable, + ) + is None + ) grants["served"] = ["tool_a", "tool_b"] reread = await get_object_permission( object_permission_id=permission_id, From bbbd0f62b3f3c4e552fb0539aef8cfb585acc7d2 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 06:34:18 +0000 Subject: [PATCH 19/20] test(proxy): audit project team ownership across endpoints, bulk paths and chaos Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management/test_project_lifecycle.py | 865 +++++++++++++++++- 1 file changed, 861 insertions(+), 4 deletions(-) diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index bf5dac0da89..b4da1a5a3f5 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -1,19 +1,38 @@ +import json +import os +import re +import signal +import threading from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor from hashlib import sha256 +from pathlib import Path from typing import Final from uuid import uuid4 import httpx +import psutil +import psycopg import pytest -from integration._support.client import JSON_OBJECT, Gateway, gateway_from_environment, object_value, string_value +from integration._support.client import ( + JSON_OBJECT, + Gateway, + delete_key_if_present, + eventually, + gateway_from_environment, + object_value, + string_value, +) from integration._support.database import read_rows, write_rows -from integration._support.process import owned_proxy +from integration._support.process import group_members, owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request, wire_server from pydantic import JsonValue from litellm.models.user import LiteLLM_UserTable from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken _OWNED_PROXY_SALT_KEY: Final = "sk-integration-salt" +_OWNERSHIP_MARKER: Final = re.compile(rb"ownership-[0-9a-f]{32}") def _project_rows(project_id: str) -> list[dict[str, JsonValue]]: @@ -40,6 +59,22 @@ def _key_permission_and_budget_ids(key: str) -> list[dict[str, JsonValue]]: ) +def _key_state_rows(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT to_jsonb(k) AS key_row, to_jsonb(b) AS budget_row FROM "LiteLLM_VerificationToken" AS k ' + 'LEFT JOIN "LiteLLM_BudgetTable" AS b ON b.budget_id = k.budget_id WHERE k.token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + +def _project_state_rows(project_id: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT to_jsonb(p) AS project_row, to_jsonb(b) AS budget_row FROM "LiteLLM_ProjectTable" AS p ' + 'LEFT JOIN "LiteLLM_BudgetTable" AS b ON b.budget_id = p.budget_id WHERE p.project_id = %s', + (project_id,), + ) + + def _object_permission_rows(permission_id: str) -> list[dict[str, JsonValue]]: return read_rows( 'SELECT * FROM "LiteLLM_ObjectPermissionTable" WHERE object_permission_id = %s', @@ -47,6 +82,13 @@ def _object_permission_rows(permission_id: str) -> list[dict[str, JsonValue]]: ) +def _object_permission_table_rows() -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT to_jsonb(p) AS row FROM "LiteLLM_ObjectPermissionTable" AS p ORDER BY object_permission_id', + (), + ) + + def _budget_rows(budget_id: str) -> list[dict[str, JsonValue]]: return read_rows( 'SELECT to_jsonb(b) AS row FROM "LiteLLM_BudgetTable" AS b WHERE budget_id = %s', @@ -67,14 +109,14 @@ def _clear_key_object_permission(key: str, permission_id: str) -> None: def _deleted_key_rows(key: str) -> list[dict[str, JsonValue]]: return read_rows( - 'SELECT token FROM "LiteLLM_DeletedVerificationToken" WHERE token = %s', + 'SELECT to_jsonb(d) AS row FROM "LiteLLM_DeletedVerificationToken" AS d WHERE d.token = %s', (sha256(key.encode()).hexdigest(),), ) def _deprecated_key_rows(key: str) -> list[dict[str, JsonValue]]: return read_rows( - 'SELECT token FROM "LiteLLM_DeprecatedVerificationToken" WHERE token = %s', + 'SELECT to_jsonb(d) AS row FROM "LiteLLM_DeprecatedVerificationToken" AS d WHERE d.token = %s', (sha256(key.encode()).hexdigest(),), ) @@ -109,6 +151,821 @@ def _discard_unexpected_key(candidate: Gateway, response: httpx.Response) -> Non candidate.post("/key/delete", {"keys": [string_value(body["key"])]}) +def _ownership_chat_reply(identity: str, stream: bool) -> Reply: + if not stream: + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4.1-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ownership ok"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 7, "completion_tokens": 2, "total_tokens": 9}, + } + ).encode() + ) + return Reply( + content_type="text/event-stream", + chunks=( + f"data: {json.dumps({'id': identity, 'object': 'chat.completion.chunk', 'created': 1, 'model': 'gpt-4.1-mini', 'choices': [{'index': 0, 'delta': {'role': 'assistant', 'content': 'ownership'}}]})}\n\n".encode(), + f"data: {json.dumps({'id': identity, 'object': 'chat.completion.chunk', 'created': 1, 'model': 'gpt-4.1-mini', 'choices': [{'index': 0, 'delta': {'content': ' ok'}, 'finish_reason': 'stop'}], 'usage': {'prompt_tokens': 7, 'completion_tokens': 2, 'total_tokens': 9}})}\n\n".encode(), + b"data: [DONE]\n\n", + ), + ) + + +def _ownership_responses_reply(identity: str, stream: bool) -> Reply: + response: Final = { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4.1-mini", + "output": [ + { + "id": f"msg_{identity}", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "ownership ok", "annotations": []}], + } + ], + "usage": {"input_tokens": 7, "output_tokens": 2, "total_tokens": 9}, + } + if not stream: + return Reply(body=json.dumps(response).encode()) + events: Final = ( + { + "type": "response.created", + "sequence_number": 0, + "response": {**response, "status": "in_progress", "output": []}, + }, + { + "type": "response.output_text.delta", + "sequence_number": 1, + "item_id": f"msg_{identity}", + "output_index": 0, + "content_index": 0, + "delta": "ownership ok", + }, + {"type": "response.completed", "sequence_number": 2, "response": response}, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _ownership_upstream(request: Request) -> Reply: + found: Final = _OWNERSHIP_MARKER.search(request.body) + if found is None: + return Reply(status=400, body=b'{"error":"missing ownership marker"}') + marker: Final = found.group(0).decode() + body: Final = JSON_OBJECT.validate_json(request.body) + stream: Final = body.get("stream") is True + if request.target.endswith("/responses"): + return _ownership_responses_reply(f"resp_{marker}", stream) + return _ownership_chat_reply(f"chatcmpl-{marker}", stream) + + +def _ownership_sse_events(text: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + JSON_OBJECT.validate_json(line[6:]) + for line in text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + + +def _ownership_request_marker(request: Request) -> str: + found: Final = _OWNERSHIP_MARKER.search(request.body) + assert found is not None, request.body + return found.group(0).decode() + + +def _ownership_request_payload( + path: str, model: str, marker: str, stream: bool +) -> tuple[dict[str, JsonValue], dict[str, str]]: + if path.endswith("/messages"): + return ( + { + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": marker}], + "stream": stream, + }, + {"anthropic-version": "2023-06-01"}, + ) + if path.endswith("/responses"): + return {"model": model, "input": marker, "stream": stream}, {} + return {"model": model, "messages": [{"role": "user", "content": marker}], "stream": stream}, {} + + +def _ownership_response_id(response: httpx.Response, path: str) -> str: + if not response.headers.get("content-type", "").startswith("text/event-stream"): + return string_value(JSON_OBJECT.validate_json(response.content)["id"]) + events: Final = _ownership_sse_events(response.text) + identities: Final = ( + tuple( + string_value(object_value(event["message"])["id"]) + for event in events + if event.get("type") == "message_start" + ) + if path.endswith("/messages") + else tuple( + string_value(object_value(event["response"])["id"]) + for event in events + if event.get("type") == "response.completed" + ) + if path.endswith("/responses") + else tuple(string_value(event["id"]) for event in events if "id" in event) + ) + unique_identities: Final = frozenset(identities) + assert len(unique_identities) == 1, response.text + return next(iter(unique_identities)) + + +def _ownership_serving_call( + candidate: Gateway, + key: str, + model: str, + index: int, + marker: str, +) -> tuple[int, str, str]: + stream: Final = index % 2 == 0 + route: Final = index % 3 + paths: Final = ("/v1/chat/completions", "/v1/responses", "/v1/messages") + path: Final = paths[route] + body, headers = _ownership_request_payload(path, model, marker, stream) + response: Final = candidate.request("POST", path, body, key=key, headers=headers) + response.read() + if response.status_code != 200: + return response.status_code, "", response.text + return response.status_code, _ownership_response_id(response, path), response.text + + +def _create_mcp_server(candidate: Gateway, server_id: str, server_name: str, alias: str) -> None: + response: Final = candidate.request( + "POST", + "/v1/mcp/server", + { + "server_id": server_id, + "server_name": server_name, + "alias": alias, + "transport": "sse", + "url": "http://127.0.0.1:9/mcp", + }, + ) + assert response.status_code == 201, response.text + + +def _delete_mcp_server(candidate: Gateway, server_id: str) -> None: + response: Final = candidate.request("DELETE", f"/v1/mcp/server/{server_id}") + assert response.status_code == 202, response.text + + +def test_service_account_generate_rejects_foreign_team_project_without_writing_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + alias: Final = f"service-account-{uuid4().hex}" + response: Final = ownership_gateway.request( + "POST", + "/key/service-account/generate", + {"team_id": team_a, "project_id": project_b, "key_alias": alias, "models": [model]}, + ) + _discard_unexpected_key(ownership_gateway, response) + assert response.status_code == 400, response.text + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND key_alias = %s', + (project_b, alias), + ) + == [] + ) + + +def test_key_generate_unowned_project_accepts_any_team(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project: Final = scenario.project(team_a, models=[model]) + write_rows('UPDATE "LiteLLM_ProjectTable" SET team_id = NULL WHERE project_id = %s', (project,)) + team_key: Final = scenario.key(team_id=team_b, project_id=project, models=[model]) + teamless_key: Final = scenario.key(project_id=project, models=[model]) + chats: Final = tuple( + ownership_gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "legacy unowned"}]}, + key=key, + ) + for key in (team_key, teamless_key) + ) + assert all(chat.status_code == 200 for chat in chats), tuple(chat.text for chat in chats) + + +@pytest.mark.parametrize( + ("path", "stream"), + ( + ("/v1/chat/completions", False), + ("/v1/chat/completions", True), + ("/v1/messages", False), + ("/v1/messages", True), + ("/v1/responses", False), + ("/v1/responses", True), + ), +) +def test_legacy_mismatched_key_keeps_serving_and_stays_editable( + ownership_gateway: Gateway, path: str, stream: bool +) -> None: + def respond(request: Request) -> Reply: + return _ownership_upstream(request) + + with wire_server(respond) as upstream: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model(api_base=f"{upstream.url}/v1") + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project, models=[model]) + write_rows( + 'UPDATE "LiteLLM_VerificationToken" SET team_id = %s WHERE token = %s', + (team_b, sha256(key.encode()).hexdigest()), + ) + marker: Final = f"ownership-{uuid4().hex}" + body, headers = _ownership_request_payload(path, model, marker, stream) + response: Final = ownership_gateway.request("POST", path, body, key=key, headers=headers) + response.read() + assert response.status_code == 200, response.text + requests: Final = upstream.drain() + assert len(requests) == 1 + expected_target: Final = "/chat/completions" if path.endswith("/chat/completions") else "/responses" + assert requests[0].target.endswith(expected_target), requests[0].target + assert marker.encode() in requests[0].body + assert _ownership_response_id(response, path) != "" + + +def test_stored_team_mismatch_allows_key_edits_and_detach(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project, models=[model]) + write_rows( + 'UPDATE "LiteLLM_VerificationToken" SET team_id = %s WHERE token = %s', + (team_b, sha256(key.encode()).hexdigest()), + ) + alias: Final = f"stored-mismatch-{uuid4().hex}" + alias_update: Final = ownership_gateway.request( + "POST", + "/key/update", + {"key": key, "key_alias": alias}, + ) + assert alias_update.status_code == 200, alias_update.text + aliased_key: Final = _key_rows(key) + assert aliased_key[0]["key_alias"] == alias + unchanged: Final = ownership_gateway.request( + "POST", + "/key/update", + {"key": key, "team_id": team_b}, + ) + assert unchanged.status_code == 200, unchanged.text + assert _key_rows(key) == aliased_key + detached: Final = ownership_gateway.request( + "POST", + "/key/update", + {"key": key, "project_id": None}, + ) + assert detached.status_code == 200, detached.text + detached_key: Final = _key_rows(key) + assert detached_key[0]["team_id"] == team_b + assert detached_key[0]["project_id"] is None + + +def test_key_regenerate_cross_team_with_object_permission_writes_nothing(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + key: Final = scenario.key( + team_id=team_a, + project_id=project_a, + models=[model], + object_permission={"vector_stores": ["before"]}, + ) + ids: Final = _key_permission_and_budget_ids(key) + assert len(ids) == 1 + permission_id: Final = string_value(ids[0]["object_permission_id"]) + scenario.cleanups.callback(_clear_key_object_permission, key, permission_id) + before: Final = ( + _object_permission_rows(permission_id), + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) + response: Final = ownership_gateway.request( + "POST", + f"/key/{key}/regenerate", + {"project_id": project_b, "object_permission": {"vector_stores": ["after"]}}, + ) + _discard_unexpected_key(ownership_gateway, response) + assert response.status_code == 400, response.text + assert ( + _object_permission_rows(permission_id), + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) == before + + +def test_key_bulk_update_cross_team_with_object_permission_preserves_permission(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key( + team_id=team_a, + project_id=project, + models=[model], + object_permission={"vector_stores": ["before"]}, + ) + ids: Final = _key_permission_and_budget_ids(key) + assert len(ids) == 1 + permission_id: Final = string_value(ids[0]["object_permission_id"]) + scenario.cleanups.callback(_clear_key_object_permission, key, permission_id) + before: Final = ( + _object_permission_rows(permission_id), + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) + response: Final = ownership_gateway.request( + "POST", + "/key/bulk_update", + { + "keys": [ + { + "key": key, + "team_id": team_b, + "object_permission": {"vector_stores": ["after"]}, + } + ], + }, + ) + assert response.status_code == 200, response.text + failed_updates: Final = JSON_OBJECT.validate_json(response.content)["failed_updates"] + assert isinstance(failed_updates, list) + assert len(failed_updates) == 1 + failed_update: Final = object_value(failed_updates[0]) + assert f"Project {project} belongs to team {team_a}" in string_value(failed_update["failed_reason"]) + assert ( + _object_permission_rows(permission_id), + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) == before + + +def test_key_update_ambiguous_mcp_permission_error_precedes_project_ownership(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + identifier: Final = f"ambiguous{uuid4().hex}" + first_id: Final = f"mcp{uuid4().hex}" + second_id: Final = f"mcp{uuid4().hex}" + _create_mcp_server(ownership_gateway, first_id, identifier, f"alias{uuid4().hex}") + scenario.cleanups.callback(_delete_mcp_server, ownership_gateway, first_id) + _create_mcp_server(ownership_gateway, second_id, f"name{uuid4().hex}", f"alias{uuid4().hex}") + scenario.cleanups.callback(_delete_mcp_server, ownership_gateway, second_id) + write_rows( + 'UPDATE "LiteLLM_MCPServerTable" SET alias = %s WHERE server_id = %s', + (identifier, second_id), + ) + team_permissions: Final = {"mcp_servers": [first_id, second_id]} + team_a: Final = scenario.team(models=[model], object_permission=team_permissions) + team_b: Final = scenario.team(models=[model], object_permission=team_permissions) + project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key( + team_id=team_a, + project_id=project, + models=[model], + object_permission={"vector_stores": ["before"]}, + ) + ids: Final = _key_permission_and_budget_ids(key) + assert len(ids) == 1 + permission_id: Final = string_value(ids[0]["object_permission_id"]) + scenario.cleanups.callback(_clear_key_object_permission, key, permission_id) + before: Final = ( + _object_permission_rows(permission_id), + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) + response: Final = ownership_gateway.request( + "POST", + "/key/update", + { + "key": key, + "team_id": team_b, + "object_permission": {"mcp_tool_permissions": {identifier: ["tool"]}}, + }, + ) + assert response.status_code == 400, response.text + assert "ambiguous" in response.text.lower(), response.text + assert "project" not in response.text.lower(), response.text + assert ( + _object_permission_rows(permission_id), + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) == before + + +def test_team_key_bulk_update_rejects_foreign_team_project(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + key: Final = scenario.key(team_id=team_a, models=[model]) + before: Final = _key_rows(key) + response: Final = ownership_gateway.request( + "POST", + "/team/key/bulk_update", + { + "team_id": team_a, + "key_ids": [sha256(key.encode()).hexdigest()], + "update_fields": {"project_id": project_b, "team_id": team_b}, + }, + ) + assert response.status_code == 422, response.text + assert "project_id" in response.text, response.text + assert "team_id" in response.text, response.text + assert _key_rows(key) == before + + +def test_bulk_update_and_regenerate_new_object_permission_is_served(ownership_gateway: Gateway) -> None: + with wire_server(_ownership_upstream) as upstream: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model(api_base=f"{upstream.url}/v1") + team: Final = scenario.team(models=[model]) + project: Final = scenario.project(team, models=[model]) + bulk_key: Final = scenario.key(team_id=team, project_id=project, models=[model]) + bulk_response: Final = ownership_gateway.request( + "POST", + "/key/bulk_update", + { + "keys": [ + { + "key": bulk_key, + "object_permission": {"vector_stores": ["bulk"]}, + } + ], + }, + ) + assert bulk_response.status_code == 200, bulk_response.text + bulk_info: Final = ownership_gateway.request("GET", "/key/info", params={"key": bulk_key}) + assert bulk_info.status_code == 200, bulk_info.text + bulk_row: Final = object_value(JSON_OBJECT.validate_json(bulk_info.content)["info"]) + bulk_permission_id: Final = string_value(bulk_row["object_permission_id"]) + assert bulk_permission_id != "" + assert len(_object_permission_rows(bulk_permission_id)) == 1 + scenario.cleanups.callback(_clear_key_object_permission, bulk_key, bulk_permission_id) + bulk_chat: Final = ownership_gateway.chat(model, key=bulk_key, text=f"ownership-{uuid4().hex}") + assert string_value(bulk_chat["id"]) != "" + regenerated: Final = ownership_gateway.request( + "POST", + "/key/generate", + {"team_id": team, "project_id": project, "models": [model]}, + ) + assert regenerated.status_code == 200, regenerated.text + regenerated_key: Final = string_value(JSON_OBJECT.validate_json(regenerated.content)["key"]) + scenario.cleanups.callback(delete_key_if_present, ownership_gateway, regenerated_key) + regeneration: Final = ownership_gateway.request( + "POST", + f"/key/{regenerated_key}/regenerate", + {"object_permission": {"vector_stores": ["regenerated"]}}, + ) + assert regeneration.status_code == 200, regeneration.text + new_key: Final = string_value(JSON_OBJECT.validate_json(regeneration.content)["key"]) + scenario.cleanups.callback(delete_key_if_present, ownership_gateway, new_key) + info: Final = ownership_gateway.request("GET", "/key/info", params={"key": new_key}) + assert info.status_code == 200, info.text + info_row: Final = object_value(JSON_OBJECT.validate_json(info.content)["info"]) + regenerated_permission_id: Final = string_value(info_row["object_permission_id"]) + assert regenerated_permission_id != "" + assert regenerated_permission_id != bulk_permission_id + assert len(_object_permission_rows(regenerated_permission_id)) == 1 + scenario.cleanups.callback(_clear_key_object_permission, new_key, regenerated_permission_id) + chat: Final = ownership_gateway.chat(model, key=new_key, text=f"ownership-{uuid4().hex}") + assert string_value(chat["id"]) != "" + + +def test_ownership_rejections_during_concurrent_traffic_burst(ownership_gateway: Gateway) -> None: + with wire_server(_ownership_upstream) as upstream: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model(api_base=f"{upstream.url}/v1") + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + serving_project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model]) + serving_key: Final = scenario.key(team_id=team_a, project_id=serving_project, models=[model]) + markers: Final = tuple(f"ownership-{uuid4().hex}" for _ in range(30)) + key_before: Final = ( + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) + serving_key_before: Final = ( + _key_rows(serving_key), + _key_permission_and_budget_ids(serving_key), + _deleted_key_rows(serving_key), + _deprecated_key_rows(serving_key), + ) + project_before: Final = _project_state_rows(project_a) + + def serve(index: int) -> tuple[int, str, str]: + return _ownership_serving_call(ownership_gateway, serving_key, model, index, markers[index]) + + def update_key() -> httpx.Response: + return ownership_gateway.request("POST", "/key/update", {"key": key, "team_id": team_b}) + + def bulk_update() -> httpx.Response: + return ownership_gateway.request( + "POST", + "/key/bulk_update", + {"keys": [{"key": key, "team_id": team_b}]}, + ) + + def move_project() -> httpx.Response: + return ownership_gateway.request( + "POST", "/project/update", {"project_id": project_a, "team_id": team_b} + ) + + with ThreadPoolExecutor(max_workers=33) as pool: + serving_futures: Final = tuple(pool.submit(serve, index) for index in range(30)) + ownership_futures: Final = ( + pool.submit(update_key), + pool.submit(bulk_update), + pool.submit(move_project), + ) + serving_results: Final = tuple(future.result(timeout=90) for future in serving_futures) + ownership_results: Final = tuple(future.result(timeout=90) for future in ownership_futures) + assert all(status == 200 for status, _, _ in serving_results), serving_results + assert all(response_id for _, response_id, _ in serving_results), serving_results + assert len({response_id for _, response_id, _ in serving_results}) == 30 + key_update_response: Final = ownership_results[0] + bulk_update_response: Final = ownership_results[1] + project_update_response: Final = ownership_results[2] + assert key_update_response.status_code == 400, key_update_response.text + assert bulk_update_response.status_code == 200, bulk_update_response.text + bulk_result: Final = JSON_OBJECT.validate_json(bulk_update_response.content) + successful_updates: Final = bulk_result["successful_updates"] + failed_updates: Final = bulk_result["failed_updates"] + assert isinstance(successful_updates, list), bulk_update_response.text + assert successful_updates == [], bulk_update_response.text + assert isinstance(failed_updates, list), bulk_update_response.text + assert len(failed_updates) == 1, bulk_update_response.text + failed_update: Final = object_value(failed_updates[0]) + assert f"Project {project_a} belongs to team {team_a}" in string_value(failed_update["failed_reason"]), ( + bulk_update_response.text + ) + assert project_update_response.status_code == 400, project_update_response.text + upstream_requests: Final = upstream.drain() + assert len(upstream_requests) == 30 + received_markers: Final = tuple(_ownership_request_marker(request) for request in upstream_requests) + assert tuple(sorted(received_markers)) == tuple(sorted(markers)) + assert ( + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) == key_before + assert ( + _key_rows(serving_key), + _key_permission_and_budget_ids(serving_key), + _deleted_key_rows(serving_key), + _deprecated_key_rows(serving_key), + ) == serving_key_before + assert _project_state_rows(project_a) == project_before + + +def test_ownership_checks_wait_for_row_locks_without_failing_readiness(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team: Final = scenario.team(models=[model]) + project: Final = scenario.project(team, models=[model], max_budget=7) + key: Final = scenario.key(team_id=team, project_id=project, models=[model]) + alias: Final = f"locked-{uuid4().hex}" + locked: Final = threading.Event() + release: Final = threading.Event() + + def hold_locks() -> None: + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute("BEGIN") + connection.execute( + 'SELECT project_id FROM "LiteLLM_ProjectTable" WHERE project_id = %s FOR UPDATE', + (project,), + ) + connection.execute( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE token = %s FOR UPDATE', + (sha256(key.encode()).hexdigest(),), + ) + locked.set() + assert release.wait(timeout=30) + connection.commit() + + with ThreadPoolExecutor(max_workers=3) as pool: + lock_future: Final = pool.submit(hold_locks) + try: + assert locked.wait(timeout=10) + key_future: Final = pool.submit( + ownership_gateway.request, + "POST", + "/key/update", + {"key": key, "key_alias": alias}, + ) + project_future: Final = pool.submit( + ownership_gateway.request, + "POST", + "/project/update", + {"project_id": project, "team_id": team, "max_budget": 19}, + ) + lock_waiters: Final = eventually( + lambda: read_rows( + "SELECT pid FROM pg_stat_activity WHERE datname = current_database() " + "AND wait_event_type = 'Lock' AND cardinality(pg_blocking_pids(pid)) > 0", + (), + ), + lambda rows: len(rows) >= 2, + seconds=10, + ) + assert len(lock_waiters) >= 2 + readiness: Final = eventually( + lambda: ownership_gateway.request("GET", "/health/readiness").status_code, + lambda status: status == 200, + seconds=10, + ) + assert readiness == 200 + finally: + release.set() + key_response: Final = key_future.result(timeout=90) + project_response: Final = project_future.result(timeout=90) + lock_future.result(timeout=30) + assert key_response.status_code == 200, key_response.text + assert project_response.status_code == 200, project_response.text + assert _key_rows(key)[0]["key_alias"] == alias + assert _project_rows(project)[0]["max_budget"] == 19.0 + + +def test_ownership_enforced_after_worker_kill(ownership_gateway: Gateway, tmp_path: Path) -> None: + started: Final = threading.Event() + release_stream: Final = threading.Event() + + def respond(request: Request) -> Reply: + started.set() + reply: Final = _ownership_upstream(request) + return Reply( + status=reply.status, + body=reply.body, + content_type=reply.content_type, + chunks=reply.chunks, + gate_after_first=release_stream, + ) + + with wire_server(respond) as upstream: + with owned_proxy_process( + ownership_gateway, + tmp_path / "worker-kill", + {"LITELLM_SALT_KEY": _OWNED_PROXY_SALT_KEY}, + workers=2, + ) as owned: + with owned.gateway.scenario() as scenario: + model: Final = scenario.model(api_base=f"{upstream.url}/v1") + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model]) + markers: Final = tuple(f"ownership-{uuid4().hex}" for _ in range(24)) + + def serve(index: int) -> tuple[int, int, str]: + try: + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": markers[index]}], + "stream": True, + }, + key=key, + ) + except httpx.TransportError as error: + return index, 0, str(error) + return index, response.status_code, response.text + + candidate_port: Final = owned.gateway.client.base_url.port + assert candidate_port is not None + workers: Final = tuple( + process + for process in group_members(owned.process.pid) + if process.pid != owned.process.pid + and any( + connection.laddr.port == candidate_port and connection.status == psutil.CONN_LISTEN + for connection in process.net_connections(kind="inet") + ) + ) + assert len(workers) == 2, tuple(process.pid for process in workers) + victim: Final = workers[0] + survivor: Final = workers[1] + with ThreadPoolExecutor(max_workers=24) as pool: + futures: Final = tuple(pool.submit(serve, index) for index in range(24)) + try: + assert started.wait(timeout=10) + os.kill(victim.pid, signal.SIGKILL) + release_stream.set() + psutil.wait_procs((victim,), timeout=10) + assert not psutil.pid_exists(victim.pid), victim.pid + finally: + release_stream.set() + burst: Final = tuple(future.result(timeout=90) for future in futures) + assert psutil.pid_exists(survivor.pid), survivor.pid + assert len(burst) == 24 + burst_requests: Final = upstream.drain() + empty_body_count: Final = sum(not request.body for request in burst_requests) + burst_markers: Final = tuple( + _ownership_request_marker(request) for request in burst_requests if request.body + ) + successful_markers: Final = tuple(markers[index] for index, status, _ in burst if status == 200) + assert burst_markers, f"empty_body_captures={empty_body_count}; burst={burst}" + assert all(burst_markers.count(marker) == 1 for marker in successful_markers), ( + f"empty_body_captures={empty_body_count}; " + f"successful_markers={successful_markers}; upstream_markers={burst_markers}; burst={burst}" + ) + cross_team: Final = owned.gateway.request( + "POST", + "/key/generate", + {"team_id": team_a, "project_id": project_b, "models": [model]}, + ) + _discard_unexpected_key(owned.gateway, cross_team) + assert cross_team.status_code == 400, cross_team.text + marker: Final = f"ownership-{uuid4().hex}" + chat: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": marker}]}, + key=key, + ) + assert chat.status_code == 200, chat.text + post_kill_requests: Final = upstream.drain() + post_kill_empty_body_count: Final = sum(not request.body for request in post_kill_requests) + post_kill_markers: Final = tuple( + _ownership_request_marker(request) for request in post_kill_requests if request.body + ) + assert post_kill_markers.count(marker) == 1, ( + f"empty_body_captures={post_kill_empty_body_count}; " + f"post_kill_markers={post_kill_markers}; response={chat.text}" + ) + + @pytest.fixture(scope="module") def ownership_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: with gateway_from_environment() as gateway: From 5e19a24738ec24ca95a2b8f0a5d4c8e3647ffc96 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 07:43:31 +0000 Subject: [PATCH 20/20] fix(proxy): drop budget and permission rows written for a rejected cross-team key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../key_management_endpoints.py | 30 ++++- .../management/test_project_lifecycle.py | 47 +++++++- .../test_key_management_endpoints.py | 112 ++++++++++++++++++ 3 files changed, 181 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 1e10834752a..e35256b56c5 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1379,6 +1379,10 @@ async def _common_key_generation_helper( ) _budget_id = getattr(_budget, "budget_id", None) + created_budget_id: Final[str | None] = ( + _budget_id if prisma_client is not None and data.soft_budget is not None else None + ) + # ADD METADATA FIELDS # Set Management Endpoint Metadata Fields for field in LiteLLM_ManagementEndpoint_MetadataFields_Premium: @@ -1484,10 +1488,16 @@ async def _common_key_generation_helper( for _op_field, _op_default_value in _default_object_permission.items(): _caller_object_permission.setdefault(_op_field, _op_default_value) + should_create_object_permission: Final = prisma_client is not None and isinstance( + data_json.get("object_permission"), dict + ) data_json = await _set_object_permission( data_json=data_json, prisma_client=prisma_client, ) + created_object_permission_id: Final[str | None] = ( + cast(str | None, data_json.get("object_permission_id")) if should_create_object_permission else None + ) _validate_key_alias_format(key_alias=data_json.get("key_alias", None)) @@ -1555,7 +1565,19 @@ async def _common_key_generation_helper( prisma_client=prisma_client, ) - response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key", llm_router=llm_router) + try: + response = await generate_key_helper_fn( + request_type="key", **data_json, table_name="key", llm_router=llm_router + ) + except KeyProjectTeamMismatchError: + if prisma_client is not None: + if created_object_permission_id is not None: + await ObjectPermissionRepository(prisma_client).table.delete( + where={"object_permission_id": created_object_permission_id} + ) + if created_budget_id is not None: + await BudgetRepository(prisma_client).table.delete(where={"budget_id": created_budget_id}) + raise response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response @@ -1831,7 +1853,7 @@ async def _check_key_project_team( if project_obj.team_id is None or project_obj.team_id == key_team_id: return - raise HTTPException( + raise KeyProjectTeamMismatchError( status_code=400, detail={ "error": ( @@ -1843,6 +1865,10 @@ async def _check_key_project_team( ) +class KeyProjectTeamMismatchError(HTTPException): + pass + + async def _check_key_project_team_on_mutation( data: UpdateKeyRequest | RegenerateKeyRequest, existing_key_row: LiteLLM_VerificationToken, diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index b4da1a5a3f5..b1765082a62 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -89,6 +89,13 @@ def _object_permission_table_rows() -> list[dict[str, JsonValue]]: ) +def _budget_table_rows() -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT to_jsonb(b) AS row FROM "LiteLLM_BudgetTable" AS b ORDER BY budget_id', + (), + ) + + def _budget_rows(budget_id: str) -> list[dict[str, JsonValue]]: return read_rows( 'SELECT to_jsonb(b) AS row FROM "LiteLLM_BudgetTable" AS b WHERE budget_id = %s', @@ -337,13 +344,25 @@ def test_service_account_generate_rejects_foreign_team_project_without_writing_k team_b: Final = scenario.team(models=[model]) project_b: Final = scenario.project(team_b, models=[model]) alias: Final = f"service-account-{uuid4().hex}" + vector_store: Final = f"vector-store-{uuid4().hex}" + budget_rows_before: Final = _budget_table_rows() + object_permission_rows_before: Final = _object_permission_table_rows() response: Final = ownership_gateway.request( "POST", "/key/service-account/generate", - {"team_id": team_a, "project_id": project_b, "key_alias": alias, "models": [model]}, + { + "team_id": team_a, + "project_id": project_b, + "key_alias": alias, + "models": [model], + "soft_budget": 3.5, + "object_permission": {"vector_stores": [vector_store]}, + }, ) _discard_unexpected_key(ownership_gateway, response) assert response.status_code == 400, response.text + assert _budget_table_rows() == budget_rows_before + assert _object_permission_table_rows() == object_permission_rows_before assert ( read_rows( 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND key_alias = %s', @@ -1082,17 +1101,33 @@ def test_key_generate_rejects_foreign_team_project(ownership_gateway: Gateway) - team_a: Final = scenario.team(models=[model]) team_b: Final = scenario.team(models=[model]) project_b: Final = scenario.project(team_b, models=[model]) + alias: Final = f"cross-team-{uuid4().hex}" + vector_store: Final = f"vector-store-{uuid4().hex}" + budget_rows_before: Final = _budget_table_rows() + object_permission_rows_before: Final = _object_permission_table_rows() cross_team: Final = ownership_gateway.request( "POST", "/key/generate", - {"team_id": team_a, "project_id": project_b, "models": [model]}, + { + "team_id": team_a, + "project_id": project_b, + "key_alias": alias, + "models": [model], + "soft_budget": 3.5, + "object_permission": {"vector_stores": [vector_store]}, + }, ) _discard_unexpected_key(ownership_gateway, cross_team) assert cross_team.status_code == 400, cross_team.text - assert read_rows( - 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND team_id = %s', - (project_b, team_a), - ) == [] + assert _budget_table_rows() == budget_rows_before + assert _object_permission_table_rows() == object_permission_rows_before + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND team_id = %s', + (project_b, team_a), + ) + == [] + ) def test_key_generate_rejects_missing_team_for_owned_project(ownership_gateway: Gateway) -> None: diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index e62be2499be..77ac29755c4 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -56,6 +56,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( _check_project_key_limits, _check_team_key_limits, _common_key_generation_helper, + KeyProjectTeamMismatchError, _effective_key_after_update, _effective_key_for_generate, _enforce_custom_key_policy, @@ -1823,6 +1824,117 @@ async def test_generate_service_account_works_with_team_id(): ) +def _identity_json_object(value: dict[str, object]) -> dict[str, object]: + return value + + +def _key_generation_prisma_client( + project_record: dict[str, str] | None, +) -> tuple[MagicMock, MagicMock, MagicMock]: + prisma_client: Final = MagicMock() + prisma_client.jsonify_object.side_effect = _identity_json_object + + budget_table: Final = MagicMock() + budget_table.create = AsyncMock(return_value=MagicMock(budget_id="created-budget-id")) + budget_table.delete = AsyncMock() + object_permission_table: Final = MagicMock() + object_permission_table.create = AsyncMock( + return_value=MagicMock(object_permission_id="created-object-permission-id") + ) + object_permission_table.delete = AsyncMock() + + prisma_client.db.litellm_budgettable = budget_table + prisma_client.db.litellm_objectpermissiontable = object_permission_table + prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock(return_value=project_record) + prisma_client.insert_data = AsyncMock() + return prisma_client, budget_table, object_permission_table + + +@pytest.mark.asyncio +async def test_rejected_key_generation_deletes_created_budget_and_default_permission( + monkeypatch: pytest.MonkeyPatch, +) -> None: + prisma_client, budget_table, object_permission_table = _key_generation_prisma_client( + {"project_id": "project-b", "team_id": "team-b"} + ) + monkeypatch.setattr(litellm, "key_generation_settings", None) + monkeypatch.setattr( + litellm, + "default_key_generate_params", + {"object_permission": {"vector_stores": ["default-vector-store"]}}, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False) + + with pytest.raises(KeyProjectTeamMismatchError) as exc: + await _common_key_generation_helper( + data=GenerateKeyRequest( + budget_id="caller-budget-id", + project_id="project-b", + soft_budget=3.5, + team_id="team-a", + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin", + ), + litellm_changed_by=None, + team_table=None, + ) + + assert exc.value.status_code == 400 + budget_table.create.assert_awaited_once() + budget_table.delete.assert_awaited_once_with(where={"budget_id": "created-budget-id"}) + object_permission_table.create.assert_awaited_once_with(data={"vector_stores": ["default-vector-store"]}) + object_permission_table.delete.assert_awaited_once_with( + where={"object_permission_id": "created-object-permission-id"} + ) + prisma_client.insert_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_key_generation_does_not_delete_rows_for_other_http_exceptions( + monkeypatch: pytest.MonkeyPatch, +) -> None: + prisma_client, budget_table, object_permission_table = _key_generation_prisma_client(None) + monkeypatch.setattr(litellm, "key_generation_settings", None) + monkeypatch.setattr( + litellm, + "default_key_generate_params", + {"object_permission": {"vector_stores": ["default-vector-store"]}}, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False) + + with pytest.raises(HTTPException) as exc: + await _common_key_generation_helper( + data=GenerateKeyRequest( + project_id="project-b", + router_settings={"weights": {"gpt-4": {"unknown-deployment": 1.0}}}, + soft_budget=3.5, + team_id="team-a", + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin", + ), + litellm_changed_by=None, + team_table=None, + ) + + assert type(exc.value) is HTTPException + assert exc.value.status_code == 400 + budget_table.create.assert_awaited_once() + budget_table.delete.assert_not_awaited() + object_permission_table.create.assert_awaited_once_with(data={"vector_stores": ["default-vector-store"]}) + object_permission_table.delete.assert_not_awaited() + + @pytest.mark.asyncio async def test_generate_key_throttle_rejected_for_non_admin(): """Security regression: a non-admin creating a key must not be able to set