feat: add ClientRequestContext type alias for client-side handlers (#1989)
This commit is contained in:
committed by
GitHub
parent
74f41ff1ae
commit
4fc49c62bd
@@ -38,9 +38,9 @@ from mcp.client.auth.extensions.client_credentials import (
|
||||
PrivateKeyJWTOAuthProvider,
|
||||
SignedJWTParameters,
|
||||
)
|
||||
from mcp.client.context import ClientRequestContext
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken
|
||||
from mcp.shared.context import RequestContext
|
||||
|
||||
# Set up logging to stderr (stdout is for conformance test output)
|
||||
logging.basicConfig(
|
||||
@@ -187,7 +187,7 @@ async def run_sse_retry(server_url: str) -> None:
|
||||
|
||||
|
||||
async def default_elicitation_callback(
|
||||
context: RequestContext[ClientSession],
|
||||
context: ClientRequestContext,
|
||||
params: types.ElicitRequestParams,
|
||||
) -> types.ElicitResult | types.ErrorData:
|
||||
"""Accept elicitation and apply defaults from the schema (SEP-1034)."""
|
||||
|
||||
@@ -2120,8 +2120,8 @@ import asyncio
|
||||
import os
|
||||
|
||||
from mcp import ClientSession, StdioServerParameters, types
|
||||
from mcp.client.context import ClientRequestContext
|
||||
from mcp.client.stdio import stdio_client
|
||||
from mcp.shared.context import RequestContext
|
||||
|
||||
# Create server parameters for stdio connection
|
||||
server_params = StdioServerParameters(
|
||||
@@ -2133,7 +2133,7 @@ server_params = StdioServerParameters(
|
||||
|
||||
# Optional: create a sampling callback
|
||||
async def handle_sampling_message(
|
||||
context: RequestContext[ClientSession, None], params: types.CreateMessageRequestParams
|
||||
context: ClientRequestContext, params: types.CreateMessageRequestParams
|
||||
) -> types.CreateMessageResult:
|
||||
print(f"Sampling request: {params.messages}")
|
||||
return types.CreateMessageResult(
|
||||
|
||||
+2
-2
@@ -2121,8 +2121,8 @@ import asyncio
|
||||
import os
|
||||
|
||||
from mcp import ClientSession, StdioServerParameters, types
|
||||
from mcp.client.context import ClientRequestContext
|
||||
from mcp.client.stdio import stdio_client
|
||||
from mcp.shared.context import RequestContext
|
||||
|
||||
# Create server parameters for stdio connection
|
||||
server_params = StdioServerParameters(
|
||||
@@ -2134,7 +2134,7 @@ server_params = StdioServerParameters(
|
||||
|
||||
# Optional: create a sampling callback
|
||||
async def handle_sampling_message(
|
||||
context: RequestContext[ClientSession], params: types.CreateMessageRequestParams
|
||||
context: ClientRequestContext, params: types.CreateMessageRequestParams
|
||||
) -> types.CreateMessageResult:
|
||||
print(f"Sampling request: {params.messages}")
|
||||
return types.CreateMessageResult(
|
||||
|
||||
+9
-8
@@ -387,6 +387,7 @@ The `RequestContext` class has been split to separate shared fields from server-
|
||||
**Before (v1):**
|
||||
|
||||
```python
|
||||
from mcp.client.session import ClientSession
|
||||
from mcp.shared.context import RequestContext, LifespanContextT, RequestT
|
||||
from mcp.shared.progress import ProgressContext
|
||||
|
||||
@@ -400,19 +401,19 @@ progress_ctx: ProgressContext[SendRequestT, SendNotificationT, SendResultT, Rece
|
||||
**After (v2):**
|
||||
|
||||
```python
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.client.context import ClientRequestContext
|
||||
from mcp.client.session import ClientSession
|
||||
from mcp.server.context import ServerRequestContext, LifespanContextT, RequestT
|
||||
from mcp.shared.progress import ProgressContext
|
||||
|
||||
# RequestContext with 1 type parameter
|
||||
ctx: RequestContext[ClientSession]
|
||||
# For client-side context (sampling, elicitation, list_roots callbacks)
|
||||
ctx: ClientRequestContext
|
||||
|
||||
# For server-specific context with lifespan and request types
|
||||
server_ctx: ServerRequestContext[LifespanContextT, RequestT]
|
||||
|
||||
# ProgressContext with 1 type parameter
|
||||
progress_ctx: ProgressContext[ClientSession]
|
||||
|
||||
# For server-specific context with lifespan and request types
|
||||
from mcp.server.context import ServerRequestContext, LifespanContextT, RequestT
|
||||
|
||||
server_ctx: ServerRequestContext[LifespanContextT, RequestT]
|
||||
```
|
||||
|
||||
### Resource URI type changed from `AnyUrl` to `str`
|
||||
|
||||
+3
-3
@@ -10,8 +10,8 @@ import asyncio
|
||||
|
||||
import click
|
||||
from mcp import ClientSession
|
||||
from mcp.client.context import ClientRequestContext
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.types import (
|
||||
CallToolResult,
|
||||
CreateMessageRequestParams,
|
||||
@@ -23,7 +23,7 @@ from mcp.types import (
|
||||
|
||||
|
||||
async def elicitation_callback(
|
||||
context: RequestContext[ClientSession],
|
||||
context: ClientRequestContext,
|
||||
params: ElicitRequestParams,
|
||||
) -> ElicitResult:
|
||||
"""Handle elicitation requests from the server."""
|
||||
@@ -38,7 +38,7 @@ async def elicitation_callback(
|
||||
|
||||
|
||||
async def sampling_callback(
|
||||
context: RequestContext[ClientSession],
|
||||
context: ClientRequestContext,
|
||||
params: CreateMessageRequestParams,
|
||||
) -> CreateMessageResult:
|
||||
"""Handle sampling requests from the server."""
|
||||
|
||||
@@ -6,8 +6,8 @@ import asyncio
|
||||
import os
|
||||
|
||||
from mcp import ClientSession, StdioServerParameters, types
|
||||
from mcp.client.context import ClientRequestContext
|
||||
from mcp.client.stdio import stdio_client
|
||||
from mcp.shared.context import RequestContext
|
||||
|
||||
# Create server parameters for stdio connection
|
||||
server_params = StdioServerParameters(
|
||||
@@ -19,7 +19,7 @@ server_params = StdioServerParameters(
|
||||
|
||||
# Optional: create a sampling callback
|
||||
async def handle_sampling_message(
|
||||
context: RequestContext[ClientSession], params: types.CreateMessageRequestParams
|
||||
context: ClientRequestContext, params: types.CreateMessageRequestParams
|
||||
) -> types.CreateMessageResult:
|
||||
print(f"Sampling request: {params.messages}")
|
||||
return types.CreateMessageResult(
|
||||
|
||||
@@ -31,14 +31,14 @@ from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from mcp import ClientSession, types
|
||||
from mcp.client.context import ClientRequestContext
|
||||
from mcp.client.sse import sse_client
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared.exceptions import MCPError, UrlElicitationRequiredError
|
||||
from mcp.types import URL_ELICITATION_REQUIRED
|
||||
|
||||
|
||||
async def handle_elicitation(
|
||||
context: RequestContext[ClientSession],
|
||||
context: ClientRequestContext,
|
||||
params: types.ElicitRequestParams,
|
||||
) -> types.ElicitResult | types.ErrorData:
|
||||
"""Handle elicitation requests from the server.
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
from mcp.client._transport import Transport
|
||||
from mcp.client.client import Client
|
||||
from mcp.client.context import ClientRequestContext
|
||||
from mcp.client.session import ClientSession
|
||||
|
||||
__all__ = ["Client", "ClientSession", "Transport"]
|
||||
__all__ = ["Client", "ClientRequestContext", "ClientSession", "Transport"]
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
"""Request context for MCP client handlers."""
|
||||
|
||||
from mcp.client.session import ClientSession
|
||||
from mcp.shared._context import RequestContext
|
||||
|
||||
ClientRequestContext = RequestContext[ClientSession]
|
||||
"""Context for handling incoming requests in a client session.
|
||||
|
||||
This context is passed to client-side callbacks (sampling, elicitation, list_roots) when the server sends requests
|
||||
to the client.
|
||||
|
||||
Attributes:
|
||||
request_id: The unique identifier for this request.
|
||||
meta: Optional metadata associated with the request.
|
||||
session: The client session handling this request.
|
||||
"""
|
||||
@@ -19,7 +19,7 @@ from typing import TYPE_CHECKING, Protocol
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import mcp.types as types
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.shared.session import RequestResponder
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
@@ -10,7 +10,7 @@ from pydantic import TypeAdapter
|
||||
import mcp.types as types
|
||||
from mcp.client.experimental import ExperimentalClientFeatures
|
||||
from mcp.client.experimental.task_handlers import ExperimentalTaskHandlers
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.shared.message import SessionMessage
|
||||
from mcp.shared.session import BaseSession, ProgressFnT, RequestResponder
|
||||
from mcp.shared.version import SUPPORTED_PROTOCOL_VERSIONS
|
||||
|
||||
@@ -7,7 +7,7 @@ from typing_extensions import TypeVar
|
||||
|
||||
from mcp.server.experimental.request_context import Experimental
|
||||
from mcp.server.session import ServerSession
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.shared.message import CloseSSEStreamCallback
|
||||
|
||||
LifespanContextT = TypeVar("LifespanContextT")
|
||||
|
||||
@@ -5,7 +5,7 @@ from typing import Generic
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from mcp.shared.context import RequestContext, SessionT
|
||||
from mcp.shared._context import RequestContext, SessionT
|
||||
from mcp.types import ProgressToken
|
||||
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ from mcp import Client
|
||||
from mcp.client.session import ClientSession
|
||||
from mcp.server.mcpserver import MCPServer
|
||||
from mcp.server.mcpserver.server import Context
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.types import ListRootsResult, Root, TextContent
|
||||
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ import pytest
|
||||
from mcp import Client
|
||||
from mcp.client.session import ClientSession
|
||||
from mcp.server.mcpserver import MCPServer
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.types import (
|
||||
CreateMessageRequestParams,
|
||||
CreateMessageResult,
|
||||
|
||||
@@ -5,7 +5,7 @@ import pytest
|
||||
|
||||
import mcp.types as types
|
||||
from mcp.client.session import DEFAULT_CLIENT_INFO, ClientSession
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.shared.message import SessionMessage
|
||||
from mcp.shared.session import RequestResponder
|
||||
from mcp.shared.version import SUPPORTED_PROTOCOL_VERSIONS
|
||||
|
||||
@@ -7,7 +7,7 @@ import mcp.types as types
|
||||
from mcp import ClientCapabilities
|
||||
from mcp.client.experimental.task_handlers import ExperimentalTaskHandlers
|
||||
from mcp.client.session import ClientSession
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.shared.message import SessionMessage
|
||||
from mcp.types import (
|
||||
LATEST_PROTOCOL_VERSION,
|
||||
|
||||
@@ -22,7 +22,7 @@ from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStre
|
||||
import mcp.types as types
|
||||
from mcp.client.experimental.task_handlers import ExperimentalTaskHandlers
|
||||
from mcp.client.session import ClientSession
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.shared.experimental.tasks.in_memory_task_store import InMemoryTaskStore
|
||||
from mcp.shared.message import SessionMessage
|
||||
from mcp.shared.session import RequestResponder
|
||||
|
||||
@@ -20,7 +20,7 @@ from mcp.client.session import ClientSession
|
||||
from mcp.server import Server
|
||||
from mcp.server.experimental.task_context import ServerTaskContext
|
||||
from mcp.server.lowlevel import NotificationOptions
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.shared.experimental.tasks.helpers import is_terminal
|
||||
from mcp.shared.experimental.tasks.in_memory_task_store import InMemoryTaskStore
|
||||
from mcp.shared.message import SessionMessage
|
||||
|
||||
@@ -9,7 +9,7 @@ from mcp import Client, types
|
||||
from mcp.client.session import ClientSession, ElicitationFnT
|
||||
from mcp.server.mcpserver import Context, MCPServer
|
||||
from mcp.server.session import ServerSession
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.types import ElicitRequestParams, ElicitResult, TextContent
|
||||
|
||||
|
||||
|
||||
@@ -33,7 +33,7 @@ from examples.snippets.servers import (
|
||||
from mcp.client.session import ClientSession
|
||||
from mcp.client.sse import sse_client
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.shared.session import RequestResponder
|
||||
from mcp.types import (
|
||||
ClientResult,
|
||||
|
||||
@@ -9,7 +9,7 @@ from mcp.client.session import ClientSession
|
||||
from mcp.server.elicitation import CancelledElicitation, DeclinedElicitation, elicit_url
|
||||
from mcp.server.mcpserver import Context, MCPServer
|
||||
from mcp.server.session import ServerSession
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.types import ElicitRequestParams, ElicitResult, TextContent
|
||||
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ from mcp.server import Server
|
||||
from mcp.server.lowlevel import NotificationOptions
|
||||
from mcp.server.models import InitializationOptions
|
||||
from mcp.server.session import ServerSession
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.shared.message import SessionMessage
|
||||
from mcp.shared.progress import progress
|
||||
from mcp.shared.session import RequestResponder
|
||||
|
||||
@@ -43,12 +43,12 @@ from mcp.server.streamable_http import (
|
||||
)
|
||||
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
|
||||
from mcp.server.transport_security import TransportSecuritySettings
|
||||
from mcp.shared._context import RequestContext
|
||||
from mcp.shared._httpx_utils import (
|
||||
MCP_DEFAULT_SSE_READ_TIMEOUT,
|
||||
MCP_DEFAULT_TIMEOUT,
|
||||
create_mcp_http_client,
|
||||
)
|
||||
from mcp.shared.context import RequestContext
|
||||
from mcp.shared.message import ClientMessageMetadata, ServerMessageMetadata, SessionMessage
|
||||
from mcp.shared.session import RequestResponder
|
||||
from mcp.types import InitializeResult, JSONRPCRequest, TextContent, TextResourceContents, Tool
|
||||
|
||||
Reference in New Issue
Block a user