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>
This commit is contained in:
yucheng 2026-10-03 00:43:46 +00:00
parent f4c161136a
commit 1d24f7fdbc
3 changed files with 78 additions and 20 deletions

View file

@ -1553,13 +1553,6 @@ async def _common_key_generation_helper(
prisma_client=prisma_client,
)
if data.project_id is not None and prisma_client is not None:
await _check_key_project_team(
project_id=data.project_id,
key_team_id=data.team_id,
prisma_client=prisma_client,
)
response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key", llm_router=llm_router)
response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response
@ -4922,6 +4915,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",
@ -5742,6 +5742,13 @@ async def _execute_virtual_key_regeneration(
prisma_client=prisma_client,
)
update_data.update(update_values)
if data is not None:
await _check_key_project_team_on_mutation(
data=data,
existing_key_row=key_in_db,
prisma_client=prisma_client,
)
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,
@ -5766,13 +5773,6 @@ async def _execute_virtual_key_regeneration(
grace_period=data.grace_period if data else None,
)
if data is not None:
await _check_key_project_team_on_mutation(
data=data,
existing_key_row=key_in_db,
prisma_client=prisma_client,
)
updated_token: Final[LiteLLM_VerificationToken | None] = await _prisma_table(
VerificationTokenRepository(prisma_client)
).update(

View file

@ -33,6 +33,20 @@ def _key_rows(key: str) -> list[dict[str, JsonValue]]:
)
def _deleted_key_rows(key: str) -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT token FROM "LiteLLM_DeletedVerificationToken" WHERE token = %s',
(sha256(key.encode()).hexdigest(),),
)
def _deprecated_key_rows(key: str) -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT token FROM "LiteLLM_DeprecatedVerificationToken" WHERE token = %s',
(sha256(key.encode()).hexdigest(),),
)
def _cli_session_token(
user_id: str,
team_id: str | None,
@ -289,13 +303,19 @@ def test_key_regenerate_rejects_foreign_project_without_changing_key(ownership_g
key: Final = string_value(JSON_OBJECT.validate_json(generated.content)["key"])
scenario.cleanups.callback(scenario.delete_key, key)
before: Final = _key_rows(key)
before_deleted: Final = _deleted_key_rows(key)
before_deprecated: Final = _deprecated_key_rows(key)
assert len(before) == 1
assert before_deleted == []
assert before_deprecated == []
response: Final = ownership_gateway.request(
"POST", f"/key/{key}/regenerate", {"project_id": project_b}
"POST", f"/key/{key}/regenerate", {"project_id": project_b, "grace_period": "1h"}
)
_discard_unexpected_key(ownership_gateway, response)
assert response.status_code == 400, response.text
assert _key_rows(key) == before
assert _deleted_key_rows(key) == before_deleted
assert _deprecated_key_rows(key) == before_deprecated
def test_key_generation_rejects_missing_project_without_writing_key(ownership_gateway: Gateway) -> None:

View file

@ -20671,6 +20671,37 @@ async def test_key_generation_organization_membership_error_precedes_project_own
assert error.value.detail == "Caller is not a member of organization_id=org-not-member"
@pytest.mark.asyncio
async def test_key_generation_premium_permission_error_precedes_project_ownership(
monkeypatch: pytest.MonkeyPatch,
) -> None:
user_api_key_cache: Final = await _cache_with_project(
_OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM
)
mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache)
monkeypatch.setattr(litellm, "default_key_generate_params", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
with pytest.raises(HTTPException) as error:
await _common_key_generation_helper(
data=GenerateKeyRequest(
project_id=_OWNED_PROJECT,
team_id=_OWNERSHIP_KEY_TEAM,
permissions={"get_spend_routes": True},
),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
),
litellm_changed_by=None,
team_table=None,
)
assert error.value.status_code == 500
assert error.value.detail == {"error": "Internal Server Error."}
mock_prisma_client.insert_data.assert_not_awaited()
@pytest.mark.asyncio
async def test_key_generation_duplicate_alias_error_precedes_project_ownership(
monkeypatch: pytest.MonkeyPatch,
@ -21175,6 +21206,12 @@ async def test_regenerate_checks_project_team_ownership(
mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock(
return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id="team-b")
)
deleted_history_table: Final = MagicMock()
deleted_history_table.create_many = AsyncMock()
mock_prisma_client.db.litellm_deletedverificationtoken = deleted_history_table
deprecated_key_table: Final = MagicMock()
deprecated_key_table.upsert = AsyncMock()
mock_prisma_client.db.litellm_deprecatedverificationtoken = deprecated_key_table
user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b")
async def regenerate() -> None:
@ -21184,10 +21221,6 @@ async def test_regenerate_checks_project_team_ownership(
new_callable=AsyncMock,
return_value="sk-newtoken1234ab12",
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
@ -21198,7 +21231,10 @@ async def test_regenerate_checks_project_team_ownership(
key_in_db=existing_key,
hashed_api_key="abc123",
key="abc123",
data=RegenerateKeyRequest(project_id=_OWNED_PROJECT),
data=RegenerateKeyRequest(
project_id=_OWNED_PROJECT,
grace_period="1h" if expected_status is not None else None,
),
user_api_key_dict=_make_regenerate_user_api_key_dict(),
litellm_changed_by=None,
user_api_key_cache=user_api_key_cache,
@ -21210,6 +21246,8 @@ async def test_regenerate_checks_project_team_ownership(
await regenerate()
assert error.value.status_code == expected_status
assert "belongs to team team-b, but the key belongs to team-a" in str(error.value.detail)
deleted_history_table.create_many.assert_not_awaited()
deprecated_key_table.upsert.assert_not_awaited()
else:
await regenerate()