mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(responses): drop citation offset recursion and test the real stream iterator
The code-quality recursive_detector gate rejected _offset_annotation because it called itself to reach the nested url_citation form. The nesting is only one level deep, so a flat _offset_citation_indices helper replaces the recursion The bridge recovery tests used handwritten stand-ins for LiteLLM's own Responses streaming iterators, which tests/unit/AGENTS.md forbids. They also skipped the part that matters most: whether the typed events the real iterator yields pass _is_recoverable_stream_event and dump into payloads that recovery can read. The tests now run SyncResponsesAPIStreamingIterator and ResponsesAPIStreamingIterator over an httpx SSE body for both output_item.done and output_text.done events, and all four cases fail when recovery is disabled
This commit is contained in:
parent
4f3973c39b
commit
9124f0981d
2 changed files with 78 additions and 31 deletions
|
|
@ -79,16 +79,21 @@ _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)}
|
||||
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 annotation.items()
|
||||
for key, value in fields.items()
|
||||
}
|
||||
|
||||
|
||||
def _offset_annotation(annotation: Mapping[str, object], offset: int) -> dict[str, object]:
|
||||
"""Shift citation indices in either the flat Responses form or the nested Chat Completions ``url_citation`` form."""
|
||||
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(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -1,7 +1,8 @@
|
|||
import json
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from openai.types.responses import ResponseOutputItemDoneEvent
|
||||
|
||||
|
|
@ -11,6 +12,8 @@ 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
|
||||
|
|
@ -244,40 +247,79 @@ def _output_item_done_event() -> dict:
|
|||
}
|
||||
|
||||
|
||||
class _SyncResponsesStream:
|
||||
def __init__(self, response: ResponsesAPIResponse):
|
||||
self.completed_response = SimpleNamespace(response=response)
|
||||
self._hidden_params = {}
|
||||
|
||||
def __iter__(self):
|
||||
return iter((_output_item_done_event(), self.completed_response))
|
||||
def _output_text_done_event() -> dict:
|
||||
return {
|
||||
"type": "response.output_text.done",
|
||||
"sequence_number": 1,
|
||||
"item_id": "msg_from_stream",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"text": "Recovered from stream",
|
||||
}
|
||||
|
||||
|
||||
class _AsyncResponsesStream:
|
||||
def __init__(self, response: ResponsesAPIResponse):
|
||||
self.completed_response = SimpleNamespace(response=response)
|
||||
self._hidden_params = {}
|
||||
|
||||
async def __aiter__(self):
|
||||
for event in (_output_item_done_event(), self.completed_response):
|
||||
yield event
|
||||
def _empty_completed_event() -> dict:
|
||||
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 test_collect_response_from_stream_recovers_output_items():
|
||||
bridge = ResponsesToCompletionBridgeHandler()
|
||||
|
||||
response = bridge._collect_response_from_stream(_SyncResponsesStream(_empty_responses_response()))
|
||||
|
||||
assert response.output == [_output_item_done_event()["item"]]
|
||||
def _sse_http_response(*events: dict) -> httpx.Response:
|
||||
body = "".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"))
|
||||
|
||||
|
||||
def _real_stream_iterator_kwargs(recoverable_event: dict) -> dict:
|
||||
"""Feed the real Responses streaming iterator an SSE body whose terminal event has empty output."""
|
||||
return {
|
||||
"response": _sse_http_response(recoverable_event, _empty_completed_event()),
|
||||
"model": "gpt-5.4",
|
||||
"responses_api_provider_config": OpenAIResponsesAPIConfig(),
|
||||
"logging_obj": _bridge_kwargs(stream=True)["logging_obj"],
|
||||
"custom_llm_provider": "chatgpt",
|
||||
}
|
||||
|
||||
|
||||
def _recovered_texts(response: ResponsesAPIResponse) -> list[str]:
|
||||
return [part["text"] for item in response.output for part in item["content"]]
|
||||
|
||||
|
||||
_RECOVERABLE_STREAM_EVENTS = 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):
|
||||
stream = SyncResponsesAPIStreamingIterator(**_real_stream_iterator_kwargs(recoverable_event))
|
||||
|
||||
response = 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():
|
||||
bridge = ResponsesToCompletionBridgeHandler()
|
||||
async def test_collect_response_from_async_stream_recovers_output_items(recoverable_event):
|
||||
stream = ResponsesAPIStreamingIterator(**_real_stream_iterator_kwargs(recoverable_event))
|
||||
|
||||
response = await bridge._collect_response_from_stream_async(_AsyncResponsesStream(_empty_responses_response()))
|
||||
response = await ResponsesToCompletionBridgeHandler()._collect_response_from_stream_async(stream)
|
||||
|
||||
assert response.output == [_output_item_done_event()["item"]]
|
||||
assert _recovered_texts(response) == ["Recovered from stream"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue