diff --git a/litellm/llms/litellm_proxy/skills/handler.py b/litellm/llms/litellm_proxy/skills/handler.py index 6eb6f40fc64..7a6dbb175fb 100644 --- a/litellm/llms/litellm_proxy/skills/handler.py +++ b/litellm/llms/litellm_proxy/skills/handler.py @@ -6,8 +6,9 @@ Used by the transformation layer and skills injection hook. """ import os +import time import uuid -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Optional, Tuple from litellm._logging import verbose_logger from litellm.llms.litellm_proxy.skills.store import LiteLLMSkillsStore @@ -21,6 +22,39 @@ from litellm.proxy.common_utils.resource_ownership import ( ALLOW_UNOWNED_SKILL_ACCESS_ENV = "LITELLM_ALLOW_UNOWNED_SKILL_ACCESS" +# Skills are looked up on every chat completion that has skills enabled +# (`LiteLLMSkillsHandler.fetch_skill_from_db` in the injection hook). Cache +# the Prisma skill row for a short window so the hot path doesn't issue a DB +# round-trip per request. Same shape as `_byok_cred_cache` and the container +# ownership cache: (value, monotonic_timestamp). `None` is cached as a true +# negative ("skill does not exist") so repeated misses also avoid the DB. +_SKILL_CACHE: Dict[str, Tuple[Optional[Any], float]] = {} +_SKILL_CACHE_TTL = 60 # seconds +_SKILL_CACHE_MAX_SIZE = 10000 + + +def _read_skill_cache(skill_id: str) -> Tuple[bool, Optional[Any]]: + """Return (hit, value). hit=False means caller must consult the DB.""" + entry = _SKILL_CACHE.get(skill_id) + if entry is None: + return False, None + value, timestamp = entry + if time.monotonic() - timestamp > _SKILL_CACHE_TTL: + _SKILL_CACHE.pop(skill_id, None) + return False, None + return True, value + + +def _write_skill_cache(skill_id: str, skill: Optional[Any]) -> None: + if len(_SKILL_CACHE) >= _SKILL_CACHE_MAX_SIZE: + _SKILL_CACHE.clear() + _SKILL_CACHE[skill_id] = (skill, time.monotonic()) + + +def _invalidate_skill_cache(skill_id: str) -> None: + """Drop the cache entry after a write so the next read sees the new row.""" + _SKILL_CACHE.pop(skill_id, None) + def _allow_unowned_skill_access() -> bool: return os.getenv(ALLOW_UNOWNED_SKILL_ACCESS_ENV, "").lower() in { @@ -190,6 +224,24 @@ class LiteLLMSkillsHandler: return [_prisma_skill_to_litellm(s) for s in skills] + @staticmethod + async def _load_skill(skill_id: str) -> Optional[Any]: + """Cache-first read of the Prisma skill row. + + Caching here keeps `fetch_skill_from_db` (called per chat completion in + the skills injection hook) off the DB. Owner-scope filtering happens + on the cached row, so the cache is per-skill and not per-caller. + """ + cached_hit, cached_skill = _read_skill_cache(skill_id) + if cached_hit: + return cached_skill + + prisma_client = await LiteLLMSkillsHandler._get_prisma_client() + store = LiteLLMSkillsStore(prisma_client) + skill = await store.find_skill(skill_id) + _write_skill_cache(skill_id, skill) + return skill + @staticmethod async def get_skill( skill_id: str, @@ -207,12 +259,9 @@ class LiteLLMSkillsHandler: Raises: 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 store.find_skill(skill_id) + skill = await LiteLLMSkillsHandler._load_skill(skill_id) if skill is None: raise ValueError(f"Skill not found: {skill_id}") @@ -247,7 +296,7 @@ class LiteLLMSkillsHandler: verbose_logger.debug(f"LiteLLMSkillsHandler: Deleting skill {skill_id}") # Check if skill exists - skill = await store.find_skill(skill_id) + skill = await LiteLLMSkillsHandler._load_skill(skill_id) if skill is None: raise ValueError(f"Skill not found: {skill_id}") @@ -259,6 +308,7 @@ class LiteLLMSkillsHandler: # Delete the skill await store.delete_skill(skill_id) + _invalidate_skill_cache(skill_id) return {"id": skill_id, "type": "skill_deleted"} diff --git a/litellm/proxy/container_endpoints/ownership.py b/litellm/proxy/container_endpoints/ownership.py index 86188c02414..e5577c4052e 100644 --- a/litellm/proxy/container_endpoints/ownership.py +++ b/litellm/proxy/container_endpoints/ownership.py @@ -1,4 +1,5 @@ import os +import time from collections import OrderedDict from typing import Any, Dict, List, Optional, Set, Tuple @@ -22,6 +23,38 @@ ALLOW_UNTRACKED_CONTAINER_ACCESS_ENV = "LITELLM_ALLOW_UNTRACKED_CONTAINER_ACCESS MAX_IN_MEMORY_CONTAINER_OWNERS = 10000 _IN_MEMORY_CONTAINER_OWNERS: "OrderedDict[str, str]" = OrderedDict() +# Short-lived cache keeps every container access check from hitting the DB +# (`_get_container_owner` is invoked on retrieve / delete / list / file-content +# paths). Mirrors the `_byok_cred_cache` pattern in mcp_server/server.py: +# (value, monotonic_timestamp) tuples, TTL'd, capped, invalidated by writes. +# A `None` value caches "untracked" so repeated negative lookups also avoid DB. +_CONTAINER_OWNER_CACHE: Dict[str, Tuple[Optional[str], float]] = {} +_CONTAINER_OWNER_CACHE_TTL = 60 # seconds +_CONTAINER_OWNER_CACHE_MAX_SIZE = 10000 + + +def _read_container_owner_cache(model_object_id: str) -> Tuple[bool, Optional[str]]: + """Return (hit, value). hit=False means caller must consult the DB.""" + entry = _CONTAINER_OWNER_CACHE.get(model_object_id) + if entry is None: + return False, None + value, timestamp = entry + if time.monotonic() - timestamp > _CONTAINER_OWNER_CACHE_TTL: + _CONTAINER_OWNER_CACHE.pop(model_object_id, None) + return False, None + return True, value + + +def _write_container_owner_cache(model_object_id: str, owner: Optional[str]) -> None: + if len(_CONTAINER_OWNER_CACHE) >= _CONTAINER_OWNER_CACHE_MAX_SIZE: + _CONTAINER_OWNER_CACHE.clear() + _CONTAINER_OWNER_CACHE[model_object_id] = (owner, time.monotonic()) + + +def _invalidate_container_owner_cache(model_object_id: str) -> None: + """Drop a cache entry after a write so the next read sees the new owner.""" + _CONTAINER_OWNER_CACHE.pop(model_object_id, None) + def _allow_untracked_container_access() -> bool: return os.getenv(ALLOW_UNTRACKED_CONTAINER_ACCESS_ENV, "").lower() in { @@ -139,6 +172,7 @@ async def record_container_owner( ): raise HTTPException(status_code=403, detail="Forbidden") _remember_container_owner(model_object_id, owner) + _invalidate_container_owner_cache(model_object_id) return response store = ContainerOwnershipStore(prisma_client) @@ -185,6 +219,7 @@ async def record_container_owner( raise HTTPException(status_code=403, detail="Forbidden") _remember_container_owner(model_object_id, owner) + _invalidate_container_owner_cache(model_object_id) return response @@ -196,15 +231,23 @@ async def _get_container_owner( original_container_id, custom_llm_provider, ) + + cached_hit, cached_value = _read_container_owner_cache(model_object_id) + if cached_hit: + return cached_value + try: prisma_client = await _get_prisma_client() if prisma_client is None: - return _IN_MEMORY_CONTAINER_OWNERS.get(model_object_id) + owner = _IN_MEMORY_CONTAINER_OWNERS.get(model_object_id) + _write_container_owner_cache(model_object_id, owner) + return owner 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) + if owner is None: + owner = _IN_MEMORY_CONTAINER_OWNERS.get(model_object_id) + _write_container_owner_cache(model_object_id, owner) + return owner except Exception as e: verbose_proxy_logger.warning( "Failed to load container ownership for container_id=%s; " @@ -212,6 +255,7 @@ async def _get_container_owner( model_object_id, e, ) + # Don't cache transient DB errors — let the next request retry. return _IN_MEMORY_CONTAINER_OWNERS.get(model_object_id) diff --git a/tests/test_litellm/containers/test_container_proxy_ownership.py b/tests/test_litellm/containers/test_container_proxy_ownership.py index 13216ad7c59..63e5e16d0d6 100644 --- a/tests/test_litellm/containers/test_container_proxy_ownership.py +++ b/tests/test_litellm/containers/test_container_proxy_ownership.py @@ -14,9 +14,11 @@ from litellm.types.containers.main import ContainerListResponse, ContainerObject @pytest.fixture(autouse=True) def clear_in_memory_container_owners(monkeypatch): ownership._IN_MEMORY_CONTAINER_OWNERS.clear() + ownership._CONTAINER_OWNER_CACHE.clear() monkeypatch.delenv(ownership.ALLOW_UNTRACKED_CONTAINER_ACCESS_ENV, raising=False) yield ownership._IN_MEMORY_CONTAINER_OWNERS.clear() + ownership._CONTAINER_OWNER_CACHE.clear() def _container(container_id: str) -> ContainerObject: @@ -1102,3 +1104,134 @@ async def test_should_forward_decoded_container_id_for_proxy_delete(monkeypatch) assert result["container_id"] == "cntr_provider" assert result["custom_llm_provider"] == "azure" assert result["model_id"] == "router-gpt" + + +# ── Cache layer ──────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_get_container_owner_uses_cache_after_first_db_hit(monkeypatch): + """Repeated access checks within the TTL window must not hit the DB. + + Greptile's P1 was that ownership reads issued a Prisma query on every + request. The cache here mirrors `_byok_cred_cache`: TTL'd, capped, and + invalidated on writes. + """ + table = AsyncMock() + fake_row = SimpleNamespace( + created_by="user-1", file_purpose=ownership.CONTAINER_OBJECT_PURPOSE + ) + table.find_first.return_value = fake_row + prisma_client = SimpleNamespace( + db=SimpleNamespace(litellm_managedobjecttable=table) + ) + monkeypatch.setattr( + ownership, + "_get_prisma_client", + AsyncMock(return_value=prisma_client), + ) + + owner_first = await ownership._get_container_owner("cntr_x", "openai") + owner_second = await ownership._get_container_owner("cntr_x", "openai") + owner_third = await ownership._get_container_owner("cntr_x", "openai") + + assert owner_first == "user-1" + assert owner_second == "user-1" + assert owner_third == "user-1" + # Single DB call across three reads — the cache absorbs the rest. + assert table.find_first.await_count == 1 + + +@pytest.mark.asyncio +async def test_get_container_owner_caches_negative_lookups(monkeypatch): + """`None` (untracked) must also be cached so repeated misses don't query.""" + table = AsyncMock() + table.find_first.return_value = None + prisma_client = SimpleNamespace( + db=SimpleNamespace(litellm_managedobjecttable=table) + ) + monkeypatch.setattr( + ownership, + "_get_prisma_client", + AsyncMock(return_value=prisma_client), + ) + + assert await ownership._get_container_owner("cntr_x", "openai") is None + assert await ownership._get_container_owner("cntr_x", "openai") is None + assert table.find_first.await_count == 1 + + +@pytest.mark.asyncio +async def test_record_container_owner_invalidates_cache(monkeypatch): + """A recorded owner must drop the cached value so the next read re-fetches. + + Otherwise a stale `None` from a prior negative lookup would survive the + create and the new owner would be invisible until the TTL elapses. + """ + # Seed the cache with a stale negative result. + ownership._write_container_owner_cache("container:openai:cntr_new", None) + cached_hit, cached_value = ownership._read_container_owner_cache( + "container:openai:cntr_new" + ) + assert cached_hit and cached_value is None + + table = AsyncMock() + table.find_unique.return_value = None + prisma_client = SimpleNamespace( + db=SimpleNamespace(litellm_managedobjecttable=table) + ) + monkeypatch.setattr( + ownership, + "_get_prisma_client", + AsyncMock(return_value=prisma_client), + ) + + await ownership.record_container_owner( + response=_container("cntr_new"), + user_api_key_dict=UserAPIKeyAuth(user_id="user-1"), + custom_llm_provider="openai", + ) + + # Invalidation drops the entry — next read goes to the DB. + cached_hit, _ = ownership._read_container_owner_cache("container:openai:cntr_new") + assert not cached_hit + + +@pytest.mark.asyncio +async def test_get_container_owner_does_not_cache_on_db_error(monkeypatch): + """DB errors must skip caching so transient failures don't pin a `None`.""" + table = AsyncMock() + table.find_first.side_effect = Exception("db unavailable") + prisma_client = SimpleNamespace( + db=SimpleNamespace(litellm_managedobjecttable=table) + ) + monkeypatch.setattr( + ownership, + "_get_prisma_client", + AsyncMock(return_value=prisma_client), + ) + + result = await ownership._get_container_owner("cntr_x", "openai") + assert result is None + cached_hit, _ = ownership._read_container_owner_cache("container:openai:cntr_x") + assert not cached_hit + + +def test_container_owner_cache_expires_after_ttl(monkeypatch): + """Entries past the TTL count as misses so writes elsewhere are eventually + visible to this process.""" + monkeypatch.setattr(ownership, "_CONTAINER_OWNER_CACHE_TTL", 0.0) + ownership._write_container_owner_cache("k", "user-1") + cached_hit, _ = ownership._read_container_owner_cache("k") + # TTL of 0 means anything in the cache is already stale. + assert not cached_hit + + +def test_container_owner_cache_evicts_when_at_capacity(monkeypatch): + """The cache must not grow unbounded; reaching capacity clears all entries.""" + monkeypatch.setattr(ownership, "_CONTAINER_OWNER_CACHE_MAX_SIZE", 2) + ownership._write_container_owner_cache("a", "user-a") + ownership._write_container_owner_cache("b", "user-b") + ownership._write_container_owner_cache("c", "user-c") + # Reaching the cap clears everything — the new write is the only survivor. + assert ownership._CONTAINER_OWNER_CACHE.keys() == {"c"} diff --git a/tests/test_litellm/llms/litellm_proxy/test_skills_ownership.py b/tests/test_litellm/llms/litellm_proxy/test_skills_ownership.py index 45ed172a845..f98a04205b4 100644 --- a/tests/test_litellm/llms/litellm_proxy/test_skills_ownership.py +++ b/tests/test_litellm/llms/litellm_proxy/test_skills_ownership.py @@ -20,6 +20,9 @@ from litellm.skills import main as skills_main @pytest.fixture(autouse=True) def clear_skill_ownership_env(monkeypatch): monkeypatch.delenv(skills_handler.ALLOW_UNOWNED_SKILL_ACCESS_ENV, raising=False) + skills_handler._SKILL_CACHE.clear() + yield + skills_handler._SKILL_CACHE.clear() def _skill(skill_id: str, created_by: str | None) -> LiteLLM_SkillsTable: @@ -433,3 +436,105 @@ async def test_should_scope_skill_injection_fetch_to_authenticated_user(monkeypa "litellm_skill_other", user_api_key_dict=auth, ) + + +# ── Cache layer ──────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_load_skill_uses_cache_after_first_db_hit(monkeypatch): + """`fetch_skill_from_db` is hit per-chat-completion; the cache absorbs + repeats so we don't issue a Prisma query on every request.""" + fake_skill = Mock(created_by="user-1", skill_id="litellm_skill_a") + store_factory = AsyncMock() + store = Mock() + store.find_skill = AsyncMock(return_value=fake_skill) + monkeypatch.setattr( + skills_handler.LiteLLMSkillsHandler, + "_get_prisma_client", + AsyncMock(return_value=store_factory), + ) + monkeypatch.setattr( + skills_handler, + "LiteLLMSkillsStore", + Mock(return_value=store), + ) + + first = await skills_handler.LiteLLMSkillsHandler._load_skill("litellm_skill_a") + second = await skills_handler.LiteLLMSkillsHandler._load_skill("litellm_skill_a") + third = await skills_handler.LiteLLMSkillsHandler._load_skill("litellm_skill_a") + + assert first is fake_skill + assert second is fake_skill + assert third is fake_skill + assert store.find_skill.await_count == 1 + + +@pytest.mark.asyncio +async def test_load_skill_caches_negative_lookups(monkeypatch): + """Missing skills must cache as `None` so repeated lookups skip the DB.""" + store_factory = AsyncMock() + store = Mock() + store.find_skill = AsyncMock(return_value=None) + monkeypatch.setattr( + skills_handler.LiteLLMSkillsHandler, + "_get_prisma_client", + AsyncMock(return_value=store_factory), + ) + monkeypatch.setattr( + skills_handler, + "LiteLLMSkillsStore", + Mock(return_value=store), + ) + + assert await skills_handler.LiteLLMSkillsHandler._load_skill("missing") is None + assert await skills_handler.LiteLLMSkillsHandler._load_skill("missing") is None + assert store.find_skill.await_count == 1 + + +@pytest.mark.asyncio +async def test_delete_skill_invalidates_cache(monkeypatch): + """After delete, the next read must consult the DB rather than the cached + pre-delete row.""" + fake_skill = Mock(created_by="user-1", skill_id="litellm_skill_a") + store = Mock() + store.find_skill = AsyncMock(return_value=fake_skill) + store.delete_skill = AsyncMock() + monkeypatch.setattr( + skills_handler.LiteLLMSkillsHandler, + "_get_prisma_client", + AsyncMock(return_value=Mock()), + ) + monkeypatch.setattr( + skills_handler, + "LiteLLMSkillsStore", + Mock(return_value=store), + ) + + # Prime the cache via the read path. + await skills_handler.LiteLLMSkillsHandler._load_skill("litellm_skill_a") + cached_hit, _ = skills_handler._read_skill_cache("litellm_skill_a") + assert cached_hit + + auth = UserAPIKeyAuth(user_id="user-1") + await skills_handler.LiteLLMSkillsHandler.delete_skill( + "litellm_skill_a", user_api_key_dict=auth + ) + + cached_hit_after, _ = skills_handler._read_skill_cache("litellm_skill_a") + assert not cached_hit_after + + +def test_skill_cache_expires_after_ttl(monkeypatch): + monkeypatch.setattr(skills_handler, "_SKILL_CACHE_TTL", 0.0) + skills_handler._write_skill_cache("k", Mock()) + cached_hit, _ = skills_handler._read_skill_cache("k") + assert not cached_hit + + +def test_skill_cache_evicts_when_at_capacity(monkeypatch): + monkeypatch.setattr(skills_handler, "_SKILL_CACHE_MAX_SIZE", 2) + skills_handler._write_skill_cache("a", Mock()) + skills_handler._write_skill_cache("b", Mock()) + skills_handler._write_skill_cache("c", Mock()) + assert skills_handler._SKILL_CACHE.keys() == {"c"}