mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
Merge 2f580d0081 into eddfb5fb20
This commit is contained in:
commit
bb86bd3598
2 changed files with 106 additions and 4 deletions
|
|
@ -5,6 +5,8 @@ Handler for transforming responses api requests to litellm.completion requests
|
|||
from collections.abc import Coroutine, Mapping
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
|
||||
LiteLLMCompletionStreamingIterator,
|
||||
|
|
@ -12,7 +14,10 @@ from litellm.responses.litellm_completion_transformation.streaming_iterator impo
|
|||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
|
||||
from litellm.responses.streaming_iterator import (
|
||||
BaseResponsesAPIStreamingIterator,
|
||||
MockResponsesAPIStreamingIterator,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseInputParam,
|
||||
ResponsesAPIOptionalRequestParams,
|
||||
|
|
@ -22,6 +27,38 @@ from litellm.types.utils import ModelResponse
|
|||
|
||||
|
||||
class LiteLLMCompletionTransformationHandler:
|
||||
@staticmethod
|
||||
def _maybe_wrap_as_fake_stream(
|
||||
responses_api_response: ResponsesAPIResponse,
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
kwargs: Mapping[str, object],
|
||||
) -> ResponsesAPIResponse | MockResponsesAPIStreamingIterator:
|
||||
"""
|
||||
An interceptor (e.g. websearch interception) can force stream=False so its
|
||||
agentic loop runs on the non-streaming path. When the caller originally asked
|
||||
for streaming, rebuild a synthetic responses stream from the final response.
|
||||
"""
|
||||
if not kwargs.get("_websearch_interception_converted_stream"):
|
||||
return responses_api_response
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
|
||||
logging_obj: Final = kwargs.get("litellm_logging_obj")
|
||||
if not isinstance(logging_obj, LiteLLMLoggingObj):
|
||||
return responses_api_response
|
||||
|
||||
raw_response: Final = httpx.Response(status_code=200, json=responses_api_response.model_dump())
|
||||
litellm_metadata: Final = kwargs.get("litellm_metadata")
|
||||
return MockResponsesAPIStreamingIterator(
|
||||
response=raw_response,
|
||||
model=model,
|
||||
responses_api_provider_config=OpenAIResponsesAPIConfig(),
|
||||
logging_obj=logging_obj,
|
||||
litellm_metadata=litellm_metadata if isinstance(litellm_metadata, dict) else None,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
def response_api_handler(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -75,7 +112,12 @@ class LiteLLMCompletionTransformationHandler:
|
|||
)
|
||||
)
|
||||
|
||||
return responses_api_response
|
||||
return self._maybe_wrap_as_fake_stream(
|
||||
responses_api_response=responses_api_response,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
elif isinstance(litellm_completion_response, litellm.CustomStreamWrapper):
|
||||
return LiteLLMCompletionStreamingIterator(
|
||||
|
|
@ -120,7 +162,12 @@ class LiteLLMCompletionTransformationHandler:
|
|||
)
|
||||
)
|
||||
|
||||
return responses_api_response
|
||||
return self._maybe_wrap_as_fake_stream(
|
||||
responses_api_response=responses_api_response,
|
||||
model=litellm_completion_request.get("model") or "",
|
||||
custom_llm_provider=litellm_completion_request.get("custom_llm_provider"),
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
elif isinstance(litellm_completion_response, litellm.CustomStreamWrapper):
|
||||
return LiteLLMCompletionStreamingIterator(
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ capture the forwarded kwargs; if the flag-setting line is removed the captured
|
|||
kwargs lack the flag and these tests fail.
|
||||
"""
|
||||
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -20,6 +20,9 @@ import pytest
|
|||
from litellm.responses.litellm_completion_transformation.handler import (
|
||||
LiteLLMCompletionTransformationHandler,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.responses.streaming_iterator import MockResponsesAPIStreamingIterator
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
|
||||
class _StopForwarding(Exception):
|
||||
|
|
@ -47,6 +50,58 @@ def test_sync_fallback_tags_skip_responses_api_bridge():
|
|||
assert captured.get("_skip_responses_api_bridge") is True
|
||||
|
||||
|
||||
def _completed_model_response() -> ModelResponse:
|
||||
return ModelResponse(
|
||||
choices=[Choices(message=Message(role="assistant", content="hello"), finish_reason="stop")],
|
||||
model="claude-sonnet-4-5",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_bridge_rebuilds_fake_stream_after_websearch_interception():
|
||||
"""Original stream=True request converted to stream=False by websearch
|
||||
interception must come back as a synthetic responses stream, not a plain
|
||||
ResponsesAPIResponse the proxy cannot async-iterate."""
|
||||
handler = LiteLLMCompletionTransformationHandler()
|
||||
logging_obj = MagicMock(spec=LiteLLMLoggingObj)
|
||||
logging_obj.dispatch_success_handlers = AsyncMock()
|
||||
logging_obj.model_call_details = {}
|
||||
|
||||
async def fake_acompletion(**kwargs):
|
||||
return _completed_model_response()
|
||||
|
||||
with patch("litellm.acompletion", fake_acompletion): # test-quality-ok: bridge forwards to module-level acompletion; no injection seam exists
|
||||
result = await handler.async_response_api_handler(
|
||||
litellm_completion_request={"model": "anthropic/claude-sonnet-4-5", "messages": []},
|
||||
request_input="hello",
|
||||
responses_api_request={},
|
||||
_websearch_interception_converted_stream=True,
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
assert isinstance(result, MockResponsesAPIStreamingIterator)
|
||||
events = [event async for event in result]
|
||||
assert any(getattr(event, "type", None) == "response.completed" for event in events)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_bridge_returns_plain_response_without_converted_stream_flag():
|
||||
handler = LiteLLMCompletionTransformationHandler()
|
||||
|
||||
async def fake_acompletion(**kwargs):
|
||||
return _completed_model_response()
|
||||
|
||||
with patch("litellm.acompletion", fake_acompletion): # test-quality-ok: bridge forwards to module-level acompletion; no injection seam exists
|
||||
result = await handler.async_response_api_handler(
|
||||
litellm_completion_request={"model": "anthropic/claude-sonnet-4-5", "messages": []},
|
||||
request_input="hello",
|
||||
responses_api_request={},
|
||||
litellm_logging_obj=MagicMock(spec=LiteLLMLoggingObj),
|
||||
)
|
||||
|
||||
assert not isinstance(result, MockResponsesAPIStreamingIterator)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_fallback_tags_skip_responses_api_bridge():
|
||||
handler = LiteLLMCompletionTransformationHandler()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue