diff --git a/src/fastmcp/utilities/mcp_config.py b/src/fastmcp/utilities/mcp_config.py index 5a2990980..6b4b80b4e 100644 --- a/src/fastmcp/utilities/mcp_config.py +++ b/src/fastmcp/utilities/mcp_config.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, Literal +from typing import TYPE_CHECKING, Annotated, Any, Literal from urllib.parse import urlparse from pydantic import AnyUrl, Field @@ -56,6 +56,12 @@ class RemoteMCPServer(FastMCPBaseModel): url: str headers: dict[str, str] = Field(default_factory=dict) transport: Literal["streamable-http", "sse", "http"] | None = None + auth: Annotated[ + str | Literal["oauth"] | None, + Field( + description='Either a string representing a Bearer token or the literal "oauth" to use OAuth authentication.' + ), + ] = None def to_transport(self) -> StreamableHttpTransport | SSETransport: from fastmcp.client.transports import SSETransport, StreamableHttpTransport @@ -66,9 +72,11 @@ class RemoteMCPServer(FastMCPBaseModel): transport = self.transport if transport == "sse": - return SSETransport(self.url, headers=self.headers) + return SSETransport(self.url, headers=self.headers, auth=self.auth) else: - return StreamableHttpTransport(self.url, headers=self.headers) + return StreamableHttpTransport( + self.url, headers=self.headers, auth=self.auth + ) class MCPConfig(FastMCPBaseModel): diff --git a/tests/utilities/test_mcp_config.py b/tests/utilities/test_mcp_config.py index ec334153d..f5dee613d 100644 --- a/tests/utilities/test_mcp_config.py +++ b/tests/utilities/test_mcp_config.py @@ -1,6 +1,8 @@ import inspect from pathlib import Path +from fastmcp.client.auth.bearer import BearerAuth +from fastmcp.client.auth.oauth import OAuthClientProvider from fastmcp.client.client import Client from fastmcp.client.transports import ( SSETransport, @@ -136,3 +138,60 @@ async def test_multi_client(tmp_path: Path): result_2 = await client.call_tool("test_2_add", {"a": 1, "b": 2}) assert result_1[0].text == "3" # type: ignore[attr-dict] assert result_2[0].text == "3" # type: ignore[attr-dict] + + +async def test_remote_config_default_no_auth(): + config = { + "mcpServers": { + "test_server": { + "url": "http://localhost:8000", + } + } + } + client = Client(config) + assert isinstance(client.transport.transport, StreamableHttpTransport) + assert client.transport.transport.auth is None + + +async def test_remote_config_with_auth_token(): + config = { + "mcpServers": { + "test_server": { + "url": "http://localhost:8000", + "auth": "test_token", + } + } + } + client = Client(config) + assert isinstance(client.transport.transport, StreamableHttpTransport) + assert isinstance(client.transport.transport.auth, BearerAuth) + assert client.transport.transport.auth.token.get_secret_value() == "test_token" + + +async def test_remote_config_sse_with_auth_token(): + config = { + "mcpServers": { + "test_server": { + "url": "http://localhost:8000/sse", + "auth": "test_token", + } + } + } + client = Client(config) + assert isinstance(client.transport.transport, SSETransport) + assert isinstance(client.transport.transport.auth, BearerAuth) + assert client.transport.transport.auth.token.get_secret_value() == "test_token" + + +async def test_remote_config_with_oauth_literal(): + config = { + "mcpServers": { + "test_server": { + "url": "http://localhost:8000", + "auth": "oauth", + } + } + } + client = Client(config) + assert isinstance(client.transport.transport, StreamableHttpTransport) + assert isinstance(client.transport.transport.auth, OAuthClientProvider)