mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(proxy): check project moves against the primary database team
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
4964d55e6b
commit
cd442a1999
3 changed files with 75 additions and 23 deletions
|
|
@ -30,6 +30,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
|
||||
|
|
@ -768,26 +769,36 @@ async def update_project(
|
|||
detail={"error": "Cannot reassign project to a team you are not an admin of"},
|
||||
)
|
||||
|
||||
if data.team_id is not None and data.team_id != existing_project.team_id:
|
||||
mismatched_key_count: Final = await bounded_db_lookup(
|
||||
prisma_client.writer_db.litellm_verificationtoken.count(
|
||||
where={
|
||||
"project_id": data.project_id,
|
||||
"OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}],
|
||||
}
|
||||
),
|
||||
name="project_key_ownership",
|
||||
if data.team_id is not None:
|
||||
current_project_record: Final = await bounded_db_lookup(
|
||||
prisma_client.writer_db.litellm_projecttable.find_unique(where={"project_id": data.project_id}),
|
||||
name="project",
|
||||
)
|
||||
if mismatched_key_count > 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
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."
|
||||
)
|
||||
},
|
||||
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 bounded_db_lookup(
|
||||
prisma_client.writer_db.litellm_verificationtoken.count(
|
||||
where={
|
||||
"project_id": data.project_id,
|
||||
"OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}],
|
||||
}
|
||||
),
|
||||
name="project_key_ownership",
|
||||
)
|
||||
if mismatched_key_count > 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
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."
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
# Validate project limits against team limits
|
||||
if target_team_obj is not None:
|
||||
|
|
|
|||
|
|
@ -1823,9 +1823,7 @@ async def _check_key_project_team(
|
|||
name="project",
|
||||
)
|
||||
project_obj: Final = (
|
||||
LiteLLM_ProjectTable.model_validate(record_to_dict(project_record))
|
||||
if project_record is not None
|
||||
else None
|
||||
LiteLLM_ProjectTable.model_validate(record_to_dict(project_record)) if project_record is not None else None
|
||||
)
|
||||
|
||||
if project_obj is None:
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ verbose_proxy_logger.setLevel(level=logging.DEBUG)
|
|||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ProjectTable,
|
||||
LiteLLM_TeamTable,
|
||||
NewProjectRequest,
|
||||
UpdateProjectRequest,
|
||||
|
|
@ -1236,7 +1237,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_prisma.db
|
||||
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)
|
||||
|
|
@ -1273,6 +1278,9 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists(
|
|||
)
|
||||
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")
|
||||
|
|
@ -1296,6 +1304,37 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists(
|
|||
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_allows_move_when_no_attached_keys_exist(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
project_id: Final = "project-without-keys"
|
||||
|
|
@ -1306,10 +1345,14 @@ async def test_update_project_allows_move_when_no_attached_keys_exist(monkeypatc
|
|||
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.db.litellm_verificationtoken.count.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_awaited_once()
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue