mirror of
https://github.com/invariantlabs-ai/invariant-gateway.git
synced 2026-10-11 15:08:44 +02:00
Add blocking and logging related tests for MCP streamable HTTP route.
This commit is contained in:
1 parent
115ae5f36b
commit
ab3fb98b67
4 files changed
+177
-82
No files matched your search
@@ -1,69 +1,14 @@
|
||||
"""This is a simple example of how to use the MCP client with Streamable HTTP transport."""
|
||||
|
||||
# pylint: disable=E1101
|
||||
# pylint: disable=W0201
|
||||
# pylint: disable=C2801
|
||||
|
||||
import asyncio
|
||||
from datetime import timedelta
|
||||
|
||||
from typing import Any, Optional
|
||||
from contextlib import AsyncExitStack
|
||||
from typing import Any
|
||||
|
||||
from mcp import ClientSession
|
||||
from mcp.client.streamable_http import streamablehttp_client
|
||||
|
||||
|
||||
class MCPClient:
|
||||
"""MCP Client for interacting with a MCP Streamable HTTP server and processing queries"""
|
||||
|
||||
def __init__(self):
|
||||
# Initialize session and client objects
|
||||
self.session: Optional[ClientSession] = None
|
||||
self.exit_stack = AsyncExitStack()
|
||||
self._streams_context = None # Initialize these to None
|
||||
self._session_context = None # so they always exist
|
||||
|
||||
async def connect_to_streamable_server(
|
||||
self, server_url: str, headers: Optional[dict] = None
|
||||
):
|
||||
"""
|
||||
Connect to an MCP server running with Streamable HTTP transport
|
||||
|
||||
Args:
|
||||
server_url: URL of the MCP server
|
||||
headers: Optional headers to include in the request
|
||||
"""
|
||||
# Store the context managers so they stay alive
|
||||
self._streams_context = streamablehttp_client(
|
||||
url=server_url,
|
||||
headers=headers or {},
|
||||
timeout=timedelta(seconds=5),
|
||||
sse_read_timeout=timedelta(seconds=10),
|
||||
)
|
||||
read_stream, write_stream, session_id = await self._streams_context.__aenter__()
|
||||
self.session_id = session_id
|
||||
|
||||
self._session_context = ClientSession(read_stream, write_stream)
|
||||
self.session: ClientSession = await self._session_context.__aenter__()
|
||||
|
||||
await self.session.initialize()
|
||||
|
||||
async def cleanup(self):
|
||||
"""Properly clean up the session and streams"""
|
||||
if self._session_context:
|
||||
await self._session_context.__aexit__(None, None, None)
|
||||
if self._streams_context:
|
||||
await self._streams_context.__aexit__(None, None, None)
|
||||
|
||||
async def process_query(self, tool_name: str, tool_args: dict) -> str:
|
||||
"""Process a query using MCP server"""
|
||||
result = await self.session.call_tool(
|
||||
tool_name, tool_args, read_timeout_seconds=timedelta(seconds=10)
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
async def run(
|
||||
gateway_url: str,
|
||||
push_to_explorer: bool,
|
||||
@@ -81,21 +26,27 @@ async def run(
|
||||
tool_args: Arguments for the tool call
|
||||
|
||||
"""
|
||||
client = MCPClient()
|
||||
try:
|
||||
await client.connect_to_streamable_server(
|
||||
server_url=gateway_url, headers=headers or {}
|
||||
streams_context = streamablehttp_client(
|
||||
url=gateway_url,
|
||||
headers=headers or {},
|
||||
timeout=timedelta(seconds=5),
|
||||
sse_read_timeout=timedelta(seconds=10),
|
||||
)
|
||||
# list tools
|
||||
listed_tools = await client.session.list_tools()
|
||||
# call tool
|
||||
if tool_name == "tools/list":
|
||||
return listed_tools
|
||||
else:
|
||||
return await client.process_query(tool_name, tool_args)
|
||||
async with streams_context as (read_stream, write_stream, _):
|
||||
async with ClientSession(read_stream, write_stream) as session:
|
||||
await session.initialize()
|
||||
# list tools
|
||||
listed_tools = await session.list_tools()
|
||||
# call tool
|
||||
if tool_name == "tools/list":
|
||||
return listed_tools
|
||||
else:
|
||||
return await session.call_tool(
|
||||
tool_name, tool_args, read_timeout_seconds=timedelta(seconds=10)
|
||||
)
|
||||
finally:
|
||||
# Sleep for a while to allow the server to process the background tasks
|
||||
# like pushing traces to the explorer
|
||||
if push_to_explorer:
|
||||
await asyncio.sleep(2)
|
||||
await client.cleanup()
|
||||
Reference in new issue
Block a user