diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index 642a78789b2..8dc5e30dd89 100644 --- a/litellm/completion_extras/litellm_responses_transformation/handler.py +++ b/litellm/completion_extras/litellm_responses_transformation/handler.py @@ -2,12 +2,13 @@ Handler for transforming /chat/completions api requests to litellm.responses requests """ -from collections.abc import AsyncIterable, Coroutine, Iterable +from collections.abc import AsyncIterable, Coroutine, Iterable, Mapping from typing import TYPE_CHECKING, Any, Final, Union +from pydantic import BaseModel from typing_extensions import TypedDict -from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.llms.openai import ResponsesAPIResponse, ResponsesAPIStreamEvents if TYPE_CHECKING: from litellm import CustomStreamWrapper, LiteLLMLoggingObj, ModelResponse @@ -54,17 +55,20 @@ class ResponsesToCompletionBridgeHandler: def _coerce_response_object( response_obj: object, hidden_params: dict | None, + stream_events: Iterable[object] = (), ) -> "ResponsesAPIResponse": if isinstance(response_obj, ResponsesAPIResponse): - response = response_obj + base_response = response_obj elif isinstance(response_obj, dict): try: - response = ResponsesAPIResponse(**response_obj) + base_response = ResponsesAPIResponse(**response_obj) except Exception: - response = ResponsesAPIResponse.model_construct(**response_obj) + base_response = ResponsesAPIResponse.model_construct(**response_obj) else: raise ValueError("Unexpected responses stream payload") + response: Final = ResponsesToCompletionBridgeHandler._recover_stream_output(base_response, stream_events) + if hidden_params: existing: Final = getattr(response, "_hidden_params", None) if not isinstance(existing, dict) or not existing: @@ -74,9 +78,38 @@ class ResponsesToCompletionBridgeHandler: existing.setdefault(key, value) return response + @staticmethod + def _recover_stream_output(response: ResponsesAPIResponse, stream_events: Iterable[object]) -> ResponsesAPIResponse: + if response.output: + return response + from .transformation import LiteLLMResponsesTransformationHandler + + parsed_chunks: Final = ( + payload + for event in stream_events + if (payload := ResponsesToCompletionBridgeHandler._stream_event_payload(event)) is not None + ) + recovered_output: Final = LiteLLMResponsesTransformationHandler.recover_output_items_from_chunks(parsed_chunks) + return response.model_copy(update={"output": recovered_output}) if recovered_output else response + + @staticmethod + def _stream_event_payload(event: object) -> Mapping[str, object] | None: + if isinstance(event, Mapping): + return event + if isinstance(event, BaseModel): + return event.model_dump() + return None + + @staticmethod + def _is_recoverable_stream_event(event: object) -> bool: + event_type: Final = event.get("type") if isinstance(event, Mapping) else getattr(event, "type", None) + return event_type in ( + ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE, + ) + def _collect_response_from_stream(self, stream_iter: Iterable[object]) -> "ResponsesAPIResponse": - for _ in stream_iter: - pass + stream_events: Final = tuple(event for event in stream_iter if self._is_recoverable_stream_event(event)) completed: Final[object] = getattr(stream_iter, "completed_response", None) response_obj: Final[object] = getattr(completed, "response", None) if completed else None @@ -84,14 +117,13 @@ class ResponsesToCompletionBridgeHandler: raise ValueError("Stream ended without a completed response") hidden_params: Final = getattr(stream_iter, "_hidden_params", None) - response: Final = self._coerce_response_object(response_obj, hidden_params) + response: Final = self._coerce_response_object(response_obj, hidden_params, stream_events) if not isinstance(response, ResponsesAPIResponse): raise ValueError("Stream completed response is invalid") return response async def _collect_response_from_stream_async(self, stream_iter: AsyncIterable[object]) -> "ResponsesAPIResponse": - async for _ in stream_iter: - pass + stream_events: Final = tuple([event async for event in stream_iter if self._is_recoverable_stream_event(event)]) completed: Final[object] = getattr(stream_iter, "completed_response", None) response_obj: Final[object] = getattr(completed, "response", None) if completed else None @@ -99,7 +131,7 @@ class ResponsesToCompletionBridgeHandler: raise ValueError("Stream ended without a completed response") hidden_params: Final = getattr(stream_iter, "_hidden_params", None) - response: Final = self._coerce_response_object(response_obj, hidden_params) + response: Final = self._coerce_response_object(response_obj, hidden_params, stream_events) if not isinstance(response, ResponsesAPIResponse): raise ValueError("Stream completed response is invalid") return response diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 71d3f1e900e..5bdeb503a3c 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -5,6 +5,7 @@ Handler for transforming /chat/completions api requests to litellm.responses req import json import os from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence +from itertools import accumulate, chain from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, TypeVar, Union, cast @@ -78,6 +79,20 @@ _RESPONSES_API_ONLY_FIELDS: Final = frozenset((*Response.model_fields, *Response ) +def _offset_citation_indices(fields: Mapping[str, object], offset: int) -> dict[str, object]: + return { + key: value + offset if key in ("start_index", "end_index") and isinstance(value, int) else value + for key, value in fields.items() + } + + +def _offset_annotation(annotation: Mapping[str, object], offset: int) -> dict[str, object]: + citation: Final = annotation.get("url_citation") + if isinstance(citation, dict): + return {**annotation, "url_citation": _offset_citation_indices(citation, offset)} + return _offset_citation_indices(annotation, offset) + + def _provider_metadata(response_fields: Mapping[str, object] | None) -> Mapping[str, object]: return MappingProxyType( { @@ -334,22 +349,30 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # Handle message items with output_text content if item_type == "message": content_list: Final = item.get("content", []) - for content_item in content_list: - if isinstance(content_item, dict): - content_type = content_item.get("type") - if content_type == "output_text": - response_text = content_item.get("text", "") - # Extract annotations from content if present - annotations = LiteLLMResponsesTransformationHandler._convert_annotations_to_chat_format( - content_item.get("annotations", None) - ) - msg = Message( - role=item.get("role", "assistant"), - content=response_text if response_text else "", - annotations=annotations, - ) - choice = Choices(message=msg, finish_reason="stop", index=index) - return choice, index + 1 + text_parts: Final = tuple( + part for part in content_list if isinstance(part, dict) and part.get("type") == "output_text" + ) + response_text_parts: Final = tuple( + text if isinstance(text := part.get("text"), str) else "" for part in text_parts + ) + offsets: Final = accumulate(map(len, response_text_parts), initial=0) + annotation_groups: Final = ( + tuple( + _offset_annotation(annotation, offset) + for annotation in self._convert_annotations_to_chat_format(part.get("annotations")) or () + ) + for part, offset in zip(text_parts, offsets) + ) + message_annotations: Final = list(chain.from_iterable(annotation_groups)) + + if text_parts: + msg: Final = Message( + role=item.get("role", "assistant"), + content="".join(response_text_parts), + annotations=message_annotations or None, + ) + choice: Final = Choices(message=msg, finish_reason="stop", index=index) + return choice, index + 1 # function_call / custom_tool_call dicts are intercepted and accumulated by # _convert_response_output_to_choices before this callback is reached @@ -833,18 +856,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return cast(list[dict[str, object]], response_output) @classmethod - def _recover_output_items_from_raw_sse(cls, raw_sse: str | None) -> list[dict[str, object]]: - if not raw_sse or not isinstance(raw_sse, str): - return [] - + def recover_output_items_from_chunks(cls, parsed_chunks: Iterable[Mapping[str, object]]) -> list[dict[str, object]]: recovered_output_items: Final[dict[int, dict[str, object]]] = {} recovered_text_only_items: Final[dict[int, dict[str, object]]] = {} - for chunk in raw_sse.splitlines(): - parsed_chunk = parse_sse_json_chunk(chunk) - if parsed_chunk is None: - continue - + for parsed_chunk in parsed_chunks: event_type = parsed_chunk.get("type") if event_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED: @@ -881,6 +897,16 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return [] + @classmethod + def _recover_output_items_from_raw_sse(cls, raw_sse: str | None) -> list[dict[str, object]]: + if not raw_sse or not isinstance(raw_sse, str): + return [] + + parsed_chunks: Final = ( + parsed_chunk for chunk in raw_sse.splitlines() if (parsed_chunk := parse_sse_json_chunk(chunk)) is not None + ) + return cls.recover_output_items_from_chunks(parsed_chunks) + @classmethod def _recover_output_items_from_logging(cls, logging_obj: "LiteLLMLoggingObj") -> list[dict[str, object]]: model_call_details: Final = getattr(logging_obj, "model_call_details", {}) or {} diff --git a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py index c5d7ca96a21..871f2a3ff70 100644 --- a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py +++ b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py @@ -1,8 +1,15 @@ +import json +from collections.abc import Iterator from datetime import datetime +from itertools import chain +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest - +from openai.types.responses import ResponseOutputItemDoneEvent +from pydantic import JsonValue +from typing_extensions import ReadOnly, TypedDict import litellm from litellm.completion_extras.litellm_responses_transformation.handler import ( @@ -10,6 +17,9 @@ from litellm.completion_extras.litellm_responses_transformation.handler import ( ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig +from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator, SyncResponsesAPIStreamingIterator +from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ModelResponse @@ -218,6 +228,165 @@ def _completed_chat_response() -> ModelResponse: ) +def _empty_responses_response() -> ResponsesAPIResponse: + return ResponsesAPIResponse.model_construct(output=[], error=None) + + +def _output_item_done_event() -> dict[str, JsonValue]: + return { + "type": "response.output_item.done", + "output_index": 0, + "item": { + "type": "message", + "id": "msg_from_stream", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "Recovered from stream", + "annotations": [], + } + ], + }, + } + + +def _output_text_done_event() -> dict[str, JsonValue]: + return { + "type": "response.output_text.done", + "sequence_number": 1, + "item_id": "msg_from_stream", + "output_index": 0, + "content_index": 0, + "text": "Recovered from stream", + } + + +def _empty_completed_event() -> dict[str, JsonValue]: + return { + "type": "response.completed", + "sequence_number": 2, + "response": { + "id": "resp_test", + "object": "response", + "created_at": 0, + "status": "completed", + "model": "gpt-5.4", + "output": [], + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + }, + } + + +def _sse_http_response(*events: dict[str, JsonValue]) -> httpx.Response: + body: Final = "".join(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n" for event in events) + return httpx.Response(200, content=body.encode(), request=httpx.Request("POST", "https://example.com/responses")) + + +class _StreamIteratorKwargs(TypedDict): + response: ReadOnly[httpx.Response] + model: ReadOnly[str] + responses_api_provider_config: ReadOnly[OpenAIResponsesAPIConfig] + logging_obj: ReadOnly[LiteLLMLogging] + custom_llm_provider: ReadOnly[str] + + +def _real_stream_iterator_kwargs(recoverable_event: dict[str, JsonValue]) -> _StreamIteratorKwargs: + return { + "response": _sse_http_response(recoverable_event, _empty_completed_event()), + "model": "gpt-5.4", + "responses_api_provider_config": OpenAIResponsesAPIConfig(), + "logging_obj": LiteLLMLogging( + litellm_call_id="test-call", + call_type="completion", + model="gpt-5.4", + messages=[{"role": "user", "content": "hi"}], + function_id="fn-id", + stream=True, + start_time=datetime(2026, 1, 1), + ), + "custom_llm_provider": "chatgpt", + } + + +def _recovered_texts(response: ResponsesAPIResponse) -> list[str]: + content_parts: Final = chain.from_iterable(item["content"] for item in response.output) + return [part["text"] for part in content_parts] + + +_RECOVERABLE_STREAM_EVENTS: Final = pytest.mark.parametrize( + "recoverable_event", + [_output_item_done_event(), _output_text_done_event()], + ids=["output_item_done", "output_text_done"], +) + + +@_RECOVERABLE_STREAM_EVENTS +def test_collect_response_from_stream_recovers_output_items(recoverable_event: dict[str, JsonValue]) -> None: + stream: Final = SyncResponsesAPIStreamingIterator(**_real_stream_iterator_kwargs(recoverable_event)) + + response: Final = ResponsesToCompletionBridgeHandler()._collect_response_from_stream(stream) + + assert _recovered_texts(response) == ["Recovered from stream"] + + +@_RECOVERABLE_STREAM_EVENTS +@pytest.mark.asyncio +async def test_collect_response_from_async_stream_recovers_output_items(recoverable_event: dict[str, JsonValue]) -> None: + stream: Final = ResponsesAPIStreamingIterator(**_real_stream_iterator_kwargs(recoverable_event)) + + response: Final = await ResponsesToCompletionBridgeHandler()._collect_response_from_stream_async(stream) + + assert _recovered_texts(response) == ["Recovered from stream"] + + +@pytest.mark.parametrize( + "terminal_payload", + [ + {"id": "resp_test", "created_at": 0, "output": []}, + {"output": []}, + ], + ids=["validated-terminal", "partial-terminal"], +) +def test_recovery_accepts_sdk_events_and_dictionary_terminals(terminal_payload: dict[str, JsonValue]) -> None: + event: Final = ResponseOutputItemDoneEvent(**_output_item_done_event(), sequence_number=1) + + response: Final = ResponsesToCompletionBridgeHandler._coerce_response_object( + terminal_payload, + {"headers": {"x-request-id": "req_test"}}, + (object(), event), + ) + + assert response.output == [event.item.model_dump()] + assert response._hidden_params["headers"] == {"x-request-id": "req_test"} + assert terminal_payload["output"] == [] + + +def test_recovery_preserves_complete_response_without_consuming_events() -> None: + terminal: Final = ResponsesAPIResponse.model_construct(output=[_output_item_done_event()["item"]]) + + def unavailable_events() -> Iterator[object]: + raise AssertionError("Complete terminal output must bypass event recovery") + yield + + response: Final = ResponsesToCompletionBridgeHandler._coerce_response_object(terminal, None, unavailable_events()) + + assert response is terminal + assert response.output == [_output_item_done_event()["item"]] + + +def test_recovery_keeps_empty_terminal_when_no_output_can_be_recovered() -> None: + terminal: Final = _empty_responses_response() + + response: Final = ResponsesToCompletionBridgeHandler._coerce_response_object(terminal, None, (object(),)) + + assert response is terminal + assert response.output == [] + + @pytest.mark.asyncio async def test_acompletion_streams_completed_model_response(): """A streaming request whose bridge call comes back already completed must still be diff --git a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index 282b84104a6..46e407b124c 100644 --- a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -610,6 +610,78 @@ def test_transform_response_recovers_empty_output_from_raw_sse(): assert result.choices[0].message.content == "Recovered from SSE" +@pytest.mark.parametrize("nested_citation", [False, True]) +@pytest.mark.parametrize("recovery_path", ["raw_sse", "bridge"]) +def test_transform_response_recovers_text_and_citation_offsets( + nested_citation: bool, recovery_path: Literal["raw_sse", "bridge"] +) -> None: + from copy import deepcopy + + from litellm.completion_extras.litellm_responses_transformation.handler import ( + ResponsesToCompletionBridgeHandler, + ) + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + + handler: Final = LiteLLMResponsesTransformationHandler() + citation: Final = {"start_index": 0, "end_index": 5, "title": "Source", "url": "https://example.com"} + annotation: Final = ( + {"type": "url_citation", "url_citation": citation} if nested_citation else {"type": "url_citation", **citation} + ) + events: Final = ( + { + "type": "response.output_text.done", + "output_index": 0, + "content_index": 1, + "item_id": "msg_from_stream", + "text": "world!", + "annotations": [annotation], + }, + { + "type": "response.output_text.done", + "output_index": 0, + "content_index": 0, + "item_id": "msg_from_stream", + "text": "Hello, ", + "annotations": [annotation], + }, + ) + original_events: Final = deepcopy(events) + raw_sse: Final = "\n".join(f"data: {json.dumps(event)}" for event in events) + raw_response: Final = ( + ResponsesToCompletionBridgeHandler._coerce_response_object(_make_empty_responses_api_response(), None, events) + if recovery_path == "bridge" + else _make_empty_responses_api_response() + ) + model_response: Final = _make_empty_model_response() + logging_obj: Final = Mock() + logging_obj.model_call_details = {"original_response": raw_sse if recovery_path == "raw_sse" else None} + + result: Final = handler.transform_response( + model="gpt-5.4", + raw_response=raw_response, + model_response=model_response, + logging_obj=logging_obj, + request_data={"model": "gpt-5.4"}, + messages=[{"role": "user", "content": "Reply with exactly: Hello, world!"}], + optional_params={}, + litellm_params={}, + encoding=Mock(), + ) + + assert len(result.choices) == 1 + assert result.choices[0].message.content == "Hello, world!" + shifted_citation: Final = {**citation, "start_index": 7, "end_index": 12} + shifted_annotation: Final = ( + {"type": "url_citation", "url_citation": shifted_citation} + if nested_citation + else {"type": "url_citation", **shifted_citation} + ) + assert result.choices[0].message.annotations == [annotation, shifted_annotation] + assert events == original_events + + def test_transform_response_recovers_output_item_done_from_raw_sse(): from litellm.completion_extras.litellm_responses_transformation.transformation import ( LiteLLMResponsesTransformationHandler,