From 1f1dd295edce3180868e3bb1321acdf699185c35 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 09:26:15 +0000 Subject: [PATCH] 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):