diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index 0da04b5c25d..22dec3e57af 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -3,6 +3,8 @@ import uuid from collections.abc import Mapping, Sequence from typing import Any, Final, cast +from openai.types.responses import ResponseFunctionToolCall, ResponseFunctionWebSearch, ResponseReasoningItem + import litellm from litellm.main import stream_chunk_builder from litellm.responses.litellm_completion_transformation.custom_tools import ( @@ -48,6 +50,12 @@ from litellm.types.llms.openai import ( WebSearchCallInProgressEvent, WebSearchCallSearchingEvent, ) +from litellm.types.responses.main import ( + CustomToolCallOutputItem, + GenericResponseOutputItem, + OutputFunctionToolCall, + OutputText, +) from litellm.types.utils import Delta as ChatCompletionDelta from litellm.types.utils import ( ModelResponse, @@ -77,9 +85,9 @@ def _output_items_with_id(items: tuple[Any, ...], item_type: str, item_id: str | ) -def _delta_has_signed_thinking_block(delta: object) -> bool: +def _delta_has_thinking_block(delta: object) -> bool: blocks: Final = getattr(delta, "thinking_blocks", None) or () - return any(isinstance(b, dict) and (b.get("signature") or b.get("data")) for b in blocks) + return any(isinstance(b, dict) and (b.get("thinking") or b.get("signature") or b.get("data")) for b in blocks) class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): @@ -131,7 +139,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._tool_item_id_by_call_id: dict[str, str] = {} # mutable-ok: filled per call id as tool call events stream self._tool_call_id_by_index: dict[int, str] = {} self._ambiguous_tool_call_indexes: set[int] = set() - self._next_tool_output_index: int = 1 # output_index=0 reserved for the message item + self._next_tool_output_index: int = 1 self._final_tool_events_queued: bool = False self._sequence_number: int = 0 self._cached_reasoning_item_id: str | None = None @@ -143,6 +151,9 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._reasoning_active = False self._reasoning_done_emitted = False self._reasoning_item_id: str | None = None + self._reasoning_output_index: int = 0 + self._reasoning_start_chunk_index: int = 0 + self._completed_reasoning_items: tuple[tuple[int, ResponseReasoningItem], ...] = () self._accumulated_reasoning_content_parts: list[str] = [] self._accumulated_provider_specific_fields: dict[str, object] = {} self._custom_tool_names: set[str] = extract_custom_tool_names(self.responses_api_request.get("tools")) @@ -198,7 +209,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return delta.content or delta.function_call or delta.tool_calls or chunk.choices[0].finish_reason is not None - def _reserve_web_search_indexes(self, provider_fields: object) -> None: + def _reserve_web_search_indexes(self, provider_fields: object, *, eager: bool = False) -> None: if not isinstance(provider_fields, dict): return calls: Final = provider_fields.get("web_search_calls") @@ -213,12 +224,24 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): if call_id: output_index = self._get_or_assign_tool_output_index(call_id) self._web_search_calls[call_id] = item - if status == "in_progress": + if not eager and status == "in_progress": self._pending_tool_events = [ event for event in self._pending_tool_events if getattr(event, "output_index", None) != output_index ] + if ( + eager + and status in ("in_progress", "searching", "completed", "failed") + and call_id not in self._queued_web_search_call_ids + ): + self._queue_web_search_events( + (search_call := ResponseFunctionWebSearch.model_validate(cast(object, item))).id.removeprefix( + "ws_" + ), + search_call, + finalize=False, + ) def _tool_call_id(self, tool_call: object) -> str: index: Final = self._normalize_tool_call_index(tool_call) @@ -336,7 +359,6 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): if web_search_call is not None: if call_id not in self._queued_web_search_call_ids: self._queue_web_search_events(call_id, web_search_call) - self._queued_web_search_call_ids.add(call_id) continue # Track if this is a new tool call that wasn't streamed @@ -399,32 +421,36 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): ) self._pending_tool_events.append(item_done_event) - def _queue_web_search_events(self, call_id: str, web_search_call: object) -> None: - from openai.types.responses import ResponseFunctionWebSearch - + def _queue_web_search_events(self, call_id: str, web_search_call: object, *, finalize: bool = True) -> None: item: Final = ( web_search_call if isinstance(web_search_call, ResponseFunctionWebSearch) else ResponseFunctionWebSearch.model_validate(web_search_call) ) output_index: Final = self._get_or_assign_tool_output_index(call_id) - self._sequence_number += 1 - added: Final = OutputItemAddedEvent( - type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, - output_index=output_index, - item=BaseLiteLLMOpenAIResponseObject( - **{ - "id": item.id, - "type": item.type, - "status": "in_progress", - "action": None, - } - ), - ) - added.__dict__["sequence_number"] = self._sequence_number - self._pending_tool_events.append(added) + if self._tool_item_id_by_call_id.get(call_id) != item.id: + self._pending_tool_events = [ + event for event in self._pending_tool_events if getattr(event, "output_index", None) != output_index + ] + self._tool_item_id_by_call_id[call_id] = item.id + self._sequence_number += 1 + added: Final = OutputItemAddedEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + output_index=output_index, + item=BaseLiteLLMOpenAIResponseObject( + **{"id": item.id, "type": item.type, "status": "in_progress", "action": None} + ), + ) + added.__dict__["sequence_number"] = self._sequence_number + self._pending_tool_events.append(added) + self._sequence_number += 1 + in_progress: Final = WebSearchCallInProgressEvent( + type=ResponsesAPIStreamEvents.WEB_SEARCH_CALL_IN_PROGRESS, output_index=output_index, item_id=item.id + ).model_copy(update={"sequence_number": self._sequence_number}) + self._pending_tool_events.append(in_progress) + if not finalize and item.status in ("in_progress", "searching"): + return for event_type, event_class in ( - (ResponsesAPIStreamEvents.WEB_SEARCH_CALL_IN_PROGRESS, WebSearchCallInProgressEvent), (ResponsesAPIStreamEvents.WEB_SEARCH_CALL_SEARCHING, WebSearchCallSearchingEvent), (ResponsesAPIStreamEvents.WEB_SEARCH_CALL_COMPLETED, WebSearchCallCompletedEvent), ): @@ -441,6 +467,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): item=BaseLiteLLMOpenAIResponseObject(**item.model_dump()), ) ) + self._queued_web_search_call_ids.add(call_id) def _adopt_response_id_from_chunk(self, chunk: ModelResponseStream) -> None: if self._cached_response_id is not None: @@ -607,7 +634,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._cached_item_id = f"msg_{uuid.uuid4()}" self.sent_message_item_added_event = True self.sent_content_part_added_event = True - if self._cached_reasoning_item_id is not None: + if self._cached_reasoning_item_id is not None or 0 in self._tool_output_index_by_call_id.values(): self._message_output_index = self._next_tool_output_index self._next_tool_output_index += 1 else: @@ -642,11 +669,11 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): for key, val in src.items(): self._accumulated_provider_specific_fields[key] = val - def create_litellm_model_response(self) -> ModelResponse | None: + def create_litellm_model_response(self, *, chunk_start: int = 0) -> ModelResponse | None: response: Final = cast( ModelResponse | None, stream_chunk_builder( - chunks=self.collected_chat_completion_chunks, + chunks=self.collected_chat_completion_chunks[chunk_start:], logging_obj=self.litellm_logging_obj, ), ) @@ -659,7 +686,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return response def _encoded_thinking_blocks(self) -> str | None: - response: Final = self.create_litellm_model_response() + response: Final = self.create_litellm_model_response(chunk_start=self._reasoning_start_chunk_index) if response is None: return None thinking_blocks: Final[Sequence[Mapping[str, object]]] = ( @@ -709,7 +736,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return ReasoningSummaryTextDoneEvent( type=ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DONE, item_id=reasoning_item_id, - output_index=0, + output_index=self._reasoning_output_index, sequence_number=sequence_number, summary_index=0, text=reasoning_content, @@ -740,7 +767,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return ReasoningSummaryPartDoneEvent( type=ResponsesAPIStreamEvents.REASONING_SUMMARY_PART_DONE, item_id=reasoning_item_id, - output_index=0, + output_index=self._reasoning_output_index, sequence_number=sequence_number, summary_index=0, part=BaseLiteLLMOpenAIResponseObject( @@ -851,7 +878,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): """ return OutputItemDoneEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, - output_index=0, + output_index=self._reasoning_output_index, sequence_number=sequence_number, item=BaseLiteLLMOpenAIResponseObject( **{ @@ -868,6 +895,28 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): ), ) + def _queue_reasoning_done_events(self) -> None: + reasoning_content: Final = "".join(self._accumulated_reasoning_content_parts) + reasoning_item_id: Final = self._reasoning_item_id or self._cached_reasoning_item_id or mint_reasoning_item_id() + self._sequence_number += 1 + text_done: Final = self.create_reasoning_summary_text_done_event( + reasoning_item_id, reasoning_content, self._sequence_number + ) + self._sequence_number += 1 + part_done: Final = self.create_reasoning_summary_part_done_event( + reasoning_item_id, reasoning_content, self._sequence_number + ) + self._sequence_number += 1 + item_done: Final = self.create_reasoning_output_item_done_event( + reasoning_item_id, reasoning_content, self._sequence_number + ) + self._completed_reasoning_items += ( + (self._reasoning_output_index, ResponseReasoningItem.model_validate(item_done.item.model_dump())), + ) + self._pending_response_events.extend((text_done, part_done, item_done)) + self._reasoning_done_emitted = True + self._reasoning_active = False + def return_default_done_events( self, litellm_complete_object: ModelResponse ) -> BaseLiteLLMOpenAIResponseObject | None: @@ -943,27 +992,37 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): else: raise StopAsyncIteration - def _ensure_output_item_for_chunk(self, chunk: ModelResponseStream) -> None: + def _ensure_output_item_for_chunk( + self, chunk: ModelResponseStream, *, allow_reasoning_resumption: bool = False + ) -> None: # Change: Never return a value, just enqueue output item events - if self.sent_output_item_added_event: + if self.sent_output_item_added_event and not allow_reasoning_resumption: return if not chunk.choices: return delta: Final = chunk.choices[0].delta - self._sequence_number += 1 - self.sent_output_item_added_event = True - # Reasoning-first - if (hasattr(delta, "reasoning_content") and delta.reasoning_content) or _delta_has_signed_thinking_block(delta): + if (hasattr(delta, "reasoning_content") and delta.reasoning_content) or _delta_has_thinking_block(delta): + if self._reasoning_active: + return + if self.sent_output_item_added_event: + self._pending_response_events.extend(self._pending_tool_events) + self._pending_tool_events.clear() + self._reasoning_output_index = self._next_tool_output_index + self._next_tool_output_index += 1 self._reasoning_active = True - if self._cached_reasoning_item_id is None: - self._cached_reasoning_item_id = mint_reasoning_item_id() + self._reasoning_done_emitted = False + self._reasoning_start_chunk_index = len(self.collected_chat_completion_chunks) + self._accumulated_reasoning_content_parts = [] + self._cached_reasoning_item_id = mint_reasoning_item_id() self._reasoning_item_id = self._cached_reasoning_item_id + self.sent_output_item_added_event = True + self._sequence_number += 1 event = OutputItemAddedEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, - output_index=0, + output_index=self._reasoning_output_index, item=BaseLiteLLMOpenAIResponseObject( **{ "id": self._cached_reasoning_item_id, @@ -977,6 +1036,11 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._pending_response_events.append(event) return + if self.sent_output_item_added_event: + return + self.sent_output_item_added_event = True + self._sequence_number += 1 + # Tool-first if hasattr(delta, "tool_calls") and delta.tool_calls: # Tool calls already handled via _queue_tool_call_delta_events @@ -1014,6 +1078,15 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): chunk = await self.litellm_custom_stream_wrapper.__anext__() if chunk is not None: chunk = cast(ModelResponseStream, chunk) + if ( + not self.sent_output_item_added_event + and chunk.choices + and chunk.choices[0].delta.tool_calls + and not chunk.choices[0].delta.content + and not getattr(chunk.choices[0].delta, "reasoning_content", None) + and not _delta_has_thinking_block(chunk.choices[0].delta) + ): + self._next_tool_output_index = 0 for src in ( getattr(chunk, "provider_specific_fields", None), getattr( @@ -1024,8 +1097,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): ): if src and isinstance(src, dict): self._merge_provider_specific_fields(src) - self._reserve_web_search_indexes(src) - self._ensure_output_item_for_chunk(chunk) + self._reserve_web_search_indexes(src, eager=True) + self._ensure_output_item_for_chunk(chunk, allow_reasoning_resumption=True) # Proceed to transformation self.collected_chat_completion_chunks.append( self._snapshot_chunk_for_stream_chunk_builder(chunk) @@ -1037,47 +1110,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): if delta and hasattr(delta, "reasoning_content") and delta.reasoning_content: self._accumulated_reasoning_content_parts.append(delta.reasoning_content) if self._is_reasoning_end(chunk): - reasoning_content = "".join(self._accumulated_reasoning_content_parts) - - # Ensure we have a valid reasoning_item_id - self._cached_reasoning_item_id = ( - self._reasoning_item_id - or self._cached_reasoning_item_id - or mint_reasoning_item_id() - ) - reasoning_item_id = self._cached_reasoning_item_id - - # Create text.done event first with its own sequence number - self._sequence_number += 1 - text_done_event = self.create_reasoning_summary_text_done_event( - reasoning_item_id=reasoning_item_id, - reasoning_content=reasoning_content, - sequence_number=self._sequence_number, - ) - - # Create part.done event second with its own sequence number - self._sequence_number += 1 - part_done_event = self.create_reasoning_summary_part_done_event( - reasoning_item_id=reasoning_item_id, - reasoning_content=reasoning_content, - sequence_number=self._sequence_number, - ) - - self._sequence_number += 1 - reasoning_output_item_done_event = self.create_reasoning_output_item_done_event( - reasoning_item_id=reasoning_item_id, - reasoning_content=reasoning_content, - sequence_number=self._sequence_number, - ) - self._pending_response_events.extend( - [ - text_done_event, - part_done_event, - reasoning_output_item_done_event, - ] - ) - self._reasoning_done_emitted = True - self._reasoning_active = False + self._queue_reasoning_done_events() response_api_chunk = self._transform_chat_completion_chunk_to_response_api_chunk(chunk) if response_api_chunk: @@ -1087,6 +1120,9 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return self._pending_response_events.pop(0) except StopAsyncIteration: + if self._reasoning_active: + self._queue_reasoning_done_events() + return self._pending_response_events.pop(0) return self.common_done_event_logic(sync_mode=False) except Exception as e: @@ -1207,7 +1243,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return ReasoningSummaryTextDeltaEvent( type=ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA, item_id=self._cached_reasoning_item_id, - output_index=0, + output_index=self._reasoning_output_index, delta=reasoning_content, ) @@ -1272,13 +1308,71 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): "message", self._cached_item_id, ) + has_streamed_layout: Final = ( + bool(self._completed_reasoning_items) or 0 in self._tool_output_index_by_call_id.values() + ) + streamed_output: Final[tuple[object, ...]] = tuple( + item + for item in message_aligned + if not has_streamed_layout + or self.sent_message_item_added_event + or not isinstance(item, GenericResponseOutputItem) + or item.type != "message" + or item.phase is not None + or any(part.text or part.annotations for part in item.content) + ) + template: Final = next( + ( + item + for item in streamed_output + if isinstance(item, GenericResponseOutputItem) and item.type == "reasoning" + ), + None, + ) + if self._completed_reasoning_items and template is not None: + indexed: Final = ( + *( + ( + index, + template.model_copy( + update={ + "id": item.id, + "encrypted_content": item.encrypted_content, + "content": [ + OutputText(type="output_text", text=part.text, annotations=[]) + for part in item.summary + if part.text + ], + } + ), + ) + for index, item in self._completed_reasoning_items + ), + *( + (self._streamed_output_index(item), item) + for item in streamed_output + if not (isinstance(item, GenericResponseOutputItem) and item.type == "reasoning") + ), + ) + return tuple(item for _, item in sorted(indexed, key=lambda entry: entry[0])) reasoning_aligned: Final = _output_items_with_id( - message_aligned, + streamed_output, "reasoning", self._cached_reasoning_item_id, ) + if 0 in self._tool_output_index_by_call_id.values(): + return tuple(sorted(reasoning_aligned, key=self._streamed_output_index)) return reasoning_aligned + def _streamed_output_index(self, item: object) -> int: + if isinstance(item, GenericResponseOutputItem) and item.type == "message": + return self._message_output_index + if isinstance(item, (ResponseFunctionToolCall, OutputFunctionToolCall, CustomToolCallOutputItem)): + return self._tool_output_index_by_call_id.get(item.call_id or "", self._next_tool_output_index) + if isinstance(item, ResponseFunctionWebSearch): + return self._tool_output_index_by_call_id.get(item.id.removeprefix("ws_"), self._next_tool_output_index) + return self._next_tool_output_index + def _emit_response_completed_event(self, litellm_model_response: ModelResponse) -> ResponseCompletedEvent | None: if litellm_model_response: # Transform the response diff --git a/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py index 7b9de4644b4..1b6b4e55f0f 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -3895,34 +3895,18 @@ class TestEnsureOutputItemContentPartAdded: def _make_iterator(self): """Create a minimal LiteLLMCompletionStreamingIterator for testing.""" + from unittest.mock import MagicMock + from litellm.responses.litellm_completion_transformation.streaming_iterator import ( LiteLLMCompletionStreamingIterator, ) - iterator = LiteLLMCompletionStreamingIterator.__new__( - LiteLLMCompletionStreamingIterator + return LiteLLMCompletionStreamingIterator( + model="test-model", + litellm_custom_stream_wrapper=MagicMock(), + request_input="test", + responses_api_request={}, ) - iterator.sent_output_item_added_event = False - iterator.sent_content_part_added_event = False - iterator._sequence_number = 0 - iterator._cached_item_id = None - iterator._cached_reasoning_item_id = None - iterator._reasoning_active = False - iterator._pending_response_events = [] - iterator._pending_tool_events = [] - iterator._tool_output_index_by_call_id = {} - iterator._tool_args_by_call_id = {} - iterator._tool_item_id_by_call_id = {} - iterator._tool_call_id_by_index = {} - iterator._ambiguous_tool_call_indexes = set() - iterator._next_tool_output_index = 1 - iterator._final_tool_events_queued = False - iterator._custom_tool_names = set() - iterator.responses_api_request = {} - iterator._namespace_tool_names = LiteLLMCompletionResponsesConfig.namespace_tool_name_map(None) - iterator._web_search_calls = {} - iterator._queued_web_search_call_ids = set() - return iterator def _make_text_chunk(self): """Create a mock ModelResponseStream with a text delta.""" diff --git a/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py index 3e800a19b5e..953f8940056 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py +++ b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py @@ -11,10 +11,12 @@ spend tracking stores, so a follow-up previous_response_id still finds the conve """ import json +from itertools import chain from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest +from openai.types.responses import ResponseReasoningItem from litellm.responses.litellm_completion_transformation.streaming_iterator import ( LiteLLMCompletionStreamingIterator, @@ -1275,3 +1277,253 @@ def test_reasoning_done_without_a_response_snapshot_preserves_summary() -> None: "type": "reasoning", "summary": [{"type": "summary_text", "text": "The response snapshot is not available yet."}], } + + +def _reasoning_block_chunk( + block: ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock, +) -> ModelResponseStream: + return ModelResponseStream( + id=CHAT_COMPLETION_ID, + model="test-model", + choices=[ + StreamingChoices( + index=0, + delta=Delta( + reasoning_content=block.get("thinking", ""), + thinking_blocks=[block], + ), + ) + ], + ) + + +@pytest.mark.parametrize("separator", ["text", "tool", "web_search", "web_search_error"]) +@pytest.mark.parametrize("reasoning_kind", ["visible", "signature-only", "redacted"]) +@pytest.mark.parametrize("finish_with_text", [True, False], ids=["text-ending", "upstream-exhausted"]) +@pytest.mark.asyncio +async def test_resumed_reasoning_items_keep_complete_replay_payloads( + separator: str, reasoning_kind: str, finish_with_text: bool +) -> None: + call_id: Final = "srvtoolu_resume" + blocks: Final = tuple( + ChatCompletionRedactedThinkingBlock(type="redacted_thinking", data=f"{phase}-redacted") + if reasoning_kind == "redacted" + else ChatCompletionThinkingBlock( + type="thinking", thinking=phase if reasoning_kind == "visible" else "", signature=f"{phase}-signature" + ) + for phase in ("first", "second") + ) + middle: Final = ( + (_chunk("interim"),) + if separator == "text" + else (_tool_call_chunk(),) + if separator == "tool" + else tuple( + ModelResponseStream( + id=CHAT_COMPLETION_ID, + model="test-model", + choices=[ + StreamingChoices( + index=0, + delta=Delta( + tool_calls=[ + { + "id": call_id, + "index": 0, + "type": "function", + "function": {"name": "web_search", "arguments": '{"query":"example"}'}, + } + ] + if status == "in_progress" + else None, + provider_specific_fields={ + "web_search_calls": [ + build_web_search_call(call_id, {"query": "example"}, {"content": []}, status=status) + ] + }, + ), + ) + ], + ) + for status in ("in_progress", "failed" if separator == "web_search_error" else "completed") + ) + ) + iterator: Final = _build_iterator( + [ + _reasoning_block_chunk(blocks[0]), + *middle, + _reasoning_block_chunk(blocks[1]), + *([_chunk("answer", finish_reason="stop")] if finish_with_text else []), + ] + ) + events: Final = [json.loads(event.model_dump_json(exclude_none=True)) async for event in iterator] + added: Final = [event for event in events if event["type"] == "response.output_item.added"] + done: Final = [ + event + for event in events + if event["type"] == "response.output_item.done" and event["item"]["type"] == "reasoning" + ] + completed: Final = events[-1]["response"] + assert len(done) == 2 + assert len({event["item"]["id"] for event in done}) == 2 + assert [event["output_index"] for event in added] == list(range(len(added))) + assert [item["id"] for item in completed["output"]] == [event["item"]["id"] for event in added] + + for event, block in zip(done, blocks): + item: Final = event["item"] + expected: Final = [block] + text: Final = block.get("thinking", "") + assert json.loads(item["encrypted_content"]) == expected + assert item["summary"] == [{"type": "summary_text", "text": text}] + assert completed["output"][event["output_index"]]["encrypted_content"] == item["encrypted_content"] + assert completed["output"][event["output_index"]]["content"] == ( + [{"type": "output_text", "text": text, "annotations": []}] if text else [] + ) + parsed: Final = ResponseReasoningItem.model_validate_json(json.dumps(item)) + assert parsed.encrypted_content == item["encrypted_content"] + replay: Final = LiteLLMCompletionResponsesConfig._transform_responses_api_input_item_to_chat_completion_message( + input_item=item, replay_reasoning=True + ) + assert replay[0]["thinking_blocks"] == expected + + announced: Final = {event["item"]["id"]: event["output_index"] for event in added} + for event in events: + if event["type"].startswith("response.reasoning_summary"): + assert event["output_index"] == announced[event["item_id"]] + + +@pytest.mark.parametrize("combine_result_and_reasoning", [True, False]) +@pytest.mark.parametrize("visible", [True, False], ids=["visible", "signature-only"]) +@pytest.mark.asyncio +async def test_reasoning_after_initial_server_tool_keeps_item_indexes( + combine_result_and_reasoning: bool, visible: bool +) -> None: + call_id: Final = "srvtoolu_first" + block: Final = ChatCompletionThinkingBlock( + type="thinking", thinking="after search" if visible else "", signature="after-search-signature" + ) + start: Final = Delta( + tool_calls=[ + { + "id": call_id, + "index": 0, + "type": "function", + "function": {"name": "web_search", "arguments": '{"query":"example"}'}, + } + ], + provider_specific_fields={ + "web_search_calls": [ + build_web_search_call(call_id, {"query": "example"}, {"content": []}, status="in_progress") + ] + }, + ) + result_fields: Final = {"web_search_calls": [build_web_search_call(call_id, {"query": "example"}, {"content": []})]} + deltas: Final = ( + start, + *(() if combine_result_and_reasoning else (Delta(provider_specific_fields=result_fields),)), + Delta( + reasoning_content=block["thinking"], + thinking_blocks=[block], + provider_specific_fields=result_fields if combine_result_and_reasoning else None, + ), + Delta(content="answer"), + ) + iterator: Final = _build_iterator( + [ + ModelResponseStream( + id=CHAT_COMPLETION_ID, + model="test-model", + choices=[ + StreamingChoices(index=0, delta=delta, finish_reason="stop" if index == len(deltas) - 1 else None) + ], + ) + for index, delta in enumerate(deltas) + ] + ) + events: Final = [json.loads(event.model_dump_json(exclude_none=True)) async for event in iterator] + added: Final = [event for event in events if event["type"] == "response.output_item.added"] + assert [event["output_index"] for event in added] == [0, 1, 2] + assert [event["item"]["type"] for event in added] == ["web_search_call", "reasoning", "message"] + assert next( + index for index, event in enumerate(events) if event["type"] == "response.web_search_call.completed" + ) < next( + index + for index, event in enumerate(events) + if event["type"] == "response.output_item.added" and event["item"]["type"] == "reasoning" + ) + completed: Final = events[-1]["response"]["output"] + assert [item["id"] for item in completed] == [event["item"]["id"] for event in added] + done: Final = next( + event + for event in events + if event["type"] == "response.output_item.done" and event["item"]["type"] == "reasoning" + ) + assert done["output_index"] == 1 + assert json.loads(done["item"]["encrypted_content"]) == [block] + assert done["item"]["encrypted_content"] == completed[1]["encrypted_content"] + + +@pytest.mark.asyncio +async def test_resumed_thinking_blocks_without_reasoning_content_preserve_text() -> None: + expected: Final = [ + {"type": "thinking", "thinking": "first", "signature": "first-signature"}, + {"type": "thinking", "thinking": "second", "signature": "second-signature"}, + ] + groups: Final = tuple( + ( + ModelResponseStream( + id=CHAT_COMPLETION_ID, + model="test-model", + choices=[ + StreamingChoices( + index=0, delta=Delta(thinking_blocks=[{"type": "thinking", "thinking": block["thinking"]}]) + ) + ], + ), + _signature_only_thinking_chunk(block["signature"]), + _chunk("interim"), + ) + for block in expected + ) + chunks: Final = (*chain.from_iterable(groups), _chunk("answer", finish_reason="stop")) + events: Final = [json.loads(event.model_dump_json(exclude_none=True)) async for event in _build_iterator(chunks)] + done: Final = [ + event["item"] + for event in events + if event["type"] == "response.output_item.done" and event["item"]["type"] == "reasoning" + ] + completed: Final = [item for item in events[-1]["response"]["output"] if item["type"] == "reasoning"] + assert len(done) == 2 + assert [json.loads(item["encrypted_content"]) for item in done] == [[block] for block in expected] + assert [item["encrypted_content"] for item in done] == [item["encrypted_content"] for item in completed] + + +@pytest.mark.parametrize("resume_reasoning", [False, True]) +@pytest.mark.asyncio +async def test_signed_tool_response_has_no_unannounced_empty_message(resume_reasoning: bool) -> None: + chunks: Final = ( + _reasoning_block_chunk( + ChatCompletionThinkingBlock(type="thinking", thinking="first", signature="first-signature") + ), + _tool_call_chunk(), + *( + ( + _reasoning_block_chunk( + ChatCompletionThinkingBlock(type="thinking", thinking="second", signature="second-signature") + ), + ) + if resume_reasoning + else () + ), + _chunk("", finish_reason="tool_calls"), + ) + events: Final = [json.loads(event.model_dump_json(exclude_none=True)) async for event in _build_iterator(chunks)] + added: Final = [event for event in events if event["type"] == "response.output_item.added"] + completed: Final = events[-1]["response"]["output"] + assert [item["type"] for item in completed] == ["reasoning", "function_call"] + ( + ["reasoning"] if resume_reasoning else [] + ) + assert [item["id"] for item in completed] == [event["item"]["id"] for event in added] + for event in events: + if event["type"] == "response.output_item.done": + assert completed[event["output_index"]]["id"] == event["item"]["id"]