mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
fixed mcp streaming ordering
This commit is contained in:
parent
88ccffccc8
commit
e33a2b486a
2 changed files with 210 additions and 90 deletions
|
|
@ -269,7 +269,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
self.should_auto_execute = self._should_auto_execute_tools()
|
||||
|
||||
# Streaming state management
|
||||
self.phase = "mcp_discovery" # mcp_discovery -> initial_response -> tool_execution -> follow_up_response -> finished
|
||||
# get_response_created: emit response.created first (OpenAI SDK expects it before other events)
|
||||
# mcp_discovery -> initial_response -> tool_execution -> follow_up_response -> finished
|
||||
self.phase = "get_response_created"
|
||||
self.finished = False
|
||||
|
||||
# Event queues and generation flags
|
||||
|
|
@ -388,101 +390,123 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
async def __anext__(self) -> ResponsesAPIStreamingResponse:
|
||||
"""
|
||||
Phase-based streaming:
|
||||
1. mcp_discovery - Emit MCP discovery events
|
||||
2. initial_response - Stream the first LLM response
|
||||
3. tool_execution - Emit tool execution events
|
||||
4. follow_up_response - Stream the follow-up response
|
||||
5. finished - End iteration
|
||||
1. get_response_created - Emit response.created first (OpenAI SDK expects it before other events)
|
||||
2. mcp_discovery - Emit MCP discovery events
|
||||
3. initial_response - Stream the first LLM response
|
||||
4. tool_execution - Emit tool execution events
|
||||
5. follow_up_response - Stream the follow-up response
|
||||
6. finished - End iteration
|
||||
"""
|
||||
|
||||
# Phase 1: MCP Discovery Events
|
||||
if self.phase == "mcp_discovery":
|
||||
# Generate MCP discovery events if not already done
|
||||
# MCP discovery events are already generated and available
|
||||
|
||||
# Emit MCP discovery events
|
||||
if self.mcp_discovery_events:
|
||||
return self.mcp_discovery_events.pop(0)
|
||||
|
||||
# All MCP discovery events emitted, move to next phase
|
||||
verbose_logger.debug(
|
||||
"MCP discovery phase complete, transitioning to initial_response"
|
||||
)
|
||||
self.phase = "initial_response"
|
||||
await self._create_initial_response_iterator()
|
||||
# Fall through to process the initial response immediately
|
||||
|
||||
# Phase 2: Initial Response Stream
|
||||
if self.phase == "initial_response":
|
||||
if self.base_iterator:
|
||||
# Check if base_iterator is actually iterable
|
||||
if hasattr(self.base_iterator, "__anext__"):
|
||||
try:
|
||||
chunk = await cast(Any, self.base_iterator).__anext__() # type: ignore[attr-defined]
|
||||
|
||||
# If auto-execution is enabled, check for completed responses
|
||||
if self.should_auto_execute and self._is_response_completed(
|
||||
chunk
|
||||
):
|
||||
# Collect the response for tool execution
|
||||
response_obj = getattr(chunk, "response", None)
|
||||
if isinstance(response_obj, ResponsesAPIResponse):
|
||||
self.collected_response = response_obj
|
||||
# Move to tool execution phase after emitting this chunk
|
||||
self.phase = "tool_execution"
|
||||
await self._generate_tool_execution_events()
|
||||
|
||||
return chunk
|
||||
except StopAsyncIteration:
|
||||
# Initial response ended, move to next phase
|
||||
if self.should_auto_execute and self.collected_response:
|
||||
self.phase = "tool_execution"
|
||||
await self._generate_tool_execution_events()
|
||||
else:
|
||||
self.phase = "finished"
|
||||
raise
|
||||
else:
|
||||
# base_iterator is not async iterable (likely a ResponsesAPIResponse)
|
||||
# Collect it for tool execution if needed
|
||||
if self.should_auto_execute and isinstance(
|
||||
self.base_iterator, ResponsesAPIResponse
|
||||
):
|
||||
self.collected_response = self.base_iterator
|
||||
self.phase = "tool_execution"
|
||||
await self._generate_tool_execution_events()
|
||||
else:
|
||||
self.phase = "finished"
|
||||
raise StopAsyncIteration
|
||||
|
||||
# Phase 3: Tool Execution Events
|
||||
if self.phase == "tool_execution":
|
||||
# Emit any queued tool execution events
|
||||
if self.tool_execution_events:
|
||||
return self.tool_execution_events.pop(0)
|
||||
|
||||
# Move to follow-up response phase
|
||||
self.phase = "follow_up_response"
|
||||
await self._create_follow_up_iterator()
|
||||
|
||||
# Phase 4: Follow-up Response Stream
|
||||
if self.phase == "follow_up_response":
|
||||
if self.follow_up_iterator:
|
||||
try:
|
||||
return await cast(Any, self.follow_up_iterator).__anext__() # type: ignore[attr-defined]
|
||||
except StopAsyncIteration:
|
||||
self.phase = "finished"
|
||||
raise
|
||||
while True:
|
||||
if self.phase == "get_response_created":
|
||||
result = await self._handle_get_response_created_phase()
|
||||
elif self.phase == "mcp_discovery":
|
||||
result = await self._handle_mcp_discovery_phase()
|
||||
elif self.phase == "initial_response":
|
||||
result = await self._handle_initial_response_phase()
|
||||
elif self.phase == "tool_execution":
|
||||
result = await self._handle_tool_execution_phase()
|
||||
elif self.phase == "follow_up_response":
|
||||
result = await self._handle_follow_up_response_phase()
|
||||
else:
|
||||
self.phase = "finished"
|
||||
raise StopAsyncIteration
|
||||
|
||||
# Phase 5: Finished
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
async def _handle_get_response_created_phase(
|
||||
self,
|
||||
) -> Optional[ResponsesAPIStreamingResponse]:
|
||||
"""Emit response.created first (OpenAI SDK expects it before other events)."""
|
||||
await self._create_initial_response_iterator()
|
||||
if self.phase == "finished":
|
||||
raise StopAsyncIteration
|
||||
if self.base_iterator and hasattr(self.base_iterator, "__anext__"):
|
||||
try:
|
||||
first_chunk = await cast(Any, self.base_iterator).__anext__() # type: ignore[attr-defined]
|
||||
self.phase = "mcp_discovery"
|
||||
return first_chunk
|
||||
except StopAsyncIteration:
|
||||
self.phase = "finished"
|
||||
raise
|
||||
self.phase = "mcp_discovery"
|
||||
return None
|
||||
|
||||
# Should not reach here
|
||||
async def _handle_mcp_discovery_phase(
|
||||
self,
|
||||
) -> Optional[ResponsesAPIStreamingResponse]:
|
||||
"""Emit MCP discovery events (after response.created)."""
|
||||
if self.mcp_discovery_events:
|
||||
return self.mcp_discovery_events.pop(0)
|
||||
verbose_logger.debug(
|
||||
"MCP discovery phase complete, transitioning to initial_response"
|
||||
)
|
||||
self.phase = "initial_response"
|
||||
return None
|
||||
|
||||
async def _handle_initial_response_phase(
|
||||
self,
|
||||
) -> Optional[ResponsesAPIStreamingResponse]:
|
||||
"""Stream the first LLM response."""
|
||||
if not self.base_iterator:
|
||||
self.phase = "finished"
|
||||
raise StopAsyncIteration
|
||||
if not hasattr(self.base_iterator, "__anext__"):
|
||||
return await self._handle_initial_response_non_iterable()
|
||||
try:
|
||||
chunk = await cast(Any, self.base_iterator).__anext__() # type: ignore[attr-defined]
|
||||
if self.should_auto_execute and self._is_response_completed(chunk):
|
||||
response_obj = getattr(chunk, "response", None)
|
||||
if isinstance(response_obj, ResponsesAPIResponse):
|
||||
self.collected_response = response_obj
|
||||
self.phase = "tool_execution"
|
||||
await self._generate_tool_execution_events()
|
||||
return chunk
|
||||
except StopAsyncIteration:
|
||||
if self.should_auto_execute and self.collected_response:
|
||||
self.phase = "tool_execution"
|
||||
await self._generate_tool_execution_events()
|
||||
return None
|
||||
self.phase = "finished"
|
||||
raise
|
||||
|
||||
async def _handle_initial_response_non_iterable(
|
||||
self,
|
||||
) -> Optional[ResponsesAPIStreamingResponse]:
|
||||
"""Handle base_iterator that is not async iterable (e.g. ResponsesAPIResponse)."""
|
||||
if self.should_auto_execute and isinstance(
|
||||
self.base_iterator, ResponsesAPIResponse
|
||||
):
|
||||
self.collected_response = self.base_iterator
|
||||
self.phase = "tool_execution"
|
||||
await self._generate_tool_execution_events()
|
||||
return None
|
||||
self.phase = "finished"
|
||||
raise StopAsyncIteration
|
||||
|
||||
async def _handle_tool_execution_phase(
|
||||
self,
|
||||
) -> Optional[ResponsesAPIStreamingResponse]:
|
||||
"""Emit tool execution events."""
|
||||
if self.tool_execution_events:
|
||||
return self.tool_execution_events.pop(0)
|
||||
self.phase = "follow_up_response"
|
||||
await self._create_follow_up_iterator()
|
||||
return None
|
||||
|
||||
async def _handle_follow_up_response_phase(
|
||||
self,
|
||||
) -> Optional[ResponsesAPIStreamingResponse]:
|
||||
"""Stream the follow-up response."""
|
||||
if not self.follow_up_iterator:
|
||||
self.phase = "finished"
|
||||
raise StopAsyncIteration
|
||||
try:
|
||||
return await cast(Any, self.follow_up_iterator).__anext__() # type: ignore[attr-defined]
|
||||
except StopAsyncIteration:
|
||||
self.phase = "finished"
|
||||
raise
|
||||
|
||||
def _is_response_completed(self, chunk: ResponsesAPIStreamingResponse) -> bool:
|
||||
"""Check if this chunk indicates the response is completed"""
|
||||
from litellm.types.llms.openai import ResponsesAPIStreamEvents
|
||||
|
|
|
|||
|
|
@ -467,15 +467,111 @@ async def test_mcp_allowed_tools_filtering():
|
|||
|
||||
print("✓ MCP allowed_tools filtering test completed successfully!")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_event_ordering_response_created_first():
|
||||
"""
|
||||
Test that response.created is emitted before mcp_list_tools events.
|
||||
OpenAI Node SDK expects response.created as the first SSE event.
|
||||
"""
|
||||
from unittest.mock import AsyncMock, 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": {}},
|
||||
})(),
|
||||
]
|
||||
|
||||
mcp_discovery_events = await 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,
|
||||
)
|
||||
|
||||
async 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={},
|
||||
),
|
||||
)
|
||||
|
||||
mock_stream_obj = mock_stream()
|
||||
|
||||
with patch(
|
||||
"litellm.responses.main.aresponses",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_stream_obj,
|
||||
):
|
||||
iterator = MCPEnhancedStreamingIterator(
|
||||
base_iterator=None,
|
||||
mcp_events=mcp_discovery_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"}],
|
||||
},
|
||||
)
|
||||
|
||||
event_types = []
|
||||
async 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"First event must be response.created for OpenAI SDK compatibility, got {event_types[0]}. "
|
||||
f"Order: {event_types}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_mcp_events_validation():
|
||||
"""
|
||||
Test that MCP streaming events are properly emitted when using streaming with MCP tools.
|
||||
|
||||
This test validates:
|
||||
1. MCP discovery events are emitted first
|
||||
2. Regular streaming response events follow
|
||||
3. Tool execution events are emitted when tools are auto-executed
|
||||
1. response.created is emitted first (OpenAI SDK requirement)
|
||||
2. MCP discovery events follow
|
||||
3. Regular streaming response events follow
|
||||
4. Tool execution events are emitted when tools are auto-executed
|
||||
"""
|
||||
from unittest.mock import AsyncMock, patch
|
||||
from litellm.types.llms.openai import ResponsesAPIStreamEvents
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue