mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
apply change to sync path
This commit is contained in:
parent
e33a2b486a
commit
c9ed2eb747
2 changed files with 123 additions and 12 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue