mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
fd51cd4a1c
commit
1f1dd295ed
5 changed files with 617 additions and 213 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue