mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-09 07:09:11 +02:00
Speed up the unit test suite, and fix the task-notification race it surfaced (#4550)
This commit is contained in:
parent
eee5e91334
commit
b9b1deacb6
21 changed files with 518 additions and 361 deletions
|
|
@ -63,6 +63,7 @@ These control how the server listens when running with an HTTP transport.
|
|||
|---|---|---|---|
|
||||
| `FASTMCP_CLIENT_INIT_TIMEOUT` | `float \| None` | None | Timeout in seconds for the client initialization handshake. Set to `0` or leave unset to disable. |
|
||||
| `FASTMCP_CLIENT_DISCONNECT_TIMEOUT` | `float` | `5` | Maximum time in seconds to wait for a clean disconnect before giving up. |
|
||||
| `FASTMCP_CLIENT_TASK_POLL_INTERVAL` | `float` | `0.5` | Ceiling in seconds for the fallback poll backoff while waiting on a [background task](/servers/tasks). Applies **only** when the server does not advertise its own `pollInterval`: in that case `Task.wait()` starts polling fast (~20ms) and doubles up to this ceiling rather than polling at a fixed cadence. When the server advertises a `pollInterval`, that interval is honored exactly and this setting is ignored. |
|
||||
| `FASTMCP_CLIENT_RAISE_FIRST_EXCEPTIONGROUP_ERROR` | `bool` | `true` | When an `ExceptionGroup` is raised, re-raise the first error directly instead of the group. Simplifies debugging but may mask secondary errors. |
|
||||
|
||||
## CLI & Display
|
||||
|
|
|
|||
|
|
@ -128,7 +128,7 @@ async def slow_task() -> str:
|
|||
return "Eventually done"
|
||||
```
|
||||
|
||||
Shorter intervals give clients faster feedback but increase server load. Longer intervals reduce load but delay status updates.
|
||||
Shorter intervals give clients faster feedback but increase server load. Longer intervals reduce load but delay status updates. FastMCP clients honor the advertised interval exactly, so this is a real load control — but note that status notifications still wake a waiting client immediately, so the interval only governs how quickly a *missed* notification is noticed.
|
||||
|
||||
### Server-Wide Default
|
||||
|
||||
|
|
|
|||
|
|
@ -14,12 +14,20 @@ from typing import TYPE_CHECKING, Generic, TypeVar
|
|||
import mcp_types
|
||||
from mcp_types import GetTaskResult, TaskStatusNotification
|
||||
|
||||
import fastmcp
|
||||
from fastmcp.client.messages import Message, MessageHandler
|
||||
from fastmcp.exceptions import ToolError
|
||||
from fastmcp.utilities.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
# Floor for the fallback poll interval in Task.wait() (seconds). When the server
|
||||
# does not advertise a pollInterval, each wait() call starts its backoff ramp
|
||||
# here so fast tasks resolve quickly even if a status notification is missed.
|
||||
# When the server does advertise one, this is only a safety floor that keeps a
|
||||
# server sending `pollInterval: 0` from spinning the client in a tight loop.
|
||||
MIN_POLL_INTERVAL = 0.02
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastmcp.client.client import CallToolResult, Client
|
||||
|
||||
|
|
@ -217,6 +225,12 @@ class Task(abc.ABC, Generic[TaskResultT]):
|
|||
with fallback to polling (reliable). Optimally wakes up immediately
|
||||
on status changes when server sends notifications/tasks/status.
|
||||
|
||||
The fallback poll cadence has two modes. If the server advertises a
|
||||
`pollInterval`, that interval is honored exactly (subject only to a
|
||||
20ms safety floor), because it is a deliberate statement about how
|
||||
much load the server wants to take. If it does not, the poll starts at
|
||||
20ms and doubles up to the `client_task_poll_interval` setting.
|
||||
|
||||
Args:
|
||||
state: Desired state ('working', 'input_required', 'completed', 'failed', 'cancelled').
|
||||
If None, waits until the task exits the 'working' state (completed, failed, cancelled, input_required, etc.)
|
||||
|
|
@ -240,7 +254,10 @@ class Task(abc.ABC, Generic[TaskResultT]):
|
|||
|
||||
start = time.time()
|
||||
in_progress_states = {"working"}
|
||||
poll_interval = 0.5 # Fallback polling interval (500ms)
|
||||
# Backoff state for the unadvertised-interval mode; resets per wait()
|
||||
# call. Notifications still short-circuit the wait via the status event,
|
||||
# so this only governs the fallback poll.
|
||||
backoff = MIN_POLL_INTERVAL
|
||||
|
||||
while True:
|
||||
# Check cached status first (updated by notifications)
|
||||
|
|
@ -260,17 +277,41 @@ class Task(abc.ABC, Generic[TaskResultT]):
|
|||
)
|
||||
|
||||
remaining = timeout - elapsed
|
||||
interval, backoff = self._next_poll_delay(backoff)
|
||||
|
||||
# Wait for notification event OR poll timeout
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
self._status_event.wait(), timeout=min(poll_interval, remaining)
|
||||
self._status_event.wait(), timeout=min(interval, remaining)
|
||||
)
|
||||
self._status_event.clear()
|
||||
except asyncio.TimeoutError:
|
||||
# Fallback: poll server (notification didn't arrive in time)
|
||||
self._status_cache = await self._client.get_task_status(self._task_id)
|
||||
|
||||
def _next_poll_delay(self, backoff: float) -> tuple[float, float]:
|
||||
"""Delay before the next fallback poll, plus the backoff for the round after.
|
||||
|
||||
Advertised interval -> honor it; no advertised interval -> ramp.
|
||||
|
||||
A server that advertises `pollInterval` (milliseconds) is making a
|
||||
deliberate statement about how much load it wants to take, so that
|
||||
interval is used verbatim as the delay with no backoff ramp. The only
|
||||
adjustment is `MIN_POLL_INTERVAL` as a safety floor, so a server sending
|
||||
a zero or negative interval cannot spin this client in a tight request
|
||||
loop.
|
||||
|
||||
When the server advertises nothing, there is no guidance to honor, so
|
||||
the poll starts at `MIN_POLL_INTERVAL` and doubles each round up to the
|
||||
`client_task_poll_interval` setting.
|
||||
"""
|
||||
cache = self._status_cache
|
||||
if cache is not None and cache.poll_interval is not None:
|
||||
return max(cache.poll_interval / 1000, MIN_POLL_INTERVAL), backoff
|
||||
|
||||
ceiling = fastmcp.settings.client_task_poll_interval
|
||||
return min(backoff, ceiling), min(backoff * 2, ceiling)
|
||||
|
||||
async def _wait_terminal(self, timeout: float = 300.0) -> GetTaskResult:
|
||||
"""Wait until task reaches a terminal state (completed, failed, cancelled).
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ This module requires fastmcp[tasks] (pydocket). It is only imported when docket
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from contextlib import suppress
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING
|
||||
|
|
@ -27,6 +28,12 @@ if TYPE_CHECKING:
|
|||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
# Initial interval for reconciling execution state against Redis (seconds). The
|
||||
# interval doubles on each idle reconcile up to the task's poll_interval, so fast
|
||||
# tasks are caught within the first ~20ms checks while long-running tasks converge
|
||||
# to roughly one sync per advertised poll interval.
|
||||
_MIN_RECONCILE_INTERVAL_SECONDS = 0.02
|
||||
|
||||
|
||||
async def subscribe_to_task_updates(
|
||||
task_id: str,
|
||||
|
|
@ -47,44 +54,97 @@ async def subscribe_to_task_updates(
|
|||
session: MCP ServerSession for sending notifications
|
||||
docket: Docket instance for subscribing to execution events
|
||||
poll_interval_ms: Poll interval in milliseconds to include in notifications
|
||||
|
||||
Note: Docket's ``execution.subscribe()`` replays the current state and a progress
|
||||
event before it subscribes to Redis pub/sub. A task that completes during that
|
||||
window has its terminal state publish lost, so no live event ever arrives — a
|
||||
common case for fast tasks. Because there is no reliable signal for when the
|
||||
subscription goes live (the replayed state event arrives two iterations early),
|
||||
we simply reconcile the execution against Redis on every idle interval until a
|
||||
terminal state is observed. The interval backs off exponentially toward the
|
||||
task's advertised poll interval, so a long-running task costs about one sync per
|
||||
poll interval while live pub/sub events still short-circuit the wait instantly.
|
||||
"""
|
||||
terminal_states = {
|
||||
ExecutionState.COMPLETED,
|
||||
ExecutionState.FAILED,
|
||||
ExecutionState.CANCELLED,
|
||||
}
|
||||
try:
|
||||
execution = await docket.get_execution(task_key)
|
||||
if execution is None:
|
||||
logger.warning(f"No execution found for task {task_id}")
|
||||
return
|
||||
|
||||
# Subscribe to state and progress events from Docket
|
||||
terminal_states = {
|
||||
ExecutionState.COMPLETED,
|
||||
ExecutionState.FAILED,
|
||||
ExecutionState.CANCELLED,
|
||||
}
|
||||
async for event in execution.subscribe():
|
||||
if event["type"] == "state":
|
||||
state = ExecutionState(event["state"])
|
||||
# Send notifications/tasks/status when state changes
|
||||
await _send_status_notification(
|
||||
session=session,
|
||||
task_id=task_id,
|
||||
task_key=task_key,
|
||||
docket=docket,
|
||||
state=state,
|
||||
poll_interval_ms=poll_interval_ms,
|
||||
)
|
||||
# Stop subscribing once the task reaches a terminal state
|
||||
if state in terminal_states:
|
||||
subscription = execution.subscribe()
|
||||
# Keep a single outstanding __anext__ across reconcile timeouts. asyncio.wait
|
||||
# returns on timeout without cancelling it, so the generator (and its pub/sub
|
||||
# subscription) stays intact — unlike wait_for, which would cancel mid-iteration.
|
||||
next_event = asyncio.ensure_future(subscription.__anext__())
|
||||
# Reconcile cadence backs off exponentially so a task that runs for a long
|
||||
# time (or that no worker ever claims) doesn't pin this loop at 50 syncs/sec
|
||||
# forever; the task's advertised poll interval is the natural ceiling.
|
||||
reconcile_backoff = _MIN_RECONCILE_INTERVAL_SECONDS
|
||||
reconcile_ceiling = max(
|
||||
poll_interval_ms / 1000, _MIN_RECONCILE_INTERVAL_SECONDS
|
||||
)
|
||||
try:
|
||||
while True:
|
||||
done, _ = await asyncio.wait({next_event}, timeout=reconcile_backoff)
|
||||
if not done:
|
||||
# No live event yet: reconcile against Redis in case a
|
||||
# terminal transition was published before pub/sub went live.
|
||||
await execution.sync()
|
||||
if execution.state in terminal_states:
|
||||
await _send_status_notification(
|
||||
session=session,
|
||||
task_id=task_id,
|
||||
task_key=task_key,
|
||||
docket=docket,
|
||||
state=execution.state,
|
||||
poll_interval_ms=poll_interval_ms,
|
||||
)
|
||||
break
|
||||
reconcile_backoff = min(reconcile_backoff * 2, reconcile_ceiling)
|
||||
continue
|
||||
|
||||
try:
|
||||
event = next_event.result()
|
||||
except StopAsyncIteration:
|
||||
break
|
||||
elif event["type"] == "progress":
|
||||
# Send notification when progress message changes
|
||||
await _send_progress_notification(
|
||||
session=session,
|
||||
task_id=task_id,
|
||||
task_key=task_key,
|
||||
docket=docket,
|
||||
execution=execution,
|
||||
poll_interval_ms=poll_interval_ms,
|
||||
)
|
||||
|
||||
if event["type"] == "state":
|
||||
state = ExecutionState(event["state"])
|
||||
# Send notifications/tasks/status when state changes
|
||||
await _send_status_notification(
|
||||
session=session,
|
||||
task_id=task_id,
|
||||
task_key=task_key,
|
||||
docket=docket,
|
||||
state=state,
|
||||
poll_interval_ms=poll_interval_ms,
|
||||
)
|
||||
# Stop subscribing once the task reaches a terminal state
|
||||
if state in terminal_states:
|
||||
break
|
||||
elif event["type"] == "progress":
|
||||
# Send notification when progress message changes
|
||||
await _send_progress_notification(
|
||||
session=session,
|
||||
task_id=task_id,
|
||||
task_key=task_key,
|
||||
docket=docket,
|
||||
execution=execution,
|
||||
poll_interval_ms=poll_interval_ms,
|
||||
)
|
||||
|
||||
next_event = asyncio.ensure_future(subscription.__anext__())
|
||||
finally:
|
||||
if not next_event.done():
|
||||
next_event.cancel()
|
||||
with suppress(asyncio.CancelledError, StopAsyncIteration):
|
||||
await next_event
|
||||
await subscription.aclose()
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Subscription task failed for {task_id}: {e}", exc_info=True)
|
||||
|
|
|
|||
|
|
@ -287,6 +287,24 @@ class Settings(BaseSettings):
|
|||
),
|
||||
] = 5
|
||||
|
||||
client_task_poll_interval: Annotated[
|
||||
float,
|
||||
Field(
|
||||
description=inspect.cleandoc(
|
||||
"""
|
||||
Ceiling, in seconds, for the fallback poll backoff while waiting on a
|
||||
background task (SEP-1686). Applies only when the server does not
|
||||
advertise its own pollInterval: in that case Task.wait() starts polling
|
||||
fast (~20ms) and doubles up to this ceiling, so quick tasks resolve
|
||||
promptly while long-running tasks don't hammer the server. When the
|
||||
server does advertise a pollInterval, that interval is honored exactly
|
||||
and this setting is ignored. Must be positive.
|
||||
"""
|
||||
),
|
||||
gt=0,
|
||||
),
|
||||
] = 0.5
|
||||
|
||||
# Transport settings
|
||||
transport: Literal["stdio", "http", "sse", "streamable-http"] = "stdio"
|
||||
|
||||
|
|
|
|||
|
|
@ -27,6 +27,11 @@ async def task_notification_server():
|
|||
await asyncio.sleep(0.05)
|
||||
return value * 2
|
||||
|
||||
@mcp.tool(task=True)
|
||||
async def instant_task(value: int) -> int:
|
||||
"""Background task that completes with no delay."""
|
||||
return value * 2
|
||||
|
||||
@mcp.tool(task=True)
|
||||
async def slow_task(duration: float = 0.2) -> str:
|
||||
"""Slow background task."""
|
||||
|
|
@ -210,6 +215,36 @@ async def test_notification_with_failed_task(task_notification_server):
|
|||
) # Error details in statusMessage per spec
|
||||
|
||||
|
||||
async def test_fast_task_completion_delivered_via_notification(
|
||||
task_notification_server,
|
||||
):
|
||||
"""A near-instant task still delivers its completion via a status notification.
|
||||
|
||||
Regression test for the Docket subscribe() setup-window race: a task that
|
||||
finishes before the pub/sub subscription goes live had its terminal state
|
||||
publish lost, so no completion notification ever reached the client and
|
||||
wait() fell back to a full poll interval. The server now reconciles the
|
||||
execution against Redis to close that gap.
|
||||
|
||||
Callbacks fire only for received notifications — client-side polling updates
|
||||
the status cache directly without invoking them — so a "completed" callback
|
||||
proves the notification path (not the poll fallback) was exercised.
|
||||
"""
|
||||
received: list[str] = []
|
||||
|
||||
async with Client(task_notification_server) as client:
|
||||
task = await client.call_tool("instant_task", {"value": 21}, task=True)
|
||||
task.on_status_change(lambda status: received.append(status.status))
|
||||
|
||||
result = await task
|
||||
assert result.data == 42
|
||||
|
||||
# Allow the completion notification to arrive and dispatch.
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
assert "completed" in received
|
||||
|
||||
|
||||
async def test_wait_returns_on_input_required(task_notification_server):
|
||||
"""wait() should return immediately when task enters input_required, not hang."""
|
||||
async with Client(task_notification_server) as client:
|
||||
|
|
|
|||
90
tests/client/tasks/test_poll_interval.py
Normal file
90
tests/client/tasks/test_poll_interval.py
Normal file
|
|
@ -0,0 +1,90 @@
|
|||
"""Fallback poll cadence for client-side task waiting.
|
||||
|
||||
Two modes: a server-advertised pollInterval is honored exactly, while an
|
||||
unadvertised one falls back to an exponential ramp up to the client setting.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from mcp_types import GetTaskResult
|
||||
from pydantic import ValidationError
|
||||
|
||||
from fastmcp import Client, FastMCP
|
||||
from fastmcp.client.tasks import MIN_POLL_INTERVAL, ToolTask
|
||||
from fastmcp.settings import Settings
|
||||
from fastmcp.utilities.tests import temporary_settings
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [0, -0.5, -1])
|
||||
def test_non_positive_poll_interval_setting_is_rejected(value: float):
|
||||
with pytest.raises(ValidationError):
|
||||
Settings(client_task_poll_interval=value)
|
||||
|
||||
|
||||
def test_positive_poll_interval_setting_is_accepted():
|
||||
settings = Settings(client_task_poll_interval=0.25)
|
||||
assert settings.client_task_poll_interval == 0.25
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def task() -> ToolTask:
|
||||
client = Client(FastMCP())
|
||||
return ToolTask(client=client, task_id="t1", tool_name="echo")
|
||||
|
||||
|
||||
def _status(poll_interval: int | None) -> GetTaskResult:
|
||||
return GetTaskResult(
|
||||
task_id="t1",
|
||||
status="working",
|
||||
created_at="2026-01-01T00:00:00+00:00",
|
||||
last_updated_at="2026-01-01T00:00:00+00:00",
|
||||
ttl=None,
|
||||
poll_interval=poll_interval,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("poll_interval", [2000, 30_000])
|
||||
def test_advertised_interval_is_used_verbatim_without_backoff(
|
||||
task: ToolTask, poll_interval: int
|
||||
):
|
||||
"""An advertised interval is the delay itself, not a ceiling to ramp toward."""
|
||||
task._status_cache = _status(poll_interval)
|
||||
expected = poll_interval / 1000
|
||||
|
||||
backoff = MIN_POLL_INTERVAL
|
||||
for _ in range(5):
|
||||
delay, backoff = task._next_poll_delay(backoff)
|
||||
assert delay == expected
|
||||
|
||||
|
||||
def test_large_advertised_interval_is_honored(task: ToolTask):
|
||||
task._status_cache = _status(24 * 60 * 60 * 1000)
|
||||
delay, _ = task._next_poll_delay(MIN_POLL_INTERVAL)
|
||||
assert delay == 24 * 60 * 60
|
||||
|
||||
|
||||
@pytest.mark.parametrize("poll_interval", [0, -1, -5000])
|
||||
def test_non_positive_advertised_interval_is_floored(
|
||||
task: ToolTask, poll_interval: int
|
||||
):
|
||||
"""A buggy or hostile server must not be able to spin the client."""
|
||||
task._status_cache = _status(poll_interval)
|
||||
delay, _ = task._next_poll_delay(MIN_POLL_INTERVAL)
|
||||
assert delay == MIN_POLL_INTERVAL
|
||||
|
||||
|
||||
def test_unadvertised_interval_ramps_up_to_setting(task: ToolTask):
|
||||
task._status_cache = _status(None)
|
||||
with temporary_settings(client_task_poll_interval=0.5):
|
||||
delays = []
|
||||
backoff = MIN_POLL_INTERVAL
|
||||
for _ in range(7):
|
||||
delay, backoff = task._next_poll_delay(backoff)
|
||||
delays.append(delay)
|
||||
|
||||
assert delays == [0.02, 0.04, 0.08, 0.16, 0.32, 0.5, 0.5]
|
||||
|
||||
|
||||
def test_missing_status_cache_ramps_from_floor(task: ToolTask):
|
||||
delay, backoff = task._next_poll_delay(MIN_POLL_INTERVAL)
|
||||
assert delay == MIN_POLL_INTERVAL
|
||||
assert backoff == MIN_POLL_INTERVAL * 2
|
||||
|
|
@ -14,6 +14,7 @@ from opentelemetry.sdk.trace import TracerProvider
|
|||
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
|
||||
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
|
||||
|
||||
from fastmcp.server.auth.providers.jwt import RSAKeyPair
|
||||
from fastmcp.utilities.tests import temporary_settings
|
||||
from tests.utilities.httpx2_mock import httpx_mock as httpx_mock
|
||||
|
||||
|
|
@ -111,6 +112,31 @@ def get_fn_name(fn: Callable[..., Any]) -> str:
|
|||
return fn.__name__ # ty: ignore[unresolved-attribute]
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def rsa_key_pair() -> RSAKeyPair:
|
||||
"""A shared RSA key pair for tests that just need *some* valid key material.
|
||||
|
||||
RSA key generation costs tens of milliseconds; hundreds of auth tests
|
||||
generating a fresh key per test adds up to real wall time for no benefit,
|
||||
since almost none of them care that the key is unique. Tests that must
|
||||
prove verification fails against a *different* key should use
|
||||
``rsa_key_pair_2`` instead of calling ``RSAKeyPair.generate()`` directly.
|
||||
Tests that specifically exercise key generation or rotation should still
|
||||
call ``RSAKeyPair.generate()`` themselves.
|
||||
"""
|
||||
return RSAKeyPair.generate()
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def rsa_key_pair_2() -> RSAKeyPair:
|
||||
"""A second shared RSA key pair, distinct from ``rsa_key_pair``.
|
||||
|
||||
For tests that sign a token with the "wrong" key to prove verification
|
||||
against ``rsa_key_pair`` fails.
|
||||
"""
|
||||
return RSAKeyPair.generate()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def worker_id(request):
|
||||
"""Get the xdist worker ID, or 'master' if not using xdist."""
|
||||
|
|
|
|||
|
|
@ -177,10 +177,11 @@ class TestComponentManagementRoutes:
|
|||
class TestAuthComponentManagementRoutes:
|
||||
"""Test the component management routes with authentication for tools, resources, and prompts."""
|
||||
|
||||
def setup_method(self):
|
||||
@pytest.fixture(autouse=True)
|
||||
def setup(self, rsa_key_pair: RSAKeyPair):
|
||||
"""Set up test fixtures."""
|
||||
# Generate a key pair and create an auth provider
|
||||
key_pair = RSAKeyPair.generate()
|
||||
# Create an auth provider from the shared test key pair
|
||||
key_pair = rsa_key_pair
|
||||
self.auth = JWTVerifier(
|
||||
public_key=key_pair.public_key,
|
||||
issuer="https://dev.example.com",
|
||||
|
|
@ -463,9 +464,10 @@ class TestComponentManagerWithPath:
|
|||
class TestComponentManagerWithPathAuth:
|
||||
"""Test component manager routes with auth when mounted at a custom path."""
|
||||
|
||||
def setup_method(self):
|
||||
# Generate a key pair and create an auth provider
|
||||
key_pair = RSAKeyPair.generate()
|
||||
@pytest.fixture(autouse=True)
|
||||
def setup(self, rsa_key_pair: RSAKeyPair):
|
||||
# Create an auth provider from the shared test key pair
|
||||
key_pair = rsa_key_pair
|
||||
self.auth = JWTVerifier(
|
||||
public_key=key_pair.public_key,
|
||||
issuer="https://dev.example.com",
|
||||
|
|
|
|||
|
|
@ -116,8 +116,8 @@ def _make_proxy(identity_assertion: IdentityAssertion | None) -> OAuthProxy:
|
|||
|
||||
|
||||
@pytest.fixture
|
||||
def idp_key() -> RSAKeyPair:
|
||||
return RSAKeyPair.generate()
|
||||
def idp_key(rsa_key_pair: RSAKeyPair) -> RSAKeyPair:
|
||||
return rsa_key_pair
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -503,9 +503,11 @@ class TestValidationMatrix:
|
|||
assert second.status_code == 401
|
||||
assert second.json()["error"] == "invalid_grant"
|
||||
|
||||
async def test_wrong_signature_rejected(self, config: IdentityAssertion):
|
||||
async def test_wrong_signature_rejected(
|
||||
self, config: IdentityAssertion, rsa_key_pair_2: RSAKeyPair
|
||||
):
|
||||
# Sign with a different key than the one served in the JWKS.
|
||||
other_key = RSAKeyPair.generate()
|
||||
other_key = rsa_key_pair_2
|
||||
proxy = _make_proxy(config)
|
||||
assertion = _mint_id_jag(other_key)
|
||||
|
||||
|
|
|
|||
|
|
@ -49,8 +49,10 @@ class TestAuth0JWTVerifier:
|
|||
scopes = verifier._extract_scopes({"permissions": "tool:whoami tool:greet"})
|
||||
assert scopes == ["tool:whoami", "tool:greet"]
|
||||
|
||||
async def test_verify_token_accepts_permissions_as_required_scopes(self):
|
||||
key_pair = RSAKeyPair.generate()
|
||||
async def test_verify_token_accepts_permissions_as_required_scopes(
|
||||
self, rsa_key_pair: RSAKeyPair
|
||||
):
|
||||
key_pair = rsa_key_pair
|
||||
verifier = Auth0JWTVerifier(
|
||||
public_key=key_pair.public_key,
|
||||
issuer=TEST_ISSUER,
|
||||
|
|
@ -66,8 +68,10 @@ class TestAuth0JWTVerifier:
|
|||
assert access_token is not None
|
||||
assert access_token.client_id == "user_123"
|
||||
|
||||
async def test_verify_token_rejects_missing_permissions(self):
|
||||
key_pair = RSAKeyPair.generate()
|
||||
async def test_verify_token_rejects_missing_permissions(
|
||||
self, rsa_key_pair: RSAKeyPair
|
||||
):
|
||||
key_pair = rsa_key_pair
|
||||
verifier = Auth0JWTVerifier(
|
||||
public_key=key_pair.public_key,
|
||||
issuer=TEST_ISSUER,
|
||||
|
|
|
|||
|
|
@ -222,10 +222,10 @@ class TestAzureProvider:
|
|||
assert verifier.required_scopes == [".default"]
|
||||
|
||||
async def test_token_accepted_with_client_id_audience(
|
||||
self, memory_storage: MemoryStore
|
||||
self, memory_storage: MemoryStore, rsa_key_pair: RSAKeyPair
|
||||
):
|
||||
"""Azure AD v2 tokens use the bare client_id as aud — must be accepted."""
|
||||
key_pair = RSAKeyPair.generate()
|
||||
key_pair = rsa_key_pair
|
||||
provider = AzureProvider(
|
||||
client_id="test_client",
|
||||
client_secret="test_secret",
|
||||
|
|
@ -252,10 +252,10 @@ class TestAzureProvider:
|
|||
assert result is not None
|
||||
|
||||
async def test_token_accepted_with_identifier_uri_audience(
|
||||
self, memory_storage: MemoryStore
|
||||
self, memory_storage: MemoryStore, rsa_key_pair: RSAKeyPair
|
||||
):
|
||||
"""Azure AD v1 tokens use the identifier_uri as aud — must be accepted."""
|
||||
key_pair = RSAKeyPair.generate()
|
||||
key_pair = rsa_key_pair
|
||||
provider = AzureProvider(
|
||||
client_id="test_client",
|
||||
client_secret="test_secret",
|
||||
|
|
@ -282,10 +282,10 @@ class TestAzureProvider:
|
|||
assert result is not None
|
||||
|
||||
async def test_token_rejected_with_wrong_audience(
|
||||
self, memory_storage: MemoryStore
|
||||
self, memory_storage: MemoryStore, rsa_key_pair: RSAKeyPair
|
||||
):
|
||||
"""Tokens for a different application must be rejected."""
|
||||
key_pair = RSAKeyPair.generate()
|
||||
key_pair = rsa_key_pair
|
||||
provider = AzureProvider(
|
||||
client_id="test_client",
|
||||
client_secret="test_secret",
|
||||
|
|
@ -888,9 +888,11 @@ class TestAzureProviderTokenIssuer:
|
|||
assert isinstance(provider._token_validator, JWTVerifier)
|
||||
assert provider._token_validator.issuer == custom_issuer
|
||||
|
||||
async def test_explicit_issuer_enforced(self, memory_storage: MemoryStore):
|
||||
async def test_explicit_issuer_enforced(
|
||||
self, memory_storage: MemoryStore, rsa_key_pair: RSAKeyPair
|
||||
):
|
||||
"""With an explicit token_issuer, wrong issuers are rejected."""
|
||||
key_pair = RSAKeyPair.generate()
|
||||
key_pair = rsa_key_pair
|
||||
expected = "https://expected.issuer.com/v2.0"
|
||||
provider = AzureProvider(
|
||||
client_id="test_client",
|
||||
|
|
@ -1109,10 +1111,10 @@ class TestAzureProviderFromB2C:
|
|||
assert "B2C_1A_SIGNUP_SIGNIN" in provider._upstream_token_endpoint
|
||||
|
||||
async def test_b2c_token_accepted_with_any_issuer(
|
||||
self, memory_storage: MemoryStore
|
||||
self, memory_storage: MemoryStore, rsa_key_pair: RSAKeyPair
|
||||
):
|
||||
"""B2C provider (issuer=None) accepts tokens from any issuer."""
|
||||
key_pair = RSAKeyPair.generate()
|
||||
key_pair = rsa_key_pair
|
||||
provider = AzureProvider.from_b2c(
|
||||
tenant_name="mytenant",
|
||||
policy_name="B2C_1_susi",
|
||||
|
|
@ -1139,10 +1141,10 @@ class TestAzureProviderFromB2C:
|
|||
assert result is not None
|
||||
|
||||
async def test_b2c_token_rejected_with_wrong_audience(
|
||||
self, memory_storage: MemoryStore
|
||||
self, memory_storage: MemoryStore, rsa_key_pair: RSAKeyPair
|
||||
):
|
||||
"""B2C provider still rejects tokens with wrong audience."""
|
||||
key_pair = RSAKeyPair.generate()
|
||||
key_pair = rsa_key_pair
|
||||
provider = AzureProvider.from_b2c(
|
||||
tenant_name="mytenant",
|
||||
policy_name="B2C_1_susi",
|
||||
|
|
|
|||
|
|
@ -425,8 +425,8 @@ class TestAzureJWTVerifier:
|
|||
assert verifier.algorithm == "RS256"
|
||||
assert verifier.required_scopes == ["access_as_user"]
|
||||
|
||||
async def test_validates_short_form_scopes(self):
|
||||
key_pair = RSAKeyPair.generate()
|
||||
async def test_validates_short_form_scopes(self, rsa_key_pair: RSAKeyPair):
|
||||
key_pair = rsa_key_pair
|
||||
verifier = AzureJWTVerifier(
|
||||
client_id="my-client-id",
|
||||
tenant_id="my-tenant-id",
|
||||
|
|
@ -446,9 +446,11 @@ class TestAzureJWTVerifier:
|
|||
assert result is not None
|
||||
assert "access_as_user" in result.scopes
|
||||
|
||||
async def test_validates_token_with_client_id_audience(self):
|
||||
async def test_validates_token_with_client_id_audience(
|
||||
self, rsa_key_pair: RSAKeyPair
|
||||
):
|
||||
"""Azure AD v2 tokens use the bare client_id GUID as audience."""
|
||||
key_pair = RSAKeyPair.generate()
|
||||
key_pair = rsa_key_pair
|
||||
verifier = AzureJWTVerifier(
|
||||
client_id="my-client-id",
|
||||
tenant_id="my-tenant-id",
|
||||
|
|
@ -467,9 +469,11 @@ class TestAzureJWTVerifier:
|
|||
assert result is not None
|
||||
assert "access_as_user" in result.scopes
|
||||
|
||||
async def test_validates_token_with_custom_identifier_uri_audience(self):
|
||||
async def test_validates_token_with_custom_identifier_uri_audience(
|
||||
self, rsa_key_pair: RSAKeyPair
|
||||
):
|
||||
"""Custom identifier_uri (e.g. Bicep deployments) accepted as audience."""
|
||||
key_pair = RSAKeyPair.generate()
|
||||
key_pair = rsa_key_pair
|
||||
verifier = AzureJWTVerifier(
|
||||
client_id="my-client-id",
|
||||
tenant_id="my-tenant-id",
|
||||
|
|
@ -489,9 +493,9 @@ class TestAzureJWTVerifier:
|
|||
assert result is not None
|
||||
assert "read" in result.scopes
|
||||
|
||||
async def test_rejects_token_with_wrong_audience(self):
|
||||
async def test_rejects_token_with_wrong_audience(self, rsa_key_pair: RSAKeyPair):
|
||||
"""Tokens for a different application must be rejected."""
|
||||
key_pair = RSAKeyPair.generate()
|
||||
key_pair = rsa_key_pair
|
||||
verifier = AzureJWTVerifier(
|
||||
client_id="my-client-id",
|
||||
tenant_id="my-tenant-id",
|
||||
|
|
|
|||
|
|
@ -124,10 +124,6 @@ class TestIntrospectionHttpClient:
|
|||
class TestJWTVerifierHttpClient:
|
||||
"""Test http_client parameter on JWTVerifier."""
|
||||
|
||||
@pytest.fixture(scope="class")
|
||||
def rsa_key_pair(self) -> RSAKeyPair:
|
||||
return RSAKeyPair.generate()
|
||||
|
||||
@pytest.fixture
|
||||
def shared_client(self) -> httpx2.AsyncClient:
|
||||
return httpx2.AsyncClient(timeout=30)
|
||||
|
|
|
|||
|
|
@ -27,11 +27,9 @@ class TestCIMDAssertionValidator:
|
|||
return CIMDAssertionValidator()
|
||||
|
||||
@pytest.fixture
|
||||
def key_pair(self):
|
||||
def key_pair(self, rsa_key_pair):
|
||||
"""Generate RSA key pair for testing."""
|
||||
from fastmcp.server.auth.providers.jwt import RSAKeyPair
|
||||
|
||||
return RSAKeyPair.generate()
|
||||
return rsa_key_pair
|
||||
|
||||
@pytest.fixture
|
||||
def jwks(self, key_pair):
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ from fastmcp import FastMCP
|
|||
from fastmcp.server.auth import RemoteAuthProvider, TokenVerifier
|
||||
from fastmcp.server.auth.auth import AccessToken
|
||||
from fastmcp.server.auth.oauth_proxy import OAuthProxy
|
||||
from fastmcp.server.auth.providers.jwt import JWTVerifier, RSAKeyPair
|
||||
from fastmcp.server.auth.providers.jwt import JWTVerifier
|
||||
from fastmcp.server.http import create_streamable_http_app
|
||||
|
||||
|
||||
|
|
@ -92,11 +92,6 @@ class _DirectClientRedirectWithIssOAuthProxy(OAuthProxy):
|
|||
class TestEnhancedAuthorizationHandler:
|
||||
"""Tests for enhanced authorization handler error responses."""
|
||||
|
||||
@pytest.fixture
|
||||
def rsa_key_pair(self) -> RSAKeyPair:
|
||||
"""Generate RSA key pair for testing."""
|
||||
return RSAKeyPair.generate()
|
||||
|
||||
@pytest.fixture
|
||||
def oauth_proxy(self, rsa_key_pair):
|
||||
"""Create OAuth proxy for testing."""
|
||||
|
|
@ -668,11 +663,6 @@ class TestEnhancedRequireAuthMiddleware:
|
|||
)
|
||||
return FastMCP("Test Server", auth=auth).http_app()
|
||||
|
||||
@pytest.fixture
|
||||
def rsa_key_pair(self) -> RSAKeyPair:
|
||||
"""Generate RSA key pair for testing."""
|
||||
return RSAKeyPair.generate()
|
||||
|
||||
@pytest.fixture
|
||||
def jwt_verifier(self, rsa_key_pair):
|
||||
"""Create JWT verifier for testing."""
|
||||
|
|
@ -899,7 +889,7 @@ class TestContentNegotiation:
|
|||
"""Tests for content negotiation in error responses."""
|
||||
|
||||
@pytest.fixture
|
||||
def oauth_proxy(self):
|
||||
def oauth_proxy(self, rsa_key_pair):
|
||||
"""Create OAuth proxy for testing."""
|
||||
return OAuthProxy(
|
||||
upstream_authorization_endpoint="https://github.com/login/oauth/authorize",
|
||||
|
|
@ -907,7 +897,7 @@ class TestContentNegotiation:
|
|||
upstream_client_id="test-client-id",
|
||||
upstream_client_secret="test-client-secret",
|
||||
token_verifier=JWTVerifier(
|
||||
public_key=RSAKeyPair.generate().public_key,
|
||||
public_key=rsa_key_pair.public_key,
|
||||
issuer="https://test.com",
|
||||
audience="https://test.com",
|
||||
base_url="https://test.com",
|
||||
|
|
|
|||
|
|
@ -81,11 +81,6 @@ class SymmetricKeyHelper:
|
|||
return token
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rsa_key_pair() -> RSAKeyPair:
|
||||
return RSAKeyPair.generate()
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def symmetric_key_helper() -> SymmetricKeyHelper:
|
||||
"""Generate a symmetric key helper for testing."""
|
||||
|
|
@ -749,13 +744,14 @@ class TestBearerTokenJWKS:
|
|||
async def test_jwks_token_validation_with_invalid_key(
|
||||
self,
|
||||
rsa_key_pair: RSAKeyPair,
|
||||
rsa_key_pair_2: RSAKeyPair,
|
||||
jwks_provider: JWTVerifier,
|
||||
mock_jwks_data: JWKSData,
|
||||
httpx_mock: HTTPXMock,
|
||||
mock_dns,
|
||||
):
|
||||
httpx_mock.add_response(json=mock_jwks_data)
|
||||
token = RSAKeyPair.generate().create_token(
|
||||
token = rsa_key_pair_2.create_token(
|
||||
subject="test-user",
|
||||
issuer="https://test.example.com",
|
||||
audience="https://api.example.com",
|
||||
|
|
|
|||
|
|
@ -13,11 +13,6 @@ from fastmcp.utilities.tests import run_server_async
|
|||
TEST_PUBLIC_IP = "93.184.216.34"
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rsa_key_pair() -> RSAKeyPair:
|
||||
return RSAKeyPair.generate()
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def bearer_token(rsa_key_pair: RSAKeyPair) -> str:
|
||||
return rsa_key_pair.create_token(
|
||||
|
|
@ -394,11 +389,14 @@ class TestBearerToken:
|
|||
assert access_token is None
|
||||
|
||||
async def test_invalid_signature_rejection(
|
||||
self, rsa_key_pair: RSAKeyPair, bearer_provider: JWTVerifier
|
||||
self,
|
||||
rsa_key_pair: RSAKeyPair,
|
||||
rsa_key_pair_2: RSAKeyPair,
|
||||
bearer_provider: JWTVerifier,
|
||||
):
|
||||
"""Test rejection of tokens with invalid signatures."""
|
||||
# Create a token with a different key pair
|
||||
other_key_pair = RSAKeyPair.generate()
|
||||
other_key_pair = rsa_key_pair_2
|
||||
token = other_key_pair.create_token(
|
||||
subject="test-user",
|
||||
issuer="https://test.example.com",
|
||||
|
|
@ -517,9 +515,10 @@ class TestFastMCPBearerAuth:
|
|||
tools = await client.list_tools() # noqa: F841
|
||||
assert "tools" not in locals()
|
||||
|
||||
async def test_token_with_bad_signature(self, mcp_server_url: str):
|
||||
rsa_key_pair = RSAKeyPair.generate()
|
||||
token = rsa_key_pair.create_token()
|
||||
async def test_token_with_bad_signature(
|
||||
self, mcp_server_url: str, rsa_key_pair_2: RSAKeyPair
|
||||
):
|
||||
token = rsa_key_pair_2.create_token()
|
||||
|
||||
with pytest.raises(MCPError):
|
||||
async with Client(mcp_server_url, auth=BearerAuth(token)) as client:
|
||||
|
|
|
|||
|
|
@ -11,11 +11,6 @@ from fastmcp.server.auth.providers.jwt import JWTVerifier, RSAKeyPair
|
|||
class TestBearerAuthBackendTokenVerifierIntegration:
|
||||
"""Test BearerAuthBackend works with TokenVerifier protocol."""
|
||||
|
||||
@pytest.fixture
|
||||
def rsa_key_pair(self) -> RSAKeyPair:
|
||||
"""Generate RSA key pair for testing."""
|
||||
return RSAKeyPair.generate()
|
||||
|
||||
@pytest.fixture
|
||||
def jwt_verifier(self, rsa_key_pair: RSAKeyPair) -> JWTVerifier:
|
||||
"""Create JWTVerifier for testing."""
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from starlette.testclient import TestClient
|
|||
from starlette.types import Receive, Scope, Send
|
||||
|
||||
from fastmcp.server import FastMCP
|
||||
from fastmcp.server.auth.providers.jwt import JWTVerifier, RSAKeyPair
|
||||
from fastmcp.server.auth.providers.jwt import JWTVerifier
|
||||
from fastmcp.server.http import HostOriginGuardMiddleware, create_streamable_http_app
|
||||
|
||||
INITIALIZE_REQUEST = {
|
||||
|
|
@ -80,11 +80,6 @@ async def _guard_status(
|
|||
class TestStreamableHTTPAppResourceMetadataURL:
|
||||
"""Test resource_metadata_url logic in create_streamable_http_app."""
|
||||
|
||||
@pytest.fixture
|
||||
def rsa_key_pair(self) -> RSAKeyPair:
|
||||
"""Generate RSA key pair for testing."""
|
||||
return RSAKeyPair.generate()
|
||||
|
||||
@pytest.fixture
|
||||
def bearer_auth_provider(self, rsa_key_pair):
|
||||
provider = JWTVerifier(
|
||||
|
|
|
|||
|
|
@ -15,13 +15,15 @@ from unittest.mock import AsyncMock, patch
|
|||
import psutil
|
||||
import pytest
|
||||
from mcp_types import TextContent
|
||||
from pydantic import ConfigDict
|
||||
|
||||
from fastmcp import FastMCP
|
||||
from fastmcp import Context, FastMCP
|
||||
from fastmcp.client.auth.bearer import BearerAuth
|
||||
from fastmcp.client.auth.oauth import OAuthClientProvider
|
||||
from fastmcp.client.client import Client
|
||||
from fastmcp.client.logging import LogMessage
|
||||
from fastmcp.client.transports import (
|
||||
FastMCPTransport,
|
||||
MCPConfigTransport,
|
||||
SSETransport,
|
||||
StdioTransport,
|
||||
|
|
@ -36,6 +38,7 @@ from fastmcp.mcp_config import (
|
|||
StdioMCPServer,
|
||||
TransformingStdioMCPServer,
|
||||
)
|
||||
from fastmcp.server.elicitation import AcceptedElicitation
|
||||
from fastmcp.tools.base import Tool as FastMCPTool
|
||||
|
||||
# These tests spawn subprocess servers via stdio which can be slow under
|
||||
|
|
@ -63,6 +66,33 @@ def gc_collect_harder():
|
|||
gc.collect()
|
||||
|
||||
|
||||
class InMemoryStdioMCPServer(StdioMCPServer):
|
||||
"""Test double for a plain (non-transforming) `StdioMCPServer` that skips
|
||||
subprocess spawning in favor of an in-memory transport.
|
||||
|
||||
`MCPConfigTransport`'s composite path calls `server_config.to_transport()`
|
||||
polymorphically for any *non-transforming* server entry (see
|
||||
`_create_proxy` in `fastmcp.client.transports.config`), so overriding
|
||||
`to_transport()` on a subclass is enough to swap in an in-memory backend
|
||||
while still exercising the real MCPConfig/MCPConfigTransport composition
|
||||
code: proxy creation, namespace-prefixed mounting, log/elicitation
|
||||
forwarding, and session handling.
|
||||
|
||||
This does NOT work for `TransformingStdioMCPServer` configs (tool
|
||||
transforms / tag filters): `_create_proxy` calls the *unbound*
|
||||
`StdioMCPServer.to_transport` for those, bypassing any subclass override,
|
||||
so transform/tag-filter tests still need a real subprocess.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True)
|
||||
|
||||
mcp: FastMCP
|
||||
command: str = "in-memory"
|
||||
|
||||
def to_transport(self) -> FastMCPTransport: # ty: ignore[invalid-method-override]
|
||||
return FastMCPTransport(mcp=self.mcp)
|
||||
|
||||
|
||||
def test_parse_single_stdio_config():
|
||||
config = {
|
||||
"mcpServers": {
|
||||
|
|
@ -310,35 +340,23 @@ def test_parse_multiple_servers():
|
|||
assert mcp_config.mcpServers["test_server_2"].env == {"TEST": "test"}
|
||||
|
||||
|
||||
async def test_multi_client(tmp_path: Path):
|
||||
server_script = inspect.cleandoc("""
|
||||
from fastmcp import FastMCP
|
||||
def _make_add_server() -> FastMCP:
|
||||
app = FastMCP()
|
||||
|
||||
mcp = FastMCP()
|
||||
@app.tool
|
||||
def add(a: int, b: int) -> int:
|
||||
return a + b
|
||||
|
||||
@mcp.tool
|
||||
def add(a: int, b: int) -> int:
|
||||
return a + b
|
||||
return app
|
||||
|
||||
if __name__ == '__main__':
|
||||
mcp.run()
|
||||
""")
|
||||
|
||||
script_path = tmp_path / "test.py"
|
||||
script_path.write_text(server_script)
|
||||
|
||||
config = {
|
||||
"mcpServers": {
|
||||
"test_1": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
"test_2": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
async def test_multi_client():
|
||||
config = MCPConfig(
|
||||
mcpServers={
|
||||
"test_1": InMemoryStdioMCPServer(mcp=_make_add_server()),
|
||||
"test_2": InMemoryStdioMCPServer(mcp=_make_add_server()),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
client = Client(config)
|
||||
|
||||
|
|
@ -352,35 +370,13 @@ async def test_multi_client(tmp_path: Path):
|
|||
assert result_2.data == 3
|
||||
|
||||
|
||||
async def test_multi_client_parallel_calls(tmp_path: Path):
|
||||
server_script = inspect.cleandoc("""
|
||||
from fastmcp import FastMCP
|
||||
|
||||
mcp = FastMCP()
|
||||
|
||||
@mcp.tool
|
||||
def add(a: int, b: int) -> int:
|
||||
return a + b
|
||||
|
||||
if __name__ == '__main__':
|
||||
mcp.run()
|
||||
""")
|
||||
|
||||
script_path = tmp_path / "test.py"
|
||||
script_path.write_text(server_script)
|
||||
|
||||
config = {
|
||||
"mcpServers": {
|
||||
"test_1": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
"test_2": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
async def test_multi_client_parallel_calls():
|
||||
config = MCPConfig(
|
||||
mcpServers={
|
||||
"test_1": InMemoryStdioMCPServer(mcp=_make_add_server()),
|
||||
"test_2": InMemoryStdioMCPServer(mcp=_make_add_server()),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
client = Client(config)
|
||||
|
||||
|
|
@ -574,41 +570,29 @@ async def test_remote_config_with_oauth_literal():
|
|||
assert isinstance(client.transport.transport.auth, OAuthClientProvider)
|
||||
|
||||
|
||||
async def test_multi_client_with_logging(tmp_path: Path, caplog):
|
||||
def _make_log_server() -> FastMCP:
|
||||
app = FastMCP()
|
||||
|
||||
@app.tool
|
||||
async def log_test(message: str, ctx: Context) -> int:
|
||||
await ctx.log(message)
|
||||
return 42
|
||||
|
||||
return app
|
||||
|
||||
|
||||
async def test_multi_client_with_logging(caplog):
|
||||
"""
|
||||
Tests that logging is properly forwarded to the ultimate client.
|
||||
"""
|
||||
caplog.set_level(logging.INFO, logger=__name__)
|
||||
|
||||
server_script = inspect.cleandoc("""
|
||||
from fastmcp import FastMCP, Context
|
||||
|
||||
mcp = FastMCP()
|
||||
|
||||
@mcp.tool
|
||||
async def log_test(message: str, ctx: Context) -> int:
|
||||
await ctx.log(message)
|
||||
return 42
|
||||
|
||||
if __name__ == '__main__':
|
||||
mcp.run()
|
||||
""")
|
||||
|
||||
script_path = tmp_path / "test.py"
|
||||
script_path.write_text(server_script)
|
||||
|
||||
config = {
|
||||
"mcpServers": {
|
||||
"test_server": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
"test_server_2": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
config = MCPConfig(
|
||||
mcpServers={
|
||||
"test_server": InMemoryStdioMCPServer(mcp=_make_log_server()),
|
||||
"test_server_2": InMemoryStdioMCPServer(mcp=_make_log_server()),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
MESSAGES = []
|
||||
|
||||
|
|
@ -852,39 +836,29 @@ async def test_single_server_config_include_tags_filtering(tmp_path: Path):
|
|||
assert "subtract" not in tool_names
|
||||
|
||||
|
||||
async def test_multi_client_with_elicitation(tmp_path: Path):
|
||||
def _make_elicit_server() -> FastMCP:
|
||||
app = FastMCP()
|
||||
|
||||
@app.tool
|
||||
async def elicit_test(ctx: Context) -> int:
|
||||
result = await ctx.elicit("Pick a number", response_type=int)
|
||||
assert isinstance(result, AcceptedElicitation)
|
||||
assert isinstance(result.data, int)
|
||||
return result.data
|
||||
|
||||
return app
|
||||
|
||||
|
||||
async def test_multi_client_with_elicitation():
|
||||
"""
|
||||
Tests that elicitation is properly forwarded to the ultimate client.
|
||||
"""
|
||||
server_script = inspect.cleandoc("""
|
||||
from fastmcp import FastMCP, Context
|
||||
|
||||
mcp = FastMCP()
|
||||
|
||||
@mcp.tool
|
||||
async def elicit_test(ctx: Context) -> int:
|
||||
result = await ctx.elicit('Pick a number', response_type=int)
|
||||
return result.data
|
||||
|
||||
if __name__ == '__main__':
|
||||
mcp.run()
|
||||
""")
|
||||
|
||||
script_path = tmp_path / "test.py"
|
||||
script_path.write_text(server_script)
|
||||
|
||||
config = {
|
||||
"mcpServers": {
|
||||
"test_server": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
"test_server_2": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
config = MCPConfig(
|
||||
mcpServers={
|
||||
"test_server": InMemoryStdioMCPServer(mcp=_make_elicit_server()),
|
||||
"test_server_2": InMemoryStdioMCPServer(mcp=_make_elicit_server()),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
async def elicitation_handler(message, response_type, params, ctx):
|
||||
return response_type(value=42)
|
||||
|
|
@ -894,41 +868,29 @@ async def test_multi_client_with_elicitation(tmp_path: Path):
|
|||
assert result.data == 42
|
||||
|
||||
|
||||
async def test_multi_server_config_transport(tmp_path: Path):
|
||||
def _make_greet_server() -> FastMCP:
|
||||
app = FastMCP()
|
||||
|
||||
@app.tool
|
||||
def greet(name: str) -> str:
|
||||
return f"Hello, {name}!"
|
||||
|
||||
return app
|
||||
|
||||
|
||||
async def test_multi_server_config_transport():
|
||||
"""
|
||||
Tests that MCPConfigTransport properly handles multi-server configurations.
|
||||
|
||||
Related to https://github.com/PrefectHQ/fastmcp/issues/2802 - verifies the
|
||||
refactored architecture creates composite servers correctly.
|
||||
"""
|
||||
server_script = inspect.cleandoc("""
|
||||
from fastmcp import FastMCP
|
||||
|
||||
mcp = FastMCP()
|
||||
|
||||
@mcp.tool
|
||||
def greet(name: str) -> str:
|
||||
return f"Hello, {name}!"
|
||||
|
||||
if __name__ == '__main__':
|
||||
mcp.run()
|
||||
""")
|
||||
|
||||
script_path = tmp_path / "greet_server.py"
|
||||
script_path.write_text(server_script)
|
||||
|
||||
config = {
|
||||
"mcpServers": {
|
||||
"server1": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
"server2": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
config = MCPConfig(
|
||||
mcpServers={
|
||||
"server1": InMemoryStdioMCPServer(mcp=_make_greet_server()),
|
||||
"server2": InMemoryStdioMCPServer(mcp=_make_greet_server()),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
# Create client with multiple servers
|
||||
client = Client(config)
|
||||
|
|
@ -993,41 +955,29 @@ async def test_multi_server_timeout_propagation():
|
|||
)
|
||||
|
||||
|
||||
async def test_multi_server_session_persistence(tmp_path: Path):
|
||||
def _make_session_server() -> FastMCP:
|
||||
app = FastMCP()
|
||||
|
||||
@app.tool
|
||||
def get_session(ctx: Context) -> str:
|
||||
return ctx.session_id
|
||||
|
||||
return app
|
||||
|
||||
|
||||
async def test_multi_server_session_persistence():
|
||||
"""Test that session IDs persist across tool calls in multi-server mode.
|
||||
|
||||
Regression test for https://github.com/PrefectHQ/fastmcp/issues/2790 —
|
||||
MCPConfigTransport was not connecting ProxyClients before mounting, so
|
||||
each tool call opened a new session with the backend server.
|
||||
"""
|
||||
server_script = inspect.cleandoc("""
|
||||
from fastmcp import FastMCP, Context
|
||||
|
||||
mcp = FastMCP()
|
||||
|
||||
@mcp.tool
|
||||
def get_session(ctx: Context) -> str:
|
||||
return ctx.session_id
|
||||
|
||||
if __name__ == '__main__':
|
||||
mcp.run()
|
||||
""")
|
||||
|
||||
script_path = tmp_path / "session_server.py"
|
||||
script_path.write_text(server_script)
|
||||
|
||||
config = {
|
||||
"mcpServers": {
|
||||
"server1": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
"server2": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
config = MCPConfig(
|
||||
mcpServers={
|
||||
"server1": InMemoryStdioMCPServer(mcp=_make_session_server()),
|
||||
"server2": InMemoryStdioMCPServer(mcp=_make_session_server()),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
client = Client(config)
|
||||
async with client:
|
||||
|
|
@ -1070,38 +1020,19 @@ async def test_single_server_config_transport():
|
|||
],
|
||||
ids=["good_first", "bad_first"],
|
||||
)
|
||||
async def test_multi_server_partial_failure(tmp_path: Path, server_order: dict):
|
||||
async def test_multi_server_partial_failure(server_order: dict):
|
||||
"""When one server fails to connect, the others should still work."""
|
||||
server_script = inspect.cleandoc("""
|
||||
from fastmcp import FastMCP
|
||||
|
||||
mcp = FastMCP()
|
||||
|
||||
@mcp.tool
|
||||
def add(a: int, b: int) -> int:
|
||||
return a + b
|
||||
|
||||
if __name__ == '__main__':
|
||||
mcp.run()
|
||||
""")
|
||||
|
||||
script_path = tmp_path / "test.py"
|
||||
script_path.write_text(server_script)
|
||||
|
||||
servers = {}
|
||||
servers: dict[str, MCPServerTypes] = {}
|
||||
for name, is_good in server_order.items():
|
||||
if is_good:
|
||||
servers[name] = {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
}
|
||||
servers[name] = InMemoryStdioMCPServer(mcp=_make_add_server())
|
||||
else:
|
||||
servers[name] = {
|
||||
"command": "this-command-does-not-exist-anywhere",
|
||||
"args": [],
|
||||
}
|
||||
servers[name] = StdioMCPServer(
|
||||
command="this-command-does-not-exist-anywhere",
|
||||
args=[],
|
||||
)
|
||||
|
||||
client = Client({"mcpServers": servers})
|
||||
client = Client(MCPConfig(mcpServers=servers))
|
||||
async with client:
|
||||
tools = await client.list_tools()
|
||||
tool_names = [t.name for t in tools]
|
||||
|
|
@ -1109,36 +1040,17 @@ async def test_multi_server_partial_failure(tmp_path: Path, server_order: dict):
|
|||
assert len(tools) == 1
|
||||
|
||||
|
||||
async def test_multi_server_partial_failure_logs_warning(tmp_path: Path, caplog):
|
||||
async def test_multi_server_partial_failure_logs_warning(caplog):
|
||||
"""A warning should be logged when a server fails to connect."""
|
||||
server_script = inspect.cleandoc("""
|
||||
from fastmcp import FastMCP
|
||||
|
||||
mcp = FastMCP()
|
||||
|
||||
@mcp.tool
|
||||
def add(a: int, b: int) -> int:
|
||||
return a + b
|
||||
|
||||
if __name__ == '__main__':
|
||||
mcp.run()
|
||||
""")
|
||||
|
||||
script_path = tmp_path / "test.py"
|
||||
script_path.write_text(server_script)
|
||||
|
||||
config = {
|
||||
"mcpServers": {
|
||||
"good_server": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
"bad_server": {
|
||||
"command": "this-command-does-not-exist-anywhere",
|
||||
"args": [],
|
||||
},
|
||||
config = MCPConfig(
|
||||
mcpServers={
|
||||
"good_server": InMemoryStdioMCPServer(mcp=_make_add_server()),
|
||||
"bad_server": StdioMCPServer(
|
||||
command="this-command-does-not-exist-anywhere",
|
||||
args=[],
|
||||
),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
async with Client(config):
|
||||
|
|
@ -1173,36 +1085,27 @@ async def test_multi_server_all_fail():
|
|||
pass
|
||||
|
||||
|
||||
async def test_multi_server_partial_failure_cleanup(tmp_path: Path):
|
||||
def _make_ping_server() -> FastMCP:
|
||||
app = FastMCP()
|
||||
|
||||
@app.tool
|
||||
def ping() -> str:
|
||||
return "pong"
|
||||
|
||||
return app
|
||||
|
||||
|
||||
async def test_multi_server_partial_failure_cleanup():
|
||||
"""Transports for failed servers should not leak into _transports."""
|
||||
server_script = inspect.cleandoc("""
|
||||
from fastmcp import FastMCP
|
||||
|
||||
mcp = FastMCP()
|
||||
|
||||
@mcp.tool
|
||||
def ping() -> str:
|
||||
return "pong"
|
||||
|
||||
if __name__ == '__main__':
|
||||
mcp.run()
|
||||
""")
|
||||
|
||||
script_path = tmp_path / "test.py"
|
||||
script_path.write_text(server_script)
|
||||
|
||||
config = {
|
||||
"mcpServers": {
|
||||
"working": {
|
||||
"command": "python",
|
||||
"args": [str(script_path)],
|
||||
},
|
||||
"broken": {
|
||||
"command": "this-command-does-not-exist-anywhere",
|
||||
"args": [],
|
||||
},
|
||||
config = MCPConfig(
|
||||
mcpServers={
|
||||
"working": InMemoryStdioMCPServer(mcp=_make_ping_server()),
|
||||
"broken": StdioMCPServer(
|
||||
command="this-command-does-not-exist-anywhere",
|
||||
args=[],
|
||||
),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
transport = MCPConfigTransport(config)
|
||||
async with transport.connect_session():
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue