mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-09 07:09:11 +02:00
Condense client kwargs
This commit is contained in:
parent
4b05e7fa2d
commit
7aa0b9b593
5 changed files with 90 additions and 96 deletions
51
.gitignore
vendored
51
.gitignore
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue