diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index a6858f0a166..b9ff051da74 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1157,6 +1157,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): tags: Optional[list] = None model_aliases: Optional[dict] = None guardrails: Optional[List[str]] = None + object_permission: Optional[LiteLLM_ObjectPermissionBase] = None class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase): @@ -1254,6 +1255,7 @@ class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase): object_permission_id: str mcp_servers: List[str] + vector_stores: List[str] class LiteLLM_TeamTable(TeamBase): @@ -1636,6 +1638,7 @@ class LiteLLM_VerificationToken(LiteLLMPydanticObjectBase): created_by: Optional[str] = None updated_at: Optional[datetime] = None updated_by: Optional[str] = None + object_permission_id: Optional[str] = None object_permission: Optional[LiteLLM_ObjectPermissionTable] = None model_config = ConfigDict(protected_namespaces=()) @@ -1757,6 +1760,7 @@ class LiteLLM_OrganizationTableUpdate(LiteLLMPydanticObjectBase): metadata: Optional[dict] = None models: Optional[List[str]] = None updated_by: Optional[str] = None + object_permission: Optional[LiteLLM_ObjectPermissionTable] = None class LiteLLM_UserTable(LiteLLMPydanticObjectBase): diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 71fead9a925..2fee9edf1c2 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -673,8 +673,9 @@ async def _set_object_permission( return data_json -def prepare_key_update_data( - data: Union[UpdateKeyRequest, RegenerateKeyRequest], existing_key_row +async def prepare_key_update_data( + data: Union[UpdateKeyRequest, RegenerateKeyRequest], + existing_key_row: LiteLLM_VerificationToken, ): data_json: dict = data.model_dump(exclude_unset=True) data_json.pop("key", None) @@ -704,6 +705,12 @@ def prepare_key_update_data( non_default_values["budget_reset_at"] = key_reset_at non_default_values["budget_duration"] = budget_duration + if "object_permission" in non_default_values: + non_default_values = await _handle_update_object_permission( + data_json=non_default_values, + existing_key_row=existing_key_row, + ) + _metadata = existing_key_row.metadata or {} # validate model_max_budget @@ -717,6 +724,60 @@ def prepare_key_update_data( return non_default_values +async def _handle_update_object_permission( + data_json: dict, + existing_key_row: LiteLLM_VerificationToken, +) -> dict: + """ + Handle the update of object permission. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise ValueError("Prisma client not found") + + ######################################################### + # Ensure `object_permission` is not added to the data_json + # We need to update the entity at the object_permission_id level in the LiteLLM_ObjectPermissionTable + ######################################################### + new_object_permission = data_json.pop("object_permission") + if new_object_permission is None: + return data_json + + # lookup existing object permission ID and update that entry + existing_object_permission_id = existing_key_row.object_permission_id + existing_object_permission = ( + await prisma_client.db.litellm_objectpermissiontable.find_unique( + where={"object_permission_id": existing_object_permission_id}, + ) + ) + existing_object_permissions_dict: dict = {} + if existing_object_permission is not None: + # update the object permission + existing_object_permissions_dict = existing_object_permission.model_dump( + exclude_unset=True, exclude_none=True + ) + existing_object_permissions_dict.update(dict(new_object_permission)) + + ######################################################### + # Commit the update to the LiteLLM_ObjectPermissionTable + ######################################################### + new_object_permission_row = ( + await prisma_client.db.litellm_objectpermissiontable.upsert( + where={"object_permission_id": existing_object_permission_id}, + data={ + "create": existing_object_permissions_dict, + "update": existing_object_permissions_dict, + }, + ) + ) + verbose_proxy_logger.debug( + f"new_object_permission_row: {new_object_permission_row}" + ) + data_json["object_permission_id"] = new_object_permission_row.object_permission_id + return data_json + + def is_different_team( data: UpdateKeyRequest, existing_key_row: LiteLLM_VerificationToken ) -> bool: @@ -853,7 +914,7 @@ async def update_key_fn( change_initiated_by=user_api_key_dict, llm_router=llm_router, ) - non_default_values = prepare_key_update_data( + non_default_values = await prepare_key_update_data( data=data, existing_key_row=existing_key_row ) @@ -1945,7 +2006,7 @@ async def regenerate_key_fn( non_default_values = {} if data is not None: # Update with any provided parameters from GenerateKeyRequest - non_default_values = prepare_key_update_data( + non_default_values = await prepare_key_update_data( data=data, existing_key_row=_key_in_db ) verbose_proxy_logger.debug("non_default_values: %s", non_default_values) @@ -2372,6 +2433,7 @@ async def _list_key_helper( {"created_at": "desc"}, {"token": "desc"}, # fallback sort ], + include={"object_permission": True}, ) verbose_proxy_logger.debug(f"Fetched {len(keys)} keys") diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index f06620ebb53..e057ff32a14 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -273,6 +273,22 @@ async def update_organization( updated_organization_row = prisma_client.jsonify_object( data.model_dump(exclude_none=True) ) + existing_organization_row = ( + await prisma_client.db.litellm_organizationtable.find_unique( + where={"organization_id": data.organization_id}, + ) + ) + + if existing_organization_row is None: + raise ValueError( + f"Organization not found for organization_id={data.organization_id}" + ) + + if data.object_permission is not None: + updated_organization_row = await handle_update_object_permission( + data_json=updated_organization_row, + existing_organization_row=existing_organization_row, + ) response = await prisma_client.db.litellm_organizationtable.update( where={"organization_id": data.organization_id}, @@ -283,6 +299,74 @@ async def update_organization( return response +async def handle_update_object_permission( + data_json: dict, + existing_organization_row: LiteLLM_OrganizationTable, +) -> dict: + """ + Handle the update of object permission for an organization. + + - Upserts the new object permission into the LiteLLM_ObjectPermissionTable + - Adds object_permission_id to data_json (this gets added in the DB) + - Pops the object_permission from data_json + - + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise ValueError("Prisma client not found") + + ######################################################### + # Ensure `object_permission` is not added to the data_json + # We need to update the entity at the object_permission_id level in the LiteLLM_ObjectPermissionTable + ######################################################### + new_object_permission: Union[dict, str] = data_json.pop("object_permission") or {} + if new_object_permission is None: + return data_json + + # lookup existing object permission ID and update that entry + existing_object_permission_id = existing_organization_row.object_permission_id + existing_object_permissions_dict = {} + + existing_object_permission = ( + await prisma_client.db.litellm_objectpermissiontable.find_unique( + where={"object_permission_id": existing_object_permission_id}, + ) + ) + + # update the object permission + if existing_object_permission is not None: + existing_object_permissions_dict = existing_object_permission.model_dump( + exclude_unset=True, exclude_none=True + ) + + if isinstance(new_object_permission, str): + new_object_permission = json.loads(new_object_permission) + + if isinstance(new_object_permission, dict): + existing_object_permissions_dict.update(new_object_permission) + + ######################################################### + # Commit the update to the LiteLLM_ObjectPermissionTable + ######################################################### + created_object_permission_row = ( + await prisma_client.db.litellm_objectpermissiontable.upsert( + where={"object_permission_id": existing_object_permission_id}, + data={ + "create": existing_object_permissions_dict, + "update": existing_object_permissions_dict, + }, + ) + ) + data_json[ + "object_permission_id" + ] = created_object_permission_row.object_permission_id + verbose_proxy_logger.debug( + f"created_object_permission_row: {created_object_permission_row}" + ) + return data_json + + @router.delete( "/organization/delete", tags=["organization management"], @@ -414,7 +498,12 @@ async def info_organization(organization_id: str): LiteLLM_OrganizationTableWithMembers ] = await prisma_client.db.litellm_organizationtable.find_unique( where={"organization_id": organization_id}, - include={"litellm_budget_table": True, "members": True, "teams": True}, + include={ + "litellm_budget_table": True, + "members": True, + "teams": True, + "object_permission": True, + }, ) if response is None: diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 2d7f52386dc..1ced6c32185 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -663,6 +663,13 @@ async def update_team( # set the budget_reset_at in DB updated_kv["budget_reset_at"] = reset_at + # Check object permission + if data.object_permission is not None: + updated_kv = await handle_update_object_permission( + data_json=updated_kv, + existing_team_row=existing_team_row, + ) + # update team metadata fields _team_metadata_fields = LiteLLM_ManagementEndpoint_MetadataFields_Premium for field in _team_metadata_fields: @@ -734,6 +741,62 @@ async def update_team( return {"team_id": team_row.team_id, "data": team_row} +async def handle_update_object_permission( + data_json: dict, existing_team_row: LiteLLM_TeamTable +) -> dict: + """ + Handle the update of object permission for a team. + + - IF there's no object_permission_id, then create a new entry in LiteLLM_ObjectPermissionTable + - IF there's an object_permission_id, then update the entry in LiteLLM_ObjectPermissionTable + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise ValueError("Prisma client not found") + + ######################################################### + # Ensure `object_permission` is not added to the data_json + # We need to update the entity at the object_permission_id level in the LiteLLM_ObjectPermissionTable + ######################################################### + new_object_permission = data_json.pop("object_permission") + if new_object_permission is None: + return data_json + + # lookup existing object permission ID and update that entry + existing_object_permission_id = existing_team_row.object_permission_id + existing_object_permission = ( + await prisma_client.db.litellm_objectpermissiontable.find_unique( + where={"object_permission_id": existing_object_permission_id}, + ) + ) + + existing_object_permissions_dict: Dict = {} + + # update the object permission + if existing_object_permission is not None: + existing_object_permissions_dict = existing_object_permission.model_dump( + exclude_unset=True, exclude_none=True + ) + existing_object_permissions_dict.update(dict(new_object_permission)) + created_object_permission_row = ( + await prisma_client.db.litellm_objectpermissiontable.upsert( + where={"object_permission_id": existing_object_permission_id}, + data={ + "create": existing_object_permissions_dict, + "update": existing_object_permissions_dict, + }, + ) + ) + data_json[ + "object_permission_id" + ] = created_object_permission_row.object_permission_id + verbose_proxy_logger.debug( + f"created_object_permission_row: {created_object_permission_row}" + ) + return data_json + + def _check_team_member_admin_add( member: Union[Member, List[Member]], premium_user: bool, @@ -1486,7 +1549,8 @@ async def team_info( team_info: Optional[ BaseModel ] = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id} + where={"team_id": team_id}, + include={"object_permission": True}, ) if team_info is None: raise Exception diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index e02d9acf3c2..91b64c09c02 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -639,7 +639,8 @@ async def test_proxy_config_update_from_db(): } -def test_prepare_key_update_data(): +@pytest.mark.asyncio +async def test_prepare_key_update_data(): from litellm.proxy.management_endpoints.key_management_endpoints import ( prepare_key_update_data, ) @@ -647,15 +648,15 @@ def test_prepare_key_update_data(): existing_key_row = MagicMock() data = UpdateKeyRequest(key="test_key", models=["gpt-4"], duration="120s") - updated_data = prepare_key_update_data(data, existing_key_row) + updated_data = await prepare_key_update_data(data, existing_key_row) assert "expires" in updated_data data = UpdateKeyRequest(key="test_key", metadata={}) - updated_data = prepare_key_update_data(data, existing_key_row) + updated_data = await prepare_key_update_data(data, existing_key_row) assert updated_data["metadata"] == {} data = UpdateKeyRequest(key="test_key", metadata=None) - updated_data = prepare_key_update_data(data, existing_key_row) + updated_data = await prepare_key_update_data(data, existing_key_row) assert updated_data["metadata"] is None diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 755b833486a..f0a2f7cab70 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -251,3 +251,217 @@ async def test_key_generation_with_object_permission(monkeypatch): ] assert len(key_insert_calls) == 1 assert key_insert_calls[0]["data"].get("object_permission_id") == "objperm123" + + +@pytest.mark.asyncio +async def test_key_update_object_permissions_existing_permission(monkeypatch): + """ + Test updating object permissions when a key already has an existing object_permission_id. + + This test verifies that when updating vector stores for a key that already has an + object_permission_id, the existing LiteLLM_ObjectPermissionTable record is updated + with the new permissions and the object_permission_id remains the same. + """ + from unittest.mock import AsyncMock, MagicMock + + import pytest + + from litellm.proxy._types import ( + LiteLLM_ObjectPermissionBase, + LiteLLM_VerificationToken, + ) + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _handle_update_object_permission, + ) + + # Mock prisma client + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + # Mock existing key with object_permission_id + existing_key_row = LiteLLM_VerificationToken( + token="test_token_hash", + object_permission_id="existing_perm_id_123", + user_id="user123", + team_id=None, + ) + + # Mock existing object permission record + existing_object_permission = MagicMock() + existing_object_permission.model_dump.return_value = { + "object_permission_id": "existing_perm_id_123", + "vector_stores": ["old_store_1", "old_store_2"], + } + + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=existing_object_permission + ) + + # Mock upsert operation + updated_permission = MagicMock() + updated_permission.object_permission_id = "existing_perm_id_123" + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( + return_value=updated_permission + ) + + # Test data with new object permission + data_json = { + "object_permission": LiteLLM_ObjectPermissionBase( + vector_stores=["new_store_1", "new_store_2", "new_store_3"] + ).model_dump(exclude_unset=True, exclude_none=True), + "user_id": "user123", + } + + # Call the function + result = await _handle_update_object_permission( + data_json=data_json, + existing_key_row=existing_key_row, + ) + + # Verify the object_permission was removed from data_json and object_permission_id was set + assert "object_permission" not in result + assert result["object_permission_id"] == "existing_perm_id_123" + + # Verify database operations were called correctly + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( + where={"object_permission_id": "existing_perm_id_123"} + ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() + + +@pytest.mark.asyncio +async def test_key_update_object_permissions_no_existing_permission(monkeypatch): + """ + Test creating object permissions when a key has no existing object_permission_id. + + This test verifies that when updating object permissions for a key that has + object_permission_id set to None, a new entry is created in the + LiteLLM_ObjectPermissionTable and the key is updated with the new object_permission_id. + """ + from unittest.mock import AsyncMock, MagicMock + + import pytest + + from litellm.proxy._types import ( + LiteLLM_ObjectPermissionBase, + LiteLLM_VerificationToken, + ) + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _handle_update_object_permission, + ) + + # Mock prisma client + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + existing_key_row_no_perm = LiteLLM_VerificationToken( + token="test_token_hash_2", + object_permission_id=None, + user_id="user456", + team_id=None, + ) + + # Mock find_unique to return None (no existing permission) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=None + ) + + # Mock upsert to create new record + new_permission = MagicMock() + new_permission.object_permission_id = "new_perm_id_456" + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( + return_value=new_permission + ) + + data_json = { + "object_permission": LiteLLM_ObjectPermissionBase( + vector_stores=["brand_new_store"] + ).model_dump(exclude_unset=True, exclude_none=True), + "user_id": "user456", + } + + result = await _handle_update_object_permission( + data_json=data_json, + existing_key_row=existing_key_row_no_perm, + ) + + # Verify new object_permission_id was set + assert "object_permission" not in result + assert result["object_permission_id"] == "new_perm_id_456" + + # Verify find_unique was called with None + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( + where={"object_permission_id": None} + ) + + # Verify upsert was called to create new record + mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() + + +@pytest.mark.asyncio +async def test_key_update_object_permissions_missing_permission_record(monkeypatch): + """ + Test creating object permissions when existing object_permission_id record is not found. + + This test verifies that when updating object permissions for a key that has an + object_permission_id but the corresponding record cannot be found in the database, + a new entry is created in the LiteLLM_ObjectPermissionTable with the new permissions. + """ + from unittest.mock import AsyncMock, MagicMock + + import pytest + + from litellm.proxy._types import ( + LiteLLM_ObjectPermissionBase, + LiteLLM_VerificationToken, + ) + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _handle_update_object_permission, + ) + + # Mock prisma client + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + existing_key_row_missing_perm = LiteLLM_VerificationToken( + token="test_token_hash_3", + object_permission_id="missing_perm_id_789", + user_id="user789", + team_id=None, + ) + + # Mock find_unique to return None (permission record not found) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=None + ) + + # Mock upsert to create new record + new_permission = MagicMock() + new_permission.object_permission_id = "recreated_perm_id_789" + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( + return_value=new_permission + ) + + data_json = { + "object_permission": LiteLLM_ObjectPermissionBase( + vector_stores=["recreated_store"] + ).model_dump(exclude_unset=True, exclude_none=True), + "user_id": "user789", + } + + result = await _handle_update_object_permission( + data_json=data_json, + existing_key_row=existing_key_row_missing_perm, + ) + + # Verify new object_permission_id was set + assert "object_permission" not in result + assert result["object_permission_id"] == "recreated_perm_id_789" + + # Verify find_unique was called with the missing permission ID + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( + where={"object_permission_id": "missing_perm_id_789"} + ) + + # Verify upsert was called to create new record + mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() diff --git a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py new file mode 100644 index 00000000000..a602a6e338a --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py @@ -0,0 +1,242 @@ +import asyncio +import json +import os +import sys +import uuid +from typing import Optional, cast +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import HTTPException +from fastapi.testclient import TestClient + +sys.path.insert( + 0, os.path.abspath("../../../") +) # Adds the parent directory to the system path + + +@pytest.mark.asyncio +async def test_organization_update_object_permissions_existing_permission(monkeypatch): + """ + Test updating object permissions when an organization already has an existing object_permission_id. + + This test verifies that when updating vector stores for an organization that already has an + object_permission_id, the existing LiteLLM_ObjectPermissionTable record is updated + with the new permissions and the object_permission_id remains the same. + """ + from unittest.mock import AsyncMock, MagicMock + + import pytest + + from litellm.proxy._types import ( + LiteLLM_ObjectPermissionBase, + LiteLLM_OrganizationTable, + ) + from litellm.proxy.management_endpoints.organization_endpoints import ( + handle_update_object_permission, + ) + + # Mock prisma client + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + # Mock existing organization with object_permission_id + existing_organization_row = LiteLLM_OrganizationTable( + organization_id="test_org_id", + object_permission_id="existing_perm_id_123", + organization_alias="test_org", + budget_id="test_budget_id", + models=["test_model_1", "test_model_2"], + created_by="test_created_by", + updated_by="test_updated_by", + ) + + # Mock existing object permission record + existing_object_permission = MagicMock() + existing_object_permission.model_dump.return_value = { + "object_permission_id": "existing_perm_id_123", + "vector_stores": ["old_store_1", "old_store_2"], + } + + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=existing_object_permission + ) + + # Mock upsert operation + updated_permission = MagicMock() + updated_permission.object_permission_id = "existing_perm_id_123" + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( + return_value=updated_permission + ) + + # Test data with new object permission + data_json = { + "object_permission": LiteLLM_ObjectPermissionBase( + vector_stores=["new_store_1", "new_store_2", "new_store_3"] + ).model_dump(exclude_unset=True, exclude_none=True), + "organization_alias": "updated_org", + } + + # Call the function + result = await handle_update_object_permission( + data_json=data_json, + existing_organization_row=existing_organization_row, + ) + + # Verify the object_permission was removed from data_json and object_permission_id was set + assert "object_permission" not in result + assert result["object_permission_id"] == "existing_perm_id_123" + + # Verify database operations were called correctly + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( + where={"object_permission_id": "existing_perm_id_123"} + ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() + + +@pytest.mark.asyncio +async def test_organization_update_object_permissions_no_existing_permission( + monkeypatch, +): + """ + Test creating object permissions when an organization has no existing object_permission_id. + + This test verifies that when updating object permissions for an organization that has + object_permission_id set to None, a new entry is created in the + LiteLLM_ObjectPermissionTable and the organization is updated with the new object_permission_id. + """ + from unittest.mock import AsyncMock, MagicMock + + import pytest + + from litellm.proxy._types import ( + LiteLLM_ObjectPermissionBase, + LiteLLM_OrganizationTable, + ) + from litellm.proxy.management_endpoints.organization_endpoints import ( + handle_update_object_permission, + ) + + # Mock prisma client + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + existing_organization_row_no_perm = LiteLLM_OrganizationTable( + organization_id="test_org_id_2", + object_permission_id=None, + organization_alias="test_org_2", + budget_id="test_budget_id_2", + models=["test_model_1", "test_model_2"], + created_by="test_created_by_2", + updated_by="test_updated_by_2", + ) + + # Mock find_unique to return None (no existing permission) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=None + ) + + # Mock upsert to create new record + new_permission = MagicMock() + new_permission.object_permission_id = "new_perm_id_456" + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( + return_value=new_permission + ) + + data_json = { + "object_permission": LiteLLM_ObjectPermissionBase( + vector_stores=["brand_new_store"] + ).model_dump(exclude_unset=True, exclude_none=True), + "organization_alias": "updated_org_2", + } + + result = await handle_update_object_permission( + data_json=data_json, + existing_organization_row=existing_organization_row_no_perm, + ) + + # Verify new object_permission_id was set + assert "object_permission" not in result + assert result["object_permission_id"] == "new_perm_id_456" + + # Verify find_unique was called with None + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( + where={"object_permission_id": None} + ) + + # Verify upsert was called to create new record + mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() + + +@pytest.mark.asyncio +async def test_organization_update_object_permissions_missing_permission_record( + monkeypatch, +): + """ + Test creating object permissions when existing object_permission_id record is not found. + + This test verifies that when updating object permissions for an organization that has an + object_permission_id but the corresponding record cannot be found in the database, + a new entry is created in the LiteLLM_ObjectPermissionTable with the new permissions. + """ + from unittest.mock import AsyncMock, MagicMock + + import pytest + + from litellm.proxy._types import ( + LiteLLM_ObjectPermissionBase, + LiteLLM_OrganizationTable, + ) + from litellm.proxy.management_endpoints.organization_endpoints import ( + handle_update_object_permission, + ) + + # Mock prisma client + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + existing_organization_row_missing_perm = LiteLLM_OrganizationTable( + organization_id="test_org_id_3", + object_permission_id="missing_perm_id_789", + organization_alias="test_org_3", + budget_id="test_budget_id_3", + models=["test_model_1", "test_model_2"], + created_by="test_created_by_3", + updated_by="test_updated_by_3", + ) + + # Mock find_unique to return None (permission record not found) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=None + ) + + # Mock upsert to create new record + new_permission = MagicMock() + new_permission.object_permission_id = "recreated_perm_id_789" + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( + return_value=new_permission + ) + + data_json = { + "object_permission": LiteLLM_ObjectPermissionBase( + vector_stores=["recreated_store"] + ).model_dump(exclude_unset=True, exclude_none=True), + "organization_alias": "updated_org_3", + } + + result = await handle_update_object_permission( + data_json=data_json, + existing_organization_row=existing_organization_row_missing_perm, + ) + + # Verify new object_permission_id was set + assert "object_permission" not in result + assert result["object_permission_id"] == "recreated_perm_id_789" + + # Verify find_unique was called with the missing permission ID + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( + where={"object_permission_id": "missing_perm_id_789"} + ) + + # Verify upsert was called to create new record + mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 6e641d2e1c0..7b9acf42e12 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -313,3 +313,205 @@ async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth): assert mock_team_create.call_count == 1 created_team_kwargs = mock_team_create.call_args.kwargs assert created_team_kwargs["data"].get("object_permission_id") == "objperm123" + + +@pytest.mark.asyncio +async def test_team_update_object_permissions_existing_permission(monkeypatch): + """ + Test updating object permissions when a team already has an existing object_permission_id. + + This test verifies that when updating vector stores for a team that already has an + object_permission_id, the existing LiteLLM_ObjectPermissionTable record is updated + with the new permissions and the object_permission_id remains the same. + """ + from unittest.mock import AsyncMock, MagicMock + + import pytest + + from litellm.proxy._types import LiteLLM_ObjectPermissionBase, LiteLLM_TeamTable + from litellm.proxy.management_endpoints.team_endpoints import ( + handle_update_object_permission, + ) + + # Mock prisma client + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + # Mock existing team with object_permission_id + existing_team_row = LiteLLM_TeamTable( + team_id="test_team_id", + object_permission_id="existing_perm_id_123", + team_alias="test_team", + ) + + # Mock existing object permission record + existing_object_permission = MagicMock() + existing_object_permission.model_dump.return_value = { + "object_permission_id": "existing_perm_id_123", + "vector_stores": ["old_store_1", "old_store_2"], + } + + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=existing_object_permission + ) + + # Mock upsert operation + updated_permission = MagicMock() + updated_permission.object_permission_id = "existing_perm_id_123" + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( + return_value=updated_permission + ) + + # Test data with new object permission + data_json = { + "object_permission": LiteLLM_ObjectPermissionBase( + vector_stores=["new_store_1", "new_store_2", "new_store_3"] + ).model_dump(exclude_unset=True, exclude_none=True), + "team_alias": "updated_team", + } + + # Call the function + result = await handle_update_object_permission( + data_json=data_json, + existing_team_row=existing_team_row, + ) + + # Verify the object_permission was removed from data_json and object_permission_id was set + assert "object_permission" not in result + assert result["object_permission_id"] == "existing_perm_id_123" + + # Verify database operations were called correctly + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( + where={"object_permission_id": "existing_perm_id_123"} + ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() + + +@pytest.mark.asyncio +async def test_team_update_object_permissions_no_existing_permission(monkeypatch): + """ + Test creating object permissions when a team has no existing object_permission_id. + + This test verifies that when updating object permissions for a team that has + object_permission_id set to None, a new entry is created in the + LiteLLM_ObjectPermissionTable and the team is updated with the new object_permission_id. + """ + from unittest.mock import AsyncMock, MagicMock + + import pytest + + from litellm.proxy._types import LiteLLM_ObjectPermissionBase, LiteLLM_TeamTable + from litellm.proxy.management_endpoints.team_endpoints import ( + handle_update_object_permission, + ) + + # Mock prisma client + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + existing_team_row_no_perm = LiteLLM_TeamTable( + team_id="test_team_id_2", + object_permission_id=None, + team_alias="test_team_2", + ) + + # Mock find_unique to return None (no existing permission) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=None + ) + + # Mock upsert to create new record + new_permission = MagicMock() + new_permission.object_permission_id = "new_perm_id_456" + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( + return_value=new_permission + ) + + data_json = { + "object_permission": LiteLLM_ObjectPermissionBase( + vector_stores=["brand_new_store"] + ).model_dump(exclude_unset=True, exclude_none=True), + "team_alias": "updated_team_2", + } + + result = await handle_update_object_permission( + data_json=data_json, + existing_team_row=existing_team_row_no_perm, + ) + + # Verify new object_permission_id was set + assert "object_permission" not in result + assert result["object_permission_id"] == "new_perm_id_456" + + # Verify find_unique was called with None + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( + where={"object_permission_id": None} + ) + + # Verify upsert was called to create new record + mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() + + +@pytest.mark.asyncio +async def test_team_update_object_permissions_missing_permission_record(monkeypatch): + """ + Test creating object permissions when existing object_permission_id record is not found. + + This test verifies that when updating object permissions for a team that has an + object_permission_id but the corresponding record cannot be found in the database, + a new entry is created in the LiteLLM_ObjectPermissionTable with the new permissions. + """ + from unittest.mock import AsyncMock, MagicMock + + import pytest + + from litellm.proxy._types import LiteLLM_ObjectPermissionBase, LiteLLM_TeamTable + from litellm.proxy.management_endpoints.team_endpoints import ( + handle_update_object_permission, + ) + + # Mock prisma client + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + existing_team_row_missing_perm = LiteLLM_TeamTable( + team_id="test_team_id_3", + object_permission_id="missing_perm_id_789", + team_alias="test_team_3", + ) + + # Mock find_unique to return None (permission record not found) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=None + ) + + # Mock upsert to create new record + new_permission = MagicMock() + new_permission.object_permission_id = "recreated_perm_id_789" + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( + return_value=new_permission + ) + + data_json = { + "object_permission": LiteLLM_ObjectPermissionBase( + vector_stores=["recreated_store"] + ).model_dump(exclude_unset=True, exclude_none=True), + "team_alias": "updated_team_3", + } + + result = await handle_update_object_permission( + data_json=data_json, + existing_team_row=existing_team_row_missing_perm, + ) + + # Verify new object_permission_id was set + assert "object_permission" not in result + assert result["object_permission_id"] == "recreated_perm_id_789" + + # Verify find_unique was called with the missing permission ID + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( + where={"object_permission_id": "missing_perm_id_789"} + ) + + # Verify upsert was called to create new record + mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() diff --git a/ui/litellm-dashboard/src/components/key_edit_view.tsx b/ui/litellm-dashboard/src/components/key_edit_view.tsx index 2c7e3d30d1e..ed337cd5631 100644 --- a/ui/litellm-dashboard/src/components/key_edit_view.tsx +++ b/ui/litellm-dashboard/src/components/key_edit_view.tsx @@ -5,6 +5,8 @@ import { KeyResponse } from "./key_team_helpers/key_list"; import { fetchTeamModels } from "../components/create_key_button"; import { modelAvailableCall } from "./networking"; import NumericalInput from "./shared/numerical_input"; +import VectorStoreSelector from "./vector_store_management/VectorStoreSelector"; + interface KeyEditViewProps { keyData: KeyResponse; onCancel: () => void; @@ -92,7 +94,8 @@ export function KeyEditView({ ...keyData, budget_duration: getBudgetDuration(keyData.budget_duration), metadata: keyData.metadata ? JSON.stringify(keyData.metadata, null, 2) : "", - guardrails: keyData.metadata?.guardrails || [] + guardrails: keyData.metadata?.guardrails || [], + vector_stores: keyData.object_permission?.vector_stores || [] }; return ( @@ -165,6 +168,15 @@ export function KeyEditView({ /> +