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>
This commit is contained in:
yucheng 2026-10-02 17:15:07 +00:00
parent 1049ac4bff
commit 7e55e90847
2 changed files with 66 additions and 15 deletions

View file

@ -139,6 +139,7 @@ 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.prisma_protocols import TableActions
from litellm.repositories.project_repository import ProjectRepository
from litellm.repositories.table_repositories import (
DeletedVerificationTokenRepository,
DeprecatedVerificationTokenRepository,
@ -1268,7 +1269,6 @@ async def _common_key_generation_helper(
project_id=data.project_id,
key_team_id=data.team_id,
prisma_client=prisma_client,
user_api_key_cache=proxy_server.user_api_key_cache,
)
# Delegated-authority ceiling (GHSA-q775-qw9r-2r4g): a non-admin caller
@ -1816,13 +1816,8 @@ async def _check_key_project_team(
project_id: str,
key_team_id: str | None,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
) -> None:
project_obj: Final = await get_project_object(
project_id=project_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
project_obj: Final = await ProjectRepository(prisma_client).find_by_id(project_id)
if project_obj is None:
raise HTTPException(
@ -1849,7 +1844,6 @@ async def _check_key_project_team_on_mutation(
data: UpdateKeyRequest | RegenerateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
) -> None:
fields_set: Final = data.model_fields_set
team_changed: Final = "team_id" in fields_set and data.team_id != existing_key_row.team_id
@ -1866,7 +1860,6 @@ async def _check_key_project_team_on_mutation(
project_id=project_id,
key_team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
@ -2905,7 +2898,6 @@ async def _process_single_key_update(
data=update_key_request,
existing_key_row=existing_key_row,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
key_request: Final = await _with_validated_object_permission(
@ -3369,7 +3361,6 @@ async def _validate_update_key_data(
data=data,
existing_key_row=existing_key_row,
prisma_client=checked_prisma_client,
user_api_key_cache=user_api_key_cache,
)
# When the caller asks to change the key's organization_id, require that
@ -5668,7 +5659,6 @@ async def _execute_virtual_key_regeneration(
data=data,
existing_key_row=key_in_db,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
_existing_key_metadata: Final = getattr(key_in_db, "metadata", None)
enforce_output_token_estimates_are_admin_only(

View file

@ -20544,6 +20544,18 @@ def _configure_key_endpoints(
) -> AsyncMock:
mock_prisma_client: Final = _make_generate_mock_prisma()
mock_prisma_client.writer_db = mock_prisma_client.db
project_obj: Final = user_api_key_cache.get_cache(
key=project_cache_key(_OWNED_PROJECT),
model_type=LiteLLM_ProjectTableCachedObj,
)
mock_prisma_client.db.litellm_projecttable = MagicMock()
mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock(
return_value=(
LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=project_obj.team_id)
if project_obj is not None
else None
)
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache)
return mock_prisma_client
@ -20578,6 +20590,54 @@ async def test_key_generation_rejects_foreign_project_team(
assert error.value.detail == expected_detail
@pytest.mark.parametrize(
("key_team_id", "expected_status"),
[(_OWNERSHIP_PROJECT_TEAM, 200), (_OWNERSHIP_KEY_TEAM, 400)],
)
@pytest.mark.asyncio
async def test_key_generation_uses_database_project_team_when_cache_is_stale(
monkeypatch: pytest.MonkeyPatch,
key_team_id: str,
expected_status: int,
) -> None:
user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_KEY_TEAM)
prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache)
prisma_client.db.litellm_projecttable = MagicMock()
prisma_client.db.litellm_projecttable.find_unique = AsyncMock(
return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM)
)
monkeypatch.setattr(litellm, "default_key_generate_params", None)
data: Final = GenerateKeyRequest(project_id=_OWNED_PROJECT, team_id=key_team_id)
user_api_key_dict: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin")
if expected_status == 200:
response: Final = await _common_key_generation_helper(
data=data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
team_table=None,
)
assert response.team_id == _OWNERSHIP_PROJECT_TEAM
return
with pytest.raises(HTTPException) as error:
await _common_key_generation_helper(
data=data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
team_table=None,
)
expected_detail: Final = {
"error": (
f"Project {_OWNED_PROJECT} belongs to team {_OWNERSHIP_PROJECT_TEAM}, but the key belongs to "
f"{_OWNERSHIP_KEY_TEAM}. A key can only be attached to a project owned by its own team."
)
}
assert error.value.status_code == 400
assert error.value.detail == expected_detail
@pytest.mark.parametrize("project_team_id", ["team-a", None])
@pytest.mark.asyncio
async def test_key_generation_accepts_same_team_and_unowned_projects(
@ -20859,7 +20919,6 @@ async def test_key_team_ownership_mutation_allows_legacy_mismatch_without_change
data=UpdateKeyRequest(key="sk-key", **request_fields),
existing_key_row=existing_key_row,
prisma_client=prisma_client,
user_api_key_cache=UserApiKeyCache(),
)
@ -20882,7 +20941,6 @@ async def test_key_team_ownership_mutation_allows_detach_with_team_change() -> N
data=UpdateKeyRequest(key="sk-key", project_id=None, team_id="team-b"),
existing_key_row=existing_key_row,
prisma_client=prisma_client,
user_api_key_cache=UserApiKeyCache(),
)
@ -20898,6 +20956,10 @@ async def test_regenerate_checks_project_team_ownership(
) -> None:
existing_key: Final = LiteLLM_VerificationToken(token="abc123", team_id=key_team_id)
mock_prisma_client: Final = _make_regenerate_mock_prisma()
mock_prisma_client.db.litellm_projecttable = MagicMock()
mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock(
return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id="team-b")
)
user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b")
async def regenerate() -> None:
@ -20949,7 +21011,6 @@ async def test_key_project_team_validation_uses_project_missing_404() -> None:
project_id=_OWNED_PROJECT,
key_team_id="team-a",
prisma_client=prisma_client,
user_api_key_cache=UserApiKeyCache(),
)
assert error.value.status_code == 404