Condense client kwargs

This commit is contained in:
Jeremiah Lowin 2025-04-06 09:23:41 -04:00
commit 7aa0b9b593
5 changed files with 90 additions and 96 deletions

51
.gitignore vendored
View file

@ -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

View file

@ -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")

View file

@ -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

View file

@ -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

View file

@ -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