Skip to content
64 changes: 64 additions & 0 deletions src/mistralai/extra/mcp/streamable_http.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
import logging
from contextlib import AsyncExitStack
from typing import Any

import httpx
from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream
from mcp.client.streamable_http import streamable_http_client # pyright: ignore[reportMissingImports]
from mcp.shared.message import SessionMessage # pyright: ignore[reportMissingImports]

from mistralai.extra.mcp.base import (
MCPClientBase,
)

from mistralai.client.types import BaseModel

logger = logging.getLogger(__name__)


class StreamableHTTPServerParams(BaseModel):
"""Parameters required for a MCPClient with Streamable HTTP transport"""

url: str
headers: dict[str, Any] | None = None
timeout: float = 30
sse_read_timeout: float = 60 * 5
Comment thread
adrienduchemin marked this conversation as resolved.
Outdated


class MCPClientStreamableHTTP(MCPClientBase):
"""MCP client that uses the Streamable HTTP transport for communication.

Credentials (for example a bearer token, or a per-request integration token)
are provided as ``headers`` and set as the default headers of the underlying
``httpx.AsyncClient``, so they are sent on every request including the
initialize call. Recent ``mcp`` releases deprecate and ignore the transport's
own ``headers`` argument, so configuring them on the client is required.
"""

_params: StreamableHTTPServerParams

def __init__(
self,
params: StreamableHTTPServerParams,
name: str | None = None,
):
super().__init__(name=name)
self._params = params

async def _get_transport(
self, exit_stack: AsyncExitStack
) -> tuple[
MemoryObjectReceiveStream[SessionMessage | Exception],
MemoryObjectSendStream[SessionMessage],
]:
http_client = await exit_stack.enter_async_context(
httpx.AsyncClient(
headers=self._params.headers,
timeout=self._params.timeout,
follow_redirects=True,
)
)
read_stream, write_stream, _ = await exit_stack.enter_async_context(
streamable_http_client(url=self._params.url, http_client=http_client)
)
return read_stream, write_stream
Loading