mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(responses): preserve citation offsets in recovered text
This commit is contained in:
parent
d0b0247e3c
commit
76e5741084
3 changed files with 88 additions and 41 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue