diff --git a/.gitignore b/.gitignore index bcb20ed19..74bd2ac8f 100644 --- a/.gitignore +++ b/.gitignore @@ -1,19 +1,62 @@ # Python-generated files __pycache__/ -*.py[oc] +*.py[cod] +*$py.class build/ dist/ wheels/ -*.egg-info +*.egg-info/ +*.egg +MANIFEST +.pytest_cache/ +.coverage +htmlcov/ +.tox/ +nosetests.xml +coverage.xml +*.cover # Virtual environments .venv -.DS_Store +venv/ +env/ +ENV/ .env +# System files +.DS_Store +# Version file src/fastmcp/_version.py -# editors +# Editors and IDEs .cursorrules .vscode/ +.idea/ +*.swp +*.swo +*~ +.project +.pydevproject +.settings/ + +# Jupyter Notebook +.ipynb_checkpoints + +# Type checking +.mypy_cache/ +.dmypy.json +dmypy.json +.pyre/ +.pytype/ + +# Local development +.python-version +.envrc +.direnv/ + +# Logs and databases +*.log +*.sqlite +*.db +*.ddb diff --git a/src/fastmcp/client/base.py b/src/fastmcp/client/base.py index 53316f706..ba95b6b1d 100644 --- a/src/fastmcp/client/base.py +++ b/src/fastmcp/client/base.py @@ -1,7 +1,7 @@ import abc import contextlib import datetime -from typing import Any, AsyncContextManager, Optional +from typing import Any, AsyncContextManager, Optional, TypedDict import mcp.types from mcp import ClientSession @@ -19,6 +19,23 @@ def _get_roots_callback(roots: list[mcp.types.Root]) -> ListRootsFnT | None: return _roots_callback +class ClientKwargs(TypedDict, total=False): + roots: list[mcp.types.Root] | None + sampling_callback: SamplingFnT | None + list_roots_callback: ListRootsFnT | None + logging_callback: LoggingFnT | None + message_handler: MessageHandlerFnT | None + read_timeout_seconds: datetime.timedelta | None + + +class SessionKwargs(TypedDict, total=False): + sampling_callback: SamplingFnT | None + list_roots_callback: ListRootsFnT | None + logging_callback: LoggingFnT | None + message_handler: MessageHandlerFnT | None + read_timeout_seconds: datetime.timedelta | None + + class BaseClient(abc.ABC): def __init__( self, @@ -48,6 +65,15 @@ class BaseClient(abc.ABC): self._message_handler = message_handler self._read_timeout_seconds = read_timeout_seconds + def _session_kwargs(self) -> SessionKwargs: + return SessionKwargs( + sampling_callback=self._sampling_callback, + list_roots_callback=self._list_roots_callback, + logging_callback=self._logging_callback, + message_handler=self._message_handler, + read_timeout_seconds=self._read_timeout_seconds, + ) + @property def transport(self): """Get the current transport connection""" @@ -71,13 +97,7 @@ class BaseClient(abc.ABC): return self._session is not None @abc.abstractmethod - def _connect( - self, - sampling_callback: SamplingFnT | None = None, - list_roots_callback: ListRootsFnT | None = None, - logging_callback: LoggingFnT | None = None, - message_handler: MessageHandlerFnT | None = None, - ) -> AsyncContextManager: + def _connect(self) -> AsyncContextManager: """Return an async context manager that handles connection lifecycle. This will be called by __aenter__ to establish the connection.""" raise NotImplementedError("Subclasses must implement this method") diff --git a/src/fastmcp/client/sse.py b/src/fastmcp/client/sse.py index de01fe0fb..578a32106 100644 --- a/src/fastmcp/client/sse.py +++ b/src/fastmcp/client/sse.py @@ -1,17 +1,10 @@ import contextlib -import datetime -import mcp.types from mcp import ClientSession from mcp.client.sse import sse_client +from typing_extensions import Unpack -from fastmcp.client.base import ( - BaseClient, - ListRootsFnT, - LoggingFnT, - MessageHandlerFnT, - SamplingFnT, -) +from fastmcp.client.base import BaseClient, ClientKwargs class SSEClient(BaseClient): @@ -19,21 +12,9 @@ class SSEClient(BaseClient): self, url: str, headers: dict[str, str] | None = None, - roots: list[mcp.types.Root] | None = None, - sampling_callback: SamplingFnT | None = None, - list_roots_callback: ListRootsFnT | None = None, - logging_callback: LoggingFnT | None = None, - message_handler: MessageHandlerFnT | None = None, - read_timeout_seconds: datetime.timedelta | None = None, + **kwargs: Unpack[ClientKwargs], ): - super().__init__( - roots=roots, - sampling_callback=sampling_callback, - list_roots_callback=list_roots_callback, - logging_callback=logging_callback, - message_handler=message_handler, - read_timeout_seconds=read_timeout_seconds, - ) + super().__init__(**kwargs) self.url = url self.headers = headers or {} @@ -45,11 +26,7 @@ class SSEClient(BaseClient): async with ClientSession( read_stream=read_stream, write_stream=write_stream, - sampling_callback=self._sampling_callback, - list_roots_callback=self._list_roots_callback, - logging_callback=self._logging_callback, - message_handler=self._message_handler, - read_timeout_seconds=self._read_timeout_seconds, + **self._session_kwargs(), ) as session: async with self._set_session(transport, session): yield self diff --git a/src/fastmcp/client/stdio.py b/src/fastmcp/client/stdio.py index 9c5f882e4..590cf1375 100644 --- a/src/fastmcp/client/stdio.py +++ b/src/fastmcp/client/stdio.py @@ -1,38 +1,19 @@ import contextlib -import datetime -import mcp.types from mcp import ClientSession, StdioServerParameters from mcp.client.stdio import stdio_client +from typing_extensions import Unpack -from fastmcp.client.base import ( - BaseClient, - ListRootsFnT, - LoggingFnT, - MessageHandlerFnT, - SamplingFnT, -) +from fastmcp.client.base import BaseClient, ClientKwargs class StdioClient(BaseClient): def __init__( self, server_script_path: str, - roots: list[mcp.types.Root] | None = None, - sampling_callback: SamplingFnT | None = None, - list_roots_callback: ListRootsFnT | None = None, - logging_callback: LoggingFnT | None = None, - message_handler: MessageHandlerFnT | None = None, - read_timeout_seconds: datetime.timedelta | None = None, + **kwargs: Unpack[ClientKwargs], ): - super().__init__( - roots=roots, - sampling_callback=sampling_callback, - list_roots_callback=list_roots_callback, - logging_callback=logging_callback, - message_handler=message_handler, - read_timeout_seconds=read_timeout_seconds, - ) + super().__init__(**kwargs) self.server_script_path = server_script_path @contextlib.asynccontextmanager @@ -54,11 +35,7 @@ class StdioClient(BaseClient): async with ClientSession( read_stream=stdio, write_stream=write, - sampling_callback=self._sampling_callback, - list_roots_callback=self._list_roots_callback, - logging_callback=self._logging_callback, - message_handler=self._message_handler, - read_timeout_seconds=self._read_timeout_seconds, + **self._session_kwargs(), ) as session: async with self._set_session(transport, session): yield self diff --git a/src/fastmcp/client/websocket.py b/src/fastmcp/client/websocket.py index c66173ec4..d3a77a47b 100644 --- a/src/fastmcp/client/websocket.py +++ b/src/fastmcp/client/websocket.py @@ -1,38 +1,19 @@ import contextlib -import datetime -import mcp.types from mcp import ClientSession from mcp.client.websocket import websocket_client +from typing_extensions import Unpack -from fastmcp.client.base import ( - BaseClient, - ListRootsFnT, - LoggingFnT, - MessageHandlerFnT, - SamplingFnT, -) +from fastmcp.client.base import BaseClient, ClientKwargs class WebSocketClient(BaseClient): def __init__( self, url: str, - roots: list[mcp.types.Root] | None = None, - sampling_callback: SamplingFnT | None = None, - list_roots_callback: ListRootsFnT | None = None, - logging_callback: LoggingFnT | None = None, - message_handler: MessageHandlerFnT | None = None, - read_timeout_seconds: datetime.timedelta | None = None, + **kwargs: Unpack[ClientKwargs], ): - super().__init__( - roots=roots, - sampling_callback=sampling_callback, - list_roots_callback=list_roots_callback, - logging_callback=logging_callback, - message_handler=message_handler, - read_timeout_seconds=read_timeout_seconds, - ) + super().__init__(**kwargs) self.url = url @contextlib.asynccontextmanager @@ -44,11 +25,7 @@ class WebSocketClient(BaseClient): async with ClientSession( read_stream=read_stream, write_stream=write_stream, - sampling_callback=self._sampling_callback, - list_roots_callback=self._list_roots_callback, - logging_callback=self._logging_callback, - message_handler=self._message_handler, - read_timeout_seconds=self._read_timeout_seconds, + **self._session_kwargs(), ) as session: async with self._set_session(transport, session): yield self