fix(cache): persist and replay streamed Responses API requests

This commit is contained in:
Noah 2026-03-31 14:59:24 +01:00
parent d251238bd7
commit 51e5175c6f
7 changed files with 1484 additions and 252 deletions

View file

@ -429,9 +429,10 @@ class Cache:
str: The final hashed cache key with the redis namespace.
"""
dynamic_cache_control: DynamicCacheControl = kwargs.get("cache", {})
metadata = kwargs.get("metadata") or {}
namespace = (
dynamic_cache_control.get("namespace")
or kwargs.get("metadata", {}).get("redis_namespace")
or metadata.get("redis_namespace")
or self.namespace
)
if namespace:

View file

@ -82,6 +82,19 @@ class CachingHandlerResponse(BaseModel):
in_memory_cache_obj = InMemoryCache()
_RESPONSES_STREAMING_CALLBACK_CALL_TYPES = {
CallTypes.aresponses.value,
CallTypes.responses.value,
}
def _should_defer_streaming_cache_hit_callbacks(
*, call_type: str, kwargs: Dict[str, Any]
) -> bool:
return (
kwargs.get("stream", False) is True
and call_type in _RESPONSES_STREAMING_CALLBACK_CALL_TYPES
)
class LLMCachingHandler:
@ -96,6 +109,7 @@ class LLMCachingHandler:
self.async_streaming_chunks: List[ModelResponse] = []
self.sync_streaming_chunks: List[ModelResponse] = []
self.request_kwargs = request_kwargs
self.preset_cache_key: Optional[str] = None
self.original_function = original_function
self.start_time = start_time
if litellm.cache is not None and isinstance(litellm.cache.cache, RedisCache):
@ -203,7 +217,10 @@ class LLMCachingHandler:
custom_llm_provider=kwargs.get("custom_llm_provider", None),
args=args,
)
if kwargs.get("stream", False) is False:
if not _should_defer_streaming_cache_hit_callbacks(
call_type=call_type,
kwargs=kwargs,
):
# LOG SUCCESS
self._async_log_cache_hit_on_callbacks(
logging_obj=logging_obj,
@ -212,11 +229,12 @@ class LLMCachingHandler:
end_time=end_time,
cache_hit=cache_hit,
)
cache_key = litellm.cache.get_cache_key(**kwargs)
if (
isinstance(cached_result, BaseModel)
or isinstance(cached_result, CustomStreamWrapper)
) and hasattr(cached_result, "_hidden_params"):
cache_key = (
self.preset_cache_key
or self.request_kwargs.get("cache_key")
or litellm.cache.get_cache_key(**self.request_kwargs)
)
if hasattr(cached_result, "_hidden_params"):
cached_result._hidden_params["cache_key"] = cache_key # type: ignore
return CachingHandlerResponse(cached_result=cached_result)
elif (
@ -262,8 +280,6 @@ class LLMCachingHandler:
kwargs: Dict[str, Any],
args: Optional[Tuple[Any, ...]] = None,
) -> CachingHandlerResponse:
from litellm.utils import CustomStreamWrapper
cached_result: Optional[Any] = None
# Check if caching should be performed BEFORE doing expensive kwargs copy
@ -279,6 +295,11 @@ class LLMCachingHandler:
args,
)
)
if new_kwargs.get("metadata") is None:
new_kwargs.pop("metadata", None)
if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs:
new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs)
self.request_kwargs = new_kwargs
print_verbose("Checking Sync Cache")
cached_result = litellm.cache.get_cache(**new_kwargs)
if cached_result is not None:
@ -319,17 +340,22 @@ class LLMCachingHandler:
is_async=False,
)
logging_obj.handle_sync_success_callbacks_for_async_calls(
result=cached_result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
if not _should_defer_streaming_cache_hit_callbacks(
call_type=call_type,
kwargs=kwargs,
):
logging_obj.handle_sync_success_callbacks_for_async_calls(
result=cached_result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
)
cache_key = (
self.preset_cache_key
or self.request_kwargs.get("cache_key")
or litellm.cache.get_cache_key(**self.request_kwargs)
)
cache_key = litellm.cache.get_cache_key(**kwargs)
if (
isinstance(cached_result, BaseModel)
or isinstance(cached_result, CustomStreamWrapper)
) and hasattr(cached_result, "_hidden_params"):
if hasattr(cached_result, "_hidden_params"):
cached_result._hidden_params["cache_key"] = cache_key # type: ignore
return CachingHandlerResponse(cached_result=cached_result)
return CachingHandlerResponse(cached_result=cached_result)
@ -596,6 +622,11 @@ class LLMCachingHandler:
args,
)
)
if new_kwargs.get("metadata") is None:
new_kwargs.pop("metadata", None)
if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs:
new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs)
self.request_kwargs = new_kwargs
cached_result: Optional[Any] = None
if call_type == CallTypes.aembedding.value:
if isinstance(new_kwargs["input"], str):
@ -620,14 +651,26 @@ class LLMCachingHandler:
if all(result is None for result in cached_result):
cached_result = None
else:
request_kwargs = new_kwargs.copy()
request_cache_key = request_kwargs.pop("cache_key", None)
if litellm.cache._supports_async() is True:
## check if dual cache is supported ##
self.preset_cache_key = (
request_cache_key or litellm.cache.get_cache_key(**request_kwargs)
)
cached_result = await litellm.cache.async_get_cache(
dynamic_cache_object=self.dual_cache, **new_kwargs
dynamic_cache_object=self.dual_cache,
cache_key=self.preset_cache_key,
**request_kwargs,
)
else: # fallback for caches that don't support async
self.preset_cache_key = (
request_cache_key or litellm.cache.get_cache_key(**request_kwargs)
)
cached_result = litellm.cache.get_cache(
dynamic_cache_object=self.dual_cache, **new_kwargs
dynamic_cache_object=self.dual_cache,
cache_key=self.preset_cache_key,
**request_kwargs,
)
return cached_result
@ -735,8 +778,27 @@ class LLMCachingHandler:
elif (call_type == "aresponses" or call_type == "responses") and isinstance(
cached_result, dict
):
# Convert cached dict back to ResponsesAPIResponse object
cached_result = ResponsesAPIResponse(**cached_result)
from litellm.responses.streaming_iterator import (
CachedResponsesAPIStreamingIterator,
)
response_obj = ResponsesAPIResponse(**cached_result)
if (
hasattr(response_obj, "_hidden_params")
and response_obj._hidden_params is not None
and isinstance(response_obj._hidden_params, dict)
):
response_obj._hidden_params["cache_hit"] = True
if kwargs.get("stream", False) is True:
cached_result = CachedResponsesAPIStreamingIterator(
response=response_obj,
logging_obj=logging_obj,
request_data=kwargs,
call_type=call_type,
)
else:
cached_result = response_obj
if (
hasattr(cached_result, "_hidden_params")

View file

@ -1,9 +1,12 @@
from __future__ import annotations
import asyncio
import json
import time
import traceback
from datetime import datetime
from typing import Any, Dict, List, Optional
from functools import lru_cache
from typing import Any, Dict, List, Literal, Optional
import httpx
@ -22,19 +25,25 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
from litellm.litellm_core_utils.thread_pool_executor import executor
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import (
OutputTextDeltaEvent,
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIRequestParams,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
ResponsesAPIStreamingResponse,
)
from litellm.types.utils import CallTypes
from litellm.utils import CustomStreamWrapper, async_post_call_success_deployment_hook
@lru_cache(maxsize=1)
def _get_openai_response_types():
from litellm.types.llms import openai as openai_types
return openai_types
def _log_background_task_failure(task: "asyncio.Task[Any]", *, task_name: str) -> None:
if task.cancelled():
return
exception = task.exception()
if exception is not None:
verbose_logger.error("%s failed: %s", task_name, exception)
class BaseResponsesAPIStreamingIterator:
"""
Base class for streaming iterators that process responses from the Responses API.
@ -46,7 +55,7 @@ class BaseResponsesAPIStreamingIterator:
self,
response: httpx.Response,
model: str,
responses_api_provider_config: BaseResponsesAPIConfig,
responses_api_provider_config: Optional[BaseResponsesAPIConfig],
logging_obj: LiteLLMLoggingObj,
litellm_metadata: Optional[Dict[str, Any]] = None,
custom_llm_provider: Optional[str] = None,
@ -58,9 +67,13 @@ class BaseResponsesAPIStreamingIterator:
self.logging_obj = logging_obj
self.finished = False
self.responses_api_provider_config = responses_api_provider_config
self.completed_response: Optional[ResponsesAPIStreamingResponse] = None
self.completed_response: Optional[Any] = None
self.start_time = getattr(logging_obj, "start_time", datetime.now())
self._failure_handled = False # Track if failure handler has been called
self._completed_response_cached = False
self._completed_response_logged = False
self._completed_response_cache_hit: Optional[bool] = None
self._persist_completed_response_before_logging = True
self._stream_created_time: float = time.time()
# track request context for hooks
@ -101,7 +114,7 @@ class BaseResponsesAPIStreamingIterator:
llm_provider=self.custom_llm_provider or "",
)
def _process_chunk(self, chunk) -> Optional[ResponsesAPIStreamingResponse]:
def _process_chunk(self, chunk) -> Optional[Any]:
"""Process a single chunk of data from the stream"""
if not chunk:
return None
@ -122,6 +135,10 @@ class BaseResponsesAPIStreamingIterator:
# Format as ResponsesAPIStreamingResponse
if isinstance(parsed_chunk, dict):
if self.responses_api_provider_config is None:
raise ValueError(
"responses_api_provider_config is required to process live streaming chunks"
)
openai_responses_api_chunk = (
self.responses_api_provider_config.transform_streaming_response(
model=self.model,
@ -144,10 +161,11 @@ class BaseResponsesAPIStreamingIterator:
if self.litellm_metadata and self.litellm_metadata.get(
"encrypted_content_affinity_enabled"
):
openai_types = _get_openai_response_types()
event_type = getattr(openai_responses_api_chunk, "type", None)
if event_type in (
ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
):
item = getattr(openai_responses_api_chunk, "item", None)
if item:
@ -168,10 +186,11 @@ class BaseResponsesAPIStreamingIterator:
# Store the completed response (also for incomplete/failed so logging still fires)
_chunk_type = getattr(openai_responses_api_chunk, "type", None)
openai_types = _get_openai_response_types()
if openai_responses_api_chunk and _chunk_type in (
ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE,
ResponsesAPIStreamEvents.RESPONSE_FAILED,
openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE,
openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED,
):
self.completed_response = openai_responses_api_chunk
# Add cost to usage object if include_cost_in_streaming_usage is True
@ -179,11 +198,11 @@ class BaseResponsesAPIStreamingIterator:
litellm.include_cost_in_streaming_usage
and self.logging_obj is not None
):
response_obj: Optional[ResponsesAPIResponse] = getattr(
response_obj: Optional[Any] = getattr(
openai_responses_api_chunk, "response", None
)
if response_obj:
usage_obj: Optional[ResponseAPIUsage] = getattr(
usage_obj: Optional[Any] = getattr(
response_obj, "usage", None
)
if usage_obj is not None:
@ -196,9 +215,13 @@ class BaseResponsesAPIStreamingIterator:
if cost is not None:
setattr(usage_obj, "cost", cost)
except Exception:
# Best-effort usage cost annotation should not break stream replay.
pass
if _chunk_type == ResponsesAPIStreamEvents.RESPONSE_FAILED:
if (
_chunk_type
== openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED
):
self._handle_logging_failed_response()
else:
self._handle_logging_completed_response()
@ -215,6 +238,59 @@ class BaseResponsesAPIStreamingIterator:
self._handle_failure(e)
raise
def _log_completed_response(self, *, is_async: bool) -> None:
if self._completed_response_logged:
return
self._completed_response_logged = True
if self._persist_completed_response_before_logging:
self._persist_completed_response_to_cache(is_async=is_async)
# Create a copy for logging to avoid modifying the response object that will be returned to the user
# The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
# to chat completion format (prompt_tokens/completion_tokens) for internal logging
# Use model_dump + model_validate instead of deepcopy to avoid pickle errors with
# Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192)
logging_response = self.completed_response
if self.completed_response is not None and hasattr(
self.completed_response, "model_dump"
):
try:
logging_response = type(self.completed_response).model_validate(
self.completed_response.model_dump()
)
except Exception:
# Fallback to original if serialization fails
pass
end_time = datetime.now()
if is_async:
asyncio.create_task(
self.logging_obj.async_success_handler(
result=logging_response,
start_time=self.start_time,
end_time=end_time,
cache_hit=self._completed_response_cache_hit,
)
)
else:
run_async_function(
async_function=self.logging_obj.async_success_handler,
result=logging_response,
start_time=self.start_time,
end_time=end_time,
cache_hit=self._completed_response_cache_hit,
)
executor.submit(
self.logging_obj.success_handler,
result=logging_response,
cache_hit=self._completed_response_cache_hit,
start_time=self.start_time,
end_time=end_time,
)
self._run_post_success_hooks(end_time=end_time)
def _handle_logging_completed_response(self):
"""Base implementation - should be overridden by subclasses"""
pass
@ -245,6 +321,88 @@ class BaseResponsesAPIStreamingIterator:
)
self._handle_failure(exception)
def _get_completed_response_object(self) -> Optional[Any]:
openai_types = _get_openai_response_types()
completed_response = self.completed_response
if isinstance(completed_response, openai_types.ResponsesAPIResponse):
return completed_response
response_obj = getattr(completed_response, "response", None)
if isinstance(response_obj, openai_types.ResponsesAPIResponse):
return response_obj
return None
def _persist_completed_response_to_cache(self, *, is_async: bool) -> None:
if self._completed_response_cached:
return
completed_response = self.completed_response
openai_types = _get_openai_response_types()
if (
getattr(completed_response, "type", None)
!= openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED
):
return
response_obj = self._get_completed_response_object()
if response_obj is None:
return
caching_handler = getattr(self.logging_obj, "_llm_caching_handler", None)
if caching_handler is None:
return
request_kwargs = getattr(caching_handler, "request_kwargs", None)
if (
not isinstance(request_kwargs, dict)
or request_kwargs.get("stream") is not True
):
return
request_kwargs = request_kwargs.copy()
preset_cache_key = getattr(caching_handler, "preset_cache_key", None)
request_cache_key = request_kwargs.pop("cache_key", None)
if preset_cache_key is None:
preset_cache_key = request_cache_key
if request_kwargs.get("metadata") is None:
request_kwargs.pop("metadata", None)
request_kwargs.pop("custom_llm_provider", None)
if preset_cache_key is not None:
request_kwargs["cache_key"] = preset_cache_key
if not caching_handler._should_store_result_in_cache(
original_function=caching_handler.original_function,
kwargs=request_kwargs,
):
return
if litellm.cache is None:
return
cached_response = response_obj.model_dump_json()
if is_async:
cache_write_task = asyncio.create_task(
litellm.cache.async_add_cache(
cached_response,
dynamic_cache_object=getattr(caching_handler, "dual_cache", None),
**request_kwargs,
)
)
cache_write_task.add_done_callback(
lambda task: _log_background_task_failure(
task,
task_name="Responses stream cache write",
)
)
else:
litellm.cache.add_cache(
cached_response,
dynamic_cache_object=getattr(caching_handler, "dual_cache", None),
**request_kwargs,
)
self._completed_response_cached = True
async def _call_post_streaming_deployment_hook(self, chunk):
"""
Allow callbacks to modify streaming chunks before returning (parity with chat).
@ -429,7 +587,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
def __aiter__(self):
return self
async def __anext__(self) -> ResponsesAPIStreamingResponse:
async def __anext__(self) -> Any:
try:
self._check_max_streaming_duration()
while True:
@ -469,40 +627,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
def _handle_logging_completed_response(self):
"""Handle logging for completed responses in async context"""
# Create a copy for logging to avoid modifying the response object that will be returned to the user
# The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
# to chat completion format (prompt_tokens/completion_tokens) for internal logging
# Use model_dump + model_validate instead of deepcopy to avoid pickle errors with
# Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192)
logging_response = self.completed_response
if self.completed_response is not None and hasattr(
self.completed_response, "model_dump"
):
try:
logging_response = type(self.completed_response).model_validate(
self.completed_response.model_dump()
)
except Exception:
# Fallback to original if serialization fails
pass
asyncio.create_task(
self.logging_obj.async_success_handler(
result=logging_response,
start_time=self.start_time,
end_time=datetime.now(),
cache_hit=None,
)
)
executor.submit(
self.logging_obj.success_handler,
result=logging_response,
cache_hit=None,
start_time=self.start_time,
end_time=datetime.now(),
)
self._run_post_success_hooks(end_time=datetime.now())
self._log_completed_response(is_async=True)
class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
@ -576,39 +701,7 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
def _handle_logging_completed_response(self):
"""Handle logging for completed responses in sync context"""
# Create a copy for logging to avoid modifying the response object that will be returned to the user
# The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
# to chat completion format (prompt_tokens/completion_tokens) for internal logging
# Use model_dump + model_validate instead of deepcopy to avoid pickle errors with
# Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192)
logging_response = self.completed_response
if self.completed_response is not None and hasattr(
self.completed_response, "model_dump"
):
try:
logging_response = type(self.completed_response).model_validate(
self.completed_response.model_dump()
)
except Exception:
# Fallback to original if serialization fails
pass
run_async_function(
async_function=self.logging_obj.async_success_handler,
result=logging_response,
start_time=self.start_time,
end_time=datetime.now(),
cache_hit=None,
)
executor.submit(
self.logging_obj.success_handler,
result=logging_response,
cache_hit=None,
start_time=self.start_time,
end_time=datetime.now(),
)
self._run_post_success_hooks(end_time=datetime.now())
self._log_completed_response(is_async=False)
class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
@ -632,90 +725,441 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
request_data: Optional[Dict[str, Any]] = None,
call_type: Optional[str] = None,
):
super().__init__(
response=response,
transformed = responses_api_provider_config.transform_response_api_response(
model=model,
responses_api_provider_config=responses_api_provider_config,
raw_response=response,
logging_obj=logging_obj,
)
super().__init__(
response=httpx.Response(200),
model=model,
responses_api_provider_config=None,
logging_obj=logging_obj,
litellm_metadata=litellm_metadata,
custom_llm_provider=custom_llm_provider,
request_data=request_data,
call_type=call_type,
)
self._set_events_from_response(transformed=transformed, logging_obj=logging_obj)
# one-time transform
transformed = (
self.responses_api_provider_config.transform_response_api_response(
model=self.model,
raw_response=response,
logging_obj=logging_obj,
)
def _set_events_from_response(
self,
transformed: Any,
logging_obj: LiteLLMLoggingObj,
) -> None:
self._events = _build_synthetic_response_events(
transformed=transformed,
logging_obj=logging_obj,
chunk_size=self.CHUNK_SIZE,
)
full_text = self._collect_text(transformed)
# build a list of 5‑char delta events
deltas = [
OutputTextDeltaEvent(
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
delta=full_text[i : i + self.CHUNK_SIZE],
item_id=transformed.id,
output_index=0,
content_index=0,
)
for i in range(0, len(full_text), self.CHUNK_SIZE)
]
# Add cost to usage object if include_cost_in_streaming_usage is True
if litellm.include_cost_in_streaming_usage and logging_obj is not None:
usage_obj: Optional[ResponseAPIUsage] = getattr(transformed, "usage", None)
if usage_obj is not None:
try:
cost: Optional[float] = logging_obj._response_cost_calculator(
result=transformed
)
if cost is not None:
setattr(usage_obj, "cost", cost)
except Exception:
# If cost calculation fails, continue without cost
pass
# append the completed event
self._events = deltas + [
ResponseCompletedEvent(
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
response=transformed,
)
]
self._idx = 0
self.completed_response = self._events[-1]
def __aiter__(self):
return self
async def __anext__(self) -> ResponsesAPIStreamingResponse:
async def __anext__(self) -> Any:
if self._idx >= len(self._events):
raise StopAsyncIteration
evt = self._events[self._idx]
self._idx += 1
openai_types = _get_openai_response_types()
if (
getattr(evt, "type", None)
== openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED
):
self.completed_response = evt
self._log_completed_response(is_async=True)
return evt
def __iter__(self):
return self
def __next__(self) -> ResponsesAPIStreamingResponse:
def __next__(self) -> Any:
if self._idx >= len(self._events):
raise StopIteration
evt = self._events[self._idx]
self._idx += 1
openai_types = _get_openai_response_types()
if (
getattr(evt, "type", None)
== openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED
):
self.completed_response = evt
self._log_completed_response(is_async=False)
return evt
def _collect_text(self, resp: ResponsesAPIResponse) -> str:
out = ""
for out_item in resp.output:
item_type = getattr(out_item, "type", None)
if item_type == "message":
for c in getattr(out_item, "content", []):
out += c.text
return out
class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
def __init__(
self,
response: Any,
logging_obj: LiteLLMLoggingObj,
request_data: Optional[Dict[str, Any]] = None,
call_type: Optional[str] = None,
):
BaseResponsesAPIStreamingIterator.__init__(
self,
response=httpx.Response(200),
model=getattr(response, "model", ""),
responses_api_provider_config=None,
logging_obj=logging_obj,
litellm_metadata=None,
custom_llm_provider="cached_response",
request_data=request_data,
call_type=call_type,
)
self._completed_response_cache_hit = True
self._persist_completed_response_before_logging = False
self._events: List[Any] = []
self._idx = 0
self._set_events_from_response(transformed=response, logging_obj=logging_obj)
def _set_events_from_response(
self,
transformed: Any,
logging_obj: LiteLLMLoggingObj,
) -> None:
self._events = _build_synthetic_response_events(
transformed=transformed,
logging_obj=logging_obj,
chunk_size=MockResponsesAPIStreamingIterator.CHUNK_SIZE,
)
self._idx = 0
self.completed_response = self._events[-1]
def __aiter__(self):
return self
async def __anext__(self) -> Any:
if self._idx >= len(self._events):
raise StopAsyncIteration
evt = self._events[self._idx]
self._idx += 1
openai_types = _get_openai_response_types()
if (
getattr(evt, "type", None)
== openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED
):
self.completed_response = evt
self._log_completed_response(is_async=True)
return evt
def __iter__(self):
return self
def __next__(self) -> Any:
if self._idx >= len(self._events):
raise StopIteration
evt = self._events[self._idx]
self._idx += 1
openai_types = _get_openai_response_types()
if (
getattr(evt, "type", None)
== openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED
):
self.completed_response = evt
self._log_completed_response(is_async=False)
return evt
def _dump_response_object(obj: Any) -> Dict[str, Any]:
if hasattr(obj, "model_dump"):
return obj.model_dump()
if isinstance(obj, dict):
return obj
return {}
def _build_response_status_event(
event_type: Literal[
"response.created",
"response.in_progress",
],
transformed: Any,
) -> Any:
openai_types = _get_openai_response_types()
in_progress_response = transformed.model_copy(
deep=True,
update={"status": "in_progress", "output": []},
)
if event_type == openai_types.ResponsesAPIStreamEvents.RESPONSE_CREATED:
return openai_types.ResponseCreatedEvent(
type=event_type, response=in_progress_response
)
return openai_types.ResponseInProgressEvent(
type=event_type, response=in_progress_response
)
def _build_content_part_done_event(
*,
item_id: str,
output_index: int,
content_index: int,
part_payload: Dict[str, Any],
) -> Optional[Any]:
openai_types = _get_openai_response_types()
part_type = part_payload.get("type")
part: Any
if part_type == "output_text":
annotations = [
openai_types.BaseLiteLLMOpenAIResponseObject(**annotation)
for annotation in part_payload.get("annotations", []) or []
]
part = openai_types.ContentPartDonePartOutputText(
type="output_text",
text=str(part_payload.get("text") or ""),
annotations=annotations,
logprobs=part_payload.get("logprobs"),
)
elif part_type == "refusal":
part = openai_types.ContentPartDonePartRefusal(
type="refusal",
refusal=str(part_payload.get("refusal") or ""),
)
elif part_type == "reasoning_text":
part = openai_types.ContentPartDonePartReasoningText(
type="reasoning_text",
reasoning=str(part_payload.get("reasoning") or ""),
)
else:
return None
return openai_types.ContentPartDoneEvent(
type=openai_types.ResponsesAPIStreamEvents.CONTENT_PART_DONE,
item_id=item_id,
output_index=output_index,
content_index=content_index,
part=part,
)
def _add_text_like_part_events(
*,
events: List[Any],
item_id: str,
output_index: int,
content_index: int,
part_payload: Dict[str, Any],
chunk_size: int,
) -> None:
openai_types = _get_openai_response_types()
part_type = part_payload.get("type")
if part_type == "output_text":
text = str(part_payload.get("text") or "")
for i in range(0, len(text), chunk_size):
events.append(
openai_types.OutputTextDeltaEvent(
type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
item_id=item_id,
output_index=output_index,
content_index=content_index,
delta=text[i : i + chunk_size],
)
)
for annotation_index, annotation in enumerate(
part_payload.get("annotations", []) or []
):
events.append(
openai_types.OutputTextAnnotationAddedEvent(
type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED,
item_id=item_id,
output_index=output_index,
content_index=content_index,
annotation_index=annotation_index,
annotation=annotation,
)
)
events.append(
openai_types.OutputTextDoneEvent(
type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE,
item_id=item_id,
output_index=output_index,
content_index=content_index,
text=text,
)
)
elif part_type == "refusal":
refusal = str(part_payload.get("refusal") or "")
for i in range(0, len(refusal), chunk_size):
events.append(
openai_types.RefusalDeltaEvent(
type=openai_types.ResponsesAPIStreamEvents.REFUSAL_DELTA,
item_id=item_id,
output_index=output_index,
content_index=content_index,
delta=refusal[i : i + chunk_size],
)
)
events.append(
openai_types.RefusalDoneEvent(
type=openai_types.ResponsesAPIStreamEvents.REFUSAL_DONE,
item_id=item_id,
output_index=output_index,
content_index=content_index,
refusal=refusal,
)
)
def _build_synthetic_response_events(
*,
transformed: Any,
logging_obj: LiteLLMLoggingObj,
chunk_size: int,
) -> List[Any]:
openai_types = _get_openai_response_types()
if litellm.include_cost_in_streaming_usage and logging_obj is not None:
usage_obj: Optional[Any] = getattr(transformed, "usage", None)
if usage_obj is not None:
try:
cost: Optional[float] = logging_obj._response_cost_calculator(
result=transformed
)
if cost is not None:
setattr(usage_obj, "cost", cost)
except Exception:
pass
events: List[Any] = [
_build_response_status_event(
openai_types.ResponsesAPIStreamEvents.RESPONSE_CREATED, transformed
),
_build_response_status_event(
openai_types.ResponsesAPIStreamEvents.RESPONSE_IN_PROGRESS, transformed
),
]
sequence_number = 0
for output_index, output_item in enumerate(
getattr(transformed, "output", []) or []
):
output_item_payload = _dump_response_object(output_item)
item_id = str(output_item_payload.get("id") or transformed.id)
item_type = output_item_payload.get("type")
events.append(
openai_types.OutputItemAddedEvent(
type=openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
output_index=output_index,
item=openai_types.BaseLiteLLMOpenAIResponseObject(
**output_item_payload
),
)
)
if item_type == "message":
for content_index, part in enumerate(
output_item_payload.get("content", []) or []
):
part_payload = _dump_response_object(part)
events.append(
openai_types.ContentPartAddedEvent(
type=openai_types.ResponsesAPIStreamEvents.CONTENT_PART_ADDED,
item_id=item_id,
output_index=output_index,
content_index=content_index,
part=openai_types.BaseLiteLLMOpenAIResponseObject(
**part_payload
),
)
)
_add_text_like_part_events(
events=events,
item_id=item_id,
output_index=output_index,
content_index=content_index,
part_payload=part_payload,
chunk_size=chunk_size,
)
done_event = _build_content_part_done_event(
item_id=item_id,
output_index=output_index,
content_index=content_index,
part_payload=part_payload,
)
if done_event is not None:
events.append(done_event)
elif item_type == "function_call":
arguments = str(output_item_payload.get("arguments") or "")
for i in range(0, len(arguments), chunk_size):
events.append(
openai_types.FunctionCallArgumentsDeltaEvent(
type=openai_types.ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA,
item_id=item_id,
output_index=output_index,
delta=arguments[i : i + chunk_size],
)
)
events.append(
openai_types.FunctionCallArgumentsDoneEvent(
type=openai_types.ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE,
item_id=item_id,
output_index=output_index,
arguments=arguments,
)
)
elif item_type == "reasoning":
for summary_index, summary in enumerate(
output_item_payload.get("summary", []) or []
):
summary_payload = _dump_response_object(summary)
summary_text = str(summary_payload.get("text") or "")
for i in range(0, len(summary_text), chunk_size):
events.append(
openai_types.ReasoningSummaryTextDeltaEvent(
type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA,
item_id=item_id,
output_index=output_index,
summary_index=summary_index,
delta=summary_text[i : i + chunk_size],
)
)
sequence_number += 1
events.append(
openai_types.ReasoningSummaryTextDoneEvent(
type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DONE,
item_id=item_id,
output_index=output_index,
sequence_number=sequence_number,
summary_index=summary_index,
text=summary_text,
)
)
sequence_number += 1
events.append(
openai_types.ReasoningSummaryPartDoneEvent(
type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_PART_DONE,
item_id=item_id,
output_index=output_index,
sequence_number=sequence_number,
summary_index=summary_index,
part=openai_types.BaseLiteLLMOpenAIResponseObject(
**summary_payload
),
)
)
sequence_number += 1
events.append(
openai_types.OutputItemDoneEvent(
type=openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
output_index=output_index,
sequence_number=sequence_number,
item=openai_types.BaseLiteLLMOpenAIResponseObject(
**output_item_payload
),
)
)
events.append(
openai_types.ResponseCompletedEvent(
type=openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
response=transformed,
)
)
return events
# ---------------------------------------------------------------------------
@ -900,8 +1344,8 @@ class ResponsesWebSocketStreaming:
# ---------------------------------------------------------------------------
_RESPONSE_CREATE_PARAMS: frozenset = (
ResponsesAPIRequestParams.__required_keys__
| ResponsesAPIRequestParams.__optional_keys__
_get_openai_response_types().ResponsesAPIRequestParams.__required_keys__
| _get_openai_response_types().ResponsesAPIRequestParams.__optional_keys__
)
_MANAGED_WS_SKIP_KWARGS: frozenset = frozenset(
@ -1034,7 +1478,7 @@ class ManagedResponsesWebSocketHandler:
@staticmethod
def _extract_output_messages(
completed_event: Dict[str, Any]
completed_event: Dict[str, Any],
) -> List[Dict[str, Any]]:
"""
Convert the output items in a ``response.completed`` event into

View file

@ -1482,6 +1482,7 @@ class ReasoningSummaryTextDeltaEvent(BaseLiteLLMOpenAIResponseObject):
type: Literal[ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA]
item_id: str
output_index: int
summary_index: int = 0
delta: str
@ -1490,7 +1491,7 @@ class ReasoningSummaryTextDoneEvent(BaseLiteLLMOpenAIResponseObject):
item_id: str
output_index: int
sequence_number: int
summary_index: int
summary_index: int = 0
text: str
@ -1499,7 +1500,7 @@ class ReasoningSummaryPartDoneEvent(BaseLiteLLMOpenAIResponseObject):
item_id: str
output_index: int
sequence_number: int
summary_index: int
summary_index: int = 0
part: BaseLiteLLMOpenAIResponseObject

View file

@ -1,6 +1,8 @@
import asyncio
from datetime import datetime
import json
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
@ -8,8 +10,16 @@ import pytest
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.responses import streaming_iterator as streaming_module
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
from litellm.types.llms.openai import ResponsesAPIStreamEvents
from litellm.responses.streaming_iterator import (
CachedResponsesAPIStreamingIterator,
ResponsesAPIStreamingIterator,
SyncResponsesAPIStreamingIterator,
)
from litellm.types.llms.openai import (
ResponseCompletedEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
from litellm.types.utils import CallTypes
@ -19,15 +29,19 @@ class _FakeLoggingObj:
self.async_success_calls = 0
self.failure_calls = 0
self.async_failure_calls = 0
self.last_success_kwargs = None
self.last_async_success_kwargs = None
self.start_time = datetime.now()
self.model_call_details = {"litellm_params": {}}
# Signature alignment with Logging handlers
def success_handler(self, *args, **kwargs):
self.success_calls += 1
self.last_success_kwargs = kwargs
async def async_success_handler(self, *args, **kwargs):
self.async_success_calls += 1
self.last_async_success_kwargs = kwargs
def failure_handler(self, *args, **kwargs):
self.failure_calls += 1
@ -36,6 +50,34 @@ class _FakeLoggingObj:
self.async_failure_calls += 1
def _make_completed_response(response_id: str = "resp_test") -> ResponseCompletedEvent:
return ResponseCompletedEvent(
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
response=ResponsesAPIResponse(
id=response_id,
created_at=int(datetime.now().timestamp()),
status="completed",
model="test-model",
object="response",
output=[
{
"type": "message",
"id": f"msg_{response_id}",
"status": "completed",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": "cached streamed response",
"annotations": [],
}
],
}
],
),
)
@pytest.mark.asyncio
async def test_responses_streaming_triggers_hooks(monkeypatch):
"""
@ -126,8 +168,12 @@ async def test_responses_streaming_calls_post_streaming_deployment_hook(monkeypa
)
# Call hook helper directly to verify chunk is modified/flagged
chunk = SimpleNamespace(type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, response=None)
chunk = await streaming_module.call_post_streaming_hooks_for_testing(iterator, chunk)
chunk = SimpleNamespace(
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, response=None
)
chunk = await streaming_module.call_post_streaming_hooks_for_testing(
iterator, chunk
)
assert getattr(chunk, "_post_streaming_hooks_ran", False) is True
assert getattr(chunk, "tagged", False) is True
@ -163,3 +209,220 @@ async def test_responses_streaming_failure_triggers_failure_handlers():
await asyncio.sleep(0.2)
assert logging_obj.failure_calls >= 1
assert logging_obj.async_failure_calls >= 1
@pytest.mark.asyncio
async def test_responses_streaming_completed_event_persists_async_cache():
logging_obj = _FakeLoggingObj()
original_cache = litellm.cache
litellm.cache = SimpleNamespace(
async_add_cache=AsyncMock(),
add_cache=MagicMock(),
)
caching_handler = SimpleNamespace(
request_kwargs={
"model": "test-model",
"input": "hello",
"stream": True,
"caching": True,
"cache_key": "stale-request-cache-key",
"metadata": None,
"custom_llm_provider": "openai",
},
preset_cache_key="responses-stream-cache-key",
original_function=litellm.aresponses,
async_set_cache=AsyncMock(),
_should_store_result_in_cache=lambda original_function, kwargs: True,
)
logging_obj._llm_caching_handler = caching_handler
iterator = ResponsesAPIStreamingIterator(
response=httpx.Response(200),
model="test-model",
responses_api_provider_config=SimpleNamespace(),
logging_obj=logging_obj,
request_data=caching_handler.request_kwargs,
call_type=CallTypes.aresponses.value,
)
iterator.completed_response = _make_completed_response()
iterator._handle_logging_completed_response()
await asyncio.sleep(0.2)
litellm.cache.async_add_cache.assert_called_once()
assert litellm.cache.async_add_cache.call_args.kwargs["stream"] is True
assert (
litellm.cache.async_add_cache.call_args.kwargs["cache_key"]
== "responses-stream-cache-key"
)
assert "metadata" not in litellm.cache.async_add_cache.call_args.kwargs
assert "custom_llm_provider" not in litellm.cache.async_add_cache.call_args.kwargs
assert (
json.loads(litellm.cache.async_add_cache.call_args.args[0])["id"]
== iterator.completed_response.response.id
)
litellm.cache = original_cache
def test_responses_streaming_completed_event_persists_sync_cache():
logging_obj = _FakeLoggingObj()
original_cache = litellm.cache
litellm.cache = SimpleNamespace(
async_add_cache=AsyncMock(),
add_cache=MagicMock(),
)
caching_handler = SimpleNamespace(
request_kwargs={
"model": "test-model",
"input": "hello",
"stream": True,
"caching": True,
"cache_key": "stale-request-cache-key",
"metadata": None,
"custom_llm_provider": "openai",
},
preset_cache_key="responses-stream-cache-key",
original_function=litellm.responses,
sync_set_cache=MagicMock(),
_should_store_result_in_cache=lambda original_function, kwargs: True,
)
logging_obj._llm_caching_handler = caching_handler
iterator = SyncResponsesAPIStreamingIterator(
response=httpx.Response(200),
model="test-model",
responses_api_provider_config=SimpleNamespace(),
logging_obj=logging_obj,
request_data=caching_handler.request_kwargs,
call_type=CallTypes.responses.value,
)
iterator.completed_response = _make_completed_response("resp_sync")
iterator._handle_logging_completed_response()
litellm.cache.add_cache.assert_called_once()
assert litellm.cache.add_cache.call_args.kwargs["stream"] is True
assert (
litellm.cache.add_cache.call_args.kwargs["cache_key"]
== "responses-stream-cache-key"
)
assert "metadata" not in litellm.cache.add_cache.call_args.kwargs
assert "custom_llm_provider" not in litellm.cache.add_cache.call_args.kwargs
assert (
json.loads(litellm.cache.add_cache.call_args.args[0])["id"]
== iterator.completed_response.response.id
)
litellm.cache = original_cache
@pytest.mark.asyncio
async def test_cached_responses_stream_async_hit_triggers_success_callbacks(
monkeypatch,
):
hook_calls = {"post_call": 0, "metadata": 0}
async def fake_post_call(request_data, response, call_type):
hook_calls["post_call"] += 1
def fake_update_metadata(**kwargs):
hook_calls["metadata"] += 1
monkeypatch.setattr(
streaming_module,
"async_post_call_success_deployment_hook",
fake_post_call,
)
monkeypatch.setattr(
streaming_module,
"update_response_metadata",
fake_update_metadata,
)
logging_obj = _FakeLoggingObj()
original_cache = litellm.cache
litellm.cache = SimpleNamespace(
async_add_cache=AsyncMock(),
add_cache=MagicMock(),
)
logging_obj._llm_caching_handler = SimpleNamespace(
request_kwargs={"model": "test-model", "input": "hello", "stream": True},
preset_cache_key="responses-stream-cache-key",
original_function=litellm.aresponses,
_should_store_result_in_cache=lambda original_function, kwargs: True,
)
iterator = CachedResponsesAPIStreamingIterator(
response=_make_completed_response("resp_cached_async").response,
logging_obj=logging_obj,
request_data={"model": "test-model", "input": "hello", "stream": True},
call_type=CallTypes.aresponses.value,
)
streamed_events = [event async for event in iterator]
await asyncio.sleep(0.2)
assert streamed_events[-1].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
assert logging_obj.success_calls == 1
assert logging_obj.async_success_calls == 1
assert logging_obj.last_success_kwargs["cache_hit"] is True
assert logging_obj.last_async_success_kwargs["cache_hit"] is True
assert hook_calls["post_call"] == 1
assert hook_calls["metadata"] == 1
litellm.cache.async_add_cache.assert_not_called()
litellm.cache.add_cache.assert_not_called()
litellm.cache = original_cache
def test_cached_responses_stream_sync_hit_triggers_success_callbacks(monkeypatch):
hook_calls = {"post_call": 0, "metadata": 0}
async def fake_post_call(request_data, response, call_type):
hook_calls["post_call"] += 1
def fake_update_metadata(**kwargs):
hook_calls["metadata"] += 1
monkeypatch.setattr(
streaming_module,
"async_post_call_success_deployment_hook",
fake_post_call,
)
monkeypatch.setattr(
streaming_module,
"update_response_metadata",
fake_update_metadata,
)
logging_obj = _FakeLoggingObj()
original_cache = litellm.cache
litellm.cache = SimpleNamespace(
async_add_cache=AsyncMock(),
add_cache=MagicMock(),
)
logging_obj._llm_caching_handler = SimpleNamespace(
request_kwargs={"model": "test-model", "input": "hello", "stream": True},
preset_cache_key="responses-stream-cache-key",
original_function=litellm.responses,
_should_store_result_in_cache=lambda original_function, kwargs: True,
)
iterator = CachedResponsesAPIStreamingIterator(
response=_make_completed_response("resp_cached_sync").response,
logging_obj=logging_obj,
request_data={"model": "test-model", "input": "hello", "stream": True},
call_type=CallTypes.responses.value,
)
streamed_events = list(iterator)
asyncio.run(asyncio.sleep(0.2))
assert streamed_events[-1].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
assert logging_obj.success_calls == 1
assert logging_obj.async_success_calls == 1
assert logging_obj.last_success_kwargs["cache_hit"] is True
assert logging_obj.last_async_success_kwargs["cache_hit"] is True
assert hook_calls["post_call"] == 1
assert hook_calls["metadata"] == 1
litellm.cache.async_add_cache.assert_not_called()
litellm.cache.add_cache.assert_not_called()
litellm.cache = original_cache

View file

@ -19,6 +19,7 @@ import pytest
import litellm
from litellm import aembedding, completion, embedding, aresponses, responses
from litellm.caching.caching import Cache
from litellm.responses.streaming_iterator import CachedResponsesAPIStreamingIterator
from unittest.mock import AsyncMock, patch, MagicMock
from litellm.caching.caching_handler import LLMCachingHandler, CachingHandlerResponse
@ -158,14 +159,20 @@ async def test_async_log_cache_hit_on_callbacks():
# Assertions
mock_logging_obj.async_success_handler.assert_called_once_with(
result=cached_result, start_time=start_time, end_time=end_time, cache_hit=cache_hit
result=cached_result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
)
# Wait for the thread to complete
await asyncio.sleep(0.5)
mock_logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once_with(
result=cached_result, start_time=start_time, end_time=end_time, cache_hit=cache_hit
result=cached_result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
)
@ -346,7 +353,7 @@ async def test_embedding_cache_model_field_consistency():
"""
# Setup cache
setup_cache()
caching_handler = LLMCachingHandler(
original_function=aembedding, request_kwargs={}, start_time=datetime.now()
)
@ -358,7 +365,7 @@ async def test_embedding_cache_model_field_consistency():
data=[
Embedding(embedding=[0.1, 0.2, 0.3], index=0, object="embedding"),
Embedding(embedding=[0.4, 0.5, 0.6], index=1, object="embedding"),
]
],
)
# Mock logging object
@ -376,14 +383,12 @@ async def test_embedding_cache_model_field_consistency():
kwargs = {
"model": original_model,
"input": ["test input 1", "test input 2"],
"caching": True
"caching": True,
}
# Step 1: Cache the embedding response
await caching_handler.async_set_cache(
result=embedding_response,
original_function=aembedding,
kwargs=kwargs
result=embedding_response, original_function=aembedding, kwargs=kwargs
)
# Step 2: Retrieve from cache
@ -400,13 +405,24 @@ async def test_embedding_cache_model_field_consistency():
assert cached_response.final_embedding_cached_response is not None
assert cached_response.final_embedding_cached_response.model == original_model
assert len(cached_response.final_embedding_cached_response.data) == 2
assert cached_response.final_embedding_cached_response.data[0].embedding == [0.1, 0.2, 0.3]
assert cached_response.final_embedding_cached_response.data[0].embedding == [
0.1,
0.2,
0.3,
]
assert cached_response.final_embedding_cached_response.data[0].index == 0
assert cached_response.final_embedding_cached_response.data[1].embedding == [0.4, 0.5, 0.6]
assert cached_response.final_embedding_cached_response.data[1].embedding == [
0.4,
0.5,
0.6,
]
assert cached_response.final_embedding_cached_response.data[1].index == 1
# Verify cache hit flag is set
assert cached_response.final_embedding_cached_response._hidden_params["cache_hit"] == True
assert (
cached_response.final_embedding_cached_response._hidden_params["cache_hit"]
== True
)
@pytest.mark.asyncio
@ -417,7 +433,7 @@ async def test_embedding_cache_model_field_with_vendor_prefix():
"""
# Setup cache
setup_cache()
caching_handler = LLMCachingHandler(
original_function=aembedding, request_kwargs={}, start_time=datetime.now()
)
@ -425,13 +441,13 @@ async def test_embedding_cache_model_field_with_vendor_prefix():
# Test with vendor-prefixed model name (like vertex_ai/text-embedding-005)
vendor_model = "vertex_ai/text-embedding-005"
actual_model = "text-embedding-005" # What the provider actually returns
# Create embedding response with the actual model name (as returned by provider)
embedding_response = EmbeddingResponse(
model=actual_model, # Provider returns this
data=[
Embedding(embedding=[0.1, 0.2, 0.3], index=0, object="embedding"),
]
],
)
# Mock logging object
@ -449,14 +465,12 @@ async def test_embedding_cache_model_field_with_vendor_prefix():
kwargs = {
"model": vendor_model, # Request uses vendor prefix
"input": ["test input"],
"caching": True
"caching": True,
}
# Cache the response
await caching_handler.async_set_cache(
result=embedding_response,
original_function=aembedding,
kwargs=kwargs
result=embedding_response, original_function=aembedding, kwargs=kwargs
)
# Retrieve from cache
@ -471,8 +485,12 @@ async def test_embedding_cache_model_field_with_vendor_prefix():
# Verify the model field matches the original provider response, not the request
assert cached_response.final_embedding_cached_response is not None
assert cached_response.final_embedding_cached_response.model == actual_model # Should be the provider's model name
assert cached_response.final_embedding_cached_response.model != vendor_model # Should NOT be the vendor-prefixed name
assert (
cached_response.final_embedding_cached_response.model == actual_model
) # Should be the provider's model name
assert (
cached_response.final_embedding_cached_response.model != vendor_model
) # Should NOT be the vendor-prefixed name
def test_extract_model_from_cached_results():
@ -485,10 +503,26 @@ def test_extract_model_from_cached_results():
# Test with valid cached results
non_null_list = [
(0, {"embedding": [0.1, 0.2], "index": 0, "object": "embedding", "model": "text-embedding-005"}),
(1, {"embedding": [0.3, 0.4], "index": 1, "object": "embedding", "model": "text-embedding-005"}),
(
0,
{
"embedding": [0.1, 0.2],
"index": 0,
"object": "embedding",
"model": "text-embedding-005",
},
),
(
1,
{
"embedding": [0.3, 0.4],
"index": 1,
"object": "embedding",
"model": "text-embedding-005",
},
),
]
model_name = caching_handler._extract_model_from_cached_results(non_null_list)
assert model_name == "text-embedding-005"
@ -497,8 +531,10 @@ def test_extract_model_from_cached_results():
(0, {"embedding": [0.1, 0.2], "index": 0, "object": "embedding"}),
(1, {"embedding": [0.3, 0.4], "index": 1, "object": "embedding"}),
]
model_name = caching_handler._extract_model_from_cached_results(non_null_list_no_model)
model_name = caching_handler._extract_model_from_cached_results(
non_null_list_no_model
)
assert model_name is None
# Test with empty list
@ -514,7 +550,7 @@ async def test_async_responses_api_caching():
"""
# Setup cache
setup_cache()
caching_handler = LLMCachingHandler(
original_function=aresponses, request_kwargs={}, start_time=datetime.now()
)
@ -537,11 +573,11 @@ async def test_async_responses_api_caching():
{
"type": "output_text",
"text": "This is a test response from the responses API.",
"annotations": []
"annotations": [],
}
]
],
}
]
],
)
# Mock logging object
@ -560,14 +596,12 @@ async def test_async_responses_api_caching():
"model": original_model,
"input": "Tell me a short story",
"max_output_tokens": 100,
"caching": True
"caching": True,
}
# Step 1: Cache the responses API response
await caching_handler.async_set_cache(
result=responses_api_response,
original_function=aresponses,
kwargs=kwargs
result=responses_api_response, original_function=aresponses, kwargs=kwargs
)
await asyncio.sleep(0.5)
@ -589,18 +623,67 @@ async def test_async_responses_api_caching():
assert cached_response.cached_result.model == original_model
assert cached_response.cached_result.status == "completed"
assert len(cached_response.cached_result.output) == 1
# Verify cache hit flag is set
assert cached_response.cached_result._hidden_params["cache_hit"] == True
@pytest.mark.asyncio
async def test_async_get_cache_updates_request_kwargs_for_streaming_responses():
"""
Ensure streamed responses retain the normalized lookup kwargs so a later
cache write can reuse the exact cache key from the read path.
"""
setup_cache()
caching_handler = LLMCachingHandler(
original_function=aresponses,
request_kwargs={"stale": True},
start_time=datetime.now(),
)
logging_obj = LiteLLMLogging(
litellm_call_id=str(datetime.now()),
call_type=CallTypes.aresponses.value,
model="gpt-4o",
messages=[],
function_id=str(uuid.uuid4()),
stream=True,
start_time=datetime.now(),
)
kwargs = {
"model": "gpt-4o",
"input": "hello",
"stream": True,
"caching": True,
}
await caching_handler._async_get_cache(
model="gpt-4o",
original_function=aresponses,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.aresponses.value,
kwargs=kwargs,
)
assert "stale" not in caching_handler.request_kwargs
assert caching_handler.request_kwargs["model"] == "gpt-4o"
assert caching_handler.request_kwargs["input"] == "hello"
assert caching_handler.request_kwargs["stream"] is True
assert caching_handler.request_kwargs["cache_key"] == litellm.cache.get_cache_key(
**caching_handler.request_kwargs
)
def test_sync_responses_api_caching():
"""
Test that synchronous responses API calls are properly cached and retrieved.
"""
# Setup cache
setup_cache()
caching_handler = LLMCachingHandler(
original_function=responses, request_kwargs={}, start_time=datetime.now()
)
@ -623,11 +706,11 @@ def test_sync_responses_api_caching():
{
"type": "output_text",
"text": "Sync response test.",
"annotations": []
"annotations": [],
}
]
],
}
]
],
)
# Mock logging object
@ -646,14 +729,11 @@ def test_sync_responses_api_caching():
"model": original_model,
"input": "Tell me another story",
"max_output_tokens": 100,
"caching": True
"caching": True,
}
# Step 1: Cache the responses API response
caching_handler.sync_set_cache(
result=responses_api_response,
kwargs=kwargs
)
caching_handler.sync_set_cache(result=responses_api_response, kwargs=kwargs)
time.sleep(0.5)
@ -673,7 +753,7 @@ def test_sync_responses_api_caching():
assert cached_response.cached_result.id == responses_api_response.id
assert cached_response.cached_result.model == original_model
assert cached_response.cached_result.status == "completed"
# Verify cache hit flag is set
assert cached_response.cached_result._hidden_params["cache_hit"] == True
@ -686,7 +766,7 @@ def test_convert_cached_responses_api_result_to_model_response():
caching_handler = LLMCachingHandler(
original_function=responses, request_kwargs={}, start_time=datetime.now()
)
logging_obj = LiteLLMLogging(
litellm_call_id=str(datetime.now()),
call_type=CallTypes.responses.value,
@ -714,11 +794,11 @@ def test_convert_cached_responses_api_result_to_model_response():
{
"type": "output_text",
"text": "Conversion test response.",
"annotations": []
"annotations": [],
}
]
],
}
]
],
}
# Convert cached result to ResponsesAPIResponse
@ -739,6 +819,318 @@ def test_convert_cached_responses_api_result_to_model_response():
assert len(result.output) == 1
def test_sync_get_cache_does_not_eagerly_log_streaming_responses_hits():
litellm.set_verbose = True
setup_cache()
caching_handler = LLMCachingHandler(
original_function=responses, request_kwargs={}, start_time=datetime.now()
)
original_model = "gpt-4o"
responses_api_response = ResponsesAPIResponse(
id="resp_stream_sync_hit",
created_at=int(time.time()),
status="completed",
model=original_model,
object="response",
output=[
{
"type": "message",
"id": "msg_stream_sync_hit",
"status": "completed",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": "Sync streamed cache hit response.",
"annotations": [],
}
],
}
],
)
logging_obj = LiteLLMLogging(
litellm_call_id=str(datetime.now()),
call_type=CallTypes.responses.value,
model=original_model,
messages=[],
function_id=str(uuid.uuid4()),
stream=True,
start_time=datetime.now(),
)
logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock()
kwargs = {
"model": original_model,
"input": "Tell me a cached story",
"stream": True,
"caching": True,
}
caching_handler.sync_set_cache(result=responses_api_response, kwargs=kwargs)
time.sleep(0.2)
cached_response = caching_handler._sync_get_cache(
model=original_model,
original_function=responses,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.responses.value,
kwargs=kwargs,
)
assert cached_response.cached_result is not None
assert isinstance(
cached_response.cached_result, CachedResponsesAPIStreamingIterator
)
logging_obj.handle_sync_success_callbacks_for_async_calls.assert_not_called()
def test_sync_get_cache_still_eagerly_logs_streaming_completion_hits():
litellm.set_verbose = True
setup_cache()
caching_handler = LLMCachingHandler(
original_function=completion, request_kwargs={}, start_time=datetime.now()
)
original_model = "gpt-4o"
logging_obj = LiteLLMLogging(
litellm_call_id=str(datetime.now()),
call_type=CallTypes.completion.value,
model=original_model,
messages=[],
function_id=str(uuid.uuid4()),
stream=True,
start_time=datetime.now(),
)
logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock()
kwargs = {
"model": original_model,
"messages": [{"role": "user", "content": "Tell me a cached joke"}],
"stream": True,
"caching": True,
}
caching_handler.sync_set_cache(result=chat_completion_response, kwargs=kwargs)
time.sleep(0.2)
cached_response = caching_handler._sync_get_cache(
model=original_model,
original_function=completion,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.completion.value,
kwargs=kwargs,
)
assert cached_response.cached_result is not None
logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once()
@pytest.mark.asyncio
async def test_async_get_cache_still_eagerly_logs_streaming_completion_hits():
litellm.set_verbose = True
setup_cache()
caching_handler = LLMCachingHandler(
original_function=completion, request_kwargs={}, start_time=datetime.now()
)
original_model = "gpt-4o"
kwargs = {
"model": original_model,
"messages": [{"role": "user", "content": "Tell me a cached joke"}],
"stream": True,
"caching": True,
}
await caching_handler.async_set_cache(
result=chat_completion_response,
original_function=litellm.acompletion,
kwargs=kwargs,
)
await asyncio.sleep(0.2)
logging_obj = LiteLLMLogging(
litellm_call_id=str(datetime.now()),
call_type=CallTypes.acompletion.value,
model=original_model,
messages=[],
function_id=str(uuid.uuid4()),
stream=True,
start_time=datetime.now(),
)
caching_handler._async_log_cache_hit_on_callbacks = MagicMock()
cached_response = await caching_handler._async_get_cache(
model=original_model,
original_function=litellm.acompletion,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.acompletion.value,
kwargs=kwargs,
)
assert cached_response is not None
assert cached_response.cached_result is not None
caching_handler._async_log_cache_hit_on_callbacks.assert_called_once()
def test_convert_cached_streaming_responses_result_to_iterator():
"""
Test that cached streaming Responses results are replayed through a synthetic
streaming iterator instead of being returned as a full response object.
"""
caching_handler = LLMCachingHandler(
original_function=responses, request_kwargs={}, start_time=datetime.now()
)
logging_obj = LiteLLMLogging(
litellm_call_id=str(datetime.now()),
call_type=CallTypes.responses.value,
model="gpt-4o",
messages=[],
function_id=str(uuid.uuid4()),
stream=True,
start_time=datetime.now(),
)
cached_result = {
"id": "resp_stream_cache_test",
"created_at": int(time.time()),
"status": "completed",
"model": "gpt-4o",
"object": "response",
"output": [
{
"type": "message",
"id": "msg_stream_cache_test",
"status": "completed",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": "Streaming cache replay test.",
"annotations": [],
}
],
}
],
}
result = caching_handler._convert_cached_result_to_model_response(
cached_result=cached_result,
call_type=CallTypes.responses.value,
kwargs={"model": "gpt-4o", "input": "test", "stream": True},
logging_obj=logging_obj,
model="gpt-4o",
args=(),
)
assert isinstance(result, CachedResponsesAPIStreamingIterator)
assert result.completed_response is not None
assert result.completed_response.response.id == cached_result["id"]
streamed_events = list(result)
assert streamed_events[0].type == "response.created"
assert streamed_events[1].type == "response.in_progress"
assert streamed_events[2].type == "response.output_item.added"
assert streamed_events[3].type == "response.content_part.added"
assert streamed_events[-4].type == "response.output_text.done"
assert streamed_events[-3].type == "response.content_part.done"
assert streamed_events[-2].type == "response.output_item.done"
assert streamed_events[-1].type == "response.completed"
assert streamed_events[-1].response.id == cached_result["id"]
assert streamed_events[-1].response.output[0].content[0].text == (
"Streaming cache replay test."
)
def test_convert_cached_streaming_reasoning_result_to_iterator():
caching_handler = LLMCachingHandler(
original_function=responses, request_kwargs={}, start_time=datetime.now()
)
logging_obj = LiteLLMLogging(
litellm_call_id=str(datetime.now()),
call_type=CallTypes.responses.value,
model="gpt-4o",
messages=[],
function_id=str(uuid.uuid4()),
stream=True,
start_time=datetime.now(),
)
cached_result = {
"id": "resp_stream_reasoning_cache_test",
"created_at": int(time.time()),
"status": "completed",
"model": "gpt-4o",
"object": "response",
"output": [
{
"type": "reasoning",
"id": "rs_stream_cache_test",
"summary": [
{
"type": "summary_text",
"text": "Cached reasoning summary.",
}
],
}
],
}
result = caching_handler._convert_cached_result_to_model_response(
cached_result=cached_result,
call_type=CallTypes.responses.value,
kwargs={"model": "gpt-4o", "input": "test", "stream": True},
logging_obj=logging_obj,
model="gpt-4o",
args=(),
)
assert isinstance(result, CachedResponsesAPIStreamingIterator)
streamed_events = list(result)
streamed_event_types = [
event.type.value if hasattr(event.type, "value") else str(event.type)
for event in streamed_events
]
assert streamed_event_types[:3] == [
"response.created",
"response.in_progress",
"response.output_item.added",
]
assert streamed_event_types[-4:] == [
"response.reasoning_summary_text.done",
"response.reasoning_summary_part.done",
"response.output_item.done",
"response.completed",
]
assert streamed_event_types.count("response.reasoning_summary_text.delta") >= 1
delta_events = [
event
for event in streamed_events
if (event.type.value if hasattr(event.type, "value") else str(event.type))
== "response.reasoning_summary_text.delta"
]
text_done_event = streamed_events[-4]
part_done_event = streamed_events[-3]
output_item_done_event = streamed_events[-2]
assert all(delta_event.summary_index == 0 for delta_event in delta_events)
assert text_done_event.text == "Cached reasoning summary."
assert text_done_event.summary_index == 0
assert part_done_event.part.type == "summary_text"
assert part_done_event.part.text == "Cached reasoning summary."
assert output_item_done_event.item.type == "reasoning"
assert output_item_done_event.item.summary[0]["text"] == "Cached reasoning summary."
@pytest.mark.asyncio
async def test_responses_api_cache_with_different_inputs():
"""
@ -747,7 +1139,7 @@ async def test_responses_api_cache_with_different_inputs():
"""
# Setup cache
setup_cache()
caching_handler = LLMCachingHandler(
original_function=aresponses, request_kwargs={}, start_time=datetime.now()
)
@ -767,21 +1159,17 @@ async def test_responses_api_cache_with_different_inputs():
"id": "msg_1",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "Response 1", "annotations": []}]
"content": [
{"type": "output_text", "text": "Response 1", "annotations": []}
],
}
]
],
)
kwargs_1 = {
"model": original_model,
"input": "First unique input",
"caching": True
}
kwargs_1 = {"model": original_model, "input": "First unique input", "caching": True}
await caching_handler.async_set_cache(
result=response_1,
original_function=aresponses,
kwargs=kwargs_1
result=response_1, original_function=aresponses, kwargs=kwargs_1
)
# Second request with different input
@ -797,21 +1185,21 @@ async def test_responses_api_cache_with_different_inputs():
"id": "msg_2",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "Response 2", "annotations": []}]
"content": [
{"type": "output_text", "text": "Response 2", "annotations": []}
],
}
]
],
)
kwargs_2 = {
"model": original_model,
"input": "Second unique input",
"caching": True
"caching": True,
}
await caching_handler.async_set_cache(
result=response_2,
original_function=aresponses,
kwargs=kwargs_2
result=response_2, original_function=aresponses, kwargs=kwargs_2
)
await asyncio.sleep(0.5)
@ -860,20 +1248,28 @@ async def test_responses_api_cache_with_different_inputs():
assert cached_2.cached_result is not None
assert cached_1.cached_result.id == "resp_1"
assert cached_2.cached_result.id == "resp_2"
# Access output content properly (could be dict or object)
output_1 = cached_1.cached_result.output[0]
if isinstance(output_1, dict):
text_1 = output_1["content"][0]["text"]
else:
text_1 = output_1.content[0].text if hasattr(output_1.content[0], 'text') else output_1.content[0]["text"]
text_1 = (
output_1.content[0].text
if hasattr(output_1.content[0], "text")
else output_1.content[0]["text"]
)
output_2 = cached_2.cached_result.output[0]
if isinstance(output_2, dict):
text_2 = output_2["content"][0]["text"]
else:
text_2 = output_2.content[0].text if hasattr(output_2.content[0], 'text') else output_2.content[0]["text"]
text_2 = (
output_2.content[0].text
if hasattr(output_2.content[0], "text")
else output_2.content[0]["text"]
)
assert text_1 == "Response 1"
assert text_2 == "Response 2"
@ -897,9 +1293,9 @@ async def test_responses_api_cache_with_different_inputs():
"role": "assistant",
"content": [
{"type": "output_text", "text": "Test", "annotations": []}
]
],
}
]
],
},
ResponsesAPIResponse,
),
@ -918,10 +1314,14 @@ async def test_responses_api_cache_with_different_inputs():
"status": "completed",
"role": "assistant",
"content": [
{"type": "output_text", "text": "Async Test", "annotations": []}
]
{
"type": "output_text",
"text": "Async Test",
"annotations": [],
}
],
}
]
],
},
ResponsesAPIResponse,
),

View file

@ -0,0 +1,61 @@
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock
import pytest
import litellm
from litellm import aresponses
from litellm._uuid import uuid
from litellm.caching.caching_handler import LLMCachingHandler
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.types.utils import CallTypes
@pytest.mark.asyncio
async def test_async_get_cache_reuses_preset_cache_key_for_responses():
caching_handler = LLMCachingHandler(
original_function=aresponses,
request_kwargs={},
start_time=datetime.now(),
)
logging_obj = LiteLLMLogging(
litellm_call_id=str(datetime.now()),
call_type=CallTypes.aresponses.value,
model="gpt-4.1-mini",
messages=[],
function_id=str(uuid.uuid4()),
stream=True,
start_time=datetime.now(),
)
original_cache = litellm.cache
mock_cache = MagicMock()
mock_cache.supported_call_types = [CallTypes.aresponses.value]
mock_cache._supports_async.return_value = True
mock_cache.get_cache_key.return_value = "responses-stream-cache-key"
mock_cache.async_get_cache = AsyncMock(return_value=None)
litellm.cache = mock_cache
kwargs = {
"model": "gpt-4.1-mini",
"input": "hello",
"stream": True,
"litellm_params": {},
}
await caching_handler._async_get_cache(
model="gpt-4.1-mini",
original_function=aresponses,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.aresponses.value,
kwargs=kwargs,
)
assert caching_handler.preset_cache_key == "responses-stream-cache-key"
mock_cache.async_get_cache.assert_awaited_once()
assert (
mock_cache.async_get_cache.call_args.kwargs["cache_key"]
== "responses-stream-cache-key"
)
litellm.cache = original_cache