fix(proxy): prevent cross-team project keys

Co-authored-by: L4XB <lukas.buck@e-mail.de>
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-02 09:26:15 +00:00
parent fd51cd4a1c
commit 1f1dd295ed
5 changed files with 617 additions and 213 deletions

View file

@ -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(

View file

@ -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
)

View file

@ -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

View file

@ -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):
"""

View file

@ -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):