mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(mcp): finish streamed responses before releasing session handlers
This commit is contained in:
parent
e62acdef76
commit
f72d75f393
2 changed files with 71 additions and 0 deletions
|
|
@ -19831,6 +19831,7 @@ async def _stream_mcp_asgi_response(handle_fn, scope: dict, receive) -> "Streami
|
|||
while True:
|
||||
chunk = await body_queue.get()
|
||||
if chunk is None:
|
||||
await handler_task
|
||||
break
|
||||
yield chunk
|
||||
finally:
|
||||
|
|
|
|||
|
|
@ -1,7 +1,16 @@
|
|||
import asyncio
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from mcp.server import Server
|
||||
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
|
||||
from starlette.applications import Starlette
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
from starlette.routing import Route
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
|
||||
from litellm.proxy.proxy_server import _stream_mcp_asgi_response
|
||||
|
||||
|
|
@ -34,3 +43,64 @@ async def test_stream_mcp_asgi_response_propagates_pre_header_http_exception():
|
|||
assert exc_info.value.headers == {
|
||||
"WWW-Authenticate": "Bearer authorization_uri=https://example.test/auth"
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streamed_initialize_preserves_session_until_delete() -> None:
|
||||
manager: Final = StreamableHTTPSessionManager(app=Server("session-test"), stateless=False)
|
||||
|
||||
async def endpoint(request: Request) -> Response:
|
||||
return await _stream_mcp_asgi_response(manager.handle_request, request.scope, request.receive)
|
||||
|
||||
app: Final = Starlette(routes=[Route("/mcp", endpoint, methods=["POST", "DELETE"])])
|
||||
async with manager.run(), httpx.AsyncClient(
|
||||
transport=httpx.ASGITransport(app=app),
|
||||
base_url="http://localhost",
|
||||
headers={"Accept": "application/json, text/event-stream"},
|
||||
) as client:
|
||||
async with asyncio.timeout(5):
|
||||
initialized: Final = await client.post(
|
||||
"/mcp",
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"method": "initialize",
|
||||
"params": {
|
||||
"protocolVersion": "2025-11-25",
|
||||
"capabilities": {},
|
||||
"clientInfo": {"name": "session-test", "version": "1"},
|
||||
},
|
||||
},
|
||||
)
|
||||
assert initialized.status_code == 200
|
||||
headers: Final = {"mcp-session-id": initialized.headers["mcp-session-id"]}
|
||||
notification: Final = await client.post(
|
||||
"/mcp", headers=headers, json={"jsonrpc": "2.0", "method": "notifications/initialized"}
|
||||
)
|
||||
assert notification.status_code == 202
|
||||
ping: Final = {"jsonrpc": "2.0", "id": 2, "method": "ping"}
|
||||
assert (await client.post("/mcp", headers=headers, json=ping)).status_code == 200
|
||||
assert (await client.delete("/mcp", headers=headers)).status_code == 200
|
||||
assert (await client.post("/mcp", headers=headers, json=ping)).status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disconnected_response_cancels_the_active_handler() -> None:
|
||||
stopped: Final = asyncio.Event()
|
||||
|
||||
async def handle(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
try:
|
||||
await send({"type": "http.response.start", "status": 200, "headers": []})
|
||||
await send({"type": "http.response.body", "body": b"first", "more_body": True})
|
||||
await asyncio.Event().wait()
|
||||
finally:
|
||||
stopped.set()
|
||||
|
||||
async def receive() -> Message:
|
||||
return {"type": "http.disconnect"}
|
||||
|
||||
async with asyncio.timeout(1):
|
||||
response: Final = await _stream_mcp_asgi_response(handle, {}, receive)
|
||||
assert await anext(response.body_iterator) == b"first"
|
||||
await response.body_iterator.aclose()
|
||||
assert stopped.is_set()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue