mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
1049ac4bff
commit
7e55e90847
2 changed files with 66 additions and 15 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue