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>
This commit is contained in:
yucheng 2026-10-03 07:43:31 +00:00
parent bbbd0f62b3
commit 5e19a24738
3 changed files with 181 additions and 8 deletions

View file

@ -1379,6 +1379,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:
@ -1484,10 +1488,16 @@ 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(
data_json=data_json,
prisma_client=prisma_client,
)
created_object_permission_id: Final[str | None] = (
cast(str | None, data_json.get("object_permission_id")) if should_create_object_permission else None
)
_validate_key_alias_format(key_alias=data_json.get("key_alias", None))
@ -1555,7 +1565,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 KeyProjectTeamMismatchError:
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
@ -1831,7 +1853,7 @@ async def _check_key_project_team(
if project_obj.team_id is None or project_obj.team_id == key_team_id:
return
raise HTTPException(
raise KeyProjectTeamMismatchError(
status_code=400,
detail={
"error": (
@ -1843,6 +1865,10 @@ async def _check_key_project_team(
)
class KeyProjectTeamMismatchError(HTTPException):
pass
async def _check_key_project_team_on_mutation(
data: UpdateKeyRequest | RegenerateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,

View file

@ -89,6 +89,13 @@ def _object_permission_table_rows() -> list[dict[str, JsonValue]]:
)
def _budget_table_rows() -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT to_jsonb(b) AS row FROM "LiteLLM_BudgetTable" AS b ORDER BY budget_id',
(),
)
def _budget_rows(budget_id: str) -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT to_jsonb(b) AS row FROM "LiteLLM_BudgetTable" AS b WHERE budget_id = %s',
@ -337,13 +344,25 @@ def test_service_account_generate_rejects_foreign_team_project_without_writing_k
team_b: Final = scenario.team(models=[model])
project_b: Final = scenario.project(team_b, models=[model])
alias: Final = f"service-account-{uuid4().hex}"
vector_store: Final = f"vector-store-{uuid4().hex}"
budget_rows_before: Final = _budget_table_rows()
object_permission_rows_before: Final = _object_permission_table_rows()
response: Final = ownership_gateway.request(
"POST",
"/key/service-account/generate",
{"team_id": team_a, "project_id": project_b, "key_alias": alias, "models": [model]},
{
"team_id": team_a,
"project_id": project_b,
"key_alias": alias,
"models": [model],
"soft_budget": 3.5,
"object_permission": {"vector_stores": [vector_store]},
},
)
_discard_unexpected_key(ownership_gateway, response)
assert response.status_code == 400, response.text
assert _budget_table_rows() == budget_rows_before
assert _object_permission_table_rows() == object_permission_rows_before
assert (
read_rows(
'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND key_alias = %s',
@ -1082,17 +1101,33 @@ def test_key_generate_rejects_foreign_team_project(ownership_gateway: Gateway) -
team_a: Final = scenario.team(models=[model])
team_b: Final = scenario.team(models=[model])
project_b: Final = scenario.project(team_b, models=[model])
alias: Final = f"cross-team-{uuid4().hex}"
vector_store: Final = f"vector-store-{uuid4().hex}"
budget_rows_before: Final = _budget_table_rows()
object_permission_rows_before: Final = _object_permission_table_rows()
cross_team: Final = ownership_gateway.request(
"POST",
"/key/generate",
{"team_id": team_a, "project_id": project_b, "models": [model]},
{
"team_id": team_a,
"project_id": project_b,
"key_alias": alias,
"models": [model],
"soft_budget": 3.5,
"object_permission": {"vector_stores": [vector_store]},
},
)
_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),
) == []
assert _budget_table_rows() == budget_rows_before
assert _object_permission_table_rows() == object_permission_rows_before
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:

View file

@ -56,6 +56,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
_check_project_key_limits,
_check_team_key_limits,
_common_key_generation_helper,
KeyProjectTeamMismatchError,
_effective_key_after_update,
_effective_key_for_generate,
_enforce_custom_key_policy,
@ -1823,6 +1824,117 @@ async def test_generate_service_account_works_with_team_id():
)
def _identity_json_object(value: dict[str, object]) -> dict[str, object]:
return value
def _key_generation_prisma_client(
project_record: dict[str, str] | None,
) -> tuple[MagicMock, MagicMock, MagicMock]:
prisma_client: Final = MagicMock()
prisma_client.jsonify_object.side_effect = _identity_json_object
budget_table: Final = MagicMock()
budget_table.create = AsyncMock(return_value=MagicMock(budget_id="created-budget-id"))
budget_table.delete = AsyncMock()
object_permission_table: Final = MagicMock()
object_permission_table.create = AsyncMock(
return_value=MagicMock(object_permission_id="created-object-permission-id")
)
object_permission_table.delete = AsyncMock()
prisma_client.db.litellm_budgettable = budget_table
prisma_client.db.litellm_objectpermissiontable = object_permission_table
prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock(return_value=project_record)
prisma_client.insert_data = AsyncMock()
return prisma_client, budget_table, object_permission_table
@pytest.mark.asyncio
async def test_rejected_key_generation_deletes_created_budget_and_default_permission(
monkeypatch: pytest.MonkeyPatch,
) -> None:
prisma_client, budget_table, object_permission_table = _key_generation_prisma_client(
{"project_id": "project-b", "team_id": "team-b"}
)
monkeypatch.setattr(litellm, "key_generation_settings", None)
monkeypatch.setattr(
litellm,
"default_key_generate_params",
{"object_permission": {"vector_stores": ["default-vector-store"]}},
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
with pytest.raises(KeyProjectTeamMismatchError) as exc:
await _common_key_generation_helper(
data=GenerateKeyRequest(
budget_id="caller-budget-id",
project_id="project-b",
soft_budget=3.5,
team_id="team-a",
),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin",
),
litellm_changed_by=None,
team_table=None,
)
assert exc.value.status_code == 400
budget_table.create.assert_awaited_once()
budget_table.delete.assert_awaited_once_with(where={"budget_id": "created-budget-id"})
object_permission_table.create.assert_awaited_once_with(data={"vector_stores": ["default-vector-store"]})
object_permission_table.delete.assert_awaited_once_with(
where={"object_permission_id": "created-object-permission-id"}
)
prisma_client.insert_data.assert_not_awaited()
@pytest.mark.asyncio
async def test_key_generation_does_not_delete_rows_for_other_http_exceptions(
monkeypatch: pytest.MonkeyPatch,
) -> None:
prisma_client, budget_table, object_permission_table = _key_generation_prisma_client(None)
monkeypatch.setattr(litellm, "key_generation_settings", None)
monkeypatch.setattr(
litellm,
"default_key_generate_params",
{"object_permission": {"vector_stores": ["default-vector-store"]}},
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
with pytest.raises(HTTPException) as exc:
await _common_key_generation_helper(
data=GenerateKeyRequest(
project_id="project-b",
router_settings={"weights": {"gpt-4": {"unknown-deployment": 1.0}}},
soft_budget=3.5,
team_id="team-a",
),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin",
),
litellm_changed_by=None,
team_table=None,
)
assert type(exc.value) is HTTPException
assert exc.value.status_code == 400
budget_table.create.assert_awaited_once()
budget_table.delete.assert_not_awaited()
object_permission_table.create.assert_awaited_once_with(data={"vector_stores": ["default-vector-store"]})
object_permission_table.delete.assert_not_awaited()
@pytest.mark.asyncio
async def test_generate_key_throttle_rejected_for_non_admin():
"""Security regression: a non-admin creating a key must not be able to set