mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(responses): preserve resumed reasoning phases
This commit is contained in:
parent
d842d266c7
commit
6f5a29fca3
3 changed files with 439 additions and 109 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue