fix(responses): preserve citation offsets in recovered text

This commit is contained in:
Daniel Phang 2026-09-26 21:50:39 -07:00
parent d0b0247e3c
commit 76e5741084
3 changed files with 88 additions and 41 deletions

View file

@ -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):

View file

@ -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

View file

@ -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():