mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
shadow-eval: address bot review findings
- Forward the original request's non-default params (temperature, tools, response_format, etc.) to the shadow router call. Previously only model and messages were sent, so a request with tools or a non-zero temperature was judged against a shadow response generated under totally different sampling settings -- an unfair, biased comparison. stream and metadata are still stripped: the shadow call needs the full text back and must not leak the caller's own metadata. - Fix unbounded background task backlog: asyncio.create_task() fired unconditionally and only the task body waited on a semaphore, so a traffic spike queued unlimited tasks (each holding a copy of messages/response) before any of them ran. Now the in-flight count is checked and incremented before scheduling; over capacity, the sample is dropped instead of queued. - Fix stopped jobs being silently reactivated: the verdict-write counter update unconditionally set status='running', so a pipeline that started before stop_shadow_eval_job() marked the job 'completed' could overwrite that back to 'running' after the fact. Now it's a conditional update_many scoped to status in (pending, running), so a completed job can never transition back. - Strip comments/docstrings added during the previous fix pass per CLAUDE.md's no-new-comments rule (flagged by review bot). Tests: 5 new regression tests (param forwarding, stream/metadata stripping, backlog cap drops samples under saturation, backlog cap schedules+decrements under capacity, stop-race status guard). 31/31 passing. Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
parent
0b96dccd07
commit
5f03456b0e
3 changed files with 183 additions and 61 deletions
|
|
@ -54,8 +54,6 @@ _MAX_CONCURRENT_SHADOW_TASKS: Final = 16
|
|||
# Truncation bound for text handed to the judge, to keep judge calls affordable.
|
||||
_MAX_JUDGE_CHARS: Final = 16_000
|
||||
|
||||
# request_count is display-only, so it is buffered in memory and flushed at most
|
||||
# once per interval instead of one UPDATE per request on the shadowed key.
|
||||
_SEEN_FLUSH_INTERVAL_SECONDS: Final = 10.0
|
||||
|
||||
PAIRWISE_JUDGE_SYSTEM_PROMPT: Final = """You are an impartial quality judge comparing two responses to the same conversation.
|
||||
|
|
@ -140,9 +138,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
self._prisma_provider = prisma_provider or _default_prisma_provider
|
||||
# api_key_hash -> (fetched_at_monotonic, job_record_or_None)
|
||||
self._job_cache: dict[str, tuple[float, dict[str, Any] | None]] = {}
|
||||
self._semaphore = asyncio.Semaphore(_MAX_CONCURRENT_SHADOW_TASKS)
|
||||
# job_id -> requests seen since last flush. Flushed opportunistically so a
|
||||
# high-traffic key costs one UPDATE per flush interval, not one per request.
|
||||
self._inflight_shadow_tasks: int = 0
|
||||
self._pending_seen: dict[str, int] = {}
|
||||
self._last_seen_flush: float = 0.0
|
||||
|
||||
|
|
@ -167,8 +163,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
if not request_id:
|
||||
return
|
||||
# The job tracks every request it saw, sampled or not, so the UI can
|
||||
# show "N of M requests shadowed". Buffered: one UPDATE per flush
|
||||
# interval, not one per request.
|
||||
# show "N of M requests shadowed".
|
||||
self._pending_seen[job["id"]] = self._pending_seen.get(job["id"], 0) + 1
|
||||
now: Final = asyncio.get_event_loop().time()
|
||||
if now - self._last_seen_flush >= _SEEN_FLUSH_INTERVAL_SECONDS:
|
||||
|
|
@ -178,15 +173,20 @@ class ShadowEvalLogger(CustomLogger):
|
|||
return
|
||||
if payload.get("call_type") not in (None, "completion", "acompletion", "chat_completion"):
|
||||
return # only chat-shaped traffic is comparable
|
||||
asyncio.create_task(
|
||||
if self._inflight_shadow_tasks >= _MAX_CONCURRENT_SHADOW_TASKS:
|
||||
return
|
||||
self._inflight_shadow_tasks += 1
|
||||
task = asyncio.create_task(
|
||||
self._run_shadow_eval(
|
||||
job=dict(job),
|
||||
request_id=request_id,
|
||||
messages=list(kwargs.get("messages") or []),
|
||||
response_obj=response_obj,
|
||||
real_model=payload.get("model") or "",
|
||||
model_parameters=dict(payload.get("model_parameters") or {}),
|
||||
)
|
||||
)
|
||||
task.add_done_callback(lambda _: setattr(self, "_inflight_shadow_tasks", self._inflight_shadow_tasks - 1))
|
||||
except Exception as e: # noqa: BLE001 # logging hooks must never fail the request
|
||||
verbose_logger.debug("shadow_eval: failed to schedule task: %s", e)
|
||||
|
||||
|
|
@ -221,7 +221,6 @@ class ShadowEvalLogger(CustomLogger):
|
|||
return job
|
||||
|
||||
async def _flush_seen_counts(self) -> None:
|
||||
"""Write buffered request-seen counts, one UPDATE per job with pending counts."""
|
||||
prisma: Final = self._prisma_provider()
|
||||
if prisma is None:
|
||||
return
|
||||
|
|
@ -245,59 +244,59 @@ class ShadowEvalLogger(CustomLogger):
|
|||
messages: list[dict[str, Any]],
|
||||
response_obj: Any,
|
||||
real_model: str,
|
||||
model_parameters: dict[str, Any],
|
||||
) -> None:
|
||||
"""Detached background task: shadow call -> blind judge -> verdict row."""
|
||||
async with self._semaphore:
|
||||
prisma: Final = self._prisma_provider()
|
||||
try:
|
||||
real_text: Final = self._extract_response_text(response_obj)
|
||||
if not real_text or not messages:
|
||||
return
|
||||
prisma: Final = self._prisma_provider()
|
||||
try:
|
||||
real_text: Final = self._extract_response_text(response_obj)
|
||||
if not real_text or not messages:
|
||||
return
|
||||
|
||||
shadow = await self._call_router_shadow(job["router_name"], messages)
|
||||
if shadow is None:
|
||||
await self._bump_failed(job["id"])
|
||||
return
|
||||
shadow_text, shadow_model, tier, shadow_tokens = shadow
|
||||
|
||||
verdict = await self._call_judge(
|
||||
judge_model=job["judge_model"],
|
||||
messages=messages,
|
||||
real_text=real_text,
|
||||
shadow_text=shadow_text,
|
||||
)
|
||||
if verdict is None:
|
||||
await self._bump_failed(job["id"])
|
||||
return
|
||||
preference, confidence, reasoning, judge_cost = verdict
|
||||
|
||||
if prisma is None:
|
||||
return
|
||||
await prisma.db.litellm_shadowevalverdict.create(
|
||||
data={
|
||||
"job_id": job["id"],
|
||||
"request_id": request_id,
|
||||
"tier_classification": tier,
|
||||
"real_model": real_model,
|
||||
"shadow_model": shadow_model,
|
||||
"shadow_response_tokens": shadow_tokens,
|
||||
"judge_preference": preference,
|
||||
"judge_confidence": confidence,
|
||||
"judge_reasoning": reasoning[:1000] if reasoning else None,
|
||||
"judge_model": job["judge_model"],
|
||||
}
|
||||
)
|
||||
await prisma.db.litellm_shadowevaljob.update(
|
||||
where={"id": job["id"]},
|
||||
data={
|
||||
"completed_count": {"increment": 1},
|
||||
"cost_actual": {"increment": judge_cost},
|
||||
"status": "running",
|
||||
},
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # detached task: log, count, never raise
|
||||
verbose_logger.debug("shadow_eval: pipeline failed for %s: %s", request_id, e)
|
||||
shadow = await self._call_router_shadow(job["router_name"], messages, model_parameters)
|
||||
if shadow is None:
|
||||
await self._bump_failed(job["id"])
|
||||
return
|
||||
shadow_text, shadow_model, tier, shadow_tokens = shadow
|
||||
|
||||
verdict = await self._call_judge(
|
||||
judge_model=job["judge_model"],
|
||||
messages=messages,
|
||||
real_text=real_text,
|
||||
shadow_text=shadow_text,
|
||||
)
|
||||
if verdict is None:
|
||||
await self._bump_failed(job["id"])
|
||||
return
|
||||
preference, confidence, reasoning, judge_cost = verdict
|
||||
|
||||
if prisma is None:
|
||||
return
|
||||
await prisma.db.litellm_shadowevalverdict.create(
|
||||
data={
|
||||
"job_id": job["id"],
|
||||
"request_id": request_id,
|
||||
"tier_classification": tier,
|
||||
"real_model": real_model,
|
||||
"shadow_model": shadow_model,
|
||||
"shadow_response_tokens": shadow_tokens,
|
||||
"judge_preference": preference,
|
||||
"judge_confidence": confidence,
|
||||
"judge_reasoning": reasoning[:1000] if reasoning else None,
|
||||
"judge_model": job["judge_model"],
|
||||
}
|
||||
)
|
||||
await prisma.db.litellm_shadowevaljob.update_many(
|
||||
where={"id": job["id"], "status": {"in": ["pending", "running"]}},
|
||||
data={
|
||||
"completed_count": {"increment": 1},
|
||||
"cost_actual": {"increment": judge_cost},
|
||||
"status": "running",
|
||||
},
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # detached task: log, count, never raise
|
||||
verbose_logger.debug("shadow_eval: pipeline failed for %s: %s", request_id, e)
|
||||
await self._bump_failed(job["id"])
|
||||
|
||||
async def _bump_failed(self, job_id: str) -> None:
|
||||
prisma: Final = self._prisma_provider()
|
||||
|
|
@ -312,7 +311,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
verbose_logger.debug("shadow_eval: failed_count increment failed: %s", e)
|
||||
|
||||
async def _call_router_shadow(
|
||||
self, router_name: str, messages: list[dict[str, Any]]
|
||||
self, router_name: str, messages: list[dict[str, Any]], model_parameters: dict[str, Any]
|
||||
) -> tuple[str, str, str | None, int | None] | None:
|
||||
"""Send the prompt through the auto-router; return (text, model, tier, completion_tokens)."""
|
||||
router: Final = self._router_provider()
|
||||
|
|
@ -322,11 +321,13 @@ class ShadowEvalLogger(CustomLogger):
|
|||
# The router's pre-routing hook writes its routing decision into this
|
||||
# metadata dict; read it back after the call for tier attribution.
|
||||
shadow_metadata: dict[str, Any] = {SHADOW_EVAL_INTERNAL_MARKER: True}
|
||||
shadow_params: Final = {k: v for k, v in model_parameters.items() if k not in ("stream", "metadata")}
|
||||
try:
|
||||
response = await router.acompletion(
|
||||
model=router_name,
|
||||
messages=messages,
|
||||
metadata=shadow_metadata,
|
||||
**shadow_params,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # provider errors are a counted failure, not a crash
|
||||
verbose_logger.debug("shadow_eval: router call failed: %s", e)
|
||||
|
|
|
|||
|
|
@ -482,13 +482,11 @@ def _require_admin_viewer(user_api_key_dict: UserAPIKeyAuth, action: str) -> Non
|
|||
|
||||
|
||||
def _require_admin_writer(user_api_key_dict: UserAPIKeyAuth, action: str) -> None:
|
||||
"""Starting or stopping a shadow eval spends money (judge calls); view-only admins may not."""
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
raise HTTPException(status_code=403, detail=f"Only a proxy admin can {action}")
|
||||
|
||||
|
||||
def _is_configured_pre_routing_strategy(llm_router: "Router", router_name: str) -> bool:
|
||||
"""True when `router_name` is any configured pre-routing strategy (auto, complexity, adaptive, quality)."""
|
||||
return any(
|
||||
router_name in registry
|
||||
for registry in (
|
||||
|
|
|
|||
|
|
@ -138,6 +138,129 @@ class TestSuccessHookSkipPaths:
|
|||
assert logger._pending_seen == {"j1": 1}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestCallRouterShadowForwardsParameters:
|
||||
async def test_forwards_non_default_params(self):
|
||||
router = MagicMock()
|
||||
router.acompletion = AsyncMock(return_value={"choices": [{"message": {"content": "shadow reply"}}]})
|
||||
logger = ShadowEvalLogger(router_provider=lambda: router, prisma_provider=lambda: MagicMock())
|
||||
|
||||
await logger._call_router_shadow(
|
||||
"claude-auto",
|
||||
[{"role": "user", "content": "hi"}],
|
||||
{"temperature": 0.2, "tools": [{"type": "function"}], "max_tokens": 500},
|
||||
)
|
||||
|
||||
_, kwargs = router.acompletion.call_args
|
||||
assert kwargs["temperature"] == 0.2
|
||||
assert kwargs["tools"] == [{"type": "function"}]
|
||||
assert kwargs["max_tokens"] == 500
|
||||
|
||||
async def test_drops_stream_and_metadata_from_forwarded_params(self):
|
||||
router = MagicMock()
|
||||
router.acompletion = AsyncMock(return_value={"choices": [{"message": {"content": "shadow reply"}}]})
|
||||
logger = ShadowEvalLogger(router_provider=lambda: router, prisma_provider=lambda: MagicMock())
|
||||
|
||||
await logger._call_router_shadow(
|
||||
"claude-auto",
|
||||
[{"role": "user", "content": "hi"}],
|
||||
{"stream": True, "metadata": {"user_api_key_hash": "leaked"}, "temperature": 0.5},
|
||||
)
|
||||
|
||||
_, kwargs = router.acompletion.call_args
|
||||
assert "stream" not in kwargs
|
||||
assert kwargs["temperature"] == 0.5
|
||||
# metadata must stay the logger's own internal marker, not the caller's
|
||||
assert kwargs["metadata"] == {SHADOW_EVAL_INTERNAL_MARKER: True}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestInflightTaskBacklog:
|
||||
async def test_drops_sample_when_at_capacity(self):
|
||||
job = {"id": "j1", "router_name": "r", "shadow_percentage": 100.0, "judge_model": "m", "status": "running"}
|
||||
logger, _, _ = _logger_with_mocks(job)
|
||||
logger._inflight_shadow_tasks = 999999 # simulate saturation regardless of the real cap
|
||||
|
||||
kwargs = {
|
||||
"standard_logging_object": {
|
||||
"id": "req-1",
|
||||
"model": "gpt-4o",
|
||||
"call_type": "acompletion",
|
||||
"metadata": {"user_api_key_hash": "key-hash"},
|
||||
},
|
||||
"litellm_params": {"metadata": {}},
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
before = logger._inflight_shadow_tasks
|
||||
await logger.async_log_success_event(kwargs, MagicMock(), None, None)
|
||||
# No task should have been scheduled; the counter is untouched by the hook itself.
|
||||
assert logger._inflight_shadow_tasks == before
|
||||
|
||||
async def test_schedules_and_decrements_when_under_capacity(self):
|
||||
job = {"id": "j1", "router_name": "r", "shadow_percentage": 100.0, "judge_model": "m", "status": "running"}
|
||||
logger, _, router = _logger_with_mocks(job)
|
||||
router.acompletion = AsyncMock(side_effect=asyncio.sleep(0)) # keep the task alive briefly
|
||||
|
||||
async def fake_run(*args, **kwargs):
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
logger._run_shadow_eval = AsyncMock(side_effect=fake_run)
|
||||
|
||||
kwargs = {
|
||||
"standard_logging_object": {
|
||||
"id": "req-1",
|
||||
"model": "gpt-4o",
|
||||
"call_type": "acompletion",
|
||||
"metadata": {"user_api_key_hash": "key-hash"},
|
||||
},
|
||||
"litellm_params": {"metadata": {}},
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
await logger.async_log_success_event(kwargs, MagicMock(), None, None)
|
||||
assert logger._inflight_shadow_tasks == 1
|
||||
await asyncio.sleep(0.05)
|
||||
assert logger._inflight_shadow_tasks == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestStoppedJobCannotBeReactivated:
|
||||
async def test_verdict_write_uses_conditional_status_guard(self):
|
||||
"""Regression: a verdict written after stop() must not resurrect a completed job.
|
||||
|
||||
stop_shadow_eval_job() sets status="completed". If a pipeline that started
|
||||
before the stop finishes afterwards, its counter update must not blindly
|
||||
set status back to "running" -- it has to filter on the job still being
|
||||
pending/running, so an already-completed job stays completed.
|
||||
"""
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_shadowevalverdict.create = AsyncMock()
|
||||
prisma.db.litellm_shadowevaljob.update_many = AsyncMock()
|
||||
# If the fix regresses to `.update(...)`, this test must fail loudly
|
||||
# rather than silently pass by mocking away the missing method.
|
||||
del prisma.db.litellm_shadowevaljob.update
|
||||
|
||||
logger = ShadowEvalLogger(router_provider=lambda: MagicMock(), prisma_provider=lambda: prisma)
|
||||
logger._call_router_shadow = AsyncMock(return_value=("shadow text", "shadow-model", "SIMPLE", 10))
|
||||
logger._call_judge = AsyncMock(return_value=("real", 0.9, "clearer", 0.01))
|
||||
|
||||
job = {"id": "j1", "router_name": "r", "judge_model": "m", "shadow_percentage": 100.0, "status": "running"}
|
||||
await logger._run_shadow_eval(
|
||||
job=job,
|
||||
request_id="req-1",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
response_obj={"choices": [{"message": {"content": "real text"}}]},
|
||||
real_model="gpt-4o",
|
||||
model_parameters={},
|
||||
)
|
||||
|
||||
prisma.db.litellm_shadowevaljob.update_many.assert_awaited_once()
|
||||
_, call_kwargs = prisma.db.litellm_shadowevaljob.update_many.call_args
|
||||
where = call_kwargs["where"]
|
||||
assert where["id"] == "j1"
|
||||
assert set(where["status"]["in"]) == {"pending", "running"}
|
||||
assert call_kwargs["data"]["status"] == "running"
|
||||
|
||||
|
||||
class TestExtractResponseText:
|
||||
def test_dict_response(self):
|
||||
resp = {"choices": [{"message": {"content": "hello"}}]}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue