From abd3af422d10b31c79ca41c61a7e86244f5cfed2 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 6 Oct 2026 20:07:30 +0000 Subject: [PATCH] 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> --- litellm/caching/caching_handler.py | 91 +- .../pass_through/messages/response_cache.py | 8 +- litellm/responses/streaming_iterator.py | 6 +- litellm/types/caching.py | 2 + tests/integration/_support/manifest.py | 1 + .../caching/response_cache_case.py | 364 ++++++++ .../caching/test_response_cache_chaos.py | 503 +++++++++++ .../test_response_cache_empty_output.py | 813 ++++++++++++++++++ tests/integration/run.py | 2 +- .../sdk/test_response_cache_sync_sdk.py | 125 +++ tests/unit/caching/test_caching_handler.py | 358 ++++++++ .../messages/test_response_cache.py | 20 + .../unit/responses/test_streaming_iterator.py | 29 +- 13 files changed, 2298 insertions(+), 24 deletions(-) create mode 100644 tests/integration/caching/response_cache_case.py create mode 100644 tests/integration/caching/test_response_cache_chaos.py create mode 100644 tests/integration/caching/test_response_cache_empty_output.py create mode 100644 tests/integration/sdk/test_response_cache_sync_sdk.py diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 9d02c84f6e8..a0bae6685de 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -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, diff --git a/litellm/llms/anthropic/pass_through/messages/response_cache.py b/litellm/llms/anthropic/pass_through/messages/response_cache.py index 5dc26c934ab..40dcbd9c8c8 100644 --- a/litellm/llms/anthropic/pass_through/messages/response_cache.py +++ b/litellm/llms/anthropic/pass_through/messages/response_cache.py @@ -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: diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 4154452af05..eb7a6556386 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -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, diff --git a/litellm/types/caching.py b/litellm/types/caching.py index 036f7c452ea..b9bcbbb93dd 100644 --- a/litellm/types/caching.py +++ b/litellm/types/caching.py @@ -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""" diff --git a/tests/integration/_support/manifest.py b/tests/integration/_support/manifest.py index fb17ebb3864..47c02470f13 100644 --- a/tests/integration/_support/manifest.py +++ b/tests/integration/_support/manifest.py @@ -18,6 +18,7 @@ OWNED_DIRECTORIES: Final = frozenset( "compatibility", "sdk", "cost_calculation", + "caching", "security", } ) diff --git a/tests/integration/caching/response_cache_case.py b/tests/integration/caching/response_cache_case.py new file mode 100644 index 00000000000..9df34721202 --- /dev/null +++ b/tests/integration/caching/response_cache_case.py @@ -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"] diff --git a/tests/integration/caching/test_response_cache_chaos.py b/tests/integration/caching/test_response_cache_chaos.py new file mode 100644 index 00000000000..2e663c97d1f --- /dev/null +++ b/tests/integration/caching/test_response_cache_chaos.py @@ -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" + ) diff --git a/tests/integration/caching/test_response_cache_empty_output.py b/tests/integration/caching/test_response_cache_empty_output.py new file mode 100644 index 00000000000..ce5a4d64ca7 --- /dev/null +++ b/tests/integration/caching/test_response_cache_empty_output.py @@ -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 diff --git a/tests/integration/run.py b/tests/integration/run.py index df408cfc809..f7b2dead197 100644 --- a/tests/integration/run.py +++ b/tests/integration/run.py @@ -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"), diff --git a/tests/integration/sdk/test_response_cache_sync_sdk.py b/tests/integration/sdk/test_response_cache_sync_sdk.py new file mode 100644 index 00000000000..54d78adbaba --- /dev/null +++ b/tests/integration/sdk/test_response_cache_sync_sdk.py @@ -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 diff --git a/tests/unit/caching/test_caching_handler.py b/tests/unit/caching/test_caching_handler.py index 1599668839a..abbd44384b1 100644 --- a/tests/unit/caching/test_caching_handler.py +++ b/tests/unit/caching/test_caching_handler.py @@ -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 diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py b/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py index a8a8eba0bf7..8690d2ee68f 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py @@ -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] + [ diff --git a/tests/unit/responses/test_streaming_iterator.py b/tests/unit/responses/test_streaming_iterator.py index 2f6dccb37f3..9e9d221ce84 100644 --- a/tests/unit/responses/test_streaming_iterator.py +++ b/tests/unit/responses/test_streaming_iterator.py @@ -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):