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):