fix(shadow_eval): copy messages before router call and raise judge output cap (#37232)

* fix(shadow_eval): copy messages before router call and raise judge output cap

* fix(shadow_eval): lead failure detail with location and pin post-failure continuation
This commit is contained in:
tin-berri 2026-08-17 17:10:52 -07:00 • committed by GitHub
parent e81cedb13a
commit 5277dab4f2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 86 additions and 3 deletions

View file

@ -9,6 +9,7 @@ across pods or stop races; the hook reads active jobs through a short-TTL cache.
import asyncio
import hashlib
import random
import traceback
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
@ -55,7 +56,7 @@ _MAX_JUDGE_PROMPT_CHARS: Final = 24_000
# The judge answers with a small JSON object; a tighter budget truncates the JSON
# mid-object and the attempt is lost to an error row.
JUDGE_MAX_OUTPUT_TOKENS: Final = 500
JUDGE_MAX_OUTPUT_TOKENS: Final = 1500
_MAX_ERROR_CHARS: Final = 500
@ -325,6 +326,14 @@ def _sample_hits(request_id: str, job_id: str, percentage: float) -> bool:
return bucket * 100.0 < percentage
def _failure_detail(e: BaseException) -> str:
"""Exception class, message, and the raising frame, so an attempt's error row names
the faulty code path without needing debug logs on the pod."""
frames: Final = traceback.extract_tb(e.__traceback__)
location: Final = f" at {frames[-1].filename.rsplit('/', 1)[-1]}:{frames[-1].lineno}" if frames else ""
return f"{type(e).__name__}{location}: {e}"
def _judge_call_cost(response: object) -> float:
"""Price a judge call, treating an unmapped judge model as free rather than fatal."""
import litellm
@ -764,7 +773,9 @@ class ShadowEvalLogger(CustomLogger):
try:
response: Final = await router.acompletion(
model=target_model,
messages=messages, # pyright: ignore[reportArgumentType] # snapshot of the SDK's own message dicts
messages=[ # mutable-ok: provider transforms rewrite messages in place, so the router gets its own copy
dict(m) for m in messages
], # pyright: ignore[reportArgumentType] # snapshot of the SDK's own message dicts
metadata=shadow_metadata,
num_retries=0,
fallbacks=[], # mutable-ok: SDK kwarg; a failed shadow is a recorded error, never a spend multiplier
@ -772,7 +783,7 @@ class ShadowEvalLogger(CustomLogger):
)
except Exception as e: # noqa: BLE001 # provider errors become error rows, not crashes
verbose_logger.debug("shadow_eval: router call failed: %s", e)
return _CallFailure(f"shadow router call failed: {e}")
return _CallFailure(f"shadow router call failed: {_failure_detail(e)}")
text: Final = _chat_final_text(response)
if not text:
return _CallFailure("shadow router returned an empty response")

View file

@ -12,10 +12,12 @@ from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.integrations.shadow_eval_logger import (
_MAX_CONCURRENT_SHADOW_TASKS,
_MAX_ERROR_CHARS,
_MAX_JUDGE_PROMPT_CHARS,
JUDGE_MAX_OUTPUT_TOKENS,
ActiveShadowEvalJob,
ShadowEvalLogger,
_failure_detail,
_judge_user_prompt,
_sample_hits,
_unmask_preference,
@ -438,6 +440,21 @@ def test_unmask_preference(raw, real_is_a, expected):
assert _unmask_preference(raw, real_is_a) == expected
def test_failure_detail_names_the_raising_frame():
try:
raise TypeError("'tuple' object does not support item assignment")
except TypeError as e:
detail = _failure_detail(e)
lineno = e.__traceback__.tb_lineno
assert detail == f"TypeError at test_shadow_eval_logger.py:{lineno}: 'tuple' object does not support item assignment"
try:
raise ValueError("p" * 5 * _MAX_ERROR_CHARS)
except ValueError as long_e:
truncated_row_error = _failure_detail(long_e)[:_MAX_ERROR_CHARS]
assert "ValueError at test_shadow_eval_logger.py:" in truncated_row_error
def test_judge_prompt_is_bounded_however_large_the_inputs():
prompt = _judge_user_prompt("c" * 200_000, "a" * 200_000, "b" * 200_000)
assert len(prompt) < _MAX_JUDGE_PROMPT_CHARS + 100
@ -476,6 +493,61 @@ class TestSuccessHookSkipChain:
assert row["error"] is None
assert prisma.db.litellm_shadowevaljob.find_many.await_count == 0
async def test_shadow_call_messages_survive_in_place_provider_rewrites(self, monkeypatch: pytest.MonkeyPatch):
"""Provider transforms (anthropic factory, cache-control hook) rewrite messages with
`messages[i] = ...`; the logger's immutable snapshot must never reach them directly."""
import litellm as litellm_module
monkeypatch.setattr(litellm_module, "completion_cost", lambda completion_response: 0.005)
prisma = _prisma()
router = _router()
inner = router.acompletion.side_effect
async def mutating_acompletion(**kwargs):
kwargs["messages"][0] = dict(kwargs["messages"][0])
return await inner(**kwargs)
router.acompletion = MagicMock(side_effect=mutating_acompletion)
logger = _logger(router=router, prisma=prisma, jobs=(_job(),))
await logger.async_log_success_event(_success_kwargs(), RESPONSE, None, None)
await _drain(logger)
row = prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"]
assert row["error"] is None
assert row["outcome"] in ("real", "shadow", "tie")
async def test_pipeline_continues_judging_after_a_failed_attempt(self, monkeypatch: pytest.MonkeyPatch):
import litellm as litellm_module
monkeypatch.setattr(litellm_module, "completion_cost", lambda completion_response: 0.005)
prisma = _prisma()
router = _router()
inner = router.acompletion.side_effect
shadow_calls = {"count": 0}
async def flaky_acompletion(**kwargs):
if kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) == SHADOW_EVAL_ROUTER_CALL_ORIGIN:
shadow_calls["count"] += 1
if shadow_calls["count"] == 1:
raise RuntimeError("provider exploded")
return await inner(**kwargs)
router.acompletion = MagicMock(side_effect=flaky_acompletion)
logger = _logger(router=router, prisma=prisma, jobs=(_job(),))
await logger.async_log_success_event(_success_kwargs(), RESPONSE, None, None)
await _drain(logger)
await logger.async_log_success_event(_success_kwargs(request_id="req-2"), RESPONSE, None, None)
await _drain(logger)
rows = [c.kwargs["data"] for c in prisma.db.litellm_shadowevalattempt.create.call_args_list]
assert [rows[0]["outcome"], rows[1]["outcome"] in ("real", "shadow")] == ["error", True]
assert "provider exploded" in rows[0]["error"]
assert rows[1]["request_id"] == "req-2"
assert rows[1]["error"] is None
assert logger._inflight_shadow_tasks == 0
@pytest.mark.parametrize(
"kwargs_mutation,job_mutation",
[