diff --git a/litellm/llms/litellm_proxy/skills/handler.py b/litellm/llms/litellm_proxy/skills/handler.py index 7093874acd..6eb6f40fc6 100644 --- a/litellm/llms/litellm_proxy/skills/handler.py +++ b/litellm/llms/litellm_proxy/skills/handler.py @@ -10,6 +10,7 @@ import uuid from typing import Any, Dict, List, Optional from litellm._logging import verbose_logger +from litellm.llms.litellm_proxy.skills.store import LiteLLMSkillsStore from litellm.proxy._types import LiteLLM_SkillsTable, NewSkillRequest, UserAPIKeyAuth from litellm.proxy.common_utils.resource_ownership import ( get_primary_resource_owner_scope, @@ -104,6 +105,7 @@ class LiteLLMSkillsHandler: LiteLLM_SkillsTable record """ prisma_client = await LiteLLMSkillsHandler._get_prisma_client() + store = LiteLLMSkillsStore(prisma_client) skill_id = f"litellm_skill_{uuid.uuid4()}" owner = get_primary_resource_owner_scope(user_api_key_dict) or user_id @@ -138,7 +140,7 @@ class LiteLLMSkillsHandler: f"LiteLLMSkillsHandler: Creating skill {skill_id} with title={data.display_title}" ) - new_skill = await prisma_client.db.litellm_skillstable.create(data=skill_data) + new_skill = await store.create_skill(skill_data) return _prisma_skill_to_litellm(new_skill) @@ -159,6 +161,7 @@ class LiteLLMSkillsHandler: List of LiteLLM_SkillsTable records """ prisma_client = await LiteLLMSkillsHandler._get_prisma_client() + store = LiteLLMSkillsStore(prisma_client) verbose_logger.debug( f"LiteLLMSkillsHandler: Listing skills with limit={limit}, offset={offset}" @@ -183,9 +186,7 @@ class LiteLLMSkillsHandler: else: find_many_kwargs["where"] = {"created_by": {"in": owner_scopes}} - skills = await prisma_client.db.litellm_skillstable.find_many( - **find_many_kwargs - ) + skills = await store.list_skills(find_many_kwargs) return [_prisma_skill_to_litellm(s) for s in skills] @@ -207,12 +208,11 @@ class LiteLLMSkillsHandler: ValueError: If skill not found """ prisma_client = await LiteLLMSkillsHandler._get_prisma_client() + store = LiteLLMSkillsStore(prisma_client) verbose_logger.debug(f"LiteLLMSkillsHandler: Getting skill {skill_id}") - skill = await prisma_client.db.litellm_skillstable.find_unique( - where={"skill_id": skill_id} - ) + skill = await store.find_skill(skill_id) if skill is None: raise ValueError(f"Skill not found: {skill_id}") @@ -242,13 +242,12 @@ class LiteLLMSkillsHandler: ValueError: If skill not found """ prisma_client = await LiteLLMSkillsHandler._get_prisma_client() + store = LiteLLMSkillsStore(prisma_client) verbose_logger.debug(f"LiteLLMSkillsHandler: Deleting skill {skill_id}") # Check if skill exists - skill = await prisma_client.db.litellm_skillstable.find_unique( - where={"skill_id": skill_id} - ) + skill = await store.find_skill(skill_id) if skill is None: raise ValueError(f"Skill not found: {skill_id}") @@ -259,7 +258,7 @@ class LiteLLMSkillsHandler: raise ValueError(f"Skill not found: {skill_id}") # Delete the skill - await prisma_client.db.litellm_skillstable.delete(where={"skill_id": skill_id}) + await store.delete_skill(skill_id) return {"id": skill_id, "type": "skill_deleted"} diff --git a/litellm/llms/litellm_proxy/skills/store.py b/litellm/llms/litellm_proxy/skills/store.py new file mode 100644 index 0000000000..6e69a7a962 --- /dev/null +++ b/litellm/llms/litellm_proxy/skills/store.py @@ -0,0 +1,22 @@ +from typing import Any, Dict, List, Optional + + +class LiteLLMSkillsStore: + def __init__(self, prisma_client: Any): + self.prisma_client = prisma_client + + @property + def _table(self) -> Any: + return self.prisma_client.db.litellm_skillstable + + async def create_skill(self, data: Dict[str, Any]) -> Any: + return await self._table.create(data=data) + + async def list_skills(self, find_many_kwargs: Dict[str, Any]) -> List[Any]: + return await self._table.find_many(**find_many_kwargs) + + async def find_skill(self, skill_id: str) -> Optional[Any]: + return await self._table.find_unique(where={"skill_id": skill_id}) + + async def delete_skill(self, skill_id: str) -> None: + await self._table.delete(where={"skill_id": skill_id}) diff --git a/litellm/proxy/container_endpoints/endpoints.py b/litellm/proxy/container_endpoints/endpoints.py index b67e1e28d8..089f22a23a 100644 --- a/litellm/proxy/container_endpoints/endpoints.py +++ b/litellm/proxy/container_endpoints/endpoints.py @@ -122,11 +122,6 @@ async def create_container( user_api_base=user_api_base, version=version, ) - return await record_container_owner( - response=response, - user_api_key_dict=user_api_key_dict, - custom_llm_provider=custom_llm_provider, - ) except Exception as e: raise await processor._handle_llm_api_exception( e=e, @@ -134,6 +129,11 @@ async def create_container( proxy_logging_obj=proxy_logging_obj, version=version, ) + return await record_container_owner( + response=response, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) @router.get( diff --git a/litellm/proxy/container_endpoints/ownership.py b/litellm/proxy/container_endpoints/ownership.py index 05fe7e6484..86188c0241 100644 --- a/litellm/proxy/container_endpoints/ownership.py +++ b/litellm/proxy/container_endpoints/ownership.py @@ -1,4 +1,5 @@ import os +from collections import OrderedDict from typing import Any, Dict, List, Optional, Set, Tuple from fastapi import HTTPException @@ -11,11 +12,15 @@ from litellm.proxy.common_utils.resource_ownership import ( is_proxy_admin, user_can_access_resource_owner, ) +from litellm.proxy.container_endpoints.ownership_store import ( + CONTAINER_OBJECT_PURPOSE, + ContainerOwnershipStore, +) from litellm.responses.utils import ResponsesAPIRequestUtils -CONTAINER_OBJECT_PURPOSE = "container" ALLOW_UNTRACKED_CONTAINER_ACCESS_ENV = "LITELLM_ALLOW_UNTRACKED_CONTAINER_ACCESS" -_IN_MEMORY_CONTAINER_OWNERS: Dict[str, str] = {} +MAX_IN_MEMORY_CONTAINER_OWNERS = 10000 +_IN_MEMORY_CONTAINER_OWNERS: "OrderedDict[str, str]" = OrderedDict() def _allow_untracked_container_access() -> bool: @@ -26,6 +31,15 @@ def _allow_untracked_container_access() -> bool: } +def _remember_container_owner(model_object_id: str, owner: str) -> None: + existing_owner = _IN_MEMORY_CONTAINER_OWNERS.get(model_object_id) + if existing_owner is not None: + _IN_MEMORY_CONTAINER_OWNERS.move_to_end(model_object_id) + _IN_MEMORY_CONTAINER_OWNERS[model_object_id] = owner + while len(_IN_MEMORY_CONTAINER_OWNERS) > MAX_IN_MEMORY_CONTAINER_OWNERS: + _IN_MEMORY_CONTAINER_OWNERS.popitem(last=False) + + def _container_model_object_id( original_container_id: str, custom_llm_provider: str, @@ -124,12 +138,11 @@ async def record_container_owner( existing_owner, user_api_key_dict ): raise HTTPException(status_code=403, detail="Forbidden") - _IN_MEMORY_CONTAINER_OWNERS[model_object_id] = owner + _remember_container_owner(model_object_id, owner) return response - existing = await prisma_client.db.litellm_managedobjecttable.find_unique( - where={"model_object_id": model_object_id} - ) + store = ContainerOwnershipStore(prisma_client) + existing = await store.find_by_model_object_id(model_object_id) if existing is not None: if getattr(existing, "file_purpose", None) != CONTAINER_OBJECT_PURPOSE: raise HTTPException(status_code=500, detail="Unable to track container") @@ -137,8 +150,8 @@ async def record_container_owner( getattr(existing, "created_by", None), user_api_key_dict ): raise HTTPException(status_code=403, detail="Forbidden") - await prisma_client.db.litellm_managedobjecttable.update( - where={"model_object_id": model_object_id}, + await store.update_owner_record( + model_object_id=model_object_id, data={ "unified_object_id": container_id, "file_object": file_object, @@ -146,7 +159,7 @@ async def record_container_owner( }, ) else: - await prisma_client.db.litellm_managedobjecttable.create( + await store.create_owner_record( data={ "unified_object_id": container_id, "model_object_id": model_object_id, @@ -165,7 +178,12 @@ async def record_container_owner( model_object_id, e, ) - _IN_MEMORY_CONTAINER_OWNERS[model_object_id] = owner + existing_owner = _IN_MEMORY_CONTAINER_OWNERS.get(model_object_id) + if existing_owner is not None and not user_can_access_resource_owner( + existing_owner, user_api_key_dict + ): + raise HTTPException(status_code=403, detail="Forbidden") + _remember_container_owner(model_object_id, owner) return response @@ -183,14 +201,9 @@ async def _get_container_owner( if prisma_client is None: return _IN_MEMORY_CONTAINER_OWNERS.get(model_object_id) - row = await prisma_client.db.litellm_managedobjecttable.find_first( - where={ - "model_object_id": model_object_id, - "file_purpose": CONTAINER_OBJECT_PURPOSE, - } - ) - if row is not None: - return getattr(row, "created_by", None) + owner = await ContainerOwnershipStore(prisma_client).get_owner(model_object_id) + if owner is not None: + return owner return _IN_MEMORY_CONTAINER_OWNERS.get(model_object_id) except Exception as e: verbose_proxy_logger.warning( @@ -283,17 +296,12 @@ async def _get_allowed_container_ids( if prisma_client is None: return in_memory_allowed_ids - rows = await prisma_client.db.litellm_managedobjecttable.find_many( - where={ - "file_purpose": CONTAINER_OBJECT_PURPOSE, - "created_by": {"in": owner_scopes}, - } + db_allowed_ids = await ContainerOwnershipStore( + prisma_client + ).list_model_object_ids_for_owners( + owner_scopes=owner_scopes, ) - return in_memory_allowed_ids | { - row.model_object_id - for row in rows - if getattr(row, "model_object_id", None) is not None - } + return in_memory_allowed_ids | db_allowed_ids except Exception as e: verbose_proxy_logger.warning( "Failed to load allowed container ids; falling back to in-process " diff --git a/litellm/proxy/container_endpoints/ownership_store.py b/litellm/proxy/container_endpoints/ownership_store.py new file mode 100644 index 0000000000..b3c406f618 --- /dev/null +++ b/litellm/proxy/container_endpoints/ownership_store.py @@ -0,0 +1,55 @@ +from typing import Any, Dict, List, Optional, Set + +CONTAINER_OBJECT_PURPOSE = "container" + + +class ContainerOwnershipStore: + def __init__(self, prisma_client: Any): + self.prisma_client = prisma_client + + @property + def _table(self) -> Any: + return self.prisma_client.db.litellm_managedobjecttable + + async def find_by_model_object_id(self, model_object_id: str) -> Optional[Any]: + return await self._table.find_unique(where={"model_object_id": model_object_id}) + + async def create_owner_record(self, data: Dict[str, Any]) -> None: + await self._table.create(data=data) + + async def update_owner_record( + self, + model_object_id: str, + data: Dict[str, Any], + ) -> None: + await self._table.update( + where={"model_object_id": model_object_id}, + data=data, + ) + + async def get_owner(self, model_object_id: str) -> Optional[str]: + row = await self._table.find_first( + where={ + "model_object_id": model_object_id, + "file_purpose": CONTAINER_OBJECT_PURPOSE, + } + ) + if row is None: + return None + return getattr(row, "created_by", None) + + async def list_model_object_ids_for_owners( + self, + owner_scopes: List[str], + ) -> Set[str]: + rows = await self._table.find_many( + where={ + "file_purpose": CONTAINER_OBJECT_PURPOSE, + "created_by": {"in": owner_scopes}, + } + ) + return { + row.model_object_id + for row in rows + if getattr(row, "model_object_id", None) is not None + } diff --git a/tests/test_litellm/containers/test_container_proxy_ownership.py b/tests/test_litellm/containers/test_container_proxy_ownership.py index fbce5d6985..19682649bb 100644 --- a/tests/test_litellm/containers/test_container_proxy_ownership.py +++ b/tests/test_litellm/containers/test_container_proxy_ownership.py @@ -241,6 +241,29 @@ async def test_should_track_container_owner_in_memory_without_prisma(monkeypatch assert provider == "openai" +@pytest.mark.asyncio +async def test_should_bound_in_memory_container_owner_tracking(monkeypatch): + monkeypatch.setattr(ownership, "MAX_IN_MEMORY_CONTAINER_OWNERS", 2) + monkeypatch.setattr( + ownership, + "_get_prisma_client", + AsyncMock(return_value=None), + ) + auth = UserAPIKeyAuth(user_id="user-1") + + for container_id in ("cntr_1", "cntr_2", "cntr_3"): + await ownership.record_container_owner( + response=_container(container_id), + user_api_key_dict=auth, + custom_llm_provider="openai", + ) + + assert list(ownership._IN_MEMORY_CONTAINER_OWNERS.keys()) == [ + "container:openai:cntr_2", + "container:openai:cntr_3", + ] + + @pytest.mark.asyncio async def test_should_deny_container_access_for_different_owner(monkeypatch): table = AsyncMock() @@ -847,6 +870,59 @@ async def test_should_record_container_owner_inside_create_endpoint(monkeypatch) ) +@pytest.mark.asyncio +async def test_should_not_route_owner_record_errors_through_llm_error_handler( + monkeypatch, +): + from litellm.proxy.container_endpoints import endpoints + + proxy_server_stub = SimpleNamespace( + general_settings={}, + llm_router=None, + proxy_config=None, + proxy_logging_obj=None, + select_data_generator=None, + user_api_base=None, + user_max_tokens=None, + user_model=None, + user_request_timeout=None, + user_temperature=None, + version="test", + ) + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_stub) + + class FakeProcessor: + def __init__(self, data): + pass + + async def base_process_llm_request(self, **kwargs): + return _container("cntr_provider") + + async def _handle_llm_api_exception(self, **kwargs): + raise AssertionError("ownership errors should not use LLM error handler") + + monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", FakeProcessor) + monkeypatch.setattr( + endpoints, + "record_container_owner", + AsyncMock(side_effect=HTTPException(status_code=403, detail="Forbidden")), + ) + + with pytest.raises(HTTPException) as exc: + await endpoints.create_container( + request=SimpleNamespace( + query_params={}, + headers={}, + json=AsyncMock(return_value={}), + body=AsyncMock(return_value=b"{}"), + ), + fastapi_response=SimpleNamespace(), + user_api_key_dict=UserAPIKeyAuth(user_id="user-1"), + ) + + assert exc.value.status_code == 403 + + @pytest.mark.asyncio async def test_should_filter_container_list_inside_list_endpoint(monkeypatch): from litellm.proxy.container_endpoints import endpoints