perf: cache the whole active shadow-eval job set, not per-key rows

The per-key TTL cache did one indexed find_first per distinct key hash
per 30s — fine at small scale, but on a proxy serving 10k active keys
that is hundreds of small reads per second across pods, all to discover
that almost every key has no job.

Cache the entire active-job set instead: one find_many per pod per TTL
(the set is admin-started and capped at one job per key, so it is
single-digit rows), served to every key as an in-memory dict hit. DB
load is now flat and constant in the number of keys, idle or active.
On a DB blip the stale snapshot is kept and the next TTL retries, so a
blip degrades freshness rather than disabling the feature. Concurrent
requests share one refresh behind a lock instead of stampeding.

Adds @@index([status]) for the status-only find_many, folded into the
unshipped shadow eval migration.

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
Abhimanyu Kapur 2026-08-08 15:46:53 -07:00
parent 7050e0c925
commit 0fc75c2bed
6 changed files with 105 additions and 41 deletions

View file

@ -23,6 +23,7 @@ CREATE TABLE IF NOT EXISTS "LiteLLM_ShadowEvalJob" (
CREATE INDEX IF NOT EXISTS "LiteLLM_ShadowEvalJob_team_id_status_idx" ON "LiteLLM_ShadowEvalJob"("team_id", "status");
CREATE INDEX IF NOT EXISTS "LiteLLM_ShadowEvalJob_api_key_id_status_idx" ON "LiteLLM_ShadowEvalJob"("api_key_id", "status");
CREATE INDEX IF NOT EXISTS "LiteLLM_ShadowEvalJob_status_idx" ON "LiteLLM_ShadowEvalJob"("status");
CREATE INDEX IF NOT EXISTS "LiteLLM_ShadowEvalJob_created_at_idx" ON "LiteLLM_ShadowEvalJob"("created_at");
CREATE TABLE IF NOT EXISTS "LiteLLM_ShadowEvalVerdict" (

View file

@ -1478,6 +1478,7 @@ model LiteLLM_ShadowEvalJob {
@@index([team_id, status])
@@index([api_key_id, status])
@@index([status])
@@index([created_at])
}

View file

@ -35,7 +35,7 @@ from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, TypeAlias
from typing import TYPE_CHECKING, Final
from pydantic import BaseModel, TypeAdapter
@ -56,7 +56,7 @@ _JSON_FENCE_RE: Final = re.compile(r"```(?:json)?\s*(.*?)\s*```", re.DOTALL | re
# success hook skips them instead of shadowing the shadow.
SHADOW_EVAL_INTERNAL_MARKER: Final = "shadow_eval_internal"
# How long the per-key active-job lookup is cached. A shadow-eval job starting
# How long the active-job snapshot is cached. A shadow-eval job starting
# or stopping takes up to this long to be noticed by running pods.
_JOB_CACHE_TTL_SECONDS: Final = 30.0
@ -251,9 +251,6 @@ def _job_is_over_spend_cap(job: ActiveShadowEvalJob) -> bool:
return job.cost_actual >= max(job.cost_estimate * _SPEND_CAP_MULTIPLIER, _SPEND_CAP_FLOOR_USD)
_JobCache: TypeAlias = "dict[str, tuple[float, ActiveShadowEvalJob | None]]" # mutable-ok: TTL cache
class ShadowEvalLogger(CustomLogger):
"""Fires blind pairwise shadow evaluations for keys with an active shadow-eval job."""
@ -266,8 +263,13 @@ class ShadowEvalLogger(CustomLogger):
resolved at call time, not at logger construction."""
self._router_provider = router_provider or _default_router_provider
self._prisma_provider = prisma_provider or _default_prisma_provider
# api_key_hash -> (fetched_at_monotonic, job_record_or_None)
self._job_cache: _JobCache = {} # mutable-ok: a TTL cache is mutable state by definition
# Snapshot of every active job, keyed by shadowed api_key_id. Refreshed as a
# whole: active jobs are admin-started and capped at one per key, so the set is
# tiny, and one find_many per pod per TTL keeps DB load flat no matter how many
# distinct keys the proxy serves.
self._jobs_by_key: dict[str, ActiveShadowEvalJob] = {} # mutable-ok: TTL cache
self._jobs_fetched_at: float | None = None
self._jobs_refresh_lock: asyncio.Lock = asyncio.Lock()
self._inflight_shadow_tasks: int = 0
self._pending_seen: dict[str, int] = {} # mutable-ok: flush buffer
self._last_seen_flush: float = 0.0
@ -352,23 +354,34 @@ class ShadowEvalLogger(CustomLogger):
#### job lookup ####
async def _get_active_job(self, api_key_hash: str) -> ActiveShadowEvalJob | None:
cached: Final = self._job_cache.get(api_key_hash)
now: Final = asyncio.get_event_loop().time()
if cached is not None and now - cached[0] < _JOB_CACHE_TTL_SECONDS:
return cached[1]
prisma: Final = self._prisma_provider()
if prisma is None:
return None
try:
record: Final = await prisma.db.litellm_shadowevaljob.find_first(
where={ # mutable-ok: Prisma filter
"api_key_id": api_key_hash,
"status": {"in": ["pending", "running"]}, # mutable-ok: Prisma filter
},
order={"created_at": "desc"}, # mutable-ok: Prisma order
)
job: Final = (
ActiveShadowEvalJob(
if self._jobs_fetched_at is None or now - self._jobs_fetched_at >= _JOB_CACHE_TTL_SECONDS:
await self._refresh_active_jobs(now)
return self._jobs_by_key.get(api_key_hash)
async def _refresh_active_jobs(self, now: float) -> None:
"""Reload the active-job set, at most once per TTL across concurrent requests.
On a DB blip the stale snapshot is kept and the next TTL retries, so a blip
degrades freshness rather than turning the feature off.
"""
async with self._jobs_refresh_lock:
if self._jobs_fetched_at is not None and now - self._jobs_fetched_at < _JOB_CACHE_TTL_SECONDS:
return
prisma: Final = self._prisma_provider()
if prisma is None:
return
try:
records: Final = await prisma.db.litellm_shadowevaljob.find_many(
where={"status": {"in": ["pending", "running"]}}, # mutable-ok: Prisma filter
order={"created_at": "desc"}, # mutable-ok: Prisma order
)
except Exception as e: # noqa: BLE001 # a DB blip must not break request logging
verbose_logger.debug("shadow_eval: active-job refresh failed: %s", e)
return
jobs_by_key: Final[dict[str, ActiveShadowEvalJob]] = {} # mutable-ok: building the new snapshot
for record in reversed(records or []):
jobs_by_key[str(record.api_key_id)] = ActiveShadowEvalJob(
id=str(record.id),
router_name=str(record.router_name),
shadow_percentage=float(record.shadow_percentage),
@ -378,14 +391,8 @@ class ShadowEvalLogger(CustomLogger):
cost_actual=float(record.cost_actual or 0.0),
ends_at=_as_utc(getattr(record, "ends_at", None)),
)
if record is not None
else None
)
except Exception as e: # noqa: BLE001 # a DB blip must not break request logging
verbose_logger.debug("shadow_eval: job lookup failed: %s", e)
return cached[1] if cached is not None else None
self._job_cache[api_key_hash] = (now, job)
return job
self._jobs_by_key = jobs_by_key # mutable-ok: atomic snapshot swap
self._jobs_fetched_at = now
async def _finalize_job(self, job: ActiveShadowEvalJob, reason: str) -> None:
"""Flip a finished job to completed, keeping the verdicts it already produced.
@ -396,8 +403,8 @@ class ShadowEvalLogger(CustomLogger):
prisma: Final = self._prisma_provider()
if prisma is None:
return
self._job_cache = { # mutable-ok: TTL cache
k: v for k, v in self._job_cache.items() if v[1] is None or v[1].id != job.id
self._jobs_by_key = { # mutable-ok: atomic snapshot swap
k: v for k, v in self._jobs_by_key.items() if v.id != job.id
}
verbose_logger.info("shadow_eval: stopping job %s: %s", job.id, reason)
try:

View file

@ -1478,6 +1478,7 @@ model LiteLLM_ShadowEvalJob {
@@index([team_id, status])
@@index([api_key_id, status])
@@index([status])
@@index([created_at])
}

View file

@ -1478,6 +1478,7 @@ model LiteLLM_ShadowEvalJob {
@@index([team_id, status])
@@index([api_key_id, status])
@@index([status])
@@index([created_at])
}

View file

@ -120,9 +120,8 @@ def _logger_with_mocks(job=None):
router = MagicMock()
logger = ShadowEvalLogger(router_provider=lambda: router, prisma_provider=lambda: prisma)
if job is not None:
# Pre-warm the cache so no DB call is needed.
loop_time = asyncio.get_event_loop().time()
logger._job_cache["key-hash"] = (loop_time, job)
logger._jobs_by_key["key-hash"] = job
logger._jobs_fetched_at = asyncio.get_event_loop().time()
return logger, prisma, router
@ -131,7 +130,7 @@ class TestSuccessHookSkipPaths:
async def test_skips_without_standard_logging_object(self):
logger, prisma, _ = _logger_with_mocks()
await logger.async_log_success_event({"messages": []}, MagicMock(), None, None)
prisma.db.litellm_shadowevaljob.find_first.assert_not_called()
prisma.db.litellm_shadowevaljob.find_many.assert_not_called()
async def test_skips_own_internal_traffic(self):
job = ActiveShadowEvalJob(id="j1", router_name="r", shadow_percentage=100.0, judge_model="m", status="running")
@ -150,7 +149,7 @@ class TestSuccessHookSkipPaths:
async def test_skips_when_no_active_job(self):
logger, prisma, _ = _logger_with_mocks()
prisma.db.litellm_shadowevaljob.find_first = AsyncMock(return_value=None)
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(return_value=[])
kwargs = {
"standard_logging_object": {
"id": "req-1",
@ -701,7 +700,7 @@ class TestPerJobSpendCap:
)
logger, prisma, _ = _logger_with_mocks(job)
prisma.db.litellm_shadowevaljob.update_many = AsyncMock()
prisma.db.litellm_shadowevaljob.find_first = AsyncMock(return_value=None)
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(return_value=[])
await logger.async_log_success_event(self._success_kwargs(), MagicMock(), None, None)
await logger.async_log_success_event(self._success_kwargs(), MagicMock(), None, None)
@ -709,6 +708,59 @@ class TestPerJobSpendCap:
assert prisma.db.litellm_shadowevaljob.update_many.await_count == 1
@pytest.mark.asyncio
class TestActiveJobSnapshot:
"""One find_many per pod per TTL serves every key, so DB load stays flat no
matter how many distinct keys send traffic through the proxy."""
@staticmethod
def _record(job_id: str, api_key_id: str) -> MagicMock:
record = MagicMock()
record.id = job_id
record.api_key_id = api_key_id
record.router_name = "r"
record.shadow_percentage = 100.0
record.judge_model = "m"
record.status = "running"
record.cost_estimate = 10.0
record.cost_actual = 0.0
record.ends_at = None
return record
async def test_one_query_serves_lookups_for_many_distinct_keys(self):
logger, prisma, _ = _logger_with_mocks()
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(return_value=[self._record("j1", "key-1")])
job_hit = await logger._get_active_job("key-1")
misses = [await logger._get_active_job(f"other-{i}") for i in range(50)]
assert job_hit is not None and job_hit.id == "j1"
assert all(m is None for m in misses)
assert prisma.db.litellm_shadowevaljob.find_many.await_count == 1
async def test_newest_job_wins_when_a_key_somehow_has_two(self):
logger, prisma, _ = _logger_with_mocks()
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(
return_value=[self._record("j-newest", "key-1"), self._record("j-oldest", "key-1")]
)
job = await logger._get_active_job("key-1")
assert job is not None and job.id == "j-newest"
async def test_db_blip_keeps_the_stale_snapshot_instead_of_disabling_the_feature(self):
logger, prisma, _ = _logger_with_mocks()
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(return_value=[self._record("j1", "key-1")])
assert (await logger._get_active_job("key-1")) is not None
logger._jobs_fetched_at = asyncio.get_event_loop().time() - 61.0
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(side_effect=RuntimeError("db down"))
job = await logger._get_active_job("key-1")
assert job is not None and job.id == "j1"
@pytest.mark.asyncio
class TestJobsStopAtTheirScheduledEnd:
"""A shadow eval samples ongoing traffic, so a job whose window has closed must
@ -766,7 +818,7 @@ class TestJobsStopAtTheirScheduledEnd:
async def test_expiry_evicts_the_cached_job_so_later_requests_do_not_rewrite_it(self):
logger, prisma, _ = _logger_with_mocks(self._job(datetime.now(timezone.utc) - timedelta(seconds=1)))
prisma.db.litellm_shadowevaljob.update_many = AsyncMock()
prisma.db.litellm_shadowevaljob.find_first = AsyncMock(return_value=None)
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(return_value=[])
await logger.async_log_success_event(TestPerJobSpendCap._success_kwargs(), MagicMock(), None, None)
await logger.async_log_success_event(TestPerJobSpendCap._success_kwargs(), MagicMock(), None, None)
@ -777,6 +829,7 @@ class TestJobsStopAtTheirScheduledEnd:
logger, prisma, _ = _logger_with_mocks()
record = MagicMock()
record.id = "j1"
record.api_key_id = "key-hash"
record.router_name = "r"
record.shadow_percentage = 100.0
record.judge_model = "m"
@ -784,7 +837,7 @@ class TestJobsStopAtTheirScheduledEnd:
record.cost_estimate = 10.0
record.cost_actual = 0.0
record.ends_at = datetime.now(timezone.utc).replace(tzinfo=None) - timedelta(seconds=1)
prisma.db.litellm_shadowevaljob.find_first = AsyncMock(return_value=record)
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(return_value=[record])
job = await logger._get_active_job("key-hash")