This commit is contained in:
dvp 2026-09-30 16:57:55 -04:00 • committed by GitHub
commit fcd780ab81
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 336 additions and 37 deletions

View file

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

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,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 {}

View file

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

View file

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