mirror of
https://github.com/invariantlabs-ai/invariant-gateway.git
synced 2026-09-30 18:29:38 +02:00
Make MCP stdio gateway fully async. With sync and async mixed behaviour for running background tasks we were running into issues.
This commit is contained in:
1 parent
a6c1124076
commit
876eb44c78
3 files changed
+198
-155
No files matched your search
@@ -17,7 +17,7 @@ MCP_SSE_SERVER_PORT = 8123
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.timeout(15)
|
||||
@pytest.mark.timeout(30)
|
||||
@pytest.mark.parametrize(
|
||||
"push_to_explorer, transport",
|
||||
[
|
||||
@@ -97,7 +97,7 @@ async def test_mcp_with_gateway(
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.timeout(15)
|
||||
@pytest.mark.timeout(30)
|
||||
@pytest.mark.parametrize("transport", ["stdio", "sse"])
|
||||
async def test_mcp_with_gateway_and_logging_guardrails(
|
||||
explorer_api_url, invariant_gateway_package_whl_file, gateway_url, transport
|
||||
@@ -205,11 +205,13 @@ async def test_mcp_with_gateway_and_logging_guardrails(
|
||||
tool_call_annotation is not None
|
||||
), "Missing 'get_last_message_from_user is called' annotation"
|
||||
assert food_annotation["extra_metadata"]["source"] == "guardrails-error"
|
||||
assert food_annotation["extra_metadata"]["guardrail"]["action"] == "log"
|
||||
assert tool_call_annotation["extra_metadata"]["source"] == "guardrails-error"
|
||||
assert tool_call_annotation["extra_metadata"]["guardrail"]["action"] == "log"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.timeout(15)
|
||||
@pytest.mark.timeout(30)
|
||||
@pytest.mark.parametrize("transport", ["stdio", "sse"])
|
||||
async def test_mcp_with_gateway_and_blocking_guardrails(
|
||||
explorer_api_url, invariant_gateway_package_whl_file, gateway_url, transport
|
||||
@@ -298,10 +300,11 @@ async def test_mcp_with_gateway_and_blocking_guardrails(
|
||||
and annotations[0]["address"] == "messages.0.tool_calls.0"
|
||||
)
|
||||
assert annotations[0]["extra_metadata"]["source"] == "guardrails-error"
|
||||
assert annotations[0]["extra_metadata"]["guardrail"]["action"] == "block"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.timeout(15)
|
||||
@pytest.mark.timeout(30)
|
||||
@pytest.mark.parametrize("transport", ["stdio", "sse"])
|
||||
async def test_mcp_sse_with_gateway_hybrid_guardrails(
|
||||
explorer_api_url, invariant_gateway_package_whl_file, gateway_url, transport
|
||||
@@ -416,4 +419,6 @@ async def test_mcp_sse_with_gateway_hybrid_guardrails(
|
||||
tool_call_annotation is not None
|
||||
), "Missing 'get_last_message_from_user is called' annotation"
|
||||
assert food_annotation["extra_metadata"]["source"] == "guardrails-error"
|
||||
assert food_annotation["extra_metadata"]["guardrail"]["action"] == "block"
|
||||
assert tool_call_annotation["extra_metadata"]["source"] == "guardrails-error"
|
||||
assert tool_call_annotation["extra_metadata"]["guardrail"]["action"] == "log"
|
||||
@@ -1,7 +1,7 @@
|
||||
"""A MCP client implementation that interacts with MCP server to make tool calls."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
|
||||
from datetime import timedelta
|
||||
from contextlib import AsyncExitStack
|
||||
from typing import Any, Optional
|
||||
@@ -69,7 +69,7 @@ class MCPClient:
|
||||
self.stdio, self.write = stdio_transport
|
||||
self.session = await self.exit_stack.enter_async_context(
|
||||
ClientSession(
|
||||
self.stdio, self.write, read_timeout_seconds=timedelta(seconds=10)
|
||||
self.stdio, self.write, read_timeout_seconds=timedelta(seconds=15)
|
||||
)
|
||||
)
|
||||
|
||||
@@ -85,10 +85,6 @@ class MCPClient:
|
||||
tool_name: Name of the tool to call
|
||||
tool_args: Arguments for the tool call
|
||||
"""
|
||||
response = await self.session.list_tools()
|
||||
if tool_name not in [tool.name for tool in response.tools]:
|
||||
raise ValueError(f"Tool '{tool_name}' not found in available tools")
|
||||
|
||||
# Execute tool call
|
||||
result = await self.session.call_tool(tool_name, tool_args)
|
||||
return result
|
||||
@@ -130,8 +126,4 @@ async def run(
|
||||
)
|
||||
return await client.call_tool(tool_name, tool_args)
|
||||
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