mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge 259821fbaa into e768ad55ce
This commit is contained in:
commit
2fc0bd092e
3 changed files with 192 additions and 12 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue