mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
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:
parent
7050e0c925
commit
0fc75c2bed
6 changed files with 105 additions and 41 deletions
|
|
@ -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" (
|
||||
|
|
|
|||
|
|
@ -1478,6 +1478,7 @@ model LiteLLM_ShadowEvalJob {
|
|||
|
||||
@@index([team_id, status])
|
||||
@@index([api_key_id, status])
|
||||
@@index([status])
|
||||
@@index([created_at])
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -1478,6 +1478,7 @@ model LiteLLM_ShadowEvalJob {
|
|||
|
||||
@@index([team_id, status])
|
||||
@@index([api_key_id, status])
|
||||
@@index([status])
|
||||
@@index([created_at])
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1478,6 +1478,7 @@ model LiteLLM_ShadowEvalJob {
|
|||
|
||||
@@index([team_id, status])
|
||||
@@index([api_key_id, status])
|
||||
@@index([status])
|
||||
@@index([created_at])
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue