This commit is contained in:
agustin18 2026-10-03 20:08:47 +08:00 • committed by GitHub
commit 2fc0bd092e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 192 additions and 12 deletions

View file

@ -13,7 +13,11 @@ 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.integrations.custom_logger import converted_stream_requested
from litellm.types.llms.openai import (
ResponseInputParam,
ResponsesAPIOptionalRequestParams,
@ -81,6 +85,18 @@ class LiteLLMCompletionTransformationHandler:
)
)
converted_stream: Final = converted_stream_requested(kwargs) or converted_stream_requested(
litellm_completion_request
)
if converted_stream and not kwargs.get("_agentic_loop_depth"):
return MockResponsesAPIStreamingIterator(
model=model,
transformed_response=responses_api_response,
logging_obj=kwargs.get("logging_obj"),
custom_llm_provider=custom_llm_provider,
litellm_metadata=kwargs.get("litellm_metadata"),
)
return responses_api_response
elif isinstance(litellm_completion_response, litellm.CustomStreamWrapper):
@ -126,6 +142,18 @@ class LiteLLMCompletionTransformationHandler:
)
)
converted_stream: Final = converted_stream_requested(kwargs) or converted_stream_requested(
litellm_completion_request
)
if converted_stream and not kwargs.get("_agentic_loop_depth"):
return MockResponsesAPIStreamingIterator(
model=litellm_completion_request.get("model") or "",
transformed_response=responses_api_response,
logging_obj=kwargs.get("logging_obj"),
custom_llm_provider=litellm_completion_request.get("custom_llm_provider"),
litellm_metadata=kwargs.get("litellm_metadata"),
)
return responses_api_response
elif isinstance(litellm_completion_response, litellm.CustomStreamWrapper):

View file

@ -334,9 +334,15 @@ class BaseResponsesAPIStreamingIterator:
# set hidden params for response headers (e.g., x-litellm-model-id)
# This matches the stream wrapper in litellm/litellm_core_utils/streaming_handler.py
_model_call_details: Final = getattr(self.logging_obj, "model_call_details", None)
_optional_params: Final = (
_typed_gets_litellm_params(_model_call_details.get)("litellm_params", {}) # mutable-ok: fallback dict
if isinstance(_model_call_details, dict)
else {} # mutable-ok: fallback empty mapping
)
_api_base: Final = get_api_base(
model=model or "",
optional_params=_typed_gets_litellm_params(self.logging_obj.model_call_details.get)("litellm_params", {}),
optional_params=_optional_params,
)
self._hidden_params: dict[str, object] = {
"model_id": _model_id_from_metadata(litellm_metadata),
@ -532,7 +538,7 @@ class BaseResponsesAPIStreamingIterator:
raise
def _log_completed_response(self, *, is_async: bool) -> None:
if self._completed_response_logged:
if self._completed_response_logged or self.logging_obj is None:
return
self._completed_response_logged = True
@ -1191,22 +1197,35 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
def __init__(
self,
response: httpx.Response,
model: str,
responses_api_provider_config: BaseResponsesAPIConfig,
logging_obj: LiteLLMLoggingObj,
response: httpx.Response | None = None,
model: str = "",
responses_api_provider_config: BaseResponsesAPIConfig | None = None,
logging_obj: LiteLLMLoggingObj | None = None,
litellm_metadata: dict[str, object] | None = None,
custom_llm_provider: str | None = None,
request_data: dict[str, object] | None = None,
call_type: str | None = None,
transformed_response: ResponsesAPIResponse | None = None,
):
transformed: Final = responses_api_provider_config.transform_response_api_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
transformed: Final[ResponsesAPIResponse | None] = (
transformed_response
if transformed_response is not None
else (
responses_api_provider_config.transform_response_api_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
if responses_api_provider_config is not None and response is not None
else None
)
)
if transformed is None:
raise ValueError(
"Either transformed_response or both responses_api_provider_config and response must be provided to MockResponsesAPIStreamingIterator"
)
super().__init__(
response=httpx.Response(200),
response=response or httpx.Response(200),
model=model,
responses_api_provider_config=None,
logging_obj=logging_obj,

View file

@ -223,3 +223,136 @@ async def test_bridged_follow_up_turn_keeps_the_addressed_response_id_off_the_pr
)
assert isinstance(response, ResponsesAPIResponse)
assert [item.type for item in response.output] == ["message"]
@pytest.mark.parametrize(
"converted_stream_flag",
[
"_websearch_interception_converted_stream",
"_code_interpreter_interception_converted_stream",
"_headroom_interception_converted_stream",
],
)
@pytest.mark.asyncio
async def test_async_fallback_wraps_converted_stream_as_synthetic_stream(converted_stream_flag):
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
from litellm.types.utils import Choices, Message, ModelResponse
handler = LiteLLMCompletionTransformationHandler()
async def fake_acompletion(**kwargs):
return ModelResponse(
id="chatcmpl-test",
created=1,
model="gpt-4o",
object="chat.completion",
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(content="Search summary", role="assistant"),
)
],
)
with patch("litellm.acompletion", fake_acompletion):
response = await handler.response_api_handler(
model="gpt-4o",
input="Search for the weather",
responses_api_request={},
custom_llm_provider="hosted_vllm",
_is_async=True,
**{converted_stream_flag: True},
)
assert isinstance(response, BaseResponsesAPIStreamingIterator)
events = [event async for event in response]
assert len(events) > 0
assert getattr(events[-1], "type", None) == "response.completed"
@pytest.mark.parametrize(
"converted_stream_flag",
[
"_websearch_interception_converted_stream",
"_code_interpreter_interception_converted_stream",
"_headroom_interception_converted_stream",
],
)
def test_sync_fallback_wraps_converted_stream_as_synthetic_stream(converted_stream_flag):
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
from litellm.types.utils import Choices, Message, ModelResponse
handler = LiteLLMCompletionTransformationHandler()
def fake_completion(**kwargs):
return ModelResponse(
id="chatcmpl-test-sync",
created=1,
model="gpt-4o",
object="chat.completion",
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(content="Sync search summary", role="assistant"),
)
],
)
with patch("litellm.completion", fake_completion):
response = handler.response_api_handler(
model="gpt-4o",
input="Search for the weather",
responses_api_request={},
custom_llm_provider="hosted_vllm",
_is_async=False,
**{converted_stream_flag: True},
)
assert isinstance(response, BaseResponsesAPIStreamingIterator)
events = list(response)
assert len(events) > 0
assert getattr(events[-1], "type", None) == "response.completed"
def test_mock_responses_streaming_iterator_validation_and_config_branches():
import httpx
from litellm.responses.streaming_iterator import MockResponsesAPIStreamingIterator
from litellm.types.llms.openai import ResponsesAPIResponse
with pytest.raises(ValueError, match="Either transformed_response or both"):
MockResponsesAPIStreamingIterator()
class _MockConfig:
def transform_response_api_response(self, **kwargs):
return ResponsesAPIResponse(
id="resp_cfg_test",
created_at=1,
status="completed",
model="test-model",
object="response",
output=[],
)
class _MockLoggingObj:
def __init__(self):
self.model_call_details = {"litellm_params": {"api_key": "fake"}}
async def async_success_handler(self, *args, **kwargs):
pass
def success_handler(self, *args, **kwargs):
pass
logging_obj = _MockLoggingObj()
iterator = MockResponsesAPIStreamingIterator(
response=httpx.Response(200),
model="gpt-4o",
responses_api_provider_config=_MockConfig(),
logging_obj=logging_obj,
)
events = list(iterator)
assert len(events) > 0
assert getattr(events[-1], "type", None) == "response.completed"