mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Add focused coverage for streamed responses cache
This commit is contained in:
parent
51e5175c6f
commit
17b70d4bc6
3 changed files with 490 additions and 1 deletions
|
|
@ -1,4 +1,5 @@
|
|||
import asyncio
|
||||
from contextlib import suppress
|
||||
from datetime import datetime
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
|
|
@ -12,6 +13,7 @@ from litellm.integrations.custom_logger import CustomLogger
|
|||
from litellm.responses import streaming_iterator as streaming_module
|
||||
from litellm.responses.streaming_iterator import (
|
||||
CachedResponsesAPIStreamingIterator,
|
||||
MockResponsesAPIStreamingIterator,
|
||||
ResponsesAPIStreamingIterator,
|
||||
SyncResponsesAPIStreamingIterator,
|
||||
)
|
||||
|
|
@ -78,6 +80,87 @@ def _make_completed_response(response_id: str = "resp_test") -> ResponseComplete
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_log_background_task_failure_logs_task_exceptions(monkeypatch):
|
||||
error_logger = MagicMock()
|
||||
monkeypatch.setattr(streaming_module.verbose_logger, "error", error_logger)
|
||||
|
||||
async def _boom():
|
||||
raise RuntimeError("boom")
|
||||
|
||||
task = asyncio.create_task(_boom())
|
||||
with suppress(RuntimeError):
|
||||
await task
|
||||
|
||||
streaming_module._log_background_task_failure(task, task_name="cache write")
|
||||
|
||||
error_logger.assert_called_once()
|
||||
assert error_logger.call_args.args == (
|
||||
"%s failed: %s",
|
||||
"cache write",
|
||||
task.exception(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_log_background_task_failure_ignores_cancelled_tasks(monkeypatch):
|
||||
error_logger = MagicMock()
|
||||
monkeypatch.setattr(streaming_module.verbose_logger, "error", error_logger)
|
||||
|
||||
task = asyncio.create_task(asyncio.sleep(1))
|
||||
task.cancel()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
streaming_module._log_background_task_failure(task, task_name="cache write")
|
||||
|
||||
error_logger.assert_not_called()
|
||||
|
||||
|
||||
def test_content_part_done_event_supports_refusal_and_reasoning_text():
|
||||
refusal_event = streaming_module._build_content_part_done_event(
|
||||
item_id="msg_1",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
part_payload={"type": "refusal", "refusal": "no"},
|
||||
)
|
||||
reasoning_event = streaming_module._build_content_part_done_event(
|
||||
item_id="msg_1",
|
||||
output_index=0,
|
||||
content_index=1,
|
||||
part_payload={"type": "reasoning_text", "reasoning": "because"},
|
||||
)
|
||||
unsupported_event = streaming_module._build_content_part_done_event(
|
||||
item_id="msg_1",
|
||||
output_index=0,
|
||||
content_index=2,
|
||||
part_payload={"type": "image"},
|
||||
)
|
||||
|
||||
assert refusal_event.part.type == "refusal"
|
||||
assert refusal_event.part.refusal == "no"
|
||||
assert reasoning_event.part.type == "reasoning_text"
|
||||
assert reasoning_event.part.reasoning == "because"
|
||||
assert unsupported_event is None
|
||||
|
||||
|
||||
def test_dump_response_object_handles_model_and_unknown_values():
|
||||
response = ResponsesAPIResponse(
|
||||
id="resp_dump",
|
||||
created_at=int(datetime.now().timestamp()),
|
||||
status="completed",
|
||||
model="gpt-4.1-mini",
|
||||
object="response",
|
||||
output=[],
|
||||
)
|
||||
|
||||
assert streaming_module._dump_response_object(response)["id"] == "resp_dump"
|
||||
assert streaming_module._dump_response_object({"type": "message"}) == {
|
||||
"type": "message"
|
||||
}
|
||||
assert streaming_module._dump_response_object(object()) == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_streaming_triggers_hooks(monkeypatch):
|
||||
"""
|
||||
|
|
@ -211,6 +294,142 @@ async def test_responses_streaming_failure_triggers_failure_handlers():
|
|||
assert logging_obj.async_failure_calls >= 1
|
||||
|
||||
|
||||
def test_process_chunk_requires_provider_config():
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200),
|
||||
model="test-model",
|
||||
responses_api_provider_config=None,
|
||||
logging_obj=_FakeLoggingObj(),
|
||||
request_data={"foo": "bar"},
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="responses_api_provider_config is required"):
|
||||
iterator._process_chunk(json.dumps({"type": "response.completed"}))
|
||||
|
||||
|
||||
def test_process_chunk_wraps_encrypted_content_with_model_id():
|
||||
openai_types = streaming_module._get_openai_response_types()
|
||||
|
||||
class _EncryptedConfig:
|
||||
def transform_streaming_response(self, **kwargs):
|
||||
return openai_types.OutputItemAddedEvent(
|
||||
type=openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
|
||||
output_index=0,
|
||||
item=openai_types.BaseLiteLLMOpenAIResponseObject(
|
||||
id="rs_123",
|
||||
type="reasoning",
|
||||
encrypted_content="ciphertext",
|
||||
),
|
||||
)
|
||||
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200),
|
||||
model="test-model",
|
||||
responses_api_provider_config=_EncryptedConfig(),
|
||||
logging_obj=_FakeLoggingObj(),
|
||||
litellm_metadata={
|
||||
"encrypted_content_affinity_enabled": True,
|
||||
"model_info": {"id": "model-123"},
|
||||
},
|
||||
request_data={"foo": "bar"},
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
event = iterator._process_chunk(json.dumps({"type": "response.output_item.added"}))
|
||||
|
||||
assert event.item.encrypted_content.startswith("litellm_enc:")
|
||||
assert event.item.encrypted_content.endswith(";ciphertext")
|
||||
|
||||
|
||||
def test_process_chunk_completed_response_updates_id_and_usage_cost(monkeypatch):
|
||||
original_include_cost = litellm.include_cost_in_streaming_usage
|
||||
litellm.include_cost_in_streaming_usage = True
|
||||
openai_types = streaming_module._get_openai_response_types()
|
||||
|
||||
class _CompletedConfig:
|
||||
def transform_streaming_response(self, **kwargs):
|
||||
return openai_types.ResponseCompletedEvent(
|
||||
type=openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response=ResponsesAPIResponse(
|
||||
id="resp_live",
|
||||
created_at=int(datetime.now().timestamp()),
|
||||
status="completed",
|
||||
model="test-model",
|
||||
object="response",
|
||||
output=[],
|
||||
usage=openai_types.ResponseAPIUsage(
|
||||
input_tokens=1,
|
||||
output_tokens=2,
|
||||
total_tokens=3,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
logging_obj = _FakeLoggingObj()
|
||||
logging_obj._response_cost_calculator = MagicMock(return_value=1.23)
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200),
|
||||
model="test-model",
|
||||
responses_api_provider_config=_CompletedConfig(),
|
||||
logging_obj=logging_obj,
|
||||
litellm_metadata={"model_info": {"id": "model-123"}},
|
||||
custom_llm_provider="openai",
|
||||
request_data={"foo": "bar"},
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
completion_handler = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
iterator, "_handle_logging_completed_response", completion_handler
|
||||
)
|
||||
|
||||
try:
|
||||
event = iterator._process_chunk(json.dumps({"type": "response.completed"}))
|
||||
finally:
|
||||
litellm.include_cost_in_streaming_usage = original_include_cost
|
||||
|
||||
assert iterator.completed_response is event
|
||||
assert event.response.id != "resp_live"
|
||||
assert event.response.id.startswith("resp_")
|
||||
assert event.response.usage.cost == 1.23
|
||||
completion_handler.assert_called_once()
|
||||
|
||||
|
||||
def test_process_chunk_failed_response_triggers_failure_logging(monkeypatch):
|
||||
openai_types = streaming_module._get_openai_response_types()
|
||||
|
||||
class _FailedConfig:
|
||||
def transform_streaming_response(self, **kwargs):
|
||||
return openai_types.ResponseFailedEvent(
|
||||
type=openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED,
|
||||
response=ResponsesAPIResponse(
|
||||
id="resp_failed",
|
||||
created_at=int(datetime.now().timestamp()),
|
||||
status="failed",
|
||||
model="test-model",
|
||||
object="response",
|
||||
output=[],
|
||||
error={"message": "provider failed"},
|
||||
),
|
||||
)
|
||||
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200),
|
||||
model="test-model",
|
||||
responses_api_provider_config=_FailedConfig(),
|
||||
logging_obj=_FakeLoggingObj(),
|
||||
request_data={"foo": "bar"},
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
failure_handler = MagicMock()
|
||||
monkeypatch.setattr(iterator, "_handle_logging_failed_response", failure_handler)
|
||||
|
||||
event = iterator._process_chunk(json.dumps({"type": "response.failed"}))
|
||||
|
||||
assert iterator.completed_response is event
|
||||
failure_handler.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_streaming_completed_event_persists_async_cache():
|
||||
logging_obj = _FakeLoggingObj()
|
||||
|
|
@ -315,6 +534,161 @@ def test_responses_streaming_completed_event_persists_sync_cache():
|
|||
litellm.cache = original_cache
|
||||
|
||||
|
||||
def test_build_synthetic_response_events_covers_annotations_function_calls_and_refusals():
|
||||
original_include_cost = litellm.include_cost_in_streaming_usage
|
||||
litellm.include_cost_in_streaming_usage = True
|
||||
logging_obj = _FakeLoggingObj()
|
||||
logging_obj._response_cost_calculator = MagicMock(side_effect=RuntimeError("boom"))
|
||||
transformed = ResponsesAPIResponse(
|
||||
id="resp_events",
|
||||
created_at=int(datetime.now().timestamp()),
|
||||
status="completed",
|
||||
model="gpt-4.1-mini",
|
||||
object="response",
|
||||
output=[
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_events",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "output_text",
|
||||
"text": "hello world",
|
||||
"annotations": [{"type": "file_citation", "file_id": "file_1"}],
|
||||
},
|
||||
{
|
||||
"type": "refusal",
|
||||
"refusal": "no thanks",
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"type": "function_call",
|
||||
"id": "fc_events",
|
||||
"call_id": "call_123",
|
||||
"name": "lookup",
|
||||
"arguments": '{"id":1}',
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
try:
|
||||
events = streaming_module._build_synthetic_response_events(
|
||||
transformed=transformed,
|
||||
logging_obj=logging_obj,
|
||||
chunk_size=5,
|
||||
)
|
||||
finally:
|
||||
litellm.include_cost_in_streaming_usage = original_include_cost
|
||||
|
||||
event_types = [
|
||||
event.type.value if hasattr(event.type, "value") else str(event.type)
|
||||
for event in events
|
||||
]
|
||||
|
||||
assert "response.output_text.annotation.added" in event_types
|
||||
assert "response.refusal.delta" in event_types
|
||||
assert "response.refusal.done" in event_types
|
||||
assert "response.function_call_arguments.delta" in event_types
|
||||
assert "response.function_call_arguments.done" in event_types
|
||||
assert event_types[-1] == "response.completed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mock_responses_streaming_iterator_async_iteration_logs_completion(
|
||||
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,
|
||||
)
|
||||
|
||||
class _MockTransformConfig:
|
||||
def transform_response_api_response(self, **kwargs):
|
||||
return _make_completed_response("resp_mock").response
|
||||
|
||||
logging_obj = _FakeLoggingObj()
|
||||
|
||||
iterator = MockResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200),
|
||||
model="test-model",
|
||||
responses_api_provider_config=_MockTransformConfig(),
|
||||
logging_obj=logging_obj,
|
||||
request_data={"model": "test-model", "stream": True},
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
streamed_events = [event async for event in iterator]
|
||||
await asyncio.sleep(0.2)
|
||||
|
||||
assert streamed_events[0].type == ResponsesAPIStreamEvents.RESPONSE_CREATED
|
||||
assert streamed_events[-1].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
assert logging_obj.success_calls == 1
|
||||
assert logging_obj.async_success_calls == 1
|
||||
assert hook_calls["post_call"] == 1
|
||||
assert hook_calls["metadata"] == 1
|
||||
|
||||
|
||||
def test_mock_responses_streaming_iterator_sync_iteration_logs_completion(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,
|
||||
)
|
||||
|
||||
class _MockTransformConfig:
|
||||
def transform_response_api_response(self, **kwargs):
|
||||
return _make_completed_response("resp_mock_sync").response
|
||||
|
||||
logging_obj = _FakeLoggingObj()
|
||||
iterator = MockResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200),
|
||||
model="test-model",
|
||||
responses_api_provider_config=_MockTransformConfig(),
|
||||
logging_obj=logging_obj,
|
||||
request_data={"model": "test-model", "stream": True},
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
streamed_events = list(iterator)
|
||||
asyncio.run(asyncio.sleep(0.2))
|
||||
|
||||
assert streamed_events[0].type == ResponsesAPIStreamEvents.RESPONSE_CREATED
|
||||
assert streamed_events[-1].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
assert logging_obj.success_calls == 1
|
||||
assert logging_obj.async_success_calls == 1
|
||||
assert hook_calls["post_call"] == 1
|
||||
assert hook_calls["metadata"] == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cached_responses_stream_async_hit_triggers_success_callbacks(
|
||||
monkeypatch,
|
||||
|
|
|
|||
|
|
@ -22,7 +22,11 @@ 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
|
||||
from litellm.caching.caching_handler import (
|
||||
LLMCachingHandler,
|
||||
CachingHandlerResponse,
|
||||
_should_defer_streaming_cache_hit_callbacks,
|
||||
)
|
||||
from litellm.caching.caching import LiteLLMCacheType
|
||||
from litellm.types.utils import CallTypes
|
||||
from litellm.types.rerank import RerankResponse
|
||||
|
|
@ -929,6 +933,37 @@ def test_sync_get_cache_still_eagerly_logs_streaming_completion_hits():
|
|||
logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once()
|
||||
|
||||
|
||||
def test_should_defer_streaming_cache_hit_callbacks_only_for_responses_streams():
|
||||
assert (
|
||||
_should_defer_streaming_cache_hit_callbacks(
|
||||
call_type=CallTypes.responses.value,
|
||||
kwargs={"stream": True},
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
_should_defer_streaming_cache_hit_callbacks(
|
||||
call_type=CallTypes.aresponses.value,
|
||||
kwargs={"stream": True},
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
_should_defer_streaming_cache_hit_callbacks(
|
||||
call_type=CallTypes.completion.value,
|
||||
kwargs={"stream": True},
|
||||
)
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
_should_defer_streaming_cache_hit_callbacks(
|
||||
call_type=CallTypes.responses.value,
|
||||
kwargs={"stream": False},
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_get_cache_still_eagerly_logs_streaming_completion_hits():
|
||||
litellm.set_verbose = True
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ 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.llms import openai as openai_types
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
|
||||
|
|
@ -59,3 +60,82 @@ async def test_async_get_cache_reuses_preset_cache_key_for_responses():
|
|||
)
|
||||
|
||||
litellm.cache = original_cache
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_get_cache_falls_back_to_sync_cache_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 = False
|
||||
mock_cache.get_cache_key.return_value = "responses-stream-cache-key"
|
||||
mock_cache.get_cache.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.get_cache.assert_called_once()
|
||||
assert mock_cache.get_cache.call_args.kwargs["cache_key"] == (
|
||||
"responses-stream-cache-key"
|
||||
)
|
||||
|
||||
litellm.cache = original_cache
|
||||
|
||||
|
||||
def test_reasoning_summary_events_default_summary_index():
|
||||
delta_event = openai_types.ReasoningSummaryTextDeltaEvent(
|
||||
type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA,
|
||||
item_id="rs_1",
|
||||
output_index=0,
|
||||
delta="abc",
|
||||
)
|
||||
text_done_event = openai_types.ReasoningSummaryTextDoneEvent(
|
||||
type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DONE,
|
||||
item_id="rs_1",
|
||||
output_index=0,
|
||||
sequence_number=1,
|
||||
text="abc",
|
||||
)
|
||||
part_done_event = openai_types.ReasoningSummaryPartDoneEvent(
|
||||
type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_PART_DONE,
|
||||
item_id="rs_1",
|
||||
output_index=0,
|
||||
sequence_number=2,
|
||||
part=openai_types.BaseLiteLLMOpenAIResponseObject(
|
||||
type="summary_text",
|
||||
text="abc",
|
||||
),
|
||||
)
|
||||
|
||||
assert delta_event.summary_index == 0
|
||||
assert text_done_event.summary_index == 0
|
||||
assert part_done_event.summary_index == 0
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue