apply change to sync path

This commit is contained in:
shivam 2026-02-27 05:04:35 -08:00
parent e33a2b486a
commit c9ed2eb747
2 changed files with 123 additions and 12 deletions

View file

@ -284,6 +284,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
mcp_events # Store the initial MCP events for backward compatibility
)
self.tool_server_map = tool_server_map
self._sync_response_created_emitted = False
# Iterator references
self.base_iterator: Optional[
@ -725,19 +726,39 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
return self
def __next__(self) -> ResponsesAPIStreamingResponse:
# First, emit any queued MCP events
if self.is_async:
raise RuntimeError("Cannot use sync iteration on async iterator")
# Emit response.created first (OpenAI SDK expects it before other events)
if not self._sync_response_created_emitted:
self._ensure_sync_base_iterator()
if self.base_iterator and hasattr(self.base_iterator, "__next__"):
try:
first_chunk = next(cast(Any, self.base_iterator)) # type: ignore[arg-type]
self._sync_response_created_emitted = True
return first_chunk
except StopIteration:
self.finished = True
raise
self._sync_response_created_emitted = True
# Then emit MCP discovery events
if self.mcp_events: # type: ignore[attr-defined]
return self.mcp_events.pop(0) # type: ignore[attr-defined]
# Then delegate to the base iterator
if not self.is_async:
try:
if self.base_iterator and hasattr(self.base_iterator, "__next__"):
return next(cast(Any, self.base_iterator)) # type: ignore[arg-type]
else:
raise StopIteration
except StopIteration:
self.finished = True
raise
else:
raise RuntimeError("Cannot use sync iteration on async iterator")
try:
if self.base_iterator and hasattr(self.base_iterator, "__next__"):
return next(cast(Any, self.base_iterator)) # type: ignore[arg-type]
raise StopIteration
except StopIteration:
self.finished = True
raise
def _ensure_sync_base_iterator(self) -> None:
"""Create base iterator synchronously when needed (for sync __next__ path)."""
if self.base_iterator is not None:
return
from litellm.litellm_core_utils.asyncify import run_async_function
run_async_function(self._create_initial_response_iterator)

View file

@ -562,6 +562,96 @@ async def test_sse_event_ordering_response_created_first():
)
def test_sse_event_ordering_sync_response_created_first():
"""
Test that sync __next__ emits response.created before mcp_list_tools events.
"""
from unittest.mock import patch
from litellm.types.llms.openai import (
ResponsesAPIStreamEvents,
ResponseCreatedEvent,
ResponsesAPIResponse,
)
from litellm.responses.mcp.mcp_streaming_iterator import (
MCPEnhancedStreamingIterator,
create_mcp_list_tools_events,
)
mock_mcp_tools = [
type("MCPTool", (), {
"name": "test_tool",
"description": "Test",
"inputSchema": {"type": "object", "properties": {}},
})(),
]
import asyncio
mcp_events = asyncio.run(
create_mcp_list_tools_events(
mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/test"}],
user_api_key_auth=None,
base_item_id="mcp_test123",
pre_processed_mcp_tools=mock_mcp_tools,
)
)
def mock_stream():
yield ResponseCreatedEvent(
type=ResponsesAPIStreamEvents.RESPONSE_CREATED,
response=ResponsesAPIResponse(
id="resp_123",
object="response",
created_at=1234567890,
status="in_progress",
error=None,
incomplete_details=None,
instructions=None,
max_output_tokens=None,
model="gpt-4o-mini",
output=[],
parallel_tool_calls=True,
previous_response_id=None,
reasoning=None,
store=True,
temperature=1.0,
text={"format": {"type": "text"}},
tool_choice="auto",
tools=[],
top_p=1.0,
truncation="disabled",
usage=None,
user=None,
metadata={},
),
)
iterator = MCPEnhancedStreamingIterator(
base_iterator=mock_stream(),
mcp_events=mcp_events,
tool_server_map={"test": "test_server"},
mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy"}],
user_api_key_auth=None,
original_request_params={
"model": "gpt-4o-mini",
"stream": True,
"tools": [{"type": "mcp", "server_url": "litellm_proxy"}],
},
)
iterator.is_async = False
event_types = []
for chunk in iterator:
event_types.append(getattr(chunk, "type", "unknown"))
if len(event_types) >= 5:
break
assert len(event_types) >= 1, "Should have at least one event"
assert event_types[0] == ResponsesAPIStreamEvents.RESPONSE_CREATED, (
f"Sync path: first event must be response.created, got {event_types[0]}. "
f"Order: {event_types}"
)
@pytest.mark.asyncio
async def test_streaming_mcp_events_validation():
"""