diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index 1f5b95e8c26..8dc5e30dd89 100644 --- a/litellm/completion_extras/litellm_responses_transformation/handler.py +++ b/litellm/completion_extras/litellm_responses_transformation/handler.py @@ -67,20 +67,7 @@ class ResponsesToCompletionBridgeHandler: else: raise ValueError("Unexpected responses stream payload") - response = base_response - if not base_response.output: - 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 - ) - if recovered_output: - response = base_response.model_copy(update={"output": recovered_output}) + response: Final = ResponsesToCompletionBridgeHandler._recover_stream_output(base_response, stream_events) if hidden_params: existing: Final = getattr(response, "_hidden_params", None) @@ -91,6 +78,20 @@ 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): diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index e260e5188bb..ca4dbcfa163 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,16 @@ _RESPONSES_API_ONLY_FIELDS: Final = frozenset((*Response.model_fields, *Response ) +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_annotation(citation, offset)} + return { + key: value + offset if key in ("start_index", "end_index") and isinstance(value, int) else value + for key, value in annotation.items() + } + + def _provider_metadata(response_fields: Mapping[str, object] | None) -> Mapping[str, object]: return MappingProxyType( { @@ -334,28 +345,29 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # Handle message items with output_text content if item_type == "message": content_list: Final = item.get("content", []) - response_text_parts: Final[list[str]] = [] - message_annotations: Final[list[ChatCompletionAnnotation]] = [] - has_output_text = False - for content_item in content_list: - if not isinstance(content_item, dict) or content_item.get("type") != "output_text": - continue - has_output_text = True - response_text = content_item.get("text", "") - response_text_parts.append(response_text if isinstance(response_text, str) else "") - annotations = LiteLLMResponsesTransformationHandler._convert_annotations_to_chat_format( - content_item.get("annotations", None) + 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 () ) - if annotations: - message_annotations.extend(annotations) + for part, offset in zip(text_parts, offsets) + ) + message_annotations: Final = list(chain.from_iterable(annotation_groups)) - if has_output_text: - msg = Message( + if text_parts: + msg: Final = Message( role=item.get("role", "assistant"), content="".join(response_text_parts), annotations=message_annotations or None, ) - choice = Choices(message=msg, finish_reason="stop", index=index) + 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 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 b09ac17d3ea..85364d1fc69 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,25 +610,51 @@ def test_transform_response_recovers_empty_output_from_raw_sse(): assert result.choices[0].message.content == "Recovered from SSE" -def test_transform_response_recovers_all_text_parts_from_raw_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, recovery_path): + 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 = LiteLLMResponsesTransformationHandler() - raw_sse = "\n".join( - [ - 'data: {"type":"response.output_text.done","output_index":0,"content_index":0,"item_id":"msg_from_stream","text":"Hello, "}', - 'data: {"type":"response.output_text.done","output_index":0,"content_index":1,"item_id":"msg_from_stream","text":"world!"}', - 'data: {"type":"response.completed","response":{"id":"resp_from_stream","object":"response","created_at":1760144904,"status":"completed","model":"gpt-5.4","output":[]}}', - "data: [DONE]", - "", - ] + citation = {"start_index": 0, "end_index": 5, "title": "Source", "url": "https://example.com"} + annotation = ( + {"type": "url_citation", "url_citation": citation} if nested_citation else {"type": "url_citation", **citation} + ) + events = ( + { + "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 = deepcopy(events) + raw_sse = "\n".join(f"data: {json.dumps(event)}" for event in events) + raw_response = ( + ResponsesToCompletionBridgeHandler._coerce_response_object(_make_empty_responses_api_response(), None, events) + if recovery_path == "bridge" + else _make_empty_responses_api_response() ) - raw_response = _make_empty_responses_api_response() model_response = _make_empty_model_response() logging_obj = Mock() - logging_obj.model_call_details = {"original_response": raw_sse} + logging_obj.model_call_details = {"original_response": raw_sse if recovery_path == "raw_sse" else None} result = handler.transform_response( model="gpt-5.4", @@ -644,6 +670,14 @@ def test_transform_response_recovers_all_text_parts_from_raw_sse(): assert len(result.choices) == 1 assert result.choices[0].message.content == "Hello, world!" + shifted_citation = {**citation, "start_index": 7, "end_index": 12} + shifted_annotation = ( + {"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():