mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge cfc7d488cd into f285229b51
This commit is contained in:
commit
fcd780ab81
4 changed files with 336 additions and 37 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue