fix(caching): never store or serve a chat completion with no choices (#44709)

* fix(caching): never store or serve a chat completion with no choices

A provider response with empty choices was written to the response cache and served on every identical request until the TTL ended, with no provider call in between. The cache now skips storing such a response and treats an already stored one as a miss, so the next request goes back to the provider and its answer replaces the entry.

* fix(caching): skip responses with no output on the Responses API and Anthropic Messages too

* fix(caching): skip streams and stored entries that carry no output

A chat or text completion stream whose chunks carried no choice is closed
by the stream wrapper with one empty choice of its own, so the assembled
response passed the choices check and was cached. The assembled stream is
now judged on its content: a stream with no text, tool call, or other
output in any choice is never stored, on the async and sync writers alike.
The Responses API stream writer and the Anthropic Messages stream writer
apply the same no-output check before storing.

A stored entry with no output read through the worker memory tier is now
evicted from that tier on the miss, so the next read reaches Redis where
the refill lands; the text completion and messages writers only write to
Redis, and the memory copy otherwise kept missing until its own TTL.

* test(integration): response cache cells for answers without output

Deterministic cells for the response cache on every unified endpoint,
streamed and not, through the OpenAI and Anthropic SDKs and raw httpx,
plus the sync SDK paths, stale entries, malformed answers, per-request
TTLs, cache delete, and chaos (Redis stopped or paused mid burst, a
worker killed, in-memory cache mode). The scripted upstream counts only
POSTs as deployment calls, since the proxy's boot-time GET /v1/models
discovery of a config deployment is not one.

* test(caching): pin the stored entry timestamp in the worker-copy test

* test(integration): drop the restating comments from the chaos cells

* test(integration): close the breaker on the first call after the Redis restart

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-06 20:07:30 +00:00 • committed by GitHub
parent 5ed7ec8511
commit abd3af422d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 2298 additions and 24 deletions

View file

@ -35,7 +35,7 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
from litellm.litellm_core_utils.logging_utils import (
_assemble_complete_response_from_streaming_chunks,
)
from litellm.types.caching import EMBEDDING_CACHE_FORMAT_VERSION, CachedEmbedding
from litellm.types.caching import CACHED_STREAM_EVENTS_KEY, EMBEDDING_CACHE_FORMAT_VERSION, CachedEmbedding
from litellm.types.integrations.custom_logger import converted_stream_requested
from litellm.types.llms.base import LiteLLMBaseModel
from litellm.types.llms.openai import ResponsesAPIResponse
@ -43,9 +43,11 @@ from litellm.types.rerank import RerankResponse
from litellm.types.utils import (
CachingDetails,
CallTypes,
Choices,
Embedding,
EmbeddingResponse,
ModelResponse,
TextChoices,
TextCompletionResponse,
TranscriptionResponse,
Usage,
@ -111,6 +113,43 @@ def _is_chat_completion_cached_dict(cached_result: dict) -> bool:
return "choices" in cached_result
def _starts_content_block(event: object) -> bool:
text: Final = event.decode("utf-8", errors="replace") if isinstance(event, bytes) else event
return isinstance(text, str) and "event: content_block_start" in text.splitlines()
def is_response_without_output(result: object) -> bool:
if isinstance(result, (ModelResponse, TextCompletionResponse)):
return not result.choices
if isinstance(result, ResponsesAPIResponse):
return not getattr(result, "output", None)
if not isinstance(result, dict):
return False
cached_stream_events: Final = result.get(CACHED_STREAM_EVENTS_KEY)
if isinstance(cached_stream_events, (list, tuple)):
return not any(_starts_content_block(event) for event in cached_stream_events)
if "choices" in result:
return not result["choices"]
if result.get("object") == "response":
return not result.get("output")
if result.get("type") == "message":
return not result.get("content")
return False
def _choice_carries_output(choice: Choices | TextChoices) -> bool:
payload: Final[Mapping[str, object]] = (
choice.message.model_dump(exclude_none=True) if isinstance(choice, Choices) else {"text": choice.text}
)
return any(key != "role" and value not in ("", [], {}) for key, value in payload.items())
def _assembled_stream_without_output(response: ModelResponse | TextCompletionResponse) -> bool:
"""The stream wrapper closes a stream whose chunks carried no choice with an empty choice of its own, so the
assembled response has choices and the no-output check on it alone would store that empty answer."""
return not any(_choice_carries_output(choice) for choice in response.choices)
def _stream_replay_requested(kwargs: Mapping[str, object]) -> bool:
if kwargs.get("stream", False) is True:
return True
@ -398,6 +437,9 @@ class LLMCachingHandler:
print_verbose("Checking Sync Cache")
with response_cache_phase("get"):
cached_result = litellm.cache.get_cache(**new_kwargs)
if is_response_without_output(cached_result):
verbose_logger.debug("LiteLLM Cache: cached response has no output, treating it as a miss")
return CachingHandlerResponse(cached_result=None)
if cached_result is not None:
if "detail" in cached_result:
# implies an error occurred
@ -836,8 +878,19 @@ class LLMCachingHandler:
cache_key=self.preset_cache_key,
**request_kwargs,
)
if is_response_without_output(cached_result):
verbose_logger.debug("LiteLLM Cache: cached response has no output, treating it as a miss")
self._forget_worker_copy()
return None
return cached_result
def _forget_worker_copy(self) -> None:
"""Drop this worker's memory copy of a stored response without output, so the next read reaches Redis,
where the refill lands: the text completion and messages writers skip the memory tier."""
if self.dual_cache is None or self.preset_cache_key is None:
return
self.dual_cache.in_memory_cache.delete_cache(self.preset_cache_key)
def _convert_cached_result_to_model_response(
self,
cached_result: Any,
@ -1063,6 +1116,9 @@ class LLMCachingHandler:
if litellm.cache is None:
return
if is_response_without_output(result):
verbose_logger.debug("LiteLLM Cache: not caching a response with no output")
return
cache: Final = litellm.cache
new_kwargs: Final = kwargs.copy()
@ -1113,6 +1169,11 @@ class LLMCachingHandler:
"""
Sync internal method to add the result to the cache
"""
if litellm.cache is None:
return
if is_response_without_output(result):
verbose_logger.debug("LiteLLM Cache: not caching a response with no output")
return
new_kwargs: Final = kwargs.copy()
new_kwargs.update(
@ -1121,8 +1182,6 @@ class LLMCachingHandler:
args,
)
)
if litellm.cache is None:
return
if self._should_store_result_in_cache(original_function=self.original_function, kwargs=new_kwargs):
with response_cache_phase("set"):
@ -1201,13 +1260,13 @@ class LLMCachingHandler:
is_async=True,
)
)
# if a complete_streaming_response is assembled, add it to the cache
if complete_streaming_response is not None:
await self.async_set_cache(
result=complete_streaming_response,
original_function=self.original_function,
kwargs=self.request_kwargs,
)
if complete_streaming_response is None or _assembled_stream_without_output(complete_streaming_response):
return
await self.async_set_cache(
result=complete_streaming_response,
original_function=self.original_function,
kwargs=self.request_kwargs,
)
def _sync_add_streaming_response_to_cache(self, processed_chunk: ModelResponse):
"""
@ -1224,12 +1283,12 @@ class LLMCachingHandler:
)
)
# if a complete_streaming_response is assembled, add it to the cache
if complete_streaming_response is not None:
self.sync_set_cache(
result=complete_streaming_response,
kwargs=self.request_kwargs,
)
if complete_streaming_response is None or _assembled_stream_without_output(complete_streaming_response):
return
self.sync_set_cache(
result=complete_streaming_response,
kwargs=self.request_kwargs,
)
def _update_litellm_logging_obj_environment(
self,

View file

@ -5,7 +5,7 @@ from typing import TYPE_CHECKING, Final, cast
import litellm
from litellm._logging import verbose_logger
from litellm.caching.caching_handler import create_cache_write_task
from litellm.caching.caching_handler import create_cache_write_task, is_response_without_output
from litellm.llms.anthropic.pass_through.messages.streaming_iterator import (
AnthropicMessagesStreamingResponse,
BaseAnthropicMessagesStreamingIterator,
@ -13,6 +13,7 @@ from litellm.llms.anthropic.pass_through.messages.streaming_iterator import (
_is_provider_error_chunk,
aclose_if_supported,
)
from litellm.types.caching import CACHED_STREAM_EVENTS_KEY
if TYPE_CHECKING:
from litellm.caching.caching_handler import LLMCachingHandler
@ -20,8 +21,6 @@ if TYPE_CHECKING:
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponseStream
CACHED_STREAM_EVENTS_KEY: Final = "litellm_cached_anthropic_sse_events"
_EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
_SSE_EVENT_BOUNDARY: Final = re.compile(r"(?<=\n\n)")
@ -114,6 +113,9 @@ class AnthropicMessagesStreamCacheWriter:
verbose_logger.exception("Anthropic Messages stream cache write failed: %s", e)
return
cached_payload: Final = {CACHED_STREAM_EVENTS_KEY: events}
if is_response_without_output(cached_payload):
verbose_logger.debug("LiteLLM Cache: not caching a stream with no content blocks")
return
dual_cache: Final = self.caching_handler.dual_cache
async def _write() -> None:

View file

@ -691,8 +691,10 @@ class BaseResponsesAPIStreamingIterator:
if getattr(completed_response, "type", None) != openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
return
from litellm.caching.caching_handler import create_cache_write_task, is_response_without_output
response_obj: Final = self._get_completed_response_object()
if response_obj is None:
if response_obj is None or is_response_without_output(response_obj):
return
caching_handler: Final[LLMCachingHandler | None] = getattr(self.logging_obj, "_llm_caching_handler", None)
@ -732,8 +734,6 @@ class BaseResponsesAPIStreamingIterator:
if cached_response is None:
return
if is_async:
from litellm.caching.caching_handler import create_cache_write_task
cache_write_task: Final = create_cache_write_task(
lambda: cache.async_add_cache(
cached_response,

View file

@ -140,6 +140,8 @@ class HealthCheckCacheParams(LiteLLMBaseModel):
EMBEDDING_CACHE_FORMAT_VERSION: Final = 2
CACHED_STREAM_EVENTS_KEY: Final = "litellm_cached_anthropic_sse_events"
class CachedEmbedding(TypedDict):
"""Type definition for cached embedding objects"""

View file

@ -18,6 +18,7 @@ OWNED_DIRECTORIES: Final = frozenset(
"compatibility",
"sdk",
"cost_calculation",
"caching",
"security",
}
)

View file

@ -0,0 +1,364 @@
"""Scripted provider answers, upstream observations and Redis helpers shared by the response cache cells.
The bodies carry ``$UNIQUE_ID`` so every upstream answer has its own id: an answer served from the
cache repeats the id of the request that filled it, an answer that reached the upstream again does not.
"""
from __future__ import annotations
import ast
import json
import os
import re
import uuid
from collections.abc import Iterator, Mapping
from contextlib import contextmanager
from typing import Final
import httpx
from pydantic import JsonValue
from redis import Redis
from tests.integration._support.client import JSON_OBJECT, Scenario, eventually, object_value, string_value
from tests.integration._support.upstream import CONTROL_URL, ScenarioHandle, delete_scenario, register_scenario
from tests.integration.cost_calculation.cost_tracking_case import JsonResponse, SseResponse
CACHE_KEY: Final = re.compile(r"^[0-9a-f]{64}$")
PROVIDER_KEY: Final = "integration-provider-key"
CHAT_MODEL: Final = "openai/gpt-4o-mini"
TEXT_MODEL: Final = "openai/gpt-3.5-turbo-instruct"
MESSAGES_MODEL: Final = "anthropic/claude-haiku-4-5"
CHAT: Final[dict[str, JsonValue]] = {
"id": "chatcmpl-$UNIQUE_ID",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "scripted"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
}
TEXT: Final[dict[str, JsonValue]] = {
"id": "cmpl-$UNIQUE_ID",
"object": "text_completion",
"created": 1,
"model": "gpt-3.5-turbo-instruct",
"choices": [{"text": "scripted", "index": 0, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
}
RESPONSE: Final[dict[str, JsonValue]] = {
"id": "resp_$UNIQUE_ID",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o-mini",
"output": [
{
"id": "msg_$UNIQUE_ID",
"type": "message",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "scripted", "annotations": []}],
}
],
"usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
}
MESSAGE: Final[dict[str, JsonValue]] = {
"id": "msg-$UNIQUE_ID",
"type": "message",
"role": "assistant",
"model": "claude-haiku-4-5",
"content": [{"type": "text", "text": "scripted"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 1, "output_tokens": 1},
}
EMBEDDING: Final[dict[str, JsonValue]] = {
"object": "list",
"data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}],
"model": "text-embedding-$UNIQUE_ID",
"usage": {"prompt_tokens": 1, "total_tokens": 1},
}
RERANK: Final[dict[str, JsonValue]] = {
"id": "rerank-$UNIQUE_ID",
"results": [{"index": 0, "relevance_score": 0.5}],
"meta": {},
}
TRANSCRIPT: Final[dict[str, JsonValue]] = {"text": "scripted $UNIQUE_ID"}
CHAT_CHUNK: Final = (
'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","created":1,"model":"gpt-4o-mini","choices":'
)
CHAT_STREAM: Final = SseResponse(
content_type="text/event-stream",
frames=(
CHAT_CHUNK + '[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}',
CHAT_CHUNK + '[{"index":0,"delta":{"content":"streamed "},"finish_reason":null}]}',
CHAT_CHUNK + '[{"index":0,"delta":{"content":"response"},"finish_reason":null}]}',
CHAT_CHUNK + '[{"index":0,"delta":{},"finish_reason":"stop"}]}',
"data: [DONE]",
),
)
CHAT_STREAM_WITHOUT_CHOICES: Final = SseResponse(
content_type="text/event-stream",
frames=(CHAT_CHUNK + "[]}", "data: [DONE]"),
)
TEXT_CHUNK: Final = (
'data: {"id":"cmpl-$UNIQUE_ID","object":"text_completion","created":1,"model":"gpt-3.5-turbo-instruct","choices":'
)
TEXT_STREAM: Final = SseResponse(
content_type="text/event-stream",
frames=(
TEXT_CHUNK + '[{"text":"streamed ","index":0,"finish_reason":null}]}',
TEXT_CHUNK + '[{"text":"response","index":0,"finish_reason":null}]}',
TEXT_CHUNK + '[{"text":"","index":0,"finish_reason":"stop"}]}',
"data: [DONE]",
),
)
TEXT_STREAM_WITHOUT_CHOICES: Final = SseResponse(
content_type="text/event-stream",
frames=(TEXT_CHUNK + "[]}", "data: [DONE]"),
)
RESPONSE_CREATED: Final = (
"event: response.created\n"
'data: {"type":"response.created","response":{"id":"resp_$UNIQUE_ID","object":"response",'
'"created_at":1,"status":"in_progress","model":"gpt-4o-mini","output":[],"usage":null}}'
)
RESPONSE_STREAM: Final = SseResponse(
content_type="text/event-stream",
frames=(
RESPONSE_CREATED,
(
"event: response.output_item.added\n"
'data: {"type":"response.output_item.added","output_index":0,'
'"item":{"type":"message","id":"msg_$UNIQUE_ID","status":"in_progress","role":"assistant","content":[]}}'
),
(
"event: response.content_part.added\n"
'data: {"type":"response.content_part.added","item_id":"msg_$UNIQUE_ID","output_index":0,'
'"content_index":0,"part":{"type":"output_text","text":"","annotations":[]}}'
),
(
"event: response.output_text.delta\n"
'data: {"type":"response.output_text.delta","item_id":"msg_$UNIQUE_ID","output_index":0,'
'"content_index":0,"delta":"streamed response"}'
),
(
"event: response.output_text.done\n"
'data: {"type":"response.output_text.done","item_id":"msg_$UNIQUE_ID","output_index":0,'
'"content_index":0,"text":"streamed response"}'
),
(
"event: response.content_part.done\n"
'data: {"type":"response.content_part.done","item_id":"msg_$UNIQUE_ID","output_index":0,'
'"content_index":0,"part":{"type":"output_text","text":"streamed response","annotations":[]}}'
),
(
"event: response.output_item.done\n"
'data: {"type":"response.output_item.done","output_index":0,'
'"item":{"type":"message","id":"msg_$UNIQUE_ID","status":"completed","role":"assistant",'
'"content":[{"type":"output_text","text":"streamed response","annotations":[]}]}}'
),
(
"event: response.completed\n"
'data: {"type":"response.completed","response":{"id":"resp_$UNIQUE_ID","object":"response",'
'"created_at":1,"status":"completed","model":"gpt-4o-mini",'
'"output":[{"id":"msg_$UNIQUE_ID","type":"message","status":"completed","role":"assistant",'
'"content":[{"type":"output_text","text":"streamed response","annotations":[]}]}],'
'"usage":{"input_tokens":1,"output_tokens":2,"total_tokens":3}}}'
),
),
)
RESPONSE_STREAM_WITHOUT_OUTPUT: Final = SseResponse(
content_type="text/event-stream",
frames=(
RESPONSE_CREATED,
(
"event: response.completed\n"
'data: {"type":"response.completed","response":{"id":"resp_$UNIQUE_ID","object":"response",'
'"created_at":1,"status":"completed","model":"gpt-4o-mini","output":[],'
'"usage":{"input_tokens":1,"output_tokens":0,"total_tokens":1}}}'
),
),
)
MESSAGE_START: Final = (
"event: message_start\n"
'data: {"type":"message_start","message":{"id":"msg_$UNIQUE_ID","type":"message","role":"assistant",'
'"model":"claude-haiku-4-5","content":[],"stop_reason":null,"stop_sequence":null,'
'"usage":{"input_tokens":1,"output_tokens":0}}}'
)
MESSAGE_END: Final = (
(
"event: message_delta\n"
'data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},'
'"usage":{"output_tokens":2}}'
),
'event: message_stop\ndata: {"type":"message_stop"}',
)
MESSAGE_STREAM: Final = SseResponse(
content_type="text/event-stream",
frames=(
MESSAGE_START,
(
"event: content_block_start\n"
'data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}'
),
(
"event: content_block_delta\n"
'data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"streamed response"}}'
),
'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}',
*MESSAGE_END,
),
)
MESSAGE_STREAM_WITHOUT_CONTENT_BLOCKS: Final = SseResponse(
content_type="text/event-stream", frames=(MESSAGE_START, *MESSAGE_END)
)
class Upstream:
"""The owned upstream's observations, drained once per read and kept, filtered per deployment."""
def __init__(self, url: str = CONTROL_URL) -> None:
self._url: Final = url.rstrip("/")
self._seen: tuple[dict[str, JsonValue], ...] = ()
def received(self, identity: str) -> tuple[dict[str, JsonValue], ...]:
with httpx.Client(timeout=10, trust_env=False) as client:
payload: Final = JSON_OBJECT.validate_json(client.get(f"{self._url}/__observations").content)
requests: Final = payload["requests"]
assert isinstance(requests, list)
self._seen = (
*self._seen,
*(object_value(item) for item in requests),
) # rebind-ok: drained observations accumulate
return tuple(item for item in self._seen if f"/{identity}/" in string_value(item.get("path", "")))
def calls(self, identity: str) -> int:
"""The LLM calls the deployment took; the proxy's boot-time ``GET /v1/models`` discovery of a config
deployment is not one."""
return sum(1 for item in self.received(identity) if item.get("method") == "POST")
def json_response(body: Mapping[str, JsonValue], status: int = 200) -> JsonResponse:
return JsonResponse(content_type="application/json", body=dict(body), status=status)
def slowed(stream: SseResponse, frame_delay_ms: int) -> SseResponse:
return stream.model_copy(update={"frame_delay_ms": frame_delay_ms})
def emptied(body: Mapping[str, JsonValue], field: str) -> dict[str, JsonValue]:
return {**body, field: []}
def scenario_id() -> str:
return f"cache-{uuid.uuid4().hex[:12]}"
def scripted(scenario: Scenario, response: JsonResponse | SseResponse) -> ScenarioHandle:
handle: Final = register_scenario(scenario_id(), response, control_url=scenario.gateway.upstream_url)
scenario.cleanups.callback(delete_scenario, handle)
return handle
@contextmanager
def scripted_in_process(response: JsonResponse | SseResponse) -> Iterator[ScenarioHandle]:
handle: Final = register_scenario(scenario_id(), response)
try:
yield handle
finally:
delete_scenario(handle)
def rescript(handle: ScenarioHandle, response: JsonResponse | SseResponse) -> None:
register_scenario(handle.scenario_id, response, control_url=handle.control_url)
def prompt() -> str:
return f"cache probe {uuid.uuid4().hex}"
def redis_store() -> Redis:
return Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]))
def entry_key(store: Redis, needle: str) -> str | None:
"""The response cache key (a bare sha256) whose stored value names ``needle``, or ``None``."""
for raw_key in store.scan_iter(count=1000):
key: Final = raw_key.decode()
if not CACHE_KEY.match(key) or store.type(key) != b"string":
continue
value: Final = store.get(key)
if value is not None and needle.encode() in value:
return key
return None
def await_entry(store: Redis, needle: str) -> str:
key: Final = eventually(lambda: entry_key(store, needle), lambda found: found is not None)
assert key is not None
return key
def _decode_entry(raw: bytes) -> tuple[dict[str, JsonValue], bool]:
"""The stored entry and whether it is JSON; the sync SDK path stores ``str(dict)`` instead."""
text: Final = raw.decode()
if text.startswith('{"'):
return object_value(json.loads(text)), True
return object_value(ast.literal_eval(text)), False
SHORT_TTL_SECONDS: Final = 2
SHORT_TTL_CACHE_CONTROL: Final[dict[str, JsonValue]] = {"ttl": SHORT_TTL_SECONDS}
STALE_ENTRY_TTL_SECONDS: Final = 60
STALE_ID_SUFFIX: Final = "-stale"
def flip_entry_to_empty(store: Redis, key: str, field: str) -> str:
"""Rewrite the stored response so ``field`` is ``[]`` and its id carries ``-stale``, keeping the encoding.
Returns the stale id. The proxy worker that wrote the entry keeps a copy in its own memory for the entry's ttl,
so a cell fills the entry with ``SHORT_TTL_CACHE_CONTROL`` and reads the rewritten entry once that copy lapsed.
"""
raw: Final = store.get(key)
assert raw is not None, key
entry, is_json = _decode_entry(raw)
response: Final = entry["response"]
stored: Final = object_value(json.loads(response) if isinstance(response, str) else response)
assert field in stored, f"{field!r} missing from the cached response {sorted(stored)}"
stale_id: Final = f"{string_value(stored['id'])}{STALE_ID_SUFFIX}"
flipped: Final[dict[str, JsonValue]] = {**stored, field: [], "id": stale_id}
rewritten: Final[dict[str, JsonValue]] = {
**entry,
"response": json.dumps(flipped) if isinstance(response, str) else flipped,
}
store.set(key, json.dumps(rewritten) if is_json else str(rewritten), ex=STALE_ENTRY_TTL_SECONDS)
return stale_id
def chat_body(model: str, text: str, **extra: JsonValue) -> dict[str, JsonValue]:
return {"model": model, "messages": [{"role": "user", "content": text}], **extra}
def message_body(model: str, text: str, **extra: JsonValue) -> dict[str, JsonValue]:
return {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": text}], **extra}
def response_id(response: httpx.Response) -> str:
assert response.status_code == 200, response.text
return string_value(JSON_OBJECT.validate_json(response.content)["id"])
def choices(response: httpx.Response) -> list[JsonValue]:
assert response.status_code == 200, response.text
found: Final = JSON_OBJECT.validate_json(response.content)["choices"]
assert isinstance(found, list), response.text
return found
def content_of_first_choice(response: httpx.Response) -> JsonValue:
return object_value(object_value(choices(response)[0])["message"])["content"]
def text_of_first_choice(response: httpx.Response) -> JsonValue:
return object_value(choices(response)[0])["text"]

View file

@ -0,0 +1,503 @@
"""Chaos rows for the response cache on an owned two-worker proxy.
Each owned proxy carries its own deployments pointed at scripted upstream scenarios, so both workers know
them from boot. A 24-request burst (chat, text completions, Responses and Anthropic Messages, two of each
streamed at 500 ms per frame so they are still in flight when the outage lands) runs while the cache store
is stopped (C1), paused (C2) or while one worker has just been SIGKILLed (C3): every request answers 200
with content and lands exactly once in the spend log. After recovery an upstream answer with no choices
still reaches the upstream again on the next identical request. C4 runs the in-memory cache mode, where
each worker keeps its own store and the cached object is read back without a JSON round trip.
"""
from __future__ import annotations
import os
import threading
import uuid
from collections.abc import Callable, Iterator, Sequence
from concurrent.futures import Future, ThreadPoolExecutor
from dataclasses import dataclass
from hashlib import sha256
from pathlib import Path
from typing import Final
import httpx
import psutil
import pytest
import yaml
from pydantic import JsonValue
from redis import Redis
from tests.integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, string_value
from tests.integration._support.database import read_rows
from tests.integration._support.process import owned_proxy_process
from tests.integration._support.redis_process import owned_redis
from tests.integration._support.upstream import ScenarioHandle
from tests.integration.cost_calculation.cost_tracking_case import JsonResponse, SseResponse
from tests.integration.caching.response_cache_case import (
CHAT,
CHAT_MODEL,
CHAT_STREAM,
MESSAGE,
MESSAGE_STREAM,
MESSAGES_MODEL,
PROVIDER_KEY,
RESPONSE,
RESPONSE_STREAM,
TEXT,
TEXT_MODEL,
TEXT_STREAM,
Upstream,
chat_body,
choices,
content_of_first_choice,
emptied,
json_response,
message_body,
prompt,
rescript,
response_id,
scripted,
slowed,
)
RecordProperty = Callable[[str, object], None]
WORKERS: Final = 2
REMOVE_FROM_ENVIRONMENT: Final = ("DATABASE_URL_READ_REPLICA",)
BURST_TIMEOUT_SECONDS: Final = 60
CHAOS_AFTER_ANSWERS: Final = 4
STREAM_FRAME_DELAY_MS: Final = 500
PER_ENDPOINT: Final = 6
STREAMS_PER_ENDPOINT: Final = 2
LOCAL_BATCH: Final = 20
SPEND_SQL: Final = 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE api_key = %s'
@dataclass(frozen=True, slots=True)
class Deployment:
name: str
litellm_model: str
handle: ScenarioHandle
def entry(self) -> dict[str, JsonValue]:
return {
"model_name": self.name,
"litellm_params": {
"model": self.litellm_model,
"api_base": self.handle.api_base(),
"api_key": PROVIDER_KEY,
},
}
@dataclass(frozen=True, slots=True)
class Fleet:
chat: Deployment
chat_stream: Deployment
text: Deployment
text_stream: Deployment
responses: Deployment
responses_stream: Deployment
messages: Deployment
messages_stream: Deployment
flip: Deployment
def burst_members(self) -> tuple[Deployment, ...]:
return (
self.chat,
self.chat_stream,
self.text,
self.text_stream,
self.responses,
self.responses_stream,
self.messages,
self.messages_stream,
)
def entries(self) -> list[dict[str, JsonValue]]:
return [deployment.entry() for deployment in (*self.burst_members(), self.flip)]
@dataclass(frozen=True, slots=True)
class Call:
label: str
path: str
body: dict[str, JsonValue]
@property
def streamed(self) -> bool:
return self.label.startswith("stream-")
@dataclass(frozen=True, slots=True)
class Answer:
call: Call
status: int | None
text: str
@property
def has_content(self) -> bool:
marker: Final = "streamed" if self.call.streamed else "scripted"
return self.status == 200 and marker in self.text
def describe(self) -> str:
return f"{self.call.label}: {self.status} {self.text[:160]!r}"
def _deployment(scenario: Scenario, kind: str, litellm_model: str, response: JsonResponse | SseResponse) -> Deployment:
return Deployment(f"chaos-{kind}-{uuid.uuid4().hex[:8]}", litellm_model, scripted(scenario, response))
def _fleet(scenario: Scenario) -> Fleet:
return Fleet(
chat=_deployment(scenario, "chat", CHAT_MODEL, json_response(CHAT)),
chat_stream=_deployment(scenario, "chat-stream", CHAT_MODEL, slowed(CHAT_STREAM, STREAM_FRAME_DELAY_MS)),
text=_deployment(scenario, "text", TEXT_MODEL, json_response(TEXT)),
text_stream=_deployment(scenario, "text-stream", TEXT_MODEL, slowed(TEXT_STREAM, STREAM_FRAME_DELAY_MS)),
responses=_deployment(scenario, "responses", CHAT_MODEL, json_response(RESPONSE)),
responses_stream=_deployment(
scenario, "responses-stream", CHAT_MODEL, slowed(RESPONSE_STREAM, STREAM_FRAME_DELAY_MS)
),
messages=_deployment(scenario, "messages", MESSAGES_MODEL, json_response(MESSAGE)),
messages_stream=_deployment(
scenario, "messages-stream", MESSAGES_MODEL, slowed(MESSAGE_STREAM, STREAM_FRAME_DELAY_MS)
),
flip=_deployment(scenario, "flip", CHAT_MODEL, json_response(emptied(CHAT, "choices"))),
)
def _config(directory: Path, fleet: Fleet, cache_params: dict[str, JsonValue] | None = None) -> Path:
"""The rig's proxy config with the fleet as its model list and, when given, another response cache."""
rig: Final = yaml.safe_load((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text())
assert isinstance(rig, dict)
litellm_settings: Final = rig["litellm_settings"]
assert isinstance(litellm_settings, dict)
config: Final = {
**rig,
"model_list": fleet.entries(),
"litellm_settings": {**litellm_settings, **({"cache_params": cache_params} if cache_params else {})},
}
path: Final = directory / f"proxy_config-{uuid.uuid4().hex[:8]}.yaml"
path.write_text(yaml.safe_dump(config, sort_keys=False))
return path
def _overrides() -> dict[str, str]:
return {"DATABASE_URL": os.environ["DATABASE_URL"]}
def _burst_calls(fleet: Fleet, per_endpoint: int, streams_per_endpoint: int) -> tuple[Call, ...]:
def calls() -> Iterator[Call]:
for index in range(per_endpoint):
streamed: Final = index < streams_per_endpoint
prefix: Final = "stream" if streamed else "plain"
extra: Final[dict[str, JsonValue]] = {"stream": True} if streamed else {}
chat: Final = fleet.chat_stream if streamed else fleet.chat
text: Final = fleet.text_stream if streamed else fleet.text
responses: Final = fleet.responses_stream if streamed else fleet.responses
messages: Final = fleet.messages_stream if streamed else fleet.messages
yield Call(f"{prefix}-chat-{index}", "/v1/chat/completions", chat_body(chat.name, prompt(), **extra))
yield Call(f"{prefix}-text-{index}", "/v1/completions", {"model": text.name, "prompt": prompt(), **extra})
yield Call(
f"{prefix}-responses-{index}", "/v1/responses", {"model": responses.name, "input": prompt(), **extra}
)
yield Call(f"{prefix}-messages-{index}", "/v1/messages", message_body(messages.name, prompt(), **extra))
return tuple(calls())
class Burst:
"""Every call submitted at once against ``target``; ``chaos_point`` is set once ``CHAOS_AFTER_ANSWERS``
have answered (or failed), while the slowed streams are still in flight."""
def __init__(self, target: str, key: str) -> None:
self._target: Final = target
self._key: Final = key
self._lock: Final = threading.Lock()
self._answers = 0 # rebind-ok: counter behind _lock
self._futures: dict[str, Future[Answer]] = {}
self.chaos_point: Final = threading.Event()
def start(self, pool: ThreadPoolExecutor, calls: Sequence[Call]) -> None:
assert not self._futures, "burst already started"
self._futures.update((call.label, pool.submit(self._send, call)) for call in calls)
assert self.chaos_point.wait(BURST_TIMEOUT_SECONDS), (
f"fewer than {CHAOS_AFTER_ANSWERS} requests answered within {BURST_TIMEOUT_SECONDS}s"
)
def _send(self, call: Call) -> Answer:
try:
with httpx.Client(base_url=self._target, timeout=BURST_TIMEOUT_SECONDS, trust_env=False) as client:
response: Final = client.post(
call.path, json=call.body, headers={"Authorization": f"Bearer {self._key}"}
)
answer = Answer(call, response.status_code, response.text)
except httpx.HTTPError as error:
answer = Answer(call, None, f"{type(error).__name__}: {error}"[:200])
with self._lock:
self._answers += 1
if self._answers >= CHAOS_AFTER_ANSWERS:
self.chaos_point.set()
return answer
def answered(self) -> int:
with self._lock:
return self._answers
def pending_streams(self) -> tuple[str, ...]:
return tuple(
label for label, future in self._futures.items() if label.startswith("stream-") and not future.done()
)
def outcomes(self) -> tuple[Answer, ...]:
return tuple(future.result(timeout=BURST_TIMEOUT_SECONDS + 30) for future in self._futures.values())
def _assert_all_answered_with_content(outcomes: Sequence[Answer]) -> None:
failures: Final = tuple(answer.describe() for answer in outcomes if not answer.has_content)
assert not failures, (
f"{len(failures)} of {len(outcomes)} burst requests lack a 200 with content:\n " + "\n ".join(failures)
)
def _plain_ids(outcomes: Sequence[Answer]) -> frozenset[str]:
"""The ids of the non-streamed chat, text and Messages answers, which the spend log records verbatim."""
return frozenset(
string_value(JSON_OBJECT.validate_json(answer.text)["id"])
for answer in outcomes
if not answer.call.streamed and answer.call.path != "/v1/responses"
) # comprehension-ok: one filter over the burst
def _spend_request_ids(key: str) -> frozenset[str]:
return frozenset(
string_value(row["request_id"]) for row in read_rows(SPEND_SQL, (sha256(key.encode()).hexdigest(),))
)
def _assert_every_burst_request_landed_once(key: str, outcomes: Sequence[Answer]) -> None:
rows: Final = eventually(lambda: _spend_request_ids(key), lambda ids: len(ids) >= len(outcomes), seconds=70)
assert len(rows) == len(outcomes), f"{len(rows)} spend rows for {len(outcomes)} burst requests: {sorted(rows)}"
missing: Final = _plain_ids(outcomes) - rows
assert not missing, f"burst answers without a spend row: {sorted(missing)}"
def _chat(target: str, key: str, body: dict[str, JsonValue]) -> httpx.Response:
with httpx.Client(base_url=target, timeout=30, trust_env=False) as client:
return client.post("/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {key}"})
def _assert_empty_answer_reaches_upstream_again(target: str, key: str, upstream: Upstream, flip: Deployment) -> None:
"""Through ``target`` after the outage: the scripted empty answer is never served again once the
upstream answers for real, and the real answer is the one served from the cache afterwards."""
request: Final = chat_body(flip.name, prompt())
empty: Final = _chat(target, key, request)
assert choices(empty) == [], empty.text
assert upstream.calls(flip.handle.scenario_id) == 1
rescript(flip.handle, json_response(CHAT))
refilled: Final = _chat(target, key, request)
assert content_of_first_choice(refilled) == "scripted", f"the empty answer was served: {refilled.text}"
assert upstream.calls(flip.handle.scenario_id) == 2
hit: Final = eventually(lambda: _chat(target, key, request), lambda r: response_id(r) == response_id(refilled))
assert content_of_first_choice(hit) == "scripted", hit.text
assert upstream.calls(flip.handle.scenario_id) >= 2
def _workers(root: psutil.Process) -> tuple[psutil.Process, ...]:
"""uvicorn's worker children of the owned proxy root, spawned through ``multiprocessing.spawn``."""
def spawned() -> Iterator[psutil.Process]:
for child in root.children():
try:
cmdline = child.cmdline()
except (psutil.NoSuchProcess, psutil.AccessDenied):
continue
if any("multiprocessing.spawn" in part for part in cmdline):
yield child
return tuple(sorted(spawned(), key=lambda process: process.pid))
def _cache_ping(target: str, key: str) -> httpx.Response:
with httpx.Client(base_url=target, timeout=30, trust_env=False) as client:
return client.get("/cache/ping", headers={"Authorization": f"Bearer {key}"})
def _cache_status(response: httpx.Response) -> str:
assert response.status_code == 200, f"/cache/ping: {response.status_code} {response.text}"
return string_value(JSON_OBJECT.validate_json(response.content)["status"])
def _proxy_url(gateway: Gateway) -> str:
return str(gateway.client.base_url).rstrip("/")
@pytest.mark.timeout(300)
def test_redis_stopped_mid_burst_answers_every_request_and_never_pins_an_empty_answer_after(
gateway: Gateway, tmp_path: Path, record_property: RecordProperty
) -> None:
upstream: Final = Upstream(gateway.upstream_url)
with gateway.scenario() as scenario, owned_redis(tmp_path) as store:
fleet: Final = _fleet(scenario)
key: Final = scenario.key(models=[deployment.name for deployment in (*fleet.burst_members(), fleet.flip)])
with (
owned_proxy_process(
gateway,
tmp_path,
{
**_overrides(),
"REDIS_HOST": store.host,
"REDIS_PORT": str(store.port),
"REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT": "0",
},
config=_config(tmp_path, fleet),
remove_environment=REMOVE_FROM_ENVIRONMENT,
workers=WORKERS,
) as owned,
ThreadPoolExecutor(PER_ENDPOINT * 4) as pool,
):
target: Final = _proxy_url(owned.gateway)
assert _cache_status(_cache_ping(target, gateway.key)) == "healthy"
burst: Final = Burst(target, key)
burst.start(pool, _burst_calls(fleet, PER_ENDPOINT, STREAMS_PER_ENDPOINT))
pending: Final = burst.pending_streams()
record_property("answered_before_outage", burst.answered())
assert pending, "no stream was in flight when Redis stopped"
store.stop()
down: Final = _cache_ping(target, gateway.key)
assert down.status_code == 503, f"/cache/ping with Redis stopped: {down.status_code} {down.text}"
assert "Service Unhealthy" in down.text, down.text
outcomes: Final = burst.outcomes()
store.start()
recovered: Final = eventually(
lambda: _cache_ping(target, gateway.key), lambda r: r.status_code == 200, seconds=60
)
assert _cache_status(recovered) == "healthy"
_assert_all_answered_with_content(outcomes)
_assert_every_burst_request_landed_once(key, outcomes)
_assert_empty_answer_reaches_upstream_again(target, key, upstream, fleet.flip)
@pytest.mark.timeout(300)
def test_redis_paused_mid_burst_answers_every_request_and_never_pins_an_empty_answer_after(
gateway: Gateway, tmp_path: Path, record_property: RecordProperty
) -> None:
upstream: Final = Upstream(gateway.upstream_url)
with gateway.scenario() as scenario, owned_redis(tmp_path) as store:
fleet: Final = _fleet(scenario)
key: Final = scenario.key(models=[deployment.name for deployment in (*fleet.burst_members(), fleet.flip)])
with (
owned_proxy_process(
gateway,
tmp_path,
{**_overrides(), "REDIS_HOST": store.host, "REDIS_PORT": str(store.port)},
config=_config(tmp_path, fleet),
remove_environment=REMOVE_FROM_ENVIRONMENT,
workers=WORKERS,
) as owned,
ThreadPoolExecutor(PER_ENDPOINT * 4) as pool,
Redis(host=store.host, port=store.port) as control,
):
target: Final = _proxy_url(owned.gateway)
assert _cache_status(_cache_ping(target, gateway.key)) == "healthy"
burst: Final = Burst(target, key)
burst.start(pool, _burst_calls(fleet, PER_ENDPOINT, STREAMS_PER_ENDPOINT))
pending: Final = burst.pending_streams()
record_property("answered_before_pause", burst.answered())
assert pending, "no stream was in flight when Redis was paused"
assert control.client_pause(2000) is True
outcomes: Final = burst.outcomes()
assert _cache_status(_cache_ping(target, gateway.key)) == "healthy"
_assert_all_answered_with_content(outcomes)
_assert_every_burst_request_landed_once(key, outcomes)
_assert_empty_answer_reaches_upstream_again(target, key, upstream, fleet.flip)
@pytest.mark.timeout(300)
def test_worker_killed_before_burst_leaves_the_survivor_serving_and_never_pins_an_empty_answer(
gateway: Gateway, tmp_path: Path, record_property: RecordProperty
) -> None:
upstream: Final = Upstream(gateway.upstream_url)
with gateway.scenario() as scenario:
fleet: Final = _fleet(scenario)
key: Final = scenario.key(models=[deployment.name for deployment in (*fleet.burst_members(), fleet.flip)])
with (
owned_proxy_process(
gateway,
tmp_path,
_overrides(),
config=_config(tmp_path, fleet),
remove_environment=REMOVE_FROM_ENVIRONMENT,
workers=WORKERS,
) as owned,
ThreadPoolExecutor(PER_ENDPOINT * 4) as pool,
):
target: Final = _proxy_url(owned.gateway)
root: Final = psutil.Process(owned.process.pid)
before: Final = _workers(root)
assert len(before) == WORKERS, [process.pid for process in before]
victim: Final = before[0]
victim.kill()
victim.wait(timeout=10)
with httpx.Client(base_url=target, timeout=15, trust_env=False) as fresh:
readiness: Final = fresh.get("/health/readiness")
assert readiness.status_code == 200, f"/health/readiness with worker {victim.pid} dead: {readiness.text}"
burst: Final = Burst(target, key)
burst.start(pool, _burst_calls(fleet, PER_ENDPOINT // 2, STREAMS_PER_ENDPOINT // 2))
outcomes: Final = burst.outcomes()
respawned: Final = eventually(
lambda: tuple(process.pid for process in _workers(root)),
lambda pids: len(pids) == WORKERS and victim.pid not in pids,
seconds=60,
)
record_property(
"worker_pids", {"before": [process.pid for process in before], "killed": victim.pid, "after": respawned}
)
_assert_all_answered_with_content(outcomes)
_assert_every_burst_request_landed_once(key, outcomes)
_assert_empty_answer_reaches_upstream_again(target, key, upstream, fleet.flip)
def _batch(target: str, key: str, body: dict[str, JsonValue], size: int) -> tuple[httpx.Response, ...]:
"""``size`` identical requests, each on a fresh connection so both workers take traffic."""
return tuple(_chat(target, key, body) for _ in range(size))
@pytest.mark.timeout(300)
def test_in_memory_cache_mode_never_serves_an_empty_answer_from_a_worker_store(
gateway: Gateway, tmp_path: Path
) -> None:
upstream: Final = Upstream(gateway.upstream_url)
with gateway.scenario() as scenario:
fleet: Final = _fleet(scenario)
key: Final = scenario.key(models=[fleet.flip.name])
with owned_proxy_process(
gateway,
tmp_path,
_overrides(),
config=_config(tmp_path, fleet, cache_params={"type": "local"}),
remove_environment=REMOVE_FROM_ENVIRONMENT,
workers=WORKERS,
) as owned:
target: Final = _proxy_url(owned.gateway)
request: Final = chat_body(fleet.flip.name, prompt())
first_batch: Final = _batch(target, key, request, LOCAL_BATCH)
assert all(choices(response) == [] for response in first_batch), [r.text for r in first_batch]
assert upstream.calls(fleet.flip.handle.scenario_id) == LOCAL_BATCH, (
"an empty answer was served from a worker store"
)
rescript(fleet.flip.handle, json_response(CHAT))
second_batch: Final = _batch(target, key, request, LOCAL_BATCH)
served_empty: Final = tuple(response.text for response in second_batch if choices(response) == [])
assert not served_empty, (
f"{len(served_empty)} of {LOCAL_BATCH} answers still empty after the upstream recovered"
)
assert all(content_of_first_choice(response) == "scripted" for response in second_batch)
after_second: Final = upstream.calls(fleet.flip.handle.scenario_id)
assert LOCAL_BATCH < after_second <= 2 * LOCAL_BATCH, after_second
third_batch: Final = _batch(target, key, request, LOCAL_BATCH)
assert all(content_of_first_choice(response) == "scripted" for response in third_batch)
assert upstream.calls(fleet.flip.handle.scenario_id) - after_second <= WORKERS, (
"the real answer is not cached"
)

View file

@ -0,0 +1,813 @@
"""The response cache never stores an answer with no output and never serves a stored one.
Every cell scripts the owned upstream for one deployment, drives the rig proxy through the client a
user would hold (OpenAI SDK, Anthropic SDK, raw httpx) and reads three things: what the caller got,
how often the upstream was called for that deployment (an answer served from the cache never reaches
it), and what landed in Redis or the spend log. An upstream answer carrying no choices, no output items
or no content blocks reaches the upstream again on the next identical request, and the real answer
that follows is the one the cache keeps. Answers with output, errors, ``no-cache`` requests, the
``/cache/delete`` workaround and the other cached call types are unchanged.
"""
from __future__ import annotations
import asyncio
import json
import uuid
from collections.abc import Callable, Iterator, Mapping
from dataclasses import dataclass
from functools import partial
from hashlib import sha256
from typing import Final, TypeAlias, TypeVar
import httpx
import pytest
from anthropic import Anthropic, AsyncAnthropic
from openai import AsyncOpenAI, OpenAI
from pydantic import JsonValue
from tests.integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value
from tests.integration._support.database import read_rows
from tests.integration.caching.response_cache_case import (
CHAT,
CHAT_MODEL,
CHAT_STREAM,
CHAT_STREAM_WITHOUT_CHOICES,
EMBEDDING,
MESSAGE,
MESSAGE_STREAM,
MESSAGE_STREAM_WITHOUT_CONTENT_BLOCKS,
MESSAGES_MODEL,
RERANK,
RESPONSE,
RESPONSE_STREAM,
RESPONSE_STREAM_WITHOUT_OUTPUT,
SHORT_TTL_CACHE_CONTROL,
TEXT,
TEXT_MODEL,
TEXT_STREAM_WITHOUT_CHOICES,
TRANSCRIPT,
Upstream,
await_entry,
chat_body,
choices,
content_of_first_choice,
emptied,
flip_entry_to_empty,
json_response,
message_body,
prompt,
redis_store,
rescript,
response_id,
scripted,
text_of_first_choice,
)
MISSING: Final = object()
EMPTY_STREAMS: Final = 6
REPLAY_ATTEMPTS: Final = 8
SETTLE_SECONDS: Final = 15
T = TypeVar("T")
_Answer: TypeAlias = tuple[frozenset[str], str, tuple[str, ...]]
SPEND_SQL: Final = (
'SELECT request_id, cache_hit, spend FROM "LiteLLM_SpendLogs" WHERE api_key = %s ORDER BY "startTime", request_id'
)
def _proxy_root(gateway: Gateway) -> str:
return str(gateway.client.base_url).rstrip("/")
def _openai(gateway: Gateway, key: str | None = None) -> OpenAI:
return OpenAI(base_url=f"{_proxy_root(gateway)}/v1", api_key=key or gateway.key, max_retries=0)
def _async_openai(gateway: Gateway, key: str | None = None) -> AsyncOpenAI:
return AsyncOpenAI(base_url=f"{_proxy_root(gateway)}/v1", api_key=key or gateway.key, max_retries=0)
def _anthropic(gateway: Gateway) -> Anthropic:
return Anthropic(base_url=_proxy_root(gateway), api_key=gateway.key, max_retries=0)
def _async_anthropic(gateway: Gateway) -> AsyncAnthropic:
return AsyncAnthropic(base_url=_proxy_root(gateway), api_key=gateway.key, max_retries=0)
def _chat(gateway: Gateway, body: Mapping[str, JsonValue], *, key: str | None = None) -> httpx.Response:
return gateway.request("POST", "/v1/chat/completions", body, key=key)
def _spend_rows(key: str) -> list[dict[str, JsonValue]]:
return read_rows(SPEND_SQL, (sha256(key.encode()).hexdigest(),))
def _assert_prompt_reached_upstream(upstream: Upstream, identity: str, text: str, times: int) -> None:
received: Final = upstream.received(identity)
assert len(received) == times, f"{identity}: {len(received)} upstream calls, expected {times}"
assert all(text in json.dumps(item["body"]) for item in received), received
def _served_without_upstream(upstream: Upstream, identity: str, call: Callable[[], T]) -> T | None:
before: Final = upstream.calls(identity)
answer: Final = call()
return None if upstream.calls(identity) > before else answer
@dataclass(frozen=True, slots=True)
class Streamed:
ids: frozenset[str]
text: str
kinds: tuple[str, ...]
reached_upstream: bool
def _streams(stream: Callable[[], _Answer], upstream: Upstream, identity: str, attempts: int) -> Iterator[Streamed]:
"""Stream up to ``attempts`` times, stopping after the first answer served without an upstream call."""
for _ in range(attempts):
before: Final = upstream.calls(identity)
ids, text, kinds = stream()
answer: Final = Streamed(ids, text, kinds, upstream.calls(identity) > before)
yield answer
if not answer.reached_upstream:
return
def _never_replayed(stream: Callable[[], _Answer], upstream: Upstream, identity: str) -> tuple[Streamed, ...]:
"""``EMPTY_STREAMS`` identical streams, every one reaching the upstream and carrying ids of its own."""
answers: Final = tuple(_streams(stream, upstream, identity, EMPTY_STREAMS))
assert len(answers) == EMPTY_STREAMS and all(answer.reached_upstream for answer in answers), (
f"a stream was replayed from the cache: {answers}"
)
distinct: Final = frozenset[str]().union(*(answer.ids for answer in answers))
assert sum(len(answer.ids) for answer in answers) == len(distinct), answers
return answers
def _replayed(stream: Callable[[], _Answer], upstream: Upstream, identity: str, produced: frozenset[str]) -> Streamed:
"""The first of up to ``REPLAY_ATTEMPTS`` identical streams served without an upstream call.
A streamed request can reach the upstream once per proxy worker before a replay is served: the proxy adds
``stream_options`` to a stream before keying it only when its own router already resolves the model, and the
worker that did not take the ``/model/new`` call resolves it once its router syncs. The replay carries ids an
upstream answer produced.
"""
answers: Final = tuple(_streams(stream, upstream, identity, REPLAY_ATTEMPTS))
replay: Final = answers[-1]
assert not replay.reached_upstream, (
f"no stream was replayed from the cache in {REPLAY_ATTEMPTS} attempts: {answers}"
)
seen: Final = produced.union(*(answer.ids for answer in answers[:-1]))
assert replay.ids <= seen, (replay.ids, seen)
return replay
def test_chat_empty_choices_reach_upstream_again_and_the_real_answer_is_cached(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario, _openai(gateway) as client:
handle: Final = scripted(scenario, json_response(emptied(CHAT, "choices")))
model: Final = scenario.model(api_base=handle.api_base())
empty: Final = client.chat.completions.create(model=model, messages=[{"role": "user", "content": text}])
assert empty.choices == [], empty.model_dump_json()
_assert_prompt_reached_upstream(upstream, handle.scenario_id, text, 1)
rescript(handle, json_response(CHAT))
real: Final = client.chat.completions.create(model=model, messages=[{"role": "user", "content": text}])
assert real.id != empty.id and real.choices[0].message.content == "scripted", real.model_dump_json()
_assert_prompt_reached_upstream(upstream, handle.scenario_id, text, 2)
await_entry(store, real.id)
hit: Final = client.chat.completions.with_raw_response.create(
model=model, messages=[{"role": "user", "content": text}]
)
served: Final = hit.parse()
assert served.id == real.id and served.choices[0].message.content == "scripted", hit.text
assert hit.headers.get("x-litellm-cache-key"), dict(hit.headers)
assert upstream.calls(handle.scenario_id) == 2
def test_async_chat_spend_rows_record_the_empty_answer_and_the_refill_as_misses(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(emptied(CHAT, "choices")))
model: Final = scenario.model(api_base=handle.api_base())
key: Final = scenario.key(models=[model])
async def drive() -> tuple[str, str, str]:
async with _async_openai(gateway, key) as client:
messages: Final = [{"role": "user", "content": text}]
empty = await client.chat.completions.create(model=model, messages=messages)
assert empty.choices == [], empty.model_dump_json()
rescript(handle, json_response(CHAT))
real = await client.chat.completions.create(model=model, messages=messages)
assert real.choices[0].message.content == "scripted", real.model_dump_json()
await_entry(store, real.id)
served = await client.chat.completions.create(model=model, messages=messages)
return empty.id, real.id, served.id
empty_id, real_id, served_id = asyncio.run(drive())
assert served_id == real_id != empty_id
assert upstream.calls(handle.scenario_id) == 2
rows: Final = eventually(lambda: _spend_rows(key), lambda found: len(found) == 3, seconds=70)
by_id: Final = {string_value(row["request_id"]): row for row in rows}
assert set(by_id) >= {empty_id, real_id}, sorted(by_id)
assert by_id[empty_id]["cache_hit"] != "True" and by_id[real_id]["cache_hit"] != "True", rows
hit_row: Final = next(row for request_id, row in by_id.items() if request_id.startswith(f"{real_id}_cache_hit"))
assert hit_row["cache_hit"] == "True" and float(str(hit_row["spend"])) == 0, rows
def test_text_completion_empty_choices_reach_upstream_again(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(emptied(TEXT, "choices")))
model: Final = scenario.model(model=TEXT_MODEL, api_base=handle.api_base())
body: Final[dict[str, JsonValue]] = {"model": model, "prompt": prompt()}
empty: Final = gateway.request("POST", "/v1/completions", body)
assert choices(empty) == [], empty.text
assert upstream.calls(handle.scenario_id) == 1
rescript(handle, json_response(TEXT))
real: Final = gateway.request("POST", "/v1/completions", body)
assert text_of_first_choice(real) == "scripted" and response_id(real) != response_id(empty), real.text
assert upstream.calls(handle.scenario_id) == 2
await_entry(store, response_id(real))
hit: Final = gateway.request("POST", "/v1/completions", body)
assert response_id(hit) == response_id(real) and text_of_first_choice(hit) == "scripted", hit.text
assert hit.headers.get("x-litellm-cache-key"), dict(hit.headers)
assert upstream.calls(handle.scenario_id) == 2
def test_responses_empty_output_reaches_upstream_again(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(emptied(RESPONSE, "output")))
model: Final = scenario.model(api_base=handle.api_base())
async def drive() -> None:
async with _async_openai(gateway) as client:
empty = await client.responses.create(model=model, input=text)
assert empty.output == [], empty.model_dump_json()
assert upstream.calls(handle.scenario_id) == 1
rescript(handle, json_response(RESPONSE))
real = await client.responses.create(model=model, input=text)
assert real.output_text == "scripted", real.model_dump_json()
assert upstream.calls(handle.scenario_id) == 2
await_entry(store, real.output[0].id)
hit = await client.responses.with_raw_response.create(model=model, input=text)
served = hit.parse()
assert served.output[0].id == real.output[0].id and served.output_text == "scripted", hit.text
assert hit.headers.get("x-litellm-cache-key"), dict(hit.headers)
assert upstream.calls(handle.scenario_id) == 2
asyncio.run(drive())
def test_messages_empty_content_reaches_upstream_again(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario, _anthropic(gateway) as client:
handle: Final = scripted(scenario, json_response(emptied(MESSAGE, "content")))
model: Final = scenario.model(model=MESSAGES_MODEL, api_base=handle.api_base())
messages: Final = [{"role": "user", "content": text}]
empty: Final = client.messages.create(model=model, max_tokens=16, messages=messages)
assert empty.content == [], empty.model_dump_json()
_assert_prompt_reached_upstream(upstream, handle.scenario_id, text, 1)
rescript(handle, json_response(MESSAGE))
real: Final = client.messages.create(model=model, max_tokens=16, messages=messages)
assert real.id != empty.id and real.content[0].type == "text", real.model_dump_json()
assert real.content[0].text == "scripted"
_assert_prompt_reached_upstream(upstream, handle.scenario_id, text, 2)
await_entry(store, real.id)
served: Final = client.messages.create(model=model, max_tokens=16, messages=messages)
assert served.id == real.id and served.content[0].type == "text", served.model_dump_json()
assert served.content[0].text == "scripted"
assert upstream.calls(handle.scenario_id) == 2
def test_stale_chat_entry_with_empty_choices_is_a_miss(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(CHAT))
model: Final = scenario.model(api_base=handle.api_base())
filled: Final = _chat(gateway, chat_body(model, text, cache=SHORT_TTL_CACHE_CONTROL))
assert content_of_first_choice(filled) == "scripted", filled.text
stale_id: Final = flip_entry_to_empty(store, await_entry(store, response_id(filled)), "choices")
request: Final = chat_body(model, text)
refilled: Final = eventually(
lambda: _chat(gateway, request), lambda answer: response_id(answer) != response_id(filled), SETTLE_SECONDS
)
assert response_id(refilled) != stale_id, f"the stale entry was served: {refilled.text}"
assert content_of_first_choice(refilled) == "scripted", refilled.text
assert upstream.calls(handle.scenario_id) == 2
await_entry(store, response_id(refilled))
hit: Final = eventually(
lambda: _chat(gateway, request), lambda answer: response_id(answer) == response_id(refilled), SETTLE_SECONDS
)
assert content_of_first_choice(hit) == "scripted", hit.text
assert upstream.calls(handle.scenario_id) == 2
def test_stale_text_completion_entry_with_empty_choices_is_a_miss(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(TEXT))
model: Final = scenario.model(model=TEXT_MODEL, api_base=handle.api_base())
short_lived: Final[dict[str, JsonValue]] = {"model": model, "prompt": text, "cache": SHORT_TTL_CACHE_CONTROL}
filled: Final = gateway.request("POST", "/v1/completions", short_lived)
assert text_of_first_choice(filled) == "scripted", filled.text
stale_id: Final = flip_entry_to_empty(store, await_entry(store, response_id(filled)), "choices")
request: Final[dict[str, JsonValue]] = {"model": model, "prompt": text}
refilled: Final = eventually(
lambda: gateway.request("POST", "/v1/completions", request),
lambda answer: response_id(answer) != response_id(filled),
SETTLE_SECONDS,
)
assert response_id(refilled) != stale_id, f"the stale entry was served: {refilled.text}"
assert text_of_first_choice(refilled) == "scripted", refilled.text
assert upstream.calls(handle.scenario_id) == 2
await_entry(store, response_id(refilled))
hit: Final = eventually(
lambda: gateway.request("POST", "/v1/completions", request),
lambda answer: response_id(answer) == response_id(refilled),
SETTLE_SECONDS,
)
assert text_of_first_choice(hit) == "scripted", hit.text
assert upstream.calls(handle.scenario_id) == 2
def test_stale_responses_entry_with_empty_output_is_a_miss(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario, _openai(gateway) as client:
handle: Final = scripted(scenario, json_response(RESPONSE))
model: Final = scenario.model(api_base=handle.api_base())
filled: Final = client.responses.create(model=model, input=text, extra_body={"cache": SHORT_TTL_CACHE_CONTROL})
assert filled.output_text == "scripted", filled.model_dump_json()
flip_entry_to_empty(store, await_entry(store, filled.output[0].id), "output")
refilled: Final = eventually(
lambda: client.responses.create(model=model, input=text),
lambda answer: not answer.output or answer.output[0].id != filled.output[0].id,
SETTLE_SECONDS,
)
assert refilled.output, f"the stale entry was served: {refilled.model_dump_json()}"
assert refilled.output_text == "scripted", refilled.model_dump_json()
assert upstream.calls(handle.scenario_id) == 2
await_entry(store, refilled.output[0].id)
hit: Final = eventually(
lambda: client.responses.create(model=model, input=text),
lambda answer: bool(answer.output) and answer.output[0].id == refilled.output[0].id,
SETTLE_SECONDS,
)
assert hit.output_text == "scripted", hit.model_dump_json()
assert upstream.calls(handle.scenario_id) == 2
def test_stale_messages_entry_with_empty_content_is_a_miss(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario, _anthropic(gateway) as client:
handle: Final = scripted(scenario, json_response(MESSAGE))
model: Final = scenario.model(model=MESSAGES_MODEL, api_base=handle.api_base())
messages: Final = [{"role": "user", "content": text}]
filled: Final = client.messages.create(
model=model, max_tokens=16, messages=messages, extra_body={"cache": SHORT_TTL_CACHE_CONTROL}
)
assert filled.content[0].type == "text" and filled.content[0].text == "scripted", filled.model_dump_json()
stale_id: Final = flip_entry_to_empty(store, await_entry(store, filled.id), "content")
refilled: Final = eventually(
lambda: client.messages.create(model=model, max_tokens=16, messages=messages),
lambda answer: answer.id != filled.id,
SETTLE_SECONDS,
)
assert refilled.id != stale_id, f"the stale entry was served: {refilled.model_dump_json()}"
assert refilled.content[0].type == "text" and refilled.content[0].text == "scripted", refilled.model_dump_json()
assert upstream.calls(handle.scenario_id) == 2
await_entry(store, refilled.id)
hit: Final = eventually(
lambda: client.messages.create(model=model, max_tokens=16, messages=messages),
lambda answer: answer.id == refilled.id,
SETTLE_SECONDS,
)
assert hit.content[0].type == "text" and hit.content[0].text == "scripted", hit.model_dump_json()
assert upstream.calls(handle.scenario_id) == 2
def _streamed_chat(gateway: Gateway, model: str, text: str) -> _Answer:
"""The chunk ids, the streamed text and the chunk kinds, after the whole stream was read."""
with (
_openai(gateway) as client,
client.chat.completions.create(
model=model, messages=[{"role": "user", "content": text}], stream=True
) as stream,
):
chunks: Final = tuple(stream)
ids: Final = frozenset(chunk.id for chunk in chunks)
streamed: Final = "".join(
choice.delta.content or "" for chunk in chunks for choice in chunk.choices
) # comprehension-ok: flattens chunk choices
return ids, streamed, tuple(chunk.object for chunk in chunks)
def test_streamed_chat_without_choices_is_never_assembled_or_cached(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
text: Final = prompt()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, CHAT_STREAM_WITHOUT_CHOICES)
model: Final = scenario.model(api_base=handle.api_base())
empty: Final = _never_replayed(partial(_streamed_chat, gateway, model, text), upstream, handle.scenario_id)
assert all(answer.text == "" for answer in empty), empty
def _streamed_text(gateway: Gateway, model: str, text: str) -> _Answer:
"""The chunk ids, the streamed text and the chunk kinds of a text completion, after the whole stream was read."""
with _openai(gateway) as client, client.completions.create(model=model, prompt=text, stream=True) as stream:
chunks: Final = tuple(stream)
ids: Final = frozenset(chunk.id for chunk in chunks)
streamed: Final = "".join(
choice.text or "" for chunk in chunks for choice in chunk.choices
) # comprehension-ok: flattens chunk choices
return ids, streamed, tuple(chunk.object for chunk in chunks)
def test_streamed_text_completion_without_choices_is_never_assembled_or_cached(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
text: Final = prompt()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, TEXT_STREAM_WITHOUT_CHOICES)
model: Final = scenario.model(model=TEXT_MODEL, api_base=handle.api_base())
empty: Final = _never_replayed(partial(_streamed_text, gateway, model, text), upstream, handle.scenario_id)
assert all(answer.text == "" for answer in empty), empty
def test_streamed_chat_with_content_is_replayed_from_the_cache(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, CHAT_STREAM)
model: Final = scenario.model(api_base=handle.api_base())
stream: Final = partial(_streamed_chat, gateway, model, text)
first_ids, first_text, _ = stream()
assert first_text == "streamed response" and len(first_ids) == 1, first_ids
assert upstream.calls(handle.scenario_id) == 1
await_entry(store, next(iter(first_ids)))
replay: Final = _replayed(stream, upstream, handle.scenario_id, first_ids)
assert replay.text == "streamed response", replay
def _streamed_response(gateway: Gateway, model: str, text: str) -> _Answer:
"""The completed response's output item ids, its text and the event kinds, after the whole stream was read."""
async def read() -> _Answer:
async with (
_async_openai(gateway) as client,
await client.responses.create(model=model, input=text, stream=True) as stream,
):
events: Final = [event async for event in stream]
completed: Final = tuple(event for event in events if event.type == "response.completed")
assert len(completed) == 1, [event.type for event in events]
ids: Final = frozenset(item.id for item in completed[0].response.output)
return ids, completed[0].response.output_text, tuple(event.type for event in events)
return asyncio.run(read())
def test_streamed_responses_completing_without_output_reach_upstream_again(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, RESPONSE_STREAM_WITHOUT_OUTPUT)
model: Final = scenario.model(api_base=handle.api_base())
stream: Final = partial(_streamed_response, gateway, model, text)
empty: Final = _never_replayed(stream, upstream, handle.scenario_id)
assert all(answer.ids == frozenset() and answer.text == "" for answer in empty), empty
rescript(handle, RESPONSE_STREAM)
real_ids, real_text, _ = stream()
assert real_text == "streamed response", real_ids
assert upstream.calls(handle.scenario_id) == EMPTY_STREAMS + 1, "the empty stream was replayed"
await_entry(store, next(iter(real_ids)))
replay: Final = _replayed(stream, upstream, handle.scenario_id, real_ids)
assert replay.text == "streamed response", replay
def test_streamed_responses_with_output_are_replayed_from_the_cache(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, RESPONSE_STREAM)
model: Final = scenario.model(api_base=handle.api_base())
stream: Final = partial(_streamed_response, gateway, model, text)
first_ids, first_text, _ = stream()
assert first_text == "streamed response" and len(first_ids) == 1, first_ids
assert all(item.startswith(f"msg_{handle.scenario_id}-") for item in first_ids), first_ids
assert upstream.calls(handle.scenario_id) == 1
await_entry(store, next(iter(first_ids)))
replay: Final = _replayed(stream, upstream, handle.scenario_id, first_ids)
assert replay.text == "streamed response", replay
def _streamed_message(gateway: Gateway, model: str, text: str) -> _Answer:
"""The message id, the streamed text and the event kinds, after the whole stream was read."""
async def read() -> _Answer:
async with _async_anthropic(gateway) as client:
stream: Final = await client.messages.create(
model=model, max_tokens=16, messages=[{"role": "user", "content": text}], stream=True
)
events: Final = [event async for event in stream]
starts: Final = tuple(event for event in events if event.type == "message_start")
assert len(starts) == 1, [event.type for event in events]
deltas: Final = tuple(event for event in events if event.type == "content_block_delta")
streamed: Final = "".join(delta.delta.text for delta in deltas if delta.delta.type == "text_delta")
return frozenset({starts[0].message.id}), streamed, tuple(event.type for event in events)
return asyncio.run(read())
def test_streamed_messages_without_content_blocks_reach_upstream_again(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, MESSAGE_STREAM_WITHOUT_CONTENT_BLOCKS)
model: Final = scenario.model(model=MESSAGES_MODEL, api_base=handle.api_base())
stream: Final = partial(_streamed_message, gateway, model, text)
empty: Final = _never_replayed(stream, upstream, handle.scenario_id)
assert all("content_block_start" not in answer.kinds and answer.text == "" for answer in empty), empty
rescript(handle, MESSAGE_STREAM)
real_ids, real_text, real_kinds = stream()
assert real_text == "streamed response" and "content_block_start" in real_kinds, real_kinds
assert upstream.calls(handle.scenario_id) == EMPTY_STREAMS + 1, "the stream without content blocks was replayed"
await_entry(store, next(iter(real_ids)))
replay: Final = _replayed(stream, upstream, handle.scenario_id, real_ids)
assert replay.text == "streamed response" and replay.kinds == real_kinds, replay
def test_streamed_messages_with_content_are_replayed_from_the_cache(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, MESSAGE_STREAM)
model: Final = scenario.model(model=MESSAGES_MODEL, api_base=handle.api_base())
stream: Final = partial(_streamed_message, gateway, model, text)
first_ids, first_text, first_kinds = stream()
assert first_text == "streamed response" and "content_block_start" in first_kinds, first_kinds
assert upstream.calls(handle.scenario_id) == 1
await_entry(store, next(iter(first_ids)))
replay: Final = _replayed(stream, upstream, handle.scenario_id, first_ids)
assert replay.text == "streamed response" and replay.kinds == first_kinds, replay
def test_embeddings_are_still_served_from_the_cache(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(EMBEDDING))
model: Final = scenario.model(model="openai/text-embedding-3-small", api_base=handle.api_base())
body: Final[dict[str, JsonValue]] = {"model": model, "input": [prompt()]}
first: Final = gateway.post("/v1/embeddings", body)
assert first["object"] == "list" and first["data"], first
assert upstream.calls(handle.scenario_id) == 1
second: Final = eventually(
lambda: _served_without_upstream(
upstream, handle.scenario_id, lambda: gateway.post("/v1/embeddings", body)
),
lambda answer: answer is not None,
SETTLE_SECONDS,
)
assert second is not None and second["data"] == first["data"], second
@pytest.mark.parametrize(
"choice",
(
pytest.param(
{"index": 0, "message": {"role": "assistant", "content": ""}, "finish_reason": "content_filter"},
id="blocked-empty-content",
),
pytest.param(
{
"index": 0,
"message": {
"role": "assistant",
"content": None,
"tool_calls": [
{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}
],
},
"finish_reason": "tool_calls",
},
id="tool-calls-only",
),
),
)
def test_chat_answers_with_a_choice_but_no_text_are_still_cached(
gateway: Gateway, choice: dict[str, JsonValue]
) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response({**CHAT, "choices": [choice]}))
model: Final = scenario.model(api_base=handle.api_base())
request: Final = chat_body(model, prompt())
first: Final = _chat(gateway, request)
assert object_value(choices(first)[0])["finish_reason"] == choice["finish_reason"], first.text
assert upstream.calls(handle.scenario_id) == 1
await_entry(store, response_id(first))
second: Final = _chat(gateway, request)
assert response_id(second) == response_id(first) and choices(second) == choices(first), second.text
assert upstream.calls(handle.scenario_id) == 1
def test_no_cache_requests_reach_upstream_every_time(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(CHAT))
model: Final = scenario.model(api_base=handle.api_base())
request: Final = chat_body(model, prompt(), cache={"no-cache": True})
ids: Final = tuple(response_id(_chat(gateway, request)) for _ in range(3))
assert len(set(ids)) == 3, ids
assert upstream.calls(handle.scenario_id) == 3
@pytest.mark.parametrize("status", (500, 429))
def test_upstream_errors_reach_the_caller_and_are_never_cached(gateway: Gateway, status: int) -> None:
upstream: Final = Upstream(gateway.upstream_url)
with gateway.scenario() as scenario:
handle: Final = scripted(
scenario, json_response({"error": {"message": "Controlled provider failure", "type": "api_error"}}, status)
)
model: Final = scenario.model(api_base=handle.api_base())
request: Final = chat_body(model, prompt())
first: Final = _chat(gateway, request)
assert first.status_code == status and "Controlled provider failure" in first.text, first.text
after_first: Final = upstream.calls(handle.scenario_id)
assert after_first >= 1
second: Final = _chat(gateway, request)
assert second.status_code == status and "Controlled provider failure" in second.text, second.text
assert upstream.calls(handle.scenario_id) > after_first, "the error was served from the cache"
def test_cache_delete_on_the_entry_key_makes_the_next_request_reach_upstream(gateway: Gateway) -> None:
"""``/cache/delete`` clears the Redis entry; a worker's in-memory copy of it serves until the entry's ttl lapses."""
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
text: Final = prompt()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(CHAT))
model: Final = scenario.model(api_base=handle.api_base())
filled: Final = _chat(gateway, chat_body(model, text, cache=SHORT_TTL_CACHE_CONTROL))
key: Final = await_entry(store, response_id(filled))
assert gateway.post("/cache/delete", {"keys": [key]}) == {"status": "success"}
assert store.get(key) is None
request: Final = chat_body(model, text)
refilled: Final = eventually(
lambda: _chat(gateway, request), lambda answer: response_id(answer) != response_id(filled), SETTLE_SECONDS
)
assert content_of_first_choice(refilled) == "scripted", refilled.text
assert upstream.calls(handle.scenario_id) == 2
await_entry(store, response_id(refilled))
hit: Final = eventually(
lambda: _chat(gateway, request), lambda answer: response_id(answer) == response_id(refilled), SETTLE_SECONDS
)
assert hit.headers.get("x-litellm-cache-key") == key, dict(hit.headers)
assert upstream.calls(handle.scenario_id) == 2
def _shape(body: Mapping[str, JsonValue], field: str, value: object) -> dict[str, JsonValue]:
if value is MISSING:
return {name: item for name, item in body.items() if name != field}
assert isinstance(value, (str, int)) or value is None
return {**body, field: value}
def _request_for(path: str, model: str, text: str) -> dict[str, JsonValue]:
if path == "/v1/responses":
return {"model": model, "input": text}
if path == "/v1/messages":
return message_body(model, text)
return chat_body(model, text)
@pytest.mark.parametrize(
("path", "body", "litellm_model", "field", "value"),
(
pytest.param("/v1/chat/completions", CHAT, CHAT_MODEL, "choices", None, id="chat-choices-null"),
pytest.param("/v1/chat/completions", CHAT, CHAT_MODEL, "choices", MISSING, id="chat-choices-missing"),
pytest.param("/v1/chat/completions", CHAT, CHAT_MODEL, "choices", "x", id="chat-choices-string"),
pytest.param("/v1/chat/completions", CHAT, CHAT_MODEL, "choices", 5, id="chat-choices-int"),
),
)
def test_malformed_upstream_answers_are_an_error_and_never_cached(
gateway: Gateway, path: str, body: Mapping[str, JsonValue], litellm_model: str, field: str, value: object
) -> None:
upstream: Final = Upstream(gateway.upstream_url)
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(_shape(body, field, value)))
control: Final = scripted(scenario, json_response(CHAT))
model: Final = scenario.model(model=litellm_model, api_base=handle.api_base())
control_model: Final = scenario.model(api_base=control.api_base())
key: Final = scenario.key(models=[model, control_model])
request: Final = _request_for(path, model, prompt())
first: Final = gateway.request("POST", path, request, key=key)
assert first.status_code == 500, first.text
assert "error" in first.text.lower(), first.text
assert upstream.calls(handle.scenario_id) == 1
second: Final = gateway.request("POST", path, request, key=key)
assert second.status_code == 500, second.text
assert upstream.calls(handle.scenario_id) == 2, "the malformed answer was served from the cache"
unrelated: Final = _chat(gateway, chat_body(control_model, prompt()), key=key)
assert content_of_first_choice(unrelated) == "scripted", unrelated.text
@pytest.mark.parametrize(
("path", "body", "litellm_model", "field", "value"),
(
pytest.param("/v1/responses", RESPONSE, CHAT_MODEL, "output", None, id="responses-output-null"),
pytest.param("/v1/responses", RESPONSE, CHAT_MODEL, "output", MISSING, id="responses-output-missing"),
pytest.param("/v1/messages", MESSAGE, MESSAGES_MODEL, "content", None, id="messages-content-null"),
pytest.param("/v1/messages", MESSAGE, MESSAGES_MODEL, "content", MISSING, id="messages-content-missing"),
),
)
def test_answers_without_the_output_field_reach_upstream_again(
gateway: Gateway, path: str, body: Mapping[str, JsonValue], litellm_model: str, field: str, value: object
) -> None:
upstream: Final = Upstream(gateway.upstream_url)
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(_shape(body, field, value)))
model: Final = scenario.model(model=litellm_model, api_base=handle.api_base())
request: Final = _request_for(path, model, prompt())
first: Final = gateway.request("POST", path, request)
assert first.status_code == 200, first.text
assert upstream.calls(handle.scenario_id) == 1
second: Final = gateway.request("POST", path, request)
assert second.status_code == 200, second.text
assert upstream.calls(handle.scenario_id) == 2, "the answer without output was served from the cache"
def test_per_request_ttl_on_an_empty_answer_does_not_pin_it(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(emptied(CHAT, "choices")))
model: Final = scenario.model(api_base=handle.api_base())
request: Final = chat_body(model, prompt(), cache={"ttl": 30})
empty: Final = _chat(gateway, request)
assert choices(empty) == [], empty.text
assert upstream.calls(handle.scenario_id) == 1
rescript(handle, json_response(CHAT))
real: Final = _chat(gateway, request)
assert content_of_first_choice(real) == "scripted", real.text
assert upstream.calls(handle.scenario_id) == 2, "the empty answer was pinned for the request's ttl"
key: Final = await_entry(store, response_id(real))
assert 0 < store.ttl(key) <= 30, store.ttl(key)
def test_rerank_results_are_still_served_from_the_cache(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(RERANK))
model: Final = scenario.model(model="cohere/rerank-v4.0", api_base=handle.api_base())
body: Final[dict[str, JsonValue]] = {"model": model, "query": prompt(), "documents": ["first", "second"]}
first: Final = gateway.post("/v1/rerank", body)
assert string_value(first["id"]).startswith(f"rerank-{handle.scenario_id}"), first
assert upstream.calls(handle.scenario_id) == 1
await_entry(store, string_value(first["id"]))
second: Final = gateway.post("/v1/rerank", body)
assert second["id"] == first["id"] and second["results"] == first["results"], second
assert upstream.calls(handle.scenario_id) == 1
def test_transcriptions_are_still_served_from_the_cache(gateway: Gateway) -> None:
upstream: Final = Upstream(gateway.upstream_url)
store: Final = redis_store()
with gateway.scenario() as scenario:
handle: Final = scripted(scenario, json_response(TRANSCRIPT))
model: Final = scenario.model(model="openai/gpt-4o-mini-transcribe", api_base=handle.api_base())
audio: Final = (f"{uuid.uuid4().hex}.wav", b"RIFF" + uuid.uuid4().bytes, "audio/wav")
first: Final = gateway.request_multipart("/v1/audio/transcriptions", {"model": model}, {"file": audio})
assert first.status_code == 200, first.text
transcript: Final = string_value(JSON_OBJECT.validate_json(first.content)["text"])
assert transcript.startswith(f"scripted {handle.scenario_id}-"), first.text
assert upstream.calls(handle.scenario_id) == 1
await_entry(store, transcript)
second: Final = gateway.request_multipart("/v1/audio/transcriptions", {"model": model}, {"file": audio})
assert second.status_code == 200 and JSON_OBJECT.validate_json(second.content)["text"] == transcript, (
second.text
)
assert upstream.calls(handle.scenario_id) == 1

View file

@ -13,7 +13,7 @@ from typing import Final
GROUPS: Final = MappingProxyType(
{
"management": ("management", "authorization", "configuration"),
"accounting": ("pricing", "spend"),
"accounting": ("pricing", "spend", "caching"),
"database": ("database",),
"providers": ("providers", "routing", "streaming", "messages_endpoint", "translation"),
"extensions": ("observability", "compatibility"),

View file

@ -0,0 +1,125 @@
"""The sync SDK paths of the response cache, which the proxy never reaches.
``litellm.completion`` and ``litellm.text_completion`` read and write the Redis cache in process. An upstream
answer with no choices is never stored, a stored entry whose choices were emptied is a miss, and the real answer
the refill brings back is the one served afterwards.
"""
from __future__ import annotations
import os
from collections.abc import Iterator
from typing import Final
import pytest
import litellm
from litellm.caching.caching import Cache
from litellm.types.caching import LiteLLMCacheType
from litellm.types.utils import ModelResponse, TextCompletionResponse
from tests.integration.caching.response_cache_case import (
CHAT,
CHAT_MODEL,
PROVIDER_KEY,
TEXT,
TEXT_MODEL,
Upstream,
await_entry,
emptied,
flip_entry_to_empty,
json_response,
prompt,
redis_store,
rescript,
scripted_in_process,
)
@pytest.fixture
def sdk_cache(monkeypatch: pytest.MonkeyPatch) -> Iterator[Cache]:
cache: Final = Cache(type=LiteLLMCacheType.REDIS, host=os.environ["REDIS_HOST"], port=os.environ["REDIS_PORT"])
monkeypatch.setattr(litellm, "cache", cache)
yield cache
monkeypatch.setattr(litellm, "cache", None)
def _completion(api_base: str, text: str) -> ModelResponse:
response: Final = litellm.completion(
model=CHAT_MODEL, messages=[{"role": "user", "content": text}], api_base=api_base, api_key=PROVIDER_KEY
)
assert isinstance(response, ModelResponse), type(response)
return response
def _text_completion(api_base: str, text: str) -> TextCompletionResponse:
response: Final = litellm.text_completion(model=TEXT_MODEL, prompt=text, api_base=api_base, api_key=PROVIDER_KEY)
assert isinstance(response, TextCompletionResponse), type(response)
return response
def _content(response: ModelResponse) -> object:
assert len(response.choices) == 1, response.model_dump_json()
return response.choices[0].model_dump()["message"]["content"]
def _text(response: TextCompletionResponse) -> object:
assert len(response.choices) == 1, response.model_dump_json()
return response.choices[0].model_dump()["text"]
@pytest.mark.usefixtures("sdk_cache")
def test_sync_completion_with_empty_choices_reaches_upstream_again_and_the_real_answer_is_cached() -> None:
upstream: Final = Upstream()
text: Final = prompt()
with scripted_in_process(json_response(emptied(CHAT, "choices"))) as handle:
empty: Final = _completion(handle.api_base(), text)
assert empty.choices == [], empty.model_dump_json()
assert upstream.calls(handle.scenario_id) == 1
rescript(handle, json_response(CHAT))
refilled: Final = _completion(handle.api_base(), text)
assert _content(refilled) == "scripted", f"the empty answer was served: {refilled.model_dump_json()}"
assert refilled.id != empty.id
assert upstream.calls(handle.scenario_id) == 2
with redis_store() as store:
await_entry(store, refilled.id)
hit: Final = _completion(handle.api_base(), text)
assert hit.id == refilled.id, hit.model_dump_json()
assert _content(hit) == "scripted"
assert upstream.calls(handle.scenario_id) == 2
@pytest.mark.usefixtures("sdk_cache")
def test_sync_completion_treats_a_stale_entry_with_empty_choices_as_a_miss() -> None:
upstream: Final = Upstream()
text: Final = prompt()
with scripted_in_process(json_response(CHAT)) as handle, redis_store() as store:
first: Final = _completion(handle.api_base(), text)
assert _content(first) == "scripted", first.model_dump_json()
stale_id: Final = flip_entry_to_empty(store, await_entry(store, first.id), "choices")
refilled: Final = _completion(handle.api_base(), text)
assert refilled.id != stale_id, f"the stale entry was served: {refilled.model_dump_json()}"
assert _content(refilled) == "scripted" and refilled.id != first.id, refilled.model_dump_json()
assert upstream.calls(handle.scenario_id) == 2
await_entry(store, refilled.id)
hit: Final = _completion(handle.api_base(), text)
assert hit.id == refilled.id, hit.model_dump_json()
assert upstream.calls(handle.scenario_id) == 2
@pytest.mark.usefixtures("sdk_cache")
def test_sync_text_completion_with_empty_choices_reaches_upstream_again() -> None:
upstream: Final = Upstream()
text: Final = prompt()
with scripted_in_process(json_response(emptied(TEXT, "choices"))) as handle:
empty: Final = _text_completion(handle.api_base(), text)
assert empty.choices == [], empty.model_dump_json()
assert upstream.calls(handle.scenario_id) == 1
rescript(handle, json_response(TEXT))
refilled: Final = _text_completion(handle.api_base(), text)
assert _text(refilled) == "scripted", f"the empty answer was served: {refilled.model_dump_json()}"
assert upstream.calls(handle.scenario_id) == 2
with redis_store() as store:
await_entry(store, refilled.id)
hit: Final = _text_completion(handle.api_base(), text)
assert hit.id == refilled.id, hit.model_dump_json()
assert upstream.calls(handle.scenario_id) == 2

View file

@ -23,17 +23,23 @@ from litellm.caching.caching_handler import (
_is_chat_completion_cached_dict,
_should_defer_streaming_cache_hit_callbacks,
)
from litellm.caching import DualCache, InMemoryCache
from litellm.caching.caching import LiteLLMCacheType
from litellm.types.caching import CACHED_STREAM_EVENTS_KEY
from litellm.types.utils import CallTypes
from litellm.types.rerank import RerankResponse
from litellm.types.utils import (
Delta,
ModelResponse,
ModelResponseStream,
StreamingChoices,
EmbeddingResponse,
TextCompletionResponse,
TranscriptionResponse,
Embedding,
)
from litellm.types.llms.openai import ResponsesAPIResponse
from collections.abc import Awaitable, Callable
from datetime import timedelta, datetime
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
@ -2313,3 +2319,355 @@ async def test_response_cache_lookup_and_write_declare_the_llm_response_target(m
assert seen == {"get": "llm_response", "set": "llm_response"}
assert current_service_target() is None
def _completion_logging_obj(call_type: str) -> LiteLLMLogging:
return LiteLLMLogging(
litellm_call_id=str(uuid.uuid4()),
call_type=call_type,
model="gpt-3.5-turbo",
messages=[],
function_id=str(uuid.uuid4()),
stream=False,
start_time=_FIXED_START,
)
_FIXED_START = datetime(2026, 1, 1)
def _unique_messages() -> list[dict[str, str]]:
return [{"role": "user", "content": f"no choices {uuid.uuid4()}"}]
async def aanthropic_messages(**kwargs: object) -> None:
return None
def _responses_api_response_without_output() -> ResponsesAPIResponse:
return ResponsesAPIResponse(
id=f"resp_{uuid.uuid4()}", created_at=0, status="incomplete", model="gpt-4o", object="response", output=[]
)
def _anthropic_message_without_content() -> dict[str, object]:
return {
"id": f"msg_{uuid.uuid4()}",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-5",
"content": [],
}
def _responses_api_response_missing_output() -> ResponsesAPIResponse:
return ResponsesAPIResponse.model_construct(
id=f"resp_{uuid.uuid4()}", created_at=0, status="incomplete", model="gpt-4o", object="response"
)
def _anthropic_message_with_null_content() -> dict[str, object]:
return {
"id": f"msg_{uuid.uuid4()}",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-5",
"content": None,
}
def _anthropic_message_missing_content() -> dict[str, object]:
return {"id": f"msg_{uuid.uuid4()}", "type": "message", "role": "assistant", "model": "claude-sonnet-5"}
def _anthropic_stream_without_content_blocks() -> dict[str, object]:
return {
CACHED_STREAM_EVENTS_KEY: [
'event: message_start\ndata: {"type": "message_start", "message": {"id": "msg_1", "content": []}}\n\n',
'event: message_delta\ndata: {"type": "message_delta", "delta": {"stop_reason": "end_turn"}}\n\n',
'event: message_stop\ndata: {"type": "message_stop"}\n\n',
]
}
EmptyResult = ModelResponse | TextCompletionResponse | ResponsesAPIResponse | dict[str, object]
EmptyResponseCase = tuple[EmptyResult, Callable[..., Awaitable[object]], str, dict[str, object]]
def _empty_response_cases() -> list[EmptyResponseCase]:
return [
(litellm.ModelResponse(choices=[]), litellm.acompletion, CallTypes.acompletion.value, {"messages": _unique_messages()}),
(
litellm.TextCompletionResponse(choices=[]),
litellm.atext_completion,
CallTypes.atext_completion.value,
{"prompt": str(uuid.uuid4())},
),
(
_responses_api_response_without_output(),
aresponses,
CallTypes.aresponses.value,
{"input": str(uuid.uuid4())},
),
(
_responses_api_response_missing_output(),
aresponses,
CallTypes.aresponses.value,
{"input": str(uuid.uuid4())},
),
(
_anthropic_message_without_content(),
aanthropic_messages,
CallTypes.aanthropic_messages.value,
{"messages": _unique_messages(), "max_tokens": 16},
),
(
_anthropic_message_with_null_content(),
aanthropic_messages,
CallTypes.aanthropic_messages.value,
{"messages": _unique_messages(), "max_tokens": 16},
),
(
_anthropic_message_missing_content(),
aanthropic_messages,
CallTypes.aanthropic_messages.value,
{"messages": _unique_messages(), "max_tokens": 16},
),
(
_anthropic_stream_without_content_blocks(),
aanthropic_messages,
CallTypes.aanthropic_messages.value,
{"messages": _unique_messages(), "max_tokens": 16, "stream": True},
),
]
@pytest.mark.asyncio
@pytest.mark.parametrize("empty_result, original_function, call_type, kwargs", _empty_response_cases())
async def test_async_set_cache_skips_response_without_output(
empty_result: EmptyResult,
original_function: Callable[..., Awaitable[object]],
call_type: str,
kwargs: dict[str, object],
):
setup_cache()
handler = LLMCachingHandler(original_function=original_function, request_kwargs={}, start_time=_FIXED_START)
await handler.async_set_cache(result=empty_result, original_function=original_function, kwargs=kwargs)
await asyncio.gather(*_PENDING_CACHE_WRITES)
assert await litellm.cache.async_get_cache(**kwargs) is None
lookup = await handler._async_get_cache(
model="gpt-3.5-turbo",
original_function=original_function,
logging_obj=_completion_logging_obj(call_type),
start_time=_FIXED_START,
call_type=call_type,
kwargs=kwargs,
)
assert lookup.cached_result is None
@pytest.mark.asyncio
@pytest.mark.parametrize("empty_result, original_function, call_type, kwargs", _empty_response_cases())
async def test_async_get_cache_treats_stored_response_without_output_as_miss(
empty_result: EmptyResult,
original_function: Callable[..., Awaitable[object]],
call_type: str,
kwargs: dict[str, object],
):
setup_cache()
stored = empty_result.model_dump_json() if hasattr(empty_result, "model_dump_json") else empty_result
await litellm.cache.async_add_cache(stored, **kwargs)
assert await litellm.cache.async_get_cache(**kwargs) is not None
handler = LLMCachingHandler(original_function=original_function, request_kwargs={}, start_time=_FIXED_START)
lookup = await handler._async_get_cache(
model="gpt-3.5-turbo",
original_function=original_function,
logging_obj=_completion_logging_obj(call_type),
start_time=_FIXED_START,
call_type=call_type,
kwargs=kwargs,
)
assert lookup.cached_result is None
@pytest.mark.asyncio
async def test_async_get_cache_heals_stored_completion_without_choices():
setup_cache()
handler = LLMCachingHandler(original_function=litellm.acompletion, request_kwargs={}, start_time=_FIXED_START)
kwargs = {"messages": _unique_messages()}
poisoned = litellm.ModelResponse(
choices=[], usage=litellm.Usage(prompt_tokens=7, completion_tokens=0, total_tokens=7)
)
await litellm.cache.async_add_cache(poisoned.model_dump_json(), **kwargs)
assert await litellm.cache.async_get_cache(**kwargs) is not None
async def lookup():
return await handler._async_get_cache(
model="gpt-3.5-turbo",
original_function=litellm.acompletion,
logging_obj=_completion_logging_obj(CallTypes.acompletion.value),
start_time=_FIXED_START,
call_type=CallTypes.acompletion.value,
kwargs=kwargs,
)
assert (await lookup()).cached_result is None
await handler.async_set_cache(result=chat_completion_response, original_function=litellm.acompletion, kwargs=kwargs)
await asyncio.gather(*_PENDING_CACHE_WRITES)
healed = (await lookup()).cached_result
assert healed is not None
assert healed.choices[0].message.content == chat_completion_response.choices[0].message.content
def test_sync_set_cache_skips_response_without_output():
setup_cache()
handler = LLMCachingHandler(original_function=completion, request_kwargs={}, start_time=_FIXED_START)
kwargs = {"messages": _unique_messages()}
handler.sync_set_cache(result=litellm.ModelResponse(choices=[]), kwargs=kwargs)
assert litellm.cache.get_cache(**kwargs) is None
lookup = handler._sync_get_cache(
model="gpt-3.5-turbo",
original_function=completion,
logging_obj=_completion_logging_obj(CallTypes.completion.value),
start_time=_FIXED_START,
call_type=CallTypes.completion.value,
kwargs=kwargs,
)
assert lookup.cached_result is None
def test_sync_get_cache_heals_stored_completion_without_choices():
setup_cache()
handler = LLMCachingHandler(original_function=completion, request_kwargs={}, start_time=_FIXED_START)
kwargs = {"messages": _unique_messages()}
litellm.cache.add_cache(litellm.ModelResponse(choices=[]).model_dump_json(), **kwargs)
assert litellm.cache.get_cache(**kwargs) is not None
def lookup():
return handler._sync_get_cache(
model="gpt-3.5-turbo",
original_function=completion,
logging_obj=_completion_logging_obj(CallTypes.completion.value),
start_time=_FIXED_START,
call_type=CallTypes.completion.value,
kwargs=kwargs,
)
assert lookup().cached_result is None
handler.sync_set_cache(result=chat_completion_response, kwargs=kwargs)
healed = lookup().cached_result
assert healed is not None
assert healed.choices[0].message.content == chat_completion_response.choices[0].message.content
@pytest.mark.asyncio
async def test_acompletion_after_a_response_without_choices_calls_the_provider_again():
setup_cache()
messages = _unique_messages()
first = await litellm.acompletion(
model="gpt-4o", messages=messages, mock_response=litellm.ModelResponse(choices=[]), caching=True
)
await asyncio.gather(*_PENDING_CACHE_WRITES)
assert first.choices == []
second = await litellm.acompletion(model="gpt-4o", messages=messages, mock_response="hi", caching=True)
assert second.choices[0].message.content == "hi"
assert second._hidden_params.get("cache_hit") is not True
def _stream_chunk(choices: list[StreamingChoices]) -> ModelResponseStream:
return ModelResponseStream(
id="chatcmpl-stream", created=0, model="gpt-4o", object="chat.completion.chunk", choices=choices
)
def _closing_chunk_without_output() -> ModelResponseStream:
return _stream_chunk([StreamingChoices(finish_reason="stop", index=0, delta=Delta())])
def _cached_content(cached: object) -> object:
assert cached is not None
stored = cached if isinstance(cached, dict) else json.loads(cached)
return stored["choices"][0]["message"]["content"]
def _content_chunk(content: str) -> ModelResponseStream:
return _stream_chunk([StreamingChoices(index=0, delta=Delta(content=content, role="assistant"))])
@pytest.mark.asyncio
async def test_async_streamed_answer_without_output_is_never_cached():
setup_cache()
kwargs = {"model": "gpt-4o", "messages": _unique_messages()}
handler = LLMCachingHandler(original_function=litellm.acompletion, request_kwargs=kwargs, start_time=_FIXED_START)
await handler._add_streaming_response_to_cache(_closing_chunk_without_output())
await asyncio.gather(*_PENDING_CACHE_WRITES)
assert litellm.cache.get_cache(**kwargs) is None
@pytest.mark.asyncio
async def test_async_streamed_answer_with_content_is_cached():
setup_cache()
kwargs = {"model": "gpt-4o", "messages": _unique_messages()}
handler = LLMCachingHandler(original_function=litellm.acompletion, request_kwargs=kwargs, start_time=_FIXED_START)
await handler._add_streaming_response_to_cache(_content_chunk("hi"))
await handler._add_streaming_response_to_cache(_closing_chunk_without_output())
await asyncio.gather(*_PENDING_CACHE_WRITES)
assert _cached_content(litellm.cache.get_cache(**kwargs)) == "hi"
def test_sync_streamed_answer_without_output_is_never_cached():
setup_cache()
kwargs = {"model": "gpt-4o", "messages": _unique_messages()}
handler = LLMCachingHandler(original_function=completion, request_kwargs=kwargs, start_time=_FIXED_START)
handler._sync_add_streaming_response_to_cache(_closing_chunk_without_output())
assert litellm.cache.get_cache(**kwargs) is None
def test_sync_streamed_answer_with_content_is_cached():
setup_cache()
kwargs = {"model": "gpt-4o", "messages": _unique_messages()}
handler = LLMCachingHandler(original_function=completion, request_kwargs=kwargs, start_time=_FIXED_START)
handler._sync_add_streaming_response_to_cache(_content_chunk("hi"))
handler._sync_add_streaming_response_to_cache(_closing_chunk_without_output())
assert _cached_content(litellm.cache.get_cache(**kwargs)) == "hi"
@pytest.mark.asyncio
async def test_async_get_cache_forgets_the_worker_copy_of_a_stored_response_without_output():
setup_cache()
handler = LLMCachingHandler(original_function=litellm.acompletion, request_kwargs={}, start_time=_FIXED_START)
handler.dual_cache = DualCache(in_memory_cache=InMemoryCache())
kwargs = {"messages": _unique_messages()}
key = litellm.cache.get_cache_key(**kwargs)
poisoned = litellm.ModelResponse(choices=[]).model_dump_json()
await handler.dual_cache.async_set_cache(key, {"timestamp": _FIXED_START.timestamp(), "response": poisoned})
assert await handler.dual_cache.async_get_cache(key) is not None
lookup = await handler._async_get_cache(
model="gpt-3.5-turbo",
original_function=litellm.acompletion,
logging_obj=_completion_logging_obj(CallTypes.acompletion.value),
start_time=_FIXED_START,
call_type=CallTypes.acompletion.value,
kwargs=kwargs,
)
assert lookup.cached_result is None
assert await handler.dual_cache.async_get_cache(key) is None

View file

@ -160,6 +160,26 @@ async def test_streaming_cache_is_not_shared_with_non_streaming(local_cache, req
assert non_streaming["content"][0]["text"] == "ALPHA"
@pytest.mark.asyncio
async def test_stream_without_content_blocks_is_not_cached(local_cache, request_kwargs, monkeypatch):
empty_events = [STREAM_EVENTS[0], STREAM_EVENTS[4], STREAM_EVENTS[5]]
fake_handler = _CountingHandler(
[_byte_stream(empty_events), _byte_stream(STREAM_EVENTS), _byte_stream([b"event: never_used\n\n"])]
)
monkeypatch.setattr(handler, "anthropic_messages_handler", fake_handler)
empty = await _collect(await litellm.anthropic_messages(**request_kwargs, stream=True))
await asyncio.sleep(0)
refilled = await _collect(await litellm.anthropic_messages(**request_kwargs, stream=True))
await asyncio.sleep(0)
replayed = await _collect(await litellm.anthropic_messages(**request_kwargs, stream=True))
assert empty == empty_events
assert refilled == STREAM_EVENTS
assert replayed == STREAM_EVENTS
assert len(fake_handler.calls) == 2
@pytest.mark.asyncio
async def test_failed_stream_is_not_cached(local_cache, request_kwargs, monkeypatch):
error_events = STREAM_EVENTS[:3] + [

View file

@ -646,7 +646,15 @@ def test_stream_cache_write_completes_when_asyncio_run_closes_the_loop(monkeypat
status="completed",
model="test-model",
object="response",
output=[],
output=[
{
"type": "message",
"id": "msg_lit6184",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "cached", "annotations": []}],
}
],
),
)
monkeypatch.setattr(litellm, "cache", _SlowWriteCache())
@ -1261,6 +1269,25 @@ def test_billed_terminal_response_copies_when_estimating_and_leaves_the_original
assert response.usage is None
def test_persist_completed_response_to_cache_skips_a_response_without_output(monkeypatch):
logging_obj: Final = _logging_obj_stub()
caching_handler: Final = Mock()
caching_handler.request_kwargs = {"stream": True}
logging_obj._llm_caching_handler = caching_handler
iterator: Final = _make_iterator(sse_events=[], logging_obj=logging_obj)
iterator.completed_response = ResponseCompletedEvent.model_construct(
type="response.completed",
response=ResponsesAPIResponse.model_construct(id="resp_empty", output=[], usage=None),
)
cache: Final = Mock()
monkeypatch.setattr(litellm, "cache", cache)
iterator._persist_completed_response_to_cache(is_async=False)
cache.add_cache.assert_not_called()
caching_handler._should_store_result_in_cache.assert_not_called()
def test_persist_completed_response_to_cache_survives_an_unserializable_response(monkeypatch):
bad_response: Final = ResponsesAPIResponse.model_construct(id="r", output=[object()], usage=None)
with pytest.raises(PydanticSerializationError):