From 0fc75c2bed5f0fe77e659625d4034d84b687cab8 Mon Sep 17 00:00:00 2001 From: Abhimanyu Kapur <38531241+akapur99@users.noreply.github.com> Date: Sat, 8 Aug 2026 15:46:53 -0700 Subject: [PATCH] perf: cache the whole active shadow-eval job set, not per-key rows MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 --- .../migration.sql | 1 + .../litellm_proxy_extras/schema.prisma | 1 + litellm/integrations/shadow_eval_logger.py | 73 ++++++++++--------- litellm/proxy/schema.prisma | 1 + schema.prisma | 1 + .../integrations/test_shadow_eval_logger.py | 69 ++++++++++++++++-- 6 files changed, 105 insertions(+), 41 deletions(-) diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260807000000_add_shadow_eval_job/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260807000000_add_shadow_eval_job/migration.sql index 05ef77e91a3..405812a38e4 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260807000000_add_shadow_eval_job/migration.sql +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260807000000_add_shadow_eval_job/migration.sql @@ -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" ( diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index ce0bdc056ab..013e82d2029 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1478,6 +1478,7 @@ model LiteLLM_ShadowEvalJob { @@index([team_id, status]) @@index([api_key_id, status]) + @@index([status]) @@index([created_at]) } diff --git a/litellm/integrations/shadow_eval_logger.py b/litellm/integrations/shadow_eval_logger.py index 635abdd39bf..93b1af00f54 100644 --- a/litellm/integrations/shadow_eval_logger.py +++ b/litellm/integrations/shadow_eval_logger.py @@ -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: diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index ce0bdc056ab..013e82d2029 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1478,6 +1478,7 @@ model LiteLLM_ShadowEvalJob { @@index([team_id, status]) @@index([api_key_id, status]) + @@index([status]) @@index([created_at]) } diff --git a/schema.prisma b/schema.prisma index ce0bdc056ab..013e82d2029 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1478,6 +1478,7 @@ model LiteLLM_ShadowEvalJob { @@index([team_id, status]) @@index([api_key_id, status]) + @@index([status]) @@index([created_at]) } diff --git a/tests/test_litellm/integrations/test_shadow_eval_logger.py b/tests/test_litellm/integrations/test_shadow_eval_logger.py index 0f327e15455..521e2314447 100644 --- a/tests/test_litellm/integrations/test_shadow_eval_logger.py +++ b/tests/test_litellm/integrations/test_shadow_eval_logger.py @@ -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")