fix(proxy): keep keys and projects on the same team (#44165)

* fix(key mgmt): a project may only be attached to keys of its own team

A project is created under exactly one team, and its budget and models are
validated against that team's. Nothing checked that a key's team matched, so
`POST /key/generate` with `team_id: team-a` and a `project_id` owned by
`team-b` answered 200 and stored exactly that — a key belonging to team-a,
recorded under team-b's project and validated against its models and budget.
The same held with no team at all.

It needs no proxy-admin rights: an admin of team-a who is a member of no other
team can issue keys under another tenant's project with nothing but the project
id, and that tenant sees the spend without having granted anything.

`_check_project_key_limits` now takes the key's team and refuses a project
owned by a different one, before the model and budget checks so the refusal
reads as what it is. A project with no owning team is left alone — nobody owns
it, so there is no boundary to cross. `/key/update` passes the key's stored
team when the request does not carry one, and now runs the check whenever the
project itself is set or changed, which a request that moves only the project
previously skipped.

Two existing cells needed the project's team passed explicitly: they measure
the model allowlist, not tenancy, and would otherwise have been asserting model
behaviour on a request the new gate refuses. A third builds an unowned project
for the same reason.

Fixes #41089

* fix(key mgmt): check project ownership on every key mutation path

/key/update ran the project check only when project_id, models or
max_budget were supplied, so a request that changed team_id alone left a
foreign project attached. /key/regenerate never ran it at all. Both now
go through one helper that validates the key as the mutation leaves it.

On /key/generate the check moves after default_key_generate_params is
applied, because that can supply team_id; it was rejecting a valid key
whose team came from the defaults.

Tests move into the mapped test file per CLAUDE.md.

* test(key mgmt): read the project from the cache instead of patching the lookup

The test-quality gate rejected the new cells: 12 TQ008 for patching
`get_project_object`, an SDK internal, and 4 TQ001 for accept controls that
could only fail by raising.

The cells now seed `UserApiKeyCache` with the project, which is the idiom
the neighbouring project cells already use and removes the patching. The
accept controls are parametrised together with the rejecting ones, so each
test function carries a real assertion. The three regenerate seams that stay
stubbed each carry a reason.

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

* fix(proxy): validate bulk key team changes

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): use behavioral assertions in project ownership tests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): read project owner from the database for key ownership checks

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): read project ownership from the primary database under the lookup deadline

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): check project moves against the primary database team

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): run ownership checks after existing validation without the lookup deadline

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): type ownership writer reads

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): make ownership tests independent of runner salt and clock

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): check key project ownership after existing validation

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): run key project ownership right before the key row write

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): reject key project ownership before permission and budget writes

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): keep project object permission validation on the stored payload

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* style(proxy): format key regeneration update payload

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): invalidate the written permission id on bulk update and regenerate

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): audit project team ownership across endpoints, bulk paths and chaos

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): drop budget and permission rows written for a rejected cross-team key

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(keys): require project-team membership when regenerate binds a key to a team-owned project

/key/{key}/regenerate and /key/regenerate accepted a team_id + project_id
pair without checking that the caller belongs to the destination team, so a
user with no team could regenerate their own key into another team and its
project. Regenerate now requires a proxy admin or a member/admin of the
project's team whenever the change lands the key on a team-owned project.

Also adapts tests to main's premium_user_check rename, drops a duplicate
import left by the merge, and replaces the created object-permission id cast
with an isinstance check.

* fix(keys): 404 a missing project before any key write, and match the project move refusal to its siblings

A missing project_id on /key/service-account/generate and on both regenerate
routes fell through to a foreign-key 500; regenerate had already written a
deleted-token history row, and service-account generate left its budget and
object-permission rows behind. The ownership check now answers 404 with the
same message /key/generate uses, before any write, and the generate rollback
covers it. The /project/update move refusal now uses ProxyException
(bad_request, param team_id) like the delete-with-keys refusal.

Tests: integration coverage for a team admin regenerating a team key into
another team's project (403), the missing-project 404s with no history rows,
and the unit test that never reached the ownership check is removed.

* test(keys): keep the existing-permission case in the regeneration cache eviction test

Parametrize the test over a key that already has an object_permission_id and one that mints a new row, so eviction under the key's existing permission id stays covered.

---------

Co-authored-by: L4XB <lukas.buck@e-mail.de>
Co-authored-by: yucheng <yucheng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-10 00:53:10 -07:00 • committed by GitHub
parent bd75b5bf3c
commit 5282aed09b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 3424 additions and 223 deletions

View file

@ -12,7 +12,7 @@ Endpoints for /project operations
import json
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Final
from typing import TYPE_CHECKING, Final, cast
from fastapi import APIRouter, Depends, HTTPException, Request
from pydantic import TypeAdapter
@ -31,6 +31,7 @@ from litellm.proxy.management_helpers.utils import (
management_endpoint_wrapper,
)
from litellm.proxy.utils import PrismaClient, handle_exception_on_proxy
from litellm.repositories.base_repository import record_to_dict
from litellm.repositories.budget_repository import BudgetRepository
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.prisma_protocols import TableActions
@ -57,6 +58,15 @@ def _project_table(prisma_client: PrismaClient) -> TableActions["prisma_models.L
return ProjectRepository(prisma_client).table
def _writer_project_table(
prisma_client: PrismaClient,
) -> TableActions["prisma_models.LiteLLM_ProjectTable"]:
return cast( # cast-ok: writer_db exposes generated Prisma tables through a dynamic wrapper
"TableActions[prisma_models.LiteLLM_ProjectTable]",
prisma_client.writer_db.litellm_projecttable,
)
def _verification_token_table(
prisma_client: PrismaClient,
) -> TableActions["prisma_models.LiteLLM_VerificationToken"]:
@ -824,6 +834,38 @@ async def update_project(
**({"max_budget": None} if "max_budget" in data.model_fields_set and data.max_budget is None else {}),
}
object_permission_data: Final = update_data.pop("object_permission", None)
object_permission_payload: Final = (
_OBJECT_PERMISSION_PAYLOAD.validate_python(object_permission_data) if object_permission_data else None
)
if data.team_id is not None:
current_project_record: Final = await _writer_project_table(prisma_client).find_unique(
where={"project_id": data.project_id}
)
current_project: Final = (
LiteLLM_ProjectTable.model_validate(record_to_dict(current_project_record))
if current_project_record is not None
else None
)
if current_project is not None and data.team_id != current_project.team_id:
mismatched_key_count: Final = await prisma_client.writer_db.litellm_verificationtoken.count(
where={
"project_id": data.project_id,
"OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}],
}
)
if mismatched_key_count > 0:
raise ProxyException(
message=(
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."
),
type="bad_request",
code=400,
param="team_id",
)
if budget_updates and existing_project.budget_id:
# Update existing budget
await _budget_table(prisma_client).update(
@ -837,23 +879,17 @@ async def update_project(
for field in budget_updates.keys():
update_data.pop(field, None)
# Handle object permissions
if "object_permission" in update_data:
object_permission_data = update_data.pop("object_permission")
if object_permission_data:
object_permission_payload: Final = _OBJECT_PERMISSION_PAYLOAD.validate_python(object_permission_data)
if existing_project.object_permission_id:
# Update existing permission
await _object_permission_table(prisma_client).update(
where={"object_permission_id": existing_project.object_permission_id},
data=object_permission_payload,
)
else:
# Create new permission
created_permission: Final = await _object_permission_table(prisma_client).create(
data=object_permission_payload,
)
update_data["object_permission_id"] = created_permission.object_permission_id
if object_permission_payload is not None:
if existing_project.object_permission_id:
await _object_permission_table(prisma_client).update(
where={"object_permission_id": existing_project.object_permission_id},
data=object_permission_payload,
)
else:
created_permission: Final = await _object_permission_table(prisma_client).create(
data=object_permission_payload,
)
update_data["object_permission_id"] = created_permission.object_permission_id
# Handle metadata fields
for field in LiteLLM_ManagementEndpoint_MetadataFields:

View file

@ -43,6 +43,7 @@ from litellm.constants import (
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.models.credentials import CredentialItem
from litellm.models.project import LiteLLM_ProjectTable
from litellm.proxy._experimental.mcp_server.db import (
rotate_mcp_server_credentials_master_key,
rotate_mcp_user_credentials_master_key,
@ -122,10 +123,11 @@ from litellm.proxy.management_helpers.access_group_key_sync import (
)
from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at
from litellm.proxy.management_helpers.object_permission_utils import ( # noqa: F401 # legacy module exports
ObjectPermissionUpsert,
_set_object_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
attach_object_permission_to_dict,
handle_update_object_permission_common,
invalidate_cached_object_permissions,
prepare_object_permission_upsert,
set_object_permission,
validate_key_mcp_servers_against_team,
validate_key_search_tools_against_team,
@ -150,11 +152,12 @@ from litellm.proxy.utils import ( # noqa: F401 # legacy module exports
hash_token_if_needed,
is_valid_api_key,
)
from litellm.repositories.base_repository import BaseRepository
from litellm.repositories.base_repository import BaseRepository, record_to_dict
from litellm.repositories.budget_repository import BudgetRepository
from litellm.repositories.config_repository import ConfigParam, ConfigRepository
from litellm.repositories.credentials_repository import CredentialsRepository
from litellm.repositories.model_repository import ModelRepository
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import (
DeletedVerificationTokenRepository,
@ -279,6 +282,19 @@ def _prisma_table(
)
def _writer_project_table(
prisma_client: PrismaClient,
) -> "TableActions[prisma_models.LiteLLM_ProjectTable]":
return cast( # cast-ok: writer_db exposes generated Prisma tables through a dynamic wrapper
"TableActions[prisma_models.LiteLLM_ProjectTable]",
prisma_client.writer_db.litellm_projecttable,
)
class KeyProjectBindingError(HTTPException):
pass
def _deleted_verification_token_table(
prisma_client: PrismaClient,
) -> "TableActions[prisma_models.LiteLLM_DeletedVerificationToken]":
@ -1361,7 +1377,8 @@ async def _common_key_generation_helper(
apply_enterprise_key_management_params,
)
data = apply_enterprise_key_management_params(data, team_table)
enterprise_data: Final[object] = apply_enterprise_key_management_params(data, team_table)
data = GenerateKeyRequest.model_validate(enterprise_data)
except Exception as e:
verbose_proxy_logger.debug(
"litellm.proxy.proxy_server.generate_key_fn(): Enterprise key management params not applied - %s", e
@ -1396,6 +1413,10 @@ async def _common_key_generation_helper(
)
_budget_id = getattr(_budget, "budget_id", None)
created_budget_id: Final[str | None] = (
_budget_id if prisma_client is not None and data.soft_budget is not None else None
)
# ADD METADATA FIELDS
# Set Management Endpoint Metadata Fields
for field in LiteLLM_ManagementEndpoint_MetadataFields_Premium:
@ -1501,10 +1522,19 @@ async def _common_key_generation_helper(
for _op_field, _op_default_value in _default_object_permission.items():
_caller_object_permission.setdefault(_op_field, _op_default_value)
should_create_object_permission: Final = prisma_client is not None and isinstance(
data_json.get("object_permission"), dict
)
data_json = await set_object_permission( # rebind-ok: pre-existing rebinding on a rename-only line
data_json=data_json,
prisma_client=prisma_client,
)
written_object_permission_id: Final = data_json.get("object_permission_id")
created_object_permission_id: Final = (
written_object_permission_id
if should_create_object_permission and isinstance(written_object_permission_id, str)
else None
)
_validate_key_alias_format(key_alias=data_json.get("key_alias", None))
@ -1572,7 +1602,19 @@ async def _common_key_generation_helper(
prisma_client=prisma_client,
)
response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key", llm_router=llm_router)
try:
response = await generate_key_helper_fn(
request_type="key", **data_json, table_name="key", llm_router=llm_router
)
except KeyProjectBindingError:
if prisma_client is not None:
if created_object_permission_id is not None:
await ObjectPermissionRepository(prisma_client).table.delete(
where={"object_permission_id": created_object_permission_id}
)
if created_budget_id is not None:
await BudgetRepository(prisma_client).table.delete(where={"budget_id": created_budget_id})
raise
response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response
@ -1834,6 +1876,85 @@ async def _check_project_key_limits(
)
async def _check_key_project_team(
project_id: str,
key_team_id: str | None,
prisma_client: PrismaClient,
) -> str | None:
project_record: Final = await _writer_project_table(prisma_client).find_unique(where={"project_id": project_id})
if project_record is None:
raise KeyProjectBindingError(
status_code=404,
detail={"error": f"Project not found, project_id={project_id}"},
)
project_obj: Final = LiteLLM_ProjectTable.model_validate(record_to_dict(project_record))
if project_obj.team_id is None or project_obj.team_id == key_team_id:
return project_obj.team_id
raise KeyProjectBindingError(
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_key_project_team_on_mutation(
data: UpdateKeyRequest | RegenerateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
prisma_client: PrismaClient,
) -> str | None:
fields_set: Final = data.model_fields_set
team_changed: Final = "team_id" in fields_set and data.team_id != existing_key_row.team_id
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 None
project_id: Final = data.project_id if "project_id" in fields_set else existing_key_row.project_id
if project_id is None:
return None
team_id: Final = data.team_id if "team_id" in fields_set else existing_key_row.team_id
return await _check_key_project_team(
project_id=project_id,
key_team_id=team_id,
prisma_client=prisma_client,
)
async def _check_caller_in_project_team(
project_team_id: str,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
) -> None:
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
return
team_table: Final = await get_team_object(
team_id=project_team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_dict.parent_otel_span,
check_db_only=True,
)
if get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) is None:
raise HTTPException(
status_code=403,
detail={
"error": (
f"Only a proxy admin or a member of team {project_team_id} can attach a key to a project "
"owned by that team."
)
},
)
def check_org_key_model_specific_limits(
keys: Sequence[LiteLLM_VerificationToken],
org_table: LiteLLM_OrganizationTable,
@ -2488,6 +2609,18 @@ async def _update_key_row_with_soft_budget(
return result
async def _update_key_row(
prisma_client: PrismaClient,
key: str,
update_values: Mapping[str, object],
) -> _KeyUpdateResult | None:
key_update_data: Final = MappingProxyType({**update_values, "token": key})
response: Final = await prisma_client.update_data(token=key, data=key_update_data)
if response is None:
return None
return cast("_KeyUpdateResult", response) # cast-ok: key update_data returns token and data
async def prepare_key_update_data(
data: UpdateKeyRequest | RegenerateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
@ -2594,27 +2727,42 @@ async def prepare_key_update_data(
return non_default_values
async def _handle_update_object_permission(
data_json: dict,
async def _prepare_key_update_object_permission(
object_permission_data: object,
existing_key_row: LiteLLM_VerificationToken,
prisma_client: PrismaClient,
) -> dict:
"""Persist the requested object permission row and swap it for its id, only after the key policy allowed the write."""
if "object_permission" not in data_json:
return data_json
) -> ObjectPermissionUpsert | None:
if object_permission_data is None:
return None
object_permission_id: Final = await handle_update_object_permission_common(
data_json=data_json,
parsed_object_permission: Final[object] = (
json.loads(object_permission_data) if isinstance(object_permission_data, str) else object_permission_data
)
permission_data: Final[dict[str, object]] = (
TypeAdapter(dict[str, object]).validate_python(parsed_object_permission)
if isinstance(parsed_object_permission, dict)
else {}
)
return await prepare_object_permission_upsert(
new_object_permission=permission_data,
existing_object_permission_id=existing_key_row.object_permission_id,
prisma_client=prisma_client,
)
# Add the object_permission_id to data_json if one was created/updated
if object_permission_id is not None:
data_json["object_permission_id"] = object_permission_id
verbose_proxy_logger.debug("updated object_permission_id: %s", object_permission_id)
return data_json
async def _write_prepared_key_update_object_permission(
data_json: Mapping[str, object],
upsert: ObjectPermissionUpsert | None,
prisma_client: PrismaClient,
) -> Mapping[str, object]:
if upsert is None:
return data_json
await ObjectPermissionRepository(prisma_client).table.upsert(
where={"object_permission_id": upsert.object_permission_id},
data={"create": upsert.record, "update": upsert.record},
)
return MappingProxyType({**data_json, "object_permission_id": upsert.object_permission_id})
def is_different_team(data: UpdateKeyRequest, existing_key_row: LiteLLM_VerificationToken) -> bool:
@ -2726,6 +2874,23 @@ def _resolve_token_to_update(data: UpdateKeyRequest, existing_key_row: LiteLLM_V
return existing_key_row.token
async def _check_single_key_update_team_permissions(
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient | None,
existing_key_row: LiteLLM_VerificationToken,
user_api_key_cache: UserApiKeyCache,
) -> None:
if prisma_client is None:
return
await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint(
user_api_key_dict=user_api_key_dict,
route=KeyManagementRoutes.KEY_UPDATE,
prisma_client=prisma_client,
existing_key_row=existing_key_row,
user_api_key_cache=user_api_key_cache,
)
async def _process_single_key_update(
update_key_request: UpdateKeyRequest,
user_api_key_dict: UserAPIKeyAuth,
@ -2789,15 +2954,12 @@ async def _process_single_key_update(
entity="key",
)
# Check team member permissions
if prisma_client is not None:
await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint(
user_api_key_dict=user_api_key_dict,
route=KeyManagementRoutes.KEY_UPDATE,
prisma_client=prisma_client,
existing_key_row=existing_key_row,
user_api_key_cache=user_api_key_cache,
)
await _check_single_key_update_team_permissions(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
existing_key_row=existing_key_row,
user_api_key_cache=user_api_key_cache,
)
check_denied_passthrough_routes_caller_permission(
update_key_request,
user_api_key_dict,
@ -2887,11 +3049,24 @@ async def _process_single_key_update(
detail={"error": "Database not connected"},
)
update_values: Final = await _handle_update_object_permission(
data_json=non_default_values,
object_permission_upsert: Final = await _prepare_key_update_object_permission(
object_permission_data=non_default_values.get("object_permission"),
existing_key_row=existing_key_row,
prisma_client=prisma_client,
)
key_update_values: Final = MappingProxyType(
{field: value for field, value in non_default_values.items() if field != "object_permission"}
)
await _check_key_project_team_on_mutation(
data=key_request,
existing_key_row=existing_key_row,
prisma_client=prisma_client,
)
update_values: Final = await _write_prepared_key_update_object_permission(
data_json=key_update_values,
upsert=object_permission_upsert,
prisma_client=prisma_client,
)
_data: Final = {**update_values, "token": key_request.key}
response: Final[Mapping[str, object] | None] = cast( # cast-ok: every update_data branch returns a str-keyed dict
"Mapping[str, object] | None",
@ -2902,7 +3077,7 @@ async def _process_single_key_update(
await invalidate_cached_object_permissions(
object_permission_ids=(
existing_key_row.object_permission_id,
non_default_values.get("object_permission_id"),
update_values.get("object_permission_id"),
),
user_api_key_cache=user_api_key_cache,
)
@ -3563,11 +3738,24 @@ async def update_key_fn(
if prisma_client is None:
raise Exception("Not connected to DB!")
update_values: Final = await _handle_update_object_permission(
data_json=non_default_values,
object_permission_upsert: Final = await _prepare_key_update_object_permission(
object_permission_data=non_default_values.get("object_permission"),
existing_key_row=existing_key_row,
prisma_client=prisma_client,
)
key_update_values: Final = MappingProxyType(
{field: value for field, value in non_default_values.items() if field != "object_permission"}
)
await _check_key_project_team_on_mutation(
data=data,
existing_key_row=existing_key_row,
prisma_client=prisma_client,
)
update_values: Final = await _write_prepared_key_update_object_permission(
data_json=key_update_values,
upsert=object_permission_upsert,
prisma_client=prisma_client,
)
changed_by: Final = user_api_key_dict.user_id or litellm_proxy_admin_name
response: Final = (
await _update_key_row_with_soft_budget(
@ -3579,7 +3767,11 @@ async def update_key_fn(
changed_by=changed_by,
)
if "soft_budget" in data.model_fields_set
else await prisma_client.update_data(token=key, data=MappingProxyType({**update_values, "token": key}))
else await _update_key_row(
prisma_client=prisma_client,
key=key,
update_values=update_values,
)
)
# Delete - key from cache, since it's been updated!
@ -3587,7 +3779,7 @@ async def update_key_fn(
await invalidate_cached_object_permissions(
object_permission_ids=(
existing_key_row.object_permission_id,
non_default_values.get("object_permission_id"),
update_values.get("object_permission_id"),
),
user_api_key_cache=user_api_key_cache,
)
@ -4846,6 +5038,13 @@ async def generate_key_helper_fn(
# the LiteLLM_VerificationToken table will increase in size if we don't do this check
return user_data
if project_id is not None:
await _check_key_project_team(
project_id=project_id,
key_team_id=team_id,
prisma_client=prisma_client,
)
## CREATE KEY
verbose_proxy_logger.debug(
"prisma_client: Creating Key= %s",
@ -5655,7 +5854,6 @@ async def _execute_virtual_key_regeneration(
new_token: Final = await get_new_token(data=data)
new_token_hash: Final = hash_token(new_token)
new_token_key_name: Final = abbreviate_api_key(api_key=new_token)
update_data = {"token": new_token_hash, "key_name": new_token_key_name}
non_default_values = {}
if data is not None:
@ -5681,12 +5879,34 @@ async def _execute_virtual_key_regeneration(
request=data if data is not None else RegenerateKeyRequest(),
),
)
update_values: Final = await _handle_update_object_permission(
data_json=non_default_values,
object_permission_upsert: Final = await _prepare_key_update_object_permission(
object_permission_data=non_default_values.get("object_permission"),
existing_key_row=key_in_db,
prisma_client=prisma_client,
)
update_data.update(update_values)
key_update_values: Final = MappingProxyType(
{field: value for field, value in non_default_values.items() if field != "object_permission"}
)
if data is not None:
project_team_id: Final = await _check_key_project_team_on_mutation(
data=data,
existing_key_row=key_in_db,
prisma_client=prisma_client,
)
if project_team_id is not None:
await _check_caller_in_project_team(
project_team_id=project_team_id,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
update_values: Final = await _write_prepared_key_update_object_permission(
data_json=key_update_values,
upsert=object_permission_upsert,
prisma_client=prisma_client,
)
update_data: Final = MappingProxyType({"token": new_token_hash, "key_name": new_token_key_name, **update_values})
jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data)
# Snapshot before the token update: the FK cascade rewrites mapping rows to the new hash,
@ -5724,7 +5944,7 @@ async def _execute_virtual_key_regeneration(
await invalidate_cached_object_permissions(
object_permission_ids=(
key_in_db.object_permission_id,
non_default_values.get("object_permission_id"),
update_values.get("object_permission_id"),
),
user_api_key_cache=user_api_key_cache,
)

File diff suppressed because it is too large Load diff

View file

@ -1,3 +1,4 @@
import asyncio
import os
import traceback
from collections.abc import Mapping
@ -10,6 +11,9 @@ from unittest import mock
from dotenv import load_dotenv
from fastapi import HTTPException, Request
from litellm.constants import PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS
from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker
load_dotenv()
import time
@ -36,6 +40,7 @@ from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.caching.caching import DualCache
from litellm.proxy._types import (
LiteLLM_ProjectTable,
NewProjectRequest,
UpdateProjectRequest,
DeleteProjectRequest,
@ -1241,6 +1246,11 @@ def _project_update_mocks(monkeypatch, stored_metadata: dict) -> mock.MagicMock:
mock_prisma.jsonify_object = lambda data: data
mock_prisma.db.litellm_projecttable.find_unique = mock.AsyncMock(return_value=existing_row)
mock_prisma.db.litellm_projecttable.update = mock.AsyncMock(return_value=mock.MagicMock())
mock_prisma.writer_db = mock.MagicMock()
mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock(
return_value={"project_id": "project-update-test", "team_id": None}
)
mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(return_value=0)
monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", True)
monkeypatch.setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma)
@ -1264,6 +1274,207 @@ def _written_project_data(mock_prisma: mock.MagicMock) -> dict:
return mock_prisma.db.litellm_projecttable.update.await_args.kwargs["data"]
@pytest.mark.asyncio
async def test_update_project_object_permission_validation_precedes_budget_write(
monkeypatch: pytest.MonkeyPatch,
) -> None:
project_id: Final = "project-object-permission-validation"
mock_prisma: Final = _project_update_mocks(monkeypatch, {})
mock_prisma.db.litellm_projecttable.find_unique.return_value.budget_id = "budget-project"
budget_table: Final = mock.MagicMock()
budget_table.update = mock.AsyncMock()
mock_prisma.db.litellm_budgettable = budget_table
def jsonify_object_permission_as_string(payload: dict[str, object]) -> dict[str, object]:
return {**payload, "object_permission": '{"vector_stores": ["replacement-store"]}'}
mock_prisma.jsonify_object = jsonify_object_permission_as_string
with pytest.raises(ProxyException) as error:
await _run_project_update(
project_id,
max_budget=50,
object_permission={"vector_stores": ["replacement-store"]},
)
assert error.value.code == "500"
assert "Input should be a valid dictionary" in error.value.message
assert "input_type=str" in error.value.message
budget_table.update.assert_not_awaited()
mock_prisma.db.litellm_projecttable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_project_rejects_move_when_attached_teamless_key_exists(
monkeypatch: pytest.MonkeyPatch,
) -> 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_projecttable.find_unique.return_value.budget_id = "budget-project"
budget_table: Final = mock.MagicMock()
budget_table.update = mock.AsyncMock()
mock_prisma.db.litellm_budgettable = budget_table
mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock(
return_value=LiteLLM_TeamTable(team_id=destination_team_id)
)
mock_prisma.db.litellm_verificationtoken.count = mock.AsyncMock(return_value=0)
mock_prisma.writer_db = mock.MagicMock()
mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock(
return_value=LiteLLM_ProjectTable(project_id=project_id, team_id="team-a")
)
async def count_teamless_keys(*, where: Mapping[str, object]) -> int:
conditions: Final = where.get("OR")
return int(isinstance(conditions, list) and {"team_id": None} in conditions)
mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(side_effect=count_teamless_keys)
with pytest.raises(ProxyException) as error:
await _run_project_update(
project_id,
team_id=destination_team_id,
max_budget=50,
)
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.writer_db.litellm_verificationtoken.count.assert_awaited_once()
mock_prisma.db.litellm_verificationtoken.count.assert_not_awaited()
budget_table.update.assert_not_awaited()
mock_prisma.db.litellm_projecttable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_project_rejects_move_when_writer_team_differs_from_stale_reader(
monkeypatch: pytest.MonkeyPatch,
) -> None:
project_id: Final = "project-replica-lag"
destination_team_id: Final = "team-a"
mock_prisma: Final = _project_update_mocks(monkeypatch, {})
mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = destination_team_id
mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock(
return_value=LiteLLM_TeamTable(team_id=destination_team_id)
)
mock_prisma.db.litellm_verificationtoken.count = mock.AsyncMock(return_value=0)
mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock(
return_value=LiteLLM_ProjectTable(project_id=project_id, team_id="team-b")
)
mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(return_value=1)
with pytest.raises(ProxyException) as error:
await _run_project_update(project_id, team_id=destination_team_id)
assert error.value.code == "400"
assert (
f"Project {project_id} has 1 key(s) that do not belong to team {destination_team_id}. "
"Detach or delete them before moving the project."
) in error.value.message
mock_prisma.writer_db.litellm_projecttable.find_unique.assert_awaited_once()
mock_prisma.writer_db.litellm_verificationtoken.count.assert_awaited_once()
mock_prisma.db.litellm_verificationtoken.count.assert_not_awaited()
mock_prisma.db.litellm_projecttable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_project_slow_writer_ownership_reads_do_not_stall_db_tracker(
monkeypatch: pytest.MonkeyPatch,
) -> None:
project_id: Final = "project-slow-writer"
source_team_id: Final = "team-source"
destination_team_id: Final = "team-destination"
mock_prisma: Final = _project_update_mocks(monkeypatch, {})
mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = source_team_id
mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock(
return_value=LiteLLM_TeamTable(team_id=destination_team_id, models=[])
)
async def slow_project_lookup(*, where: Mapping[str, object]) -> LiteLLM_ProjectTable:
assert where == {"project_id": project_id}
await asyncio.sleep(0)
await asyncio.sleep(0)
return LiteLLM_ProjectTable(project_id=project_id, team_id=source_team_id)
async def slow_key_count(*, where: Mapping[str, object]) -> int:
assert where == {
"project_id": project_id,
"OR": [{"team_id": {"not": destination_team_id}}, {"team_id": None}],
}
await asyncio.sleep(0)
await asyncio.sleep(0)
return 0
mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock(side_effect=slow_project_lookup)
mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(side_effect=slow_key_count)
monkeypatch.setattr("litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.0)
db_lookup_stall_tracker.clear()
try:
await _run_project_update(project_id, team_id=destination_team_id)
assert mock_prisma.writer_db.litellm_projecttable.find_unique.await_count == 1
assert mock_prisma.writer_db.litellm_verificationtoken.count.await_count == 1
mock_prisma.db.litellm_projecttable.update.assert_awaited_once()
assert db_lookup_stall_tracker.stalled_within(PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS) is False
finally:
db_lookup_stall_tracker.clear()
@pytest.mark.asyncio
async def test_update_project_team_limit_error_precedes_mismatched_key_guard(
monkeypatch: pytest.MonkeyPatch,
) -> None:
project_id: Final = "project-team-limit-precedence"
source_team_id: Final = "team-source"
destination_team_id: Final = "team-destination"
mock_prisma: Final = _project_update_mocks(monkeypatch, {})
mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = source_team_id
mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock(
return_value=LiteLLM_TeamTable(
team_id=destination_team_id,
models=["allowed-model"],
)
)
mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock(
return_value=LiteLLM_ProjectTable(project_id=project_id, team_id=source_team_id)
)
mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(return_value=1)
with pytest.raises(ProxyException, match="not in team's allowed models") as error:
await _run_project_update(project_id, team_id=destination_team_id, models=["disallowed-model"])
assert error.value.code == "400"
mock_prisma.writer_db.litellm_verificationtoken.count.assert_not_awaited()
mock_prisma.db.litellm_projecttable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_project_allows_move_when_no_attached_keys_exist(monkeypatch: pytest.MonkeyPatch) -> None:
project_id: Final = "project-without-keys"
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)
mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock(
return_value=LiteLLM_ProjectTable(project_id=project_id, team_id="team-a")
)
await _run_project_update(project_id, team_id=destination_team_id)
mock_prisma.writer_db.litellm_verificationtoken.count.assert_awaited_once()
mock_prisma.db.litellm_verificationtoken.count.assert_not_awaited()
mock_prisma.db.litellm_projecttable.update.assert_awaited_once()
@pytest.mark.asyncio
async def test_update_project_clears_model_itpm_limit_sent_as_an_empty_map(monkeypatch):
"""