mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
bd75b5bf3c
commit
5282aed09b
5 changed files with 3424 additions and 223 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue