fix(proxy): keep deployment labels on cache-hit post_call guardrail rejections (#42780)

* fix(proxy): keep deployment labels on cache-hit post_call rejections

A post-call failure on a response served from the litellm cache set no first_api_call_start_time, so the failure hook flagged it as rejected before routing and dropped the model_id and provider labels

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): read the cache hit from caching_details in the failure hook

model_call_details[cache_hit] is stamped inside the enqueued success handler, so a post-call failure can observe it too early; logging_obj.caching_details is set synchronously before the cached response returns

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): cover cache-hit guardrail reject deployment labels across endpoints, modes and chaos

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): bound the worker-kill reject count by in-flight losses

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): assert provider and model labels on the cache-hit regression test

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yucheng <yucheng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-23 22:48:04 -07:00 • committed by GitHub
parent 251fdf0308
commit 1519032d90
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 1069 additions and 1 deletions

View file

@ -974,6 +974,14 @@ def _failure_usage_to_lift(
_EMPTY_LIFT: Final = MappingProxyType({})
def _reached_deployment(litellm_logging_obj: Logging) -> bool:
"""A provider handoff or a cached response both mean the router selected a deployment."""
caching_details: Final = litellm_logging_obj.caching_details
return litellm_logging_obj.model_call_details.get("first_api_call_start_time") is not None or (
caching_details is not None and caching_details.get("cache_hit") is True
)
def _stamp_deployment_attribution(
litellm_params: dict[str, object], model_group: str | None, team_id: str | None, dispatched: bool
) -> Mapping[str, object]:
@ -3325,7 +3333,7 @@ class ProxyLogging:
_litellm_params,
request_data.get("model"),
user_api_key_dict.team_id,
dispatched=litellm_logging_obj.model_call_details.get("first_api_call_start_time") is not None,
dispatched=_reached_deployment(litellm_logging_obj),
)
litellm_logging_obj.update_environment_variables(

View file

@ -0,0 +1,652 @@
import asyncio
import json
import subprocess
import uuid
from collections.abc import Callable, Generator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import anthropic
import httpx
import openai
import pytest
import yaml
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from prometheus_client.parser import text_string_to_metric_families
GUARDRAIL_PATH: Final = "/beta/litellm_basic_guardrail_api"
DEPLOYMENT_FAILURE: Final = "litellm_deployment_failure_responses_total"
DEPLOYMENT_REQUESTS: Final = "litellm_deployment_total_requests_total"
DEPLOYMENT_STATE: Final = "litellm_deployment_state"
PROXY_FAILED: Final = "litellm_proxy_failed_requests_metric_total"
def _chat_sse(marker: str) -> tuple[bytes, ...]:
chunk: Final = {
"id": "chatcmpl_" + marker,
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
}
frames: Final = (
{**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "content": ""}, "finish_reason": None}]},
{**chunk, "choices": [{"index": 0, "delta": {"content": "provider control"}, "finish_reason": None}]},
{**chunk, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]},
)
return tuple(f"data: {json.dumps(frame)}".encode() for frame in frames) + (b"data: [DONE]",)
def _provider_body(target: str, marker: str, streamed: bool) -> Reply:
match target:
case "/v1/chat/completions":
if streamed:
return Reply(content_type="text/event-stream", chunks=_chat_sse(marker))
body: dict = {
"id": "chatcmpl_" + marker,
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "provider control " + marker},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
}
case "/v1/messages":
body = {
"id": "msg_" + marker,
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-5-20250929",
"content": [{"type": "text", "text": "provider control " + marker}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 11, "output_tokens": 4},
}
case "/v1/responses":
body = {
"id": "resp_" + marker,
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o-mini",
"output": [
{
"type": "message",
"id": "msg_" + marker,
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "provider control " + marker, "annotations": []}],
}
],
"usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15},
}
case "/v1/embeddings":
body = {
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
"model": "text-embedding-3-small",
"usage": {"prompt_tokens": 3, "total_tokens": 3},
}
case _:
return Reply(status=404, body=json.dumps({"error": "unexpected provider target " + target}).encode())
return Reply(body=json.dumps(body).encode())
def _provider(marker: str) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
streamed: Final = b'"stream":true' in request.body.replace(b" ", b"")
return _provider_body(request.target.split("?", 1)[0], marker, streamed)
return respond
def _blocking_sink(request: Request) -> Reply:
assert request.target == GUARDRAIL_PATH, request.target
return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic block"}).encode())
def _failing_sink(status: int) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert request.target == GUARDRAIL_PATH, request.target
return Reply(status=status, body=json.dumps({"error": "synthetic guardrail outage"}).encode())
return respond
def _first_call_pass_sink() -> Callable[[Request], Reply]:
calls: list[int] = [] # mutable-ok: the wire handler must remember call order across requests
def respond(request: Request) -> Reply:
assert request.target == GUARDRAIL_PATH, request.target
calls.append(1)
action: dict = (
{"action": "NONE"} if len(calls) == 1 else {"action": "BLOCKED", "blocked_reason": "synthetic block"}
)
return Reply(body=json.dumps(action).encode())
return respond
def _guardrail_config(
tmp_path: Path,
name: str,
sink_url: str,
*,
mode: str = "post_call",
default_on: bool = False,
local_cache: bool = False,
ttl: int | None = None,
) -> Path:
config: dict = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"]["callbacks"] = ["prometheus"]
if local_cache:
config["litellm_settings"]["cache_params"] = {"type": "local"}
if ttl is not None:
config["litellm_settings"]["cache_params"]["ttl"] = ttl
config["guardrails"] = [
{
"guardrail_name": name,
"litellm_params": {
"guardrail": "generic_guardrail_api",
"mode": mode,
"default_on": default_on,
"api_base": sink_url,
"api_key": "synthetic-guardrail-key",
},
}
]
path: Final = tmp_path / "guardrail.yaml"
path.write_text(yaml.safe_dump(config))
return path
@dataclass(frozen=True, slots=True)
class Rig:
candidate: Gateway
scenario: Scenario
model_name: str
deployment_id: str
guardrail_name: str
policy: Wire
provider: Wire
process: subprocess.Popen[bytes]
@contextmanager
def _rig(
gateway: Gateway,
tmp_path: Path,
marker: str,
*,
sink: Callable[[Request], Reply] = _blocking_sink,
mode: str = "post_call",
default_on: bool = False,
local_cache: bool = False,
ttl: int | None = None,
workers: int = 1,
upstream_model: str = "openai/gpt-4o-mini",
api_base_suffix: str = "/v1",
env: Mapping[str, str] | None = None,
) -> Generator[Rig, None, None]:
identity: Final = "guardrail-" + marker
with wire_server(sink) as policy, wire_server(_provider(marker)) as provider:
config: Final = _guardrail_config(
tmp_path, identity, policy.url, mode=mode, default_on=default_on, local_cache=local_cache, ttl=ttl
)
prom_dir: Final = tmp_path / "prom"
prom_dir.mkdir()
with (
owned_proxy_process(
gateway,
tmp_path,
{"PROMETHEUS_MULTIPROC_DIR": str(prom_dir), **(env or {})},
config=config,
workers=workers,
) as owned,
owned.gateway.scenario() as scenario,
):
model: Final = scenario.model(
model=upstream_model, api_base=provider.url + api_base_suffix, api_key="synthetic-provider-key"
)
entries: Final = owned.gateway.get("/model/info")["data"]
assert isinstance(entries, list)
entry: Final = next(item for item in entries if object_value(item)["model_name"] == model)
yield Rig(
owned.gateway,
scenario,
model,
string_value(object_value(object_value(entry)["model_info"])["id"]),
identity,
policy,
provider,
owned.process,
)
def _metric_samples(candidate: Gateway, model_name: str) -> tuple:
response: Final = candidate.client.request(
"GET", "/metrics", headers={"Authorization": f"Bearer {candidate.key}"}, follow_redirects=True
)
assert response.status_code == 200, f"GET /metrics: {response.status_code} {response.text[:300]}"
return tuple(
sample
for family in text_string_to_metric_families(response.text)
for sample in family.samples
if sample.labels.get("requested_model") == model_name
or (sample.name == DEPLOYMENT_STATE and sample.labels.get("model_id") != "")
)
def _count(samples: tuple, name: str, model_id: str) -> float:
return float(
sum(sample.value for sample in samples if sample.name == name and sample.labels.get("model_id") == model_id)
)
def _populated_failures(samples: tuple, rig: Rig, api_provider: str) -> float:
return float(
sum(
sample.value
for sample in samples
if sample.name == DEPLOYMENT_FAILURE
and sample.labels.get("model_id") == rig.deployment_id
and sample.labels.get("api_provider") == api_provider
and sample.labels.get("litellm_model_name") != ""
)
)
def _expect_metrics(
rig: Rig,
populated: float,
blank: float,
*,
api_provider: str = "openai",
pf_id: str | None = None,
pf_populated: float | None = None,
pf_blank: float | None = None,
) -> tuple:
expected_id: Final = rig.deployment_id if pf_id is None else pf_id
expected_pf_populated: Final = populated if pf_populated is None else pf_populated
expected_pf_blank: Final = blank if pf_blank is None else pf_blank
def read() -> tuple:
samples: Final = _metric_samples(rig.candidate, rig.model_name)
satisfied: Final = (
_populated_failures(samples, rig, api_provider) == populated
and _count(samples, DEPLOYMENT_FAILURE, "") == blank
and _count(samples, PROXY_FAILED, expected_id) == expected_pf_populated
and _count(samples, PROXY_FAILED, "") == expected_pf_blank
)
return samples if satisfied else ()
return eventually(read, bool, seconds=70)
def _spend_rows(call_id: str) -> tuple[dict, ...]:
rows: Final = eventually(
lambda: read_rows(
'SELECT request_id, custom_llm_provider, model_id, status FROM "LiteLLM_SpendLogs" '
"WHERE request_id = %s OR request_id LIKE %s",
(call_id, call_id + "\\_%"),
),
lambda values: len(values) >= 1,
seconds=70,
)
return tuple(dict(row) for row in rows)
def _assert_spend(call_id: str, rig: Rig, api_provider: str = "openai") -> None:
rows: Final = _spend_rows(call_id)
failures: Final = tuple(row for row in rows if row["status"] == "failure")
assert len(failures) == 1, rows
assert (failures[0]["custom_llm_provider"], failures[0]["model_id"]) == (api_provider, rig.deployment_id), rows
def _call_id(reject: httpx.Response) -> str:
return reject.headers["x-litellm-call-id"]
def _chat_body(model: str, text: str, guardrail: str | None, stream: bool = False) -> dict:
body: dict = {"model": model, "messages": [{"role": "user", "content": text}]}
if stream:
body["stream"] = True
if guardrail is not None:
body["guardrails"] = [guardrail]
return body
def test_cache_hit_post_call_reject_keeps_deployment_labels(gateway: Gateway, tmp_path: Path) -> None:
"""H1: warm then identical post_call-rejected cache hit keeps populated deployment labels."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control h1 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code == 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0)
_assert_spend(_call_id(reject), rig)
def test_cache_hit_post_call_reject_keeps_deployment_labels_openai_sdk(gateway: Gateway, tmp_path: Path) -> None:
"""H2: same as H1 through the openai AsyncOpenAI client."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control h2 " + marker
sdk: Final = openai.AsyncOpenAI(
base_url=str(rig.candidate.client.base_url) + "/v1",
api_key=rig.candidate.key,
http_client=httpx.AsyncClient(trust_env=False, timeout=15),
)
async def run() -> int:
await sdk.chat.completions.create(model=rig.model_name, messages=[{"role": "user", "content": text}])
try:
await sdk.chat.completions.create(
model=rig.model_name,
messages=[{"role": "user", "content": text}],
extra_body={"guardrails": [rig.guardrail_name]},
)
return 200
except openai.BadRequestError:
return 400
assert asyncio.run(run()) == 400
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0)
def test_cache_hit_post_call_reject_streaming(gateway: Gateway, tmp_path: Path) -> None:
"""H3: streamed responses are not cached; the reject call hits upstream again and no failure hook fires."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control h3 " + marker
warm: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None, stream=True)
)
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name, stream=True)
)
assert reject.status_code == 200, reject.text
assert rig.provider.received.qsize() == 2, rig.provider.drain()
samples: Final = _metric_samples(rig.candidate, rig.model_name)
assert _populated_failures(samples, rig, "openai") == 0, samples
assert _count(samples, DEPLOYMENT_FAILURE, "") == 0, samples
def test_cache_hit_post_call_reject_keeps_deployment_labels_anthropic(gateway: Gateway, tmp_path: Path) -> None:
"""H4: /v1/messages cache hit reject through the anthropic SDK."""
marker: Final = uuid.uuid4().hex
with _rig(
gateway, tmp_path, marker, upstream_model="anthropic/claude-sonnet-4-5-20250929", api_base_suffix=""
) as rig:
text: Final = "cache hit control h4 " + marker
sdk: Final = anthropic.Anthropic(
base_url=str(rig.candidate.client.base_url),
api_key=rig.candidate.key,
http_client=httpx.Client(trust_env=False, timeout=15),
)
sdk.messages.create(model=rig.model_name, max_tokens=16, messages=[{"role": "user", "content": text}])
raised: bool = False # mutable-ok: a flag set inside the except block cannot be Final
try:
sdk.messages.create(
model=rig.model_name,
max_tokens=16,
messages=[{"role": "user", "content": text}],
extra_body={"guardrails": [rig.guardrail_name]},
)
except anthropic.BadRequestError:
raised = True
assert raised, "cache-hit post_call guardrail did not reject /v1/messages"
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0, api_provider="anthropic", pf_id="None")
def test_cache_hit_post_call_reject_keeps_deployment_labels_responses(gateway: Gateway, tmp_path: Path) -> None:
"""H5: /v1/responses cache hit reject."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control h5 " + marker
warm: Final = rig.candidate.request("POST", "/v1/responses", {"model": rig.model_name, "input": text})
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST",
"/v1/responses",
{"model": rig.model_name, "input": text, "guardrails": [rig.guardrail_name]},
)
assert reject.status_code == 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0, pf_id="None")
_assert_spend(_call_id(reject), rig)
def test_cache_hit_post_call_reject_embeddings(gateway: Gateway, tmp_path: Path) -> None:
"""H6: post_call guardrails do not run on embeddings; the cached response returns 200 unguarded."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control h6 " + marker
warm: Final = rig.candidate.request("POST", "/v1/embeddings", {"model": rig.model_name, "input": text})
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST",
"/v1/embeddings",
{"model": rig.model_name, "input": text, "guardrails": [rig.guardrail_name]},
)
assert reject.status_code == 200, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
samples: Final = _metric_samples(rig.candidate, rig.model_name)
assert _populated_failures(samples, rig, "openai") == 0, samples
assert _count(samples, DEPLOYMENT_FAILURE, "") == 0, samples
def test_cache_hit_during_call_reject_keeps_deployment_labels(gateway: Gateway, tmp_path: Path) -> None:
"""H7: during_call guardrail reject on a cache hit."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, mode="during_call") as rig:
text: Final = "cache hit control h7 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code == 400, reject.text
_expect_metrics(rig, 1, 0)
def test_pre_call_reject_on_cache_hit_stays_blank(gateway: Gateway, tmp_path: Path) -> None:
"""C1: pre_call reject never reaches the deployment; labels stay blank on both legs."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, mode="pre_call") as rig:
text: Final = "cache hit control c1 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code == 400, reject.text
_expect_metrics(rig, 0, 1, pf_populated=1, pf_blank=0)
def test_post_call_reject_without_cache_keeps_deployment_labels(gateway: Gateway, tmp_path: Path) -> None:
"""C2: a real provider call rejected post_call keeps populated labels on both legs."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "non cache control c2 " + marker
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code == 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0)
_assert_spend(_call_id(reject), rig)
def test_cache_hit_post_call_reject_default_on(gateway: Gateway, tmp_path: Path) -> None:
"""C3: default_on post_call guardrail rejects the cached response (sink passes the warm call)."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, sink=_first_call_pass_sink(), default_on=True) as rig:
text: Final = "cache hit control c3 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert reject.status_code == 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0)
def test_cache_hit_post_call_reject_key_metadata_guardrails(gateway: Gateway, tmp_path: Path) -> None:
"""C4: guardrail attached via key metadata guardrails on a cache hit."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, sink=_first_call_pass_sink()) as rig:
key: Final = rig.candidate.post("/key/generate", {"metadata": {"guardrails": [rig.guardrail_name]}})["key"]
text: Final = "cache hit control c4 " + marker
warm: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None), key=key
)
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None), key=key
)
assert reject.status_code == 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
assert rig.policy.received.qsize() == 2
_expect_metrics(rig, 1, 0)
def test_cache_hit_post_call_reject_local_cache(gateway: Gateway, tmp_path: Path) -> None:
"""C5: same cache-hit reject with cache_params type local."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, local_cache=True) as rig:
text: Final = "cache hit control c5 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code == 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0)
@pytest.mark.parametrize("status", (500, 403))
def test_cache_hit_post_call_guardrail_outage_keeps_deployment_labels(
gateway: Gateway, tmp_path: Path, status: int
) -> None:
"""S1/S2: guardrail sink answers 500/403 on the cache-hit call; failure hook still counts as dispatched."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, sink=_failing_sink(status)) as rig:
text: Final = "cache hit control s " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code >= 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 0, 0, pf_populated=1, pf_blank=0)
def test_two_identical_cache_hit_rejects_increment_populated_series(gateway: Gateway, tmp_path: Path) -> None:
"""E1: two identical cache-hit rejects count +2 on the populated series, two spend rows."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control e1 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
rejects: Final = tuple(
rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name))
for _ in range(2)
)
assert all(response.status_code == 400 for response in rejects), [r.text for r in rejects]
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 2, 0)
def test_two_identical_cache_hit_rejects_write_matching_spend_rows(gateway: Gateway, tmp_path: Path) -> None:
"""E1b: both cache-hit rejects land a failure spend row."""
pytest.skip("BUG: roughly one in four back-to-back cache-hit rejects never lands its LiteLLM_SpendLogs row")
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control e1b " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
rejects: Final = tuple(
rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name))
for _ in range(2)
)
assert all(response.status_code == 400 for response in rejects), [r.text for r in rejects]
for response in rejects:
_assert_spend(_call_id(response), rig)
def test_cache_hit_reject_after_ttl_expiry_is_a_miss(gateway: Gateway, tmp_path: Path) -> None:
"""E2: cache_params ttl=1; post-expiry the same body misses, hits upstream again, labels populated."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, ttl=1) as rig:
text: Final = "cache hit control e2 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
assert rig.provider.received.qsize() == 1
rejects: list[int] = [] # mutable-ok: the poll helper must remember how many rejects it issued
def expired_miss() -> int:
rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name))
rejects.append(1)
return rig.provider.received.qsize()
eventually(lambda: expired_miss() == 2, bool, seconds=70)
_expect_metrics(rig, len(rejects), 0)
def test_cache_hit_reject_metrics_aggregate_across_workers(gateway: Gateway, tmp_path: Path) -> None:
"""E3: workers=2, 8 cache-hit rejects, aggregated /metrics shows +8 on the populated series."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, workers=2) as rig:
text: Final = "cache hit control e3 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
rejects: Final = tuple(
rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name))
for _ in range(8)
)
assert all(response.status_code == 400 for response in rejects), [r.text for r in rejects]
_expect_metrics(rig, 8, 0)
def test_cache_hit_reject_deployment_metric_set_diff(gateway: Gateway, tmp_path: Path) -> None:
"""E4: exact expected label sets on litellm_deployment_* and litellm_proxy_failed_requests_metric."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control e4 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code == 400, reject.text
samples: Final = _expect_metrics(rig, 1, 0)
blank: Final = tuple(sample for sample in samples if sample.labels.get("model_id") == "")
assert blank == (), blank
states: Final = tuple(
sample.value
for sample in samples
if sample.name == DEPLOYMENT_STATE
and sample.labels.get("model_id") == rig.deployment_id
and sample.labels.get("api_base") == ""
)
assert states == (1.0,), states

View file

@ -0,0 +1,348 @@
import json
import signal
import socket
import subprocess
import threading
import uuid
from collections.abc import Callable, Generator
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from pathlib import Path
from typing import Final
import httpx
import psutil
from integration._support.client import Gateway, eventually, object_value, string_value
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
from prometheus_client.parser import text_string_to_metric_families
from test_cache_hit_guardrail_metrics import (
DEPLOYMENT_FAILURE,
GUARDRAIL_PATH,
PROXY_FAILED,
Rig,
_blocking_sink,
_chat_body,
_guardrail_config,
_provider,
_rig,
)
BURST: Final = 10
@contextmanager
def _redis(port: int) -> Generator[subprocess.Popen[bytes], None, None]:
process: Final = subprocess.Popen(["redis-server", "--port", str(port), "--save", ""], stdout=subprocess.DEVNULL)
try:
yield process
finally:
process.kill()
process.wait(timeout=10)
def _free_port() -> int:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
return reserve.getsockname()[1]
def _deployment_id(candidate: Gateway, model_name: str) -> str:
entries: Final = candidate.get("/model/info")["data"]
assert isinstance(entries, list)
entry: Final = next(item for item in entries if object_value(item)["model_name"] == model_name)
return string_value(object_value(object_value(entry)["model_info"])["id"])
def _stall_sink(stall: threading.Event, release: threading.Event) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert request.target == GUARDRAIL_PATH, request.target
if stall.is_set():
release.wait(timeout=60)
return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic block"}).encode())
return respond
def _samples(candidate: Gateway, model_names: tuple[str, ...]) -> tuple:
response: Final = candidate.client.request(
"GET", "/metrics", headers={"Authorization": f"Bearer {candidate.key}"}, follow_redirects=True
)
assert response.status_code == 200, f"GET /metrics: {response.status_code}"
return tuple(
sample
for family in text_string_to_metric_families(response.text)
for sample in family.samples
if sample.labels.get("requested_model") in model_names
)
def _populated(samples: tuple, deployment_id: str) -> float:
return float(
sum(
sample.value
for sample in samples
if sample.name == DEPLOYMENT_FAILURE and sample.labels.get("model_id") == deployment_id
)
)
def _blank(samples: tuple) -> float:
return float(
sum(
sample.value
for sample in samples
if sample.name == DEPLOYMENT_FAILURE and sample.labels.get("model_id") == ""
)
)
def _proxy_failed(samples: tuple) -> float:
return float(sum(sample.value for sample in samples if sample.name == PROXY_FAILED))
def _burst_bodies(rig: Rig, marker: str, anthropic_name: str | None) -> tuple[tuple[str, dict], ...]:
chat: Final = tuple(
("/v1/chat/completions", _chat_body(rig.model_name, f"burst {marker} {index}", rig.guardrail_name))
for index in range(BURST)
)
responses: Final = tuple(
(
"/v1/responses",
{"model": rig.model_name, "input": f"burst {marker} r{index}", "guardrails": [rig.guardrail_name]},
)
for index in range(BURST)
)
messages: Final = (
tuple(
(
"/v1/messages",
{
"model": anthropic_name,
"max_tokens": 16,
"messages": [{"role": "user", "content": f"burst {marker} m{index}"}],
"guardrails": [rig.guardrail_name],
},
)
for index in range(BURST)
)
if anthropic_name is not None
else ()
)
return chat + responses + messages
def _warm(rig: Rig, bodies: tuple[tuple[str, dict], ...]) -> None:
for path, body in bodies:
warmed: Final = dict(body)
warmed.pop("guardrails", None)
response: Final = rig.candidate.request("POST", path, warmed)
assert response.status_code == 200, f"warm {path}: {response.status_code} {response.text}"
def _fire(rig: Rig, bodies: tuple[tuple[str, dict], ...]) -> tuple[tuple[int, str | None], ...]:
def call(item: tuple[str, dict]) -> tuple[int, str | None]:
path, body = item
try:
response: Final = rig.candidate.request("POST", path, body)
return response.status_code, response.headers.get("x-litellm-call-id")
except httpx.HTTPError:
return -1, None
with ThreadPoolExecutor(max_workers=8) as pool:
return tuple(pool.map(call, bodies))
def _expect_counted_within(
rig: Rig, model_names: tuple[str, ...], deployment_ids: tuple[str, ...], low: int, high: int
) -> None:
def converged() -> tuple:
samples: Final = _samples(rig.candidate, model_names)
populated: Final = sum(_populated(samples, deployment) for deployment in deployment_ids)
if low <= populated <= high and _blank(samples) == 0:
return samples
return ()
eventually(converged, bool, seconds=70)
def _expect_exactly_once(rig: Rig, model_names: tuple[str, ...], deployment_ids: tuple[str, ...], four_xx: int) -> None:
_expect_counted_within(rig, model_names, deployment_ids, four_xx, four_xx)
def test_burst_cache_hit_rejects_count_exactly_once(gateway: Gateway, tmp_path: Path) -> None:
"""X0: 30 mixed-endpoint cache-hit rejects across two deployments, each counted once."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
anthropic_name: Final = rig.scenario.model(
model="anthropic/claude-sonnet-4-5-20250929", api_base=rig.provider.url, api_key="synthetic-provider-key"
)
anthropic_id: Final = _deployment_id(rig.candidate, anthropic_name)
bodies: Final = _burst_bodies(rig, marker, anthropic_name)
_warm(rig, bodies)
outcomes: Final = _fire(rig, bodies)
rejected: Final = sum(1 for status, _ in outcomes if status >= 400)
assert all(status == 400 for status, _ in outcomes), outcomes
_expect_exactly_once(rig, (rig.model_name, anthropic_name), (rig.deployment_id, anthropic_id), rejected)
def test_stalled_guardrail_sink_recovers_and_counts(gateway: Gateway, tmp_path: Path) -> None:
"""X1: guardrail sink stalls mid-burst; requests fail exactly once, then recovery counts again."""
marker: Final = uuid.uuid4().hex
stall: Final = threading.Event()
release: Final = threading.Event()
with _rig(gateway, tmp_path, marker, sink=_stall_sink(stall, release)) as rig:
bodies: Final = _burst_bodies(rig, marker, None)
_warm(rig, bodies)
stall.set()
with ThreadPoolExecutor(max_workers=8) as pool:
futures: Final = tuple(
pool.submit(lambda b: rig.candidate.request("POST", b[0], b[1]), body) for body in bodies
)
eventually(lambda: rig.policy.received.qsize() >= 5, bool, seconds=30)
release.set()
outcomes: Final = tuple(
(future.result().status_code, future.result().headers.get("x-litellm-call-id")) for future in futures
)
assert all(status >= 400 for status, _ in outcomes), outcomes
blocked: Final = sum(1 for status, _ in outcomes if status == 400)
outages: Final = sum(1 for status, _ in outcomes if status >= 500)
assert blocked + outages == len(bodies), outcomes
samples: Final = eventually(
lambda: _samples(rig.candidate, (rig.model_name,)),
lambda observed: _proxy_failed(observed) == blocked + outages,
seconds=70,
)
assert _proxy_failed(samples) == blocked + outages, (samples, outcomes)
follow_up: Final = rig.candidate.request(
"POST",
"/v1/chat/completions",
_chat_body(rig.model_name, "post stall unrelated " + marker, None),
)
assert follow_up.status_code == 200, follow_up.text
_expect_exactly_once(rig, (rig.model_name,), (rig.deployment_id,), blocked)
def test_redis_outage_keeps_serving_in_memory_hits(gateway: Gateway, tmp_path: Path) -> None:
"""X2: the redis cache keeps an in-memory shadow, so a redis kill does not stop cache-hit rejects."""
marker: Final = uuid.uuid4().hex
port: Final = _free_port()
with _redis(port) as redis_one:
with _rig(gateway, tmp_path, marker, env={"REDIS_HOST": "127.0.0.1", "REDIS_PORT": str(port)}) as rig:
bodies: Final = _burst_bodies(rig, marker, None)[:BURST]
_warm(rig, bodies)
reject: Final = rig.candidate.request("POST", *bodies[0])
assert reject.status_code == 400, reject.text
warmed_hits: Final = rig.provider.received.qsize()
redis_one.kill()
redis_one.wait(timeout=10)
outcomes: Final = _fire(rig, bodies[1:])
assert all(status == 400 for status, _ in outcomes), outcomes
assert rig.provider.received.qsize() == warmed_hits, (
"redis outage reached the provider",
warmed_hits,
rig.provider.received.qsize(),
)
with _redis(port):
recovered: Final = rig.candidate.request(
"POST",
"/v1/chat/completions",
_chat_body(rig.model_name, "x2 rehit " + marker, rig.guardrail_name),
)
assert recovered.status_code == 400, recovered.text
_expect_exactly_once(rig, (rig.model_name,), (rig.deployment_id,), 1 + len(bodies))
def test_worker_kill_mid_burst_keeps_counting(gateway: Gateway, tmp_path: Path) -> None:
"""X3: workers=2, SIGKILL one uvicorn child mid-burst; survivors keep rejecting; the count is answered plus at most the in-flight requests the killed worker had already counted."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, workers=2) as rig:
bodies: Final = _burst_bodies(rig, marker, None)
_warm(rig, bodies)
children: Final = psutil.Process(rig.process.pid).children(recursive=True)
assert children, "no uvicorn worker children found"
with ThreadPoolExecutor(max_workers=8) as pool:
futures: Final = tuple(
pool.submit(lambda b: rig.candidate.request("POST", b[0], b[1]), body) for body in bodies
)
eventually(lambda: rig.policy.received.qsize() >= 3, bool, seconds=30)
children[0].send_signal(signal.SIGKILL)
statuses: list[int] = [] # mutable-ok: collect per-request outcomes from concurrent futures
for future in futures:
try:
statuses.append(future.result().status_code)
except httpx.HTTPError:
statuses.append(-1)
answered: Final = sum(1 for status in statuses if status >= 0)
transport_lost: Final = sum(1 for status in statuses if status == -1)
assert all(status == 400 for status in statuses if status >= 0), (
statuses,
transport_lost,
)
_expect_counted_within(rig, (rig.model_name,), (rig.deployment_id,), answered, answered + transport_lost)
def test_proxy_restart_mid_burst_keeps_counting(gateway: Gateway, tmp_path: Path) -> None:
"""X4: restart the owned proxy between the two halves; pre-restart count asserted, then recounted."""
marker: Final = uuid.uuid4().hex
prom_dir: Final = tmp_path / "prom"
prom_dir.mkdir()
with wire_server(_blocking_sink) as policy, wire_server(_provider(marker)) as provider:
config: Final = _guardrail_config(tmp_path, "guardrail-" + marker, policy.url)
bodies: Final = tuple(
(
"/v1/chat/completions",
_chat_body("pending-model", f"burst {marker} {index}", "guardrail-" + marker),
)
for index in range(BURST)
)
with owned_proxy_process(
gateway, tmp_path, {"PROMETHEUS_MULTIPROC_DIR": str(prom_dir)}, config=config
) as owned_one:
model: Final = "restart-" + marker
owned_one.gateway.post(
"/model/new",
{
"model_name": model,
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_base": provider.url + "/v1",
"api_key": "synthetic-provider-key",
},
},
)
deployment: Final = _deployment_id(owned_one.gateway, model)
named: Final = tuple((path, {**body, "model": model}) for path, body in bodies)
first_half, second_half = named[: BURST // 2], named[BURST // 2 :]
for path, body in named:
warmed: Final = dict(body)
warmed.pop("guardrails", None)
assert owned_one.gateway.request("POST", path, warmed).status_code == 200
outcomes_one: Final = tuple(owned_one.gateway.request("POST", path, body) for path, body in first_half)
assert all(response.status_code == 400 for response in outcomes_one), [r.text for r in outcomes_one]
pre: Final = eventually(
lambda: (
_populated(_samples(owned_one.gateway, (model,)), deployment),
_blank(_samples(owned_one.gateway, (model,))),
),
lambda observed: observed[0] == len(first_half) and observed[1] == 0,
seconds=70,
)
with owned_proxy_process(
gateway, tmp_path, {"PROMETHEUS_MULTIPROC_DIR": str(prom_dir)}, config=config
) as owned_two:
outcomes_two: Final = tuple(owned_two.gateway.request("POST", path, body) for path, body in second_half)
assert all(response.status_code == 400 for response in outcomes_two), (
pre,
[(r.status_code, r.text[:200]) for r in outcomes_two],
)
post: Final = eventually(
lambda: (
_populated(_samples(owned_two.gateway, (model,)), deployment),
_blank(_samples(owned_two.gateway, (model,))),
),
lambda observed: observed[0] == len(named) and observed[1] == 0,
seconds=70,
)
assert post[0] == len(named), (pre, post, outcomes_two)
owned_two.gateway.post("/model/delete", {"id": deployment})

View file

@ -18,6 +18,7 @@ from litellm.exceptions import GuardrailRaisedException
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import ProxyErrorTypes, UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.types.utils import CachingDetails
@pytest.fixture(autouse=True)
@ -249,6 +250,65 @@ async def test_post_call_failure_hook_keeps_router_stamped_metadata_for_post_cal
assert kwargs["standard_logging_object"]["model_id"] == "routed-deployment"
@pytest.mark.asyncio
async def test_post_call_failure_hook_keeps_deployment_attribution_for_cache_hit_post_call_failures(
proxy_logging, make_user_api_key_auth, monkeypatch
):
"""A post-call guardrail blocks a response served from the litellm cache. No provider call was made,
so ``first_api_call_start_time`` is unset, but the router did pick the deployment: the pre-routing
flag must stay off so ``litellm_deployment_failure_responses`` keeps its model_id and provider labels."""
from litellm.proxy import proxy_server
recorded: list[dict] = []
class _RecordingLogger(CustomLogger):
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
recorded.append(kwargs)
monkeypatch.setattr(
proxy_server,
"llm_router",
litellm.Router(
model_list=[
{
"model_name": "internal-model",
"litellm_params": {"model": "openai/gpt-4.1", "api_key": "sk-test"},
"model_info": {"id": "routed-deployment"},
}
]
),
)
monkeypatch.setattr(litellm, "callbacks", [_RecordingLogger()])
proxy_logging.alert_types = []
request_data = {
"litellm_call_id": "cache-hit-post-call-guardrail",
"model": "internal-model",
"messages": [{"role": "user", "content": "hi"}],
"metadata": {"model_info": {"id": "routed-deployment"}},
}
logging_obj, request_data = litellm.utils.function_setup(
original_function="acompletion", rules_obj=litellm.utils.Rules(), start_time=datetime.now(), **request_data
)
logging_obj.caching_details = CachingDetails(cache_hit=True, cache_duration_ms=1.0)
request_data["litellm_logging_obj"] = logging_obj
await proxy_logging.post_call_failure_hook(
request_data=request_data,
original_exception=GuardrailRaisedException(guardrail_name="g", message="response blocked"),
user_api_key_dict=make_user_api_key_auth(request_route="/chat/completions"),
route="/chat/completions",
)
assert len(recorded) == 1
kwargs = recorded[0]
assert PROXY_REJECTED_BEFORE_ROUTING_KEY not in kwargs["litellm_params"], kwargs["litellm_params"]
assert kwargs["standard_logging_object"]["model_id"] == "routed-deployment"
assert kwargs["standard_logging_object"]["custom_llm_provider"] == "openai"
assert kwargs["model"] == "internal-model"
assert kwargs["litellm_params"]["custom_llm_provider"] == "openai"
@pytest.mark.asyncio
async def test_post_call_failure_hook_flags_pre_routing_reject_despite_caller_model_info(
proxy_logging, make_user_api_key_auth, monkeypatch