perf(proxy): cache container/skill ownership reads on the hot path

Container ownership and skill rows are looked up on every retrieve /
delete / list / file-content / chat-completion-with-skill call. The new
stores wrapped raw Prisma queries with no cache, putting one DB
round-trip on each request. Add an in-process TTL'd cache mirroring the
_byok_cred_cache pattern in mcp_server/server.py: per-key (value,
monotonic_timestamp), 60s TTL, 10000-entry cap with full-clear on
overflow, invalidated by every write. Negative results (`None`) are
cached too so untracked-resource checks also skip the DB.

Tests cover: cache-after-first-hit, negative caching, write
invalidation, no-caching-on-DB-error, TTL expiry, capacity eviction.
56 tests pass.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
user 2026-05-02 03:26:58 +00:00
parent 22ced8d507
commit 6194028f79
No known key found for this signature in database
4 changed files with 342 additions and 10 deletions

View file

@ -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"}

View file

@ -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)

View file

@ -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"}

View file

@ -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"}