feat: standardize timeout values to floats in seconds (#1766)
This commit is contained in:
@@ -29,7 +29,6 @@ import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from datetime import timedelta
|
||||
from urllib.parse import ParseResult, parse_qs, urlparse
|
||||
|
||||
import httpx
|
||||
@@ -263,8 +262,8 @@ async def _run_session(server_url: str, oauth_auth: OAuthClientProvider) -> None
|
||||
async with streamablehttp_client(
|
||||
url=server_url,
|
||||
auth=oauth_auth,
|
||||
timeout=timedelta(seconds=30),
|
||||
sse_read_timeout=timedelta(seconds=60),
|
||||
timeout=30.0,
|
||||
sse_read_timeout=60.0,
|
||||
) as (read_stream, write_stream, _):
|
||||
async with ClientSession(read_stream, write_stream) as session:
|
||||
# Initialize the session
|
||||
|
||||
@@ -207,7 +207,7 @@ class SimpleAuthClient:
|
||||
async with sse_client(
|
||||
url=self.server_url,
|
||||
auth=oauth_auth,
|
||||
timeout=60,
|
||||
timeout=60.0,
|
||||
) as (read_stream, write_stream):
|
||||
await self._run_session(read_stream, write_stream, None)
|
||||
else:
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import logging
|
||||
from datetime import timedelta
|
||||
from typing import Any, Protocol, overload
|
||||
|
||||
import anyio.lowlevel
|
||||
@@ -113,7 +112,7 @@ class ClientSession(
|
||||
self,
|
||||
read_stream: MemoryObjectReceiveStream[SessionMessage | Exception],
|
||||
write_stream: MemoryObjectSendStream[SessionMessage],
|
||||
read_timeout_seconds: timedelta | None = None,
|
||||
read_timeout_seconds: float | None = None,
|
||||
sampling_callback: SamplingFnT | None = None,
|
||||
elicitation_callback: ElicitationFnT | None = None,
|
||||
list_roots_callback: ListRootsFnT | None = None,
|
||||
@@ -369,7 +368,7 @@ class ClientSession(
|
||||
self,
|
||||
name: str,
|
||||
arguments: dict[str, Any] | None = None,
|
||||
read_timeout_seconds: timedelta | None = None,
|
||||
read_timeout_seconds: float | None = None,
|
||||
progress_callback: ProgressFnT | None = None,
|
||||
*,
|
||||
meta: dict[str, Any] | None = None,
|
||||
|
||||
@@ -12,7 +12,6 @@ import contextlib
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from datetime import timedelta
|
||||
from types import TracebackType
|
||||
from typing import Any, TypeAlias, overload
|
||||
|
||||
@@ -41,11 +40,11 @@ class SseServerParameters(BaseModel):
|
||||
# Optional headers to include in requests.
|
||||
headers: dict[str, Any] | None = None
|
||||
|
||||
# HTTP timeout for regular operations.
|
||||
timeout: float = 5
|
||||
# HTTP timeout for regular operations (in seconds).
|
||||
timeout: float = 5.0
|
||||
|
||||
# Timeout for SSE read operations.
|
||||
sse_read_timeout: float = 60 * 5
|
||||
# Timeout for SSE read operations (in seconds).
|
||||
sse_read_timeout: float = 300.0
|
||||
|
||||
|
||||
class StreamableHttpParameters(BaseModel):
|
||||
@@ -57,11 +56,11 @@ class StreamableHttpParameters(BaseModel):
|
||||
# Optional headers to include in requests.
|
||||
headers: dict[str, Any] | None = None
|
||||
|
||||
# HTTP timeout for regular operations.
|
||||
timeout: timedelta = timedelta(seconds=30)
|
||||
# HTTP timeout for regular operations (in seconds).
|
||||
timeout: float = 30.0
|
||||
|
||||
# Timeout for SSE read operations.
|
||||
sse_read_timeout: timedelta = timedelta(seconds=60 * 5)
|
||||
# Timeout for SSE read operations (in seconds).
|
||||
sse_read_timeout: float = 300.0
|
||||
|
||||
# Close the client session when the transport closes.
|
||||
terminate_on_close: bool = True
|
||||
@@ -76,7 +75,7 @@ ServerParameters: TypeAlias = StdioServerParameters | SseServerParameters | Stre
|
||||
class ClientSessionParameters:
|
||||
"""Parameters for establishing a client session to an MCP server."""
|
||||
|
||||
read_timeout_seconds: timedelta | None = None
|
||||
read_timeout_seconds: float | None = None
|
||||
sampling_callback: SamplingFnT | None = None
|
||||
elicitation_callback: ElicitationFnT | None = None
|
||||
list_roots_callback: ListRootsFnT | None = None
|
||||
@@ -197,7 +196,7 @@ class ClientSessionGroup:
|
||||
self,
|
||||
name: str,
|
||||
arguments: dict[str, Any],
|
||||
read_timeout_seconds: timedelta | None = None,
|
||||
read_timeout_seconds: float | None = None,
|
||||
progress_callback: ProgressFnT | None = None,
|
||||
*,
|
||||
meta: dict[str, Any] | None = None,
|
||||
@@ -210,7 +209,7 @@ class ClientSessionGroup:
|
||||
name: str,
|
||||
*,
|
||||
args: dict[str, Any],
|
||||
read_timeout_seconds: timedelta | None = None,
|
||||
read_timeout_seconds: float | None = None,
|
||||
progress_callback: ProgressFnT | None = None,
|
||||
meta: dict[str, Any] | None = None,
|
||||
) -> types.CallToolResult: ...
|
||||
@@ -219,7 +218,7 @@ class ClientSessionGroup:
|
||||
self,
|
||||
name: str,
|
||||
arguments: dict[str, Any] | None = None,
|
||||
read_timeout_seconds: timedelta | None = None,
|
||||
read_timeout_seconds: float | None = None,
|
||||
progress_callback: ProgressFnT | None = None,
|
||||
*,
|
||||
meta: dict[str, Any] | None = None,
|
||||
@@ -314,8 +313,8 @@ class ClientSessionGroup:
|
||||
httpx_client = create_mcp_http_client(
|
||||
headers=server_params.headers,
|
||||
timeout=httpx.Timeout(
|
||||
server_params.timeout.total_seconds(),
|
||||
read=server_params.sse_read_timeout.total_seconds(),
|
||||
server_params.timeout,
|
||||
read=server_params.sse_read_timeout,
|
||||
),
|
||||
)
|
||||
await session_stack.enter_async_context(httpx_client)
|
||||
|
||||
@@ -31,8 +31,8 @@ def _extract_session_id_from_endpoint(endpoint_url: str) -> str | None:
|
||||
async def sse_client(
|
||||
url: str,
|
||||
headers: dict[str, Any] | None = None,
|
||||
timeout: float = 5,
|
||||
sse_read_timeout: float = 60 * 5,
|
||||
timeout: float = 5.0,
|
||||
sse_read_timeout: float = 300.0,
|
||||
httpx_client_factory: McpHttpClientFactory = create_mcp_http_client,
|
||||
auth: httpx.Auth | None = None,
|
||||
on_session_created: Callable[[str], None] | None = None,
|
||||
@@ -46,8 +46,8 @@ async def sse_client(
|
||||
Args:
|
||||
url: The SSE endpoint URL.
|
||||
headers: Optional headers to include in requests.
|
||||
timeout: HTTP timeout for regular operations.
|
||||
sse_read_timeout: Timeout for SSE read operations.
|
||||
timeout: HTTP timeout for regular operations (in seconds).
|
||||
sse_read_timeout: Timeout for SSE read operations (in seconds).
|
||||
auth: Optional HTTPX authentication handler.
|
||||
on_session_created: Optional callback invoked with the session ID when received.
|
||||
"""
|
||||
|
||||
@@ -100,8 +100,8 @@ class StreamableHTTPTransport:
|
||||
self,
|
||||
url: str,
|
||||
headers: dict[str, str] | None = None,
|
||||
timeout: float | timedelta = 30,
|
||||
sse_read_timeout: float | timedelta = 60 * 5,
|
||||
timeout: float = 30.0,
|
||||
sse_read_timeout: float = 300.0,
|
||||
auth: httpx.Auth | None = None,
|
||||
) -> None: ...
|
||||
|
||||
@@ -118,8 +118,8 @@ class StreamableHTTPTransport:
|
||||
Args:
|
||||
url: The endpoint URL.
|
||||
headers: Optional headers to include in requests.
|
||||
timeout: HTTP timeout for regular operations.
|
||||
sse_read_timeout: Timeout for SSE read operations.
|
||||
timeout: HTTP timeout for regular operations (in seconds).
|
||||
sse_read_timeout: Timeout for SSE read operations (in seconds).
|
||||
auth: Optional HTTPX authentication handler.
|
||||
"""
|
||||
# Check for deprecated parameters and issue runtime warning
|
||||
|
||||
@@ -6,7 +6,6 @@ from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncGenerator
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import timedelta
|
||||
from typing import Any
|
||||
|
||||
import anyio
|
||||
@@ -49,7 +48,7 @@ async def create_client_server_memory_streams() -> AsyncGenerator[tuple[MessageS
|
||||
@asynccontextmanager
|
||||
async def create_connected_server_and_client_session(
|
||||
server: Server[Any] | FastMCP,
|
||||
read_timeout_seconds: timedelta | None = None,
|
||||
read_timeout_seconds: float | None = None,
|
||||
sampling_callback: SamplingFnT | None = None,
|
||||
list_roots_callback: ListRootsFnT | None = None,
|
||||
logging_callback: LoggingFnT | None = None,
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from contextlib import AsyncExitStack
|
||||
from datetime import timedelta
|
||||
from types import TracebackType
|
||||
from typing import Any, Generic, Protocol, TypeVar
|
||||
|
||||
@@ -189,7 +188,7 @@ class BaseSession(
|
||||
receive_request_type: type[ReceiveRequestT],
|
||||
receive_notification_type: type[ReceiveNotificationT],
|
||||
# If none, reading will never time out
|
||||
read_timeout_seconds: timedelta | None = None,
|
||||
read_timeout_seconds: float | None = None,
|
||||
) -> None:
|
||||
self._read_stream = read_stream
|
||||
self._write_stream = write_stream
|
||||
@@ -241,7 +240,7 @@ class BaseSession(
|
||||
self,
|
||||
request: SendRequestT,
|
||||
result_type: type[ReceiveResultT],
|
||||
request_read_timeout_seconds: timedelta | None = None,
|
||||
request_read_timeout_seconds: float | None = None,
|
||||
metadata: MessageMetadata = None,
|
||||
progress_callback: ProgressFnT | None = None,
|
||||
) -> ReceiveResultT:
|
||||
@@ -283,9 +282,9 @@ class BaseSession(
|
||||
# request read timeout takes precedence over session read timeout
|
||||
timeout = None
|
||||
if request_read_timeout_seconds is not None: # pragma: no cover
|
||||
timeout = request_read_timeout_seconds.total_seconds()
|
||||
timeout = request_read_timeout_seconds
|
||||
elif self._session_read_timeout_seconds is not None: # pragma: no cover
|
||||
timeout = self._session_read_timeout_seconds.total_seconds()
|
||||
timeout = self._session_read_timeout_seconds
|
||||
|
||||
try:
|
||||
with anyio.fail_after(timeout):
|
||||
|
||||
@@ -273,7 +273,7 @@ class TestClientSessionGroup:
|
||||
"mcp.client.session_group.mcp.stdio_client",
|
||||
),
|
||||
(
|
||||
SseServerParameters(url="http://test.com/sse", timeout=10),
|
||||
SseServerParameters(url="http://test.com/sse", timeout=10.0),
|
||||
"sse",
|
||||
"mcp.client.session_group.sse_client",
|
||||
), # url, headers, timeout, sse_read_timeout
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
"""Test to reproduce issue #88: Random error thrown on response."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
from datetime import timedelta
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
@@ -93,11 +92,9 @@ async def test_notification_validation_error(tmp_path: Path):
|
||||
assert not slow_request_lock.is_set()
|
||||
|
||||
# Second call should timeout (slow operation with minimal timeout)
|
||||
# Use 10ms timeout to trigger quickly without waiting
|
||||
# Use very small timeout to trigger quickly without waiting
|
||||
with pytest.raises(McpError) as exc_info:
|
||||
await session.call_tool(
|
||||
"slow", read_timeout_seconds=timedelta(microseconds=1)
|
||||
) # artificial timeout that always fails
|
||||
await session.call_tool("slow", read_timeout_seconds=0.000001) # artificial timeout that always fails
|
||||
assert "Timed out while waiting" in str(exc_info.value)
|
||||
|
||||
# release the slow request not to have hanging process
|
||||
|
||||
@@ -270,12 +270,10 @@ async def test_response_id_non_numeric_string_no_match():
|
||||
async def make_request(client_session: ClientSession):
|
||||
try:
|
||||
# Use a short timeout since we expect this to fail
|
||||
from datetime import timedelta
|
||||
|
||||
await client_session.send_request(
|
||||
ClientRequest(types.PingRequest()),
|
||||
types.EmptyResult,
|
||||
request_read_timeout_seconds=timedelta(seconds=0.5),
|
||||
request_read_timeout_seconds=0.5,
|
||||
)
|
||||
pytest.fail("Expected timeout") # pragma: no cover
|
||||
except McpError as e:
|
||||
|
||||
@@ -9,7 +9,6 @@ import multiprocessing
|
||||
import socket
|
||||
import time
|
||||
from collections.abc import Generator
|
||||
from datetime import timedelta
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
@@ -2370,8 +2369,8 @@ async def test_streamable_http_transport_deprecated_params_ignored(basic_server:
|
||||
transport = StreamableHTTPTransport( # pyright: ignore[reportDeprecated]
|
||||
url=f"{basic_server_url}/mcp",
|
||||
headers={"X-Should-Be-Ignored": "ignored"},
|
||||
timeout=999,
|
||||
sse_read_timeout=timedelta(seconds=999),
|
||||
timeout=999.0,
|
||||
sse_read_timeout=999.0,
|
||||
auth=None,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user