From dbee2f000c2a594b6c6b312a9b8aab786c982b90 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 30 Sep 2026 12:44:24 +0200 Subject: [PATCH 01/42] Updated mcp to 2.x (websocket optional is part of the core now) --- .../packages/core/agent_framework/_agents.py | 64 +++---- python/packages/core/agent_framework/_mcp.py | 2 +- .../packages/core/agent_framework/_skills.py | 20 +- python/packages/core/pyproject.toml | 2 +- .../packages/foundry_hosting/pyproject.toml | 2 +- .../_agent_tool.py | 24 +-- .../_conversion.py | 20 +- .../_workflow_tool.py | 35 ++-- python/packages/hosting-mcp/pyproject.toml | 2 +- .../tests/hosting_mcp/test_agent_tool.py | 78 ++++++-- .../tests/hosting_mcp/test_conversion.py | 4 +- .../tests/hosting_mcp/test_integration.py | 20 +- .../tests/hosting_mcp/test_workflow_tool.py | 52 +++-- python/packages/lab/pyproject.toml | 2 +- python/pyproject.toml | 4 +- python/samples/02-agents/mcp/README.md | 3 +- .../02-agents/mcp/mcp_long_running_task.py | 181 ------------------ .../mcp/mcp_progressive_disclosure.py | 26 +-- .../02-agents/mcp/mcp_sampling_approval.py | 4 +- ...foundry_chat_client_with_toolbox_skills.py | 12 +- .../skills/mcp_based_skill/mcp_based_skill.py | 2 +- python/samples/04-hosting/mcp/agent_app.py | 33 ++-- python/samples/04-hosting/mcp/fastmcp_app.py | 21 +- python/samples/04-hosting/mcp/manual_app.py | 52 ++--- python/samples/04-hosting/mcp/pyproject.toml | 2 +- python/samples/04-hosting/mcp/session_app.py | 45 +++-- python/samples/04-hosting/mcp/workflow_app.py | 29 +-- .../tests/test_dependency_bounds_runtime.py | 4 +- python/uv.lock | 91 ++++----- 29 files changed, 361 insertions(+), 475 deletions(-) delete mode 100644 python/samples/02-agents/mcp/mcp_long_running_task.py diff --git a/python/packages/core/agent_framework/_agents.py b/python/packages/core/agent_framework/_agents.py index 5215c14880a..c4597013ad2 100644 --- a/python/packages/core/agent_framework/_agents.py +++ b/python/packages/core/agent_framework/_agents.py @@ -1881,8 +1881,9 @@ def as_mcp_server( """ try: from mcp import types + from mcp.server import ServerRequestContext from mcp.server.lowlevel import Server - from mcp.shared.exceptions import McpError + from mcp.shared.exceptions import MCPError except ModuleNotFoundError as exc: raise ModuleNotFoundError( "`mcp` is required to use `Agent.as_mcp_server()`. Please install `mcp`." @@ -1899,47 +1900,40 @@ def as_mcp_server( if kwargs: server_args.update(kwargs) - server: Server[Any] = Server(**server_args) - agent_tool = self.as_tool(name=self._get_agent_name()) - async def _log(level: types.LoggingLevel, data: Any) -> None: - """Log a message to the server and logger.""" + def _log(level: types.LoggingLevel, data: Any) -> None: + """Log a message to logger.""" # Log to the local logger logger.log(LOG_LEVEL_MAPPING[level], data) - if server and server.request_context and server.request_context.session: - try: - await server.request_context.session.send_log_message(level=level, data=data) - except Exception as e: - logger.error("Failed to send log message to server: %s", e) - - @server.list_tools() - async def _list_tools() -> list[types.Tool]: + + async def _list_tools( # ruff:ignore[unused-async] + _ctx: ServerRequestContext[dict[str, Any]], _params: types.PaginatedRequestParams | None + ) -> types.ListToolsResult: """List all tools in the agent.""" schema = agent_tool.parameters() tool = types.Tool( name=agent_tool.name, description=agent_tool.description, - inputSchema=schema, + input_schema=schema, ) - await _log(level="debug", data=f"Agent tool: {agent_tool}") - return [tool] + _log(level="debug", data=f"Agent tool: {agent_tool}") + return types.ListToolsResult(tools=[tool]) - @server.call_tool() async def _call_tool( - name: str, arguments: dict[str, Any] - ) -> Sequence[types.TextContent | types.ImageContent | types.AudioContent | types.EmbeddedResource]: + _ctx: ServerRequestContext[dict[str, Any]], params: types.CallToolRequestParams + ) -> types.CallToolResult: """Call a tool in the agent.""" - await _log(level="debug", data=f"Calling tool with args: {arguments}") + arguments = params.arguments or {} + name = params.name + _log(level="debug", data=f"Calling tool with args: {arguments}") if name != agent_tool.name: - raise McpError( - error=types.ErrorData( - code=types.INTERNAL_ERROR, - message=f"Tool {name} not found", - ), + raise MCPError( + code=types.INTERNAL_ERROR, + message=f"Tool {name} not found", ) # Create an instance of the input model with the arguments @@ -1949,17 +1943,15 @@ async def _call_tool( ) result = await agent_tool.invoke(arguments=args_instance) except Exception as e: - raise McpError( - error=types.ErrorData( - code=types.INTERNAL_ERROR, - message=f"Error calling tool {name}: {e}", - ), + raise MCPError( + code=types.INTERNAL_ERROR, + message=f"Error calling tool {name}: {e}", ) from e # Convert result to MCP content. # Currently only text items are forwarded over MCP; rich content # (images, audio) is not yet supported in the MCP server path. - mcp_content: list[types.TextContent | types.ImageContent | types.EmbeddedResource] = [] + mcp_content: list[types.ContentBlock] = [] for c in result: if c.type == "text" and c.text: mcp_content.append(types.TextContent(type="text", text=c.text)) @@ -1968,15 +1960,9 @@ async def _call_tool( "MCP server does not yet forward rich content (images, audio) " "in tool results. Rich content items will be omitted." ) - return mcp_content or [types.TextContent(type="text", text="")] - - @server.set_logging_level() - async def _set_logging_level(level: types.LoggingLevel) -> None: - """Set the logging level for the server.""" - logger.setLevel(LOG_LEVEL_MAPPING[level]) - # emit this log with the new minimum level - await _log(level=level, data=f"Log level set to {level}") + return types.CallToolResult(content=mcp_content) + server: Server[Any] = Server(**server_args, on_list_tools=_list_tools, on_call_tool=_call_tool) return server def _get_agent_name(self) -> str: diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 56f68cf3c3b..a89f3f26cb9 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -4539,7 +4539,7 @@ def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: reason = f"The optional dependency `{missing_name}` is not installed." raise ModuleNotFoundError( f"`MCPWebsocketTool` requires websocket transport support. {reason} " - "Please install `mcp[ws]` and update your dependencies." + "Please install `mcp` and update your dependencies." ) from ex # Support MCP releases from before and after the transport gained its deprecation marker. diff --git a/python/packages/core/agent_framework/_skills.py b/python/packages/core/agent_framework/_skills.py index 3bef5b39424..27073038b04 100644 --- a/python/packages/core/agent_framework/_skills.py +++ b/python/packages/core/agent_framework/_skills.py @@ -96,7 +96,6 @@ if TYPE_CHECKING: from mcp.client.session import ClientSession from mcp.types import ReadResourceResult - from pydantic import AnyUrl from ._agents import SupportsAgentRun from ._sessions import AgentSession, SessionContext @@ -4440,13 +4439,6 @@ async def get_skills(self, context: SkillsSourceContext) -> list[Skill]: # region MCP Skills -def _mcp_any_url(uri: str) -> AnyUrl: - """Convert a string URI to a :class:`pydantic.AnyUrl` for MCP client calls.""" - from pydantic import AnyUrl as _AnyUrl - - return _AnyUrl(uri) - - def _is_mcp_resource_not_found(ex: Exception) -> bool: """Return ``True`` when *ex* is an :class:`McpError` indicating a missing resource. @@ -4465,7 +4457,7 @@ def _is_mcp_resource_not_found(ex: Exception) -> bool: token or crashing server is not silently mistaken for "the server has no skills." """ - from mcp.shared.exceptions import McpError as _McpError + from mcp.shared.exceptions import MCPError as _McpError if not isinstance(ex, _McpError): return False @@ -4506,7 +4498,7 @@ def _mcp_first_blob(result: ReadResourceResult) -> tuple[bytes, str | None] | No # binascii.Error (invalid base64) subclasses ValueError. logger.warning("Failed to base64-decode blob resource content from the MCP server.", exc_info=True) return None - return data, content.mimeType + return data, content.mime_type return None @@ -4768,7 +4760,7 @@ async def get_content(self) -> str: if self._content is not None: return self._content - result = await self._session_provider().read_resource(_mcp_any_url(self._skill_md_uri)) + result = await self._session_provider().read_resource(self._skill_md_uri) text = _mcp_join_text(result) if not text: raise ValueError(f"The MCP server returned no text content for SKILL.md resource '{self._skill_md_uri}'.") @@ -4798,7 +4790,7 @@ async def get_resource(self, name: str) -> SkillResource | None: uri = self._skill_root_uri + normalized try: - result = await self._session_provider().read_resource(_mcp_any_url(uri)) + result = await self._session_provider().read_resource(uri) except Exception as ex: if _is_mcp_resource_not_found(ex): logger.debug("MCP resource '%s' not available: %s", uri, ex) @@ -5162,7 +5154,7 @@ async def _download(self, entry: _McpSkillIndexEntry) -> tuple[bytes, str | None while reading the archive resource is re-raised. """ try: - result = await self._session_provider().read_resource(_mcp_any_url(cast(str, entry.url))) + result = await self._session_provider().read_resource(cast(str, entry.url)) except Exception as ex: if _is_mcp_resource_not_found(ex): logger.debug("Archive resource '%s' for skill '%s' not available: %s", entry.url, entry.name, ex) @@ -5535,7 +5527,7 @@ async def _try_read_index(self) -> _McpSkillIndex | None: absent, empty, or malformed. """ try: - result = await self._session_provider().read_resource(_mcp_any_url(self._INDEX_URI)) + result = await self._session_provider().read_resource(self._INDEX_URI) except Exception as ex: if _is_mcp_resource_not_found(ex): logger.debug("No skill://index.json resource available on MCP server: %s", ex) diff --git a/python/packages/core/pyproject.toml b/python/packages/core/pyproject.toml index 847a29cdf84..4209795772e 100644 --- a/python/packages/core/pyproject.toml +++ b/python/packages/core/pyproject.toml @@ -34,7 +34,7 @@ dependencies = [ [project.optional-dependencies] all = [ - "mcp>=1.24.0,<2", + "mcp>=2.2.0,<3", "agent-framework-a2a>=1.0.0b261002,<2", "agent-framework-ag-ui>=1.5.0,<2", "agent-framework-anthropic>=1.0.0b261002,<2", diff --git a/python/packages/foundry_hosting/pyproject.toml b/python/packages/foundry_hosting/pyproject.toml index aa7af15075c..a00a555027c 100644 --- a/python/packages/foundry_hosting/pyproject.toml +++ b/python/packages/foundry_hosting/pyproject.toml @@ -28,7 +28,7 @@ dependencies = [ "azure-ai-agentserver-responses>=2.2.0b1,<3", "azure-ai-agentserver-invocations>=1.1.0,<2", "httpx>=0.28,<1", - "mcp>=1.24.0,<2", + "mcp>=2.2.0,<3", ] [tool.uv] diff --git a/python/packages/hosting-mcp/agent_framework_hosting_mcp/_agent_tool.py b/python/packages/hosting-mcp/agent_framework_hosting_mcp/_agent_tool.py index 8905bbda4fa..4a551e5a5ba 100644 --- a/python/packages/hosting-mcp/agent_framework_hosting_mcp/_agent_tool.py +++ b/python/packages/hosting-mcp/agent_framework_hosting_mcp/_agent_tool.py @@ -11,7 +11,7 @@ from agent_framework import AgentResponse, Message, SupportsAgentRun from agent_framework._telemetry import mark_feature_used from agent_framework_hosting import AgentRunArgs, AgentState -from mcp import types +from mcp import MCPError, types from ._conversion import mcp_from_run, mcp_to_run from ._feature_usage import FeatureIndex @@ -86,11 +86,11 @@ def __init__( name for name in (*self.parameters, *self.chat_option_parameters) if name in required_names ) - async def list_tools(self) -> list[types.Tool]: + async def list_tools(self) -> types.ListToolsResult: """Return the native MCP tool definition for the target agent.""" mark_feature_used(FeatureIndex.HOSTING_MCP) target = await self.state.get_target() - return [self._tool_for_target(target)] + return types.ListToolsResult(tools=[self._tool_for_target(target)]) def _tool_for_target(self, target: AgentT) -> types.Tool: """Create the native MCP tool definition for a resolved target.""" @@ -111,7 +111,7 @@ def _tool_for_target(self, target: AgentT) -> types.Tool: return types.Tool( name=tool_name, description=self._description if self._description is not None else target.description or "", - inputSchema={ + input_schema={ "type": "object", "properties": properties, "required": [self.argument_name, *self.required_parameters], @@ -135,7 +135,7 @@ async def call_tool( self, name: str, arguments: Mapping[str, Any] | None, - ) -> list[types.ContentBlock]: + ) -> types.CallToolResult: """Run the target agent for a native MCP ``call_tool`` handler. Args: @@ -143,16 +143,16 @@ async def call_tool( arguments: Native MCP tool arguments. Returns: - Native MCP content blocks for the completed tool result. + Native MCP ``CallToolResult`` for the completed tool result. Raises: - ValueError: If the tool name or configured session id is invalid. + MCPError: If the tool name or configured session id is invalid. """ mark_feature_used(FeatureIndex.HOSTING_MCP) target = await self.state.get_target() tool = self._tool_for_target(target) if name != tool.name: - raise ValueError(f"Unknown MCP tool: {name}") + raise MCPError(types.INVALID_PARAMS, f"Unknown MCP tool: {name}") run = self.mcp_to_run(arguments) if self.session_id_parameter is None: @@ -164,11 +164,13 @@ async def call_tool( stream=False, ), ) - return self.mcp_from_run(result) + return types.CallToolResult(content=self.mcp_from_run(result)) session_id = arguments.get(self.session_id_parameter) if arguments else None if not isinstance(session_id, str) or not session_id: - raise ValueError(f"MCP tool argument '{self.session_id_parameter}' must be a non-empty string.") + raise MCPError( + types.INVALID_PARAMS, f"MCP tool argument '{self.session_id_parameter}' must be a non-empty string." + ) session = await self.state.get_or_create_session(session_id) result = cast( "AgentResponse[Any]", @@ -180,4 +182,4 @@ async def call_tool( ), ) await self.state.set_session(session_id, session) - return self.mcp_from_run(result) + return types.CallToolResult(content=self.mcp_from_run(result)) diff --git a/python/packages/hosting-mcp/agent_framework_hosting_mcp/_conversion.py b/python/packages/hosting-mcp/agent_framework_hosting_mcp/_conversion.py index b91d30570b7..91d8381e8c2 100644 --- a/python/packages/hosting-mcp/agent_framework_hosting_mcp/_conversion.py +++ b/python/packages/hosting-mcp/agent_framework_hosting_mcp/_conversion.py @@ -11,7 +11,7 @@ from agent_framework import AgentResponse, ChatOptions, Content, Message from agent_framework_hosting import AgentRunArgs -from mcp import types +from mcp import MCPError, types logger = logging.getLogger("agent_framework.hosting.mcp") @@ -39,14 +39,14 @@ def mcp_to_run( Arguments corresponding to ``Agent.run(...)``. Raises: - ValueError: If the selected argument is missing or is not a string. + MCPError: If the selected argument is missing or is not a string. """ if arguments is None or argument_name not in arguments: - raise ValueError(f"MCP tool arguments must include a '{argument_name}' string.") + raise MCPError(types.INVALID_PARAMS, f"MCP tool arguments must include a '{argument_name}' string.") message_value = arguments[argument_name] if not isinstance(message_value, str): - raise ValueError(f"MCP tool argument '{argument_name}' must be a string.") + raise MCPError(types.INVALID_PARAMS, f"MCP tool argument '{argument_name}' must be a string.") options = {name: arguments[name] for name in chat_option_arguments if name in arguments} return AgentRunArgs( @@ -98,8 +98,8 @@ def mcp_from_run( types.ResourceLink( type="resource_link", name=name, - uri=content.uri, # pyright: ignore[reportArgumentType] - mimeType=content.media_type, + uri=content.uri, + mime_type=content.media_type, _meta=metadata, ) ) @@ -117,7 +117,7 @@ def mcp_from_run( types.ImageContent( type="image", data=encoded, - mimeType=content.media_type, + mime_type=content.media_type, _meta=metadata, ) ) @@ -126,7 +126,7 @@ def mcp_from_run( types.AudioContent( type="audio", data=encoded, - mimeType=content.media_type, + mime_type=content.media_type, _meta=metadata, ) ) @@ -140,9 +140,9 @@ def mcp_from_run( types.EmbeddedResource( type="resource", resource=types.BlobResourceContents( - uri=resource_uri, # pyright: ignore[reportArgumentType] + uri=resource_uri, blob=encoded, - mimeType=content.media_type, + mime_type=content.media_type, ), _meta=metadata, ) diff --git a/python/packages/hosting-mcp/agent_framework_hosting_mcp/_workflow_tool.py b/python/packages/hosting-mcp/agent_framework_hosting_mcp/_workflow_tool.py index b3018e5a31f..a6325041fee 100644 --- a/python/packages/hosting-mcp/agent_framework_hosting_mcp/_workflow_tool.py +++ b/python/packages/hosting-mcp/agent_framework_hosting_mcp/_workflow_tool.py @@ -11,8 +11,8 @@ from agent_framework import AgentResponse, Message, Workflow, WorkflowRunResult from agent_framework._telemetry import mark_feature_used from agent_framework_hosting import WorkflowState -from mcp import types -from pydantic import TypeAdapter +from mcp import MCPError, types +from pydantic import TypeAdapter, ValidationError from ._conversion import mcp_from_run from ._feature_usage import FeatureIndex @@ -53,11 +53,11 @@ def __init__( self._description = description self.argument_name = argument_name - async def list_tools(self) -> list[types.Tool]: + async def list_tools(self) -> types.ListToolsResult: """Return the native MCP tool definition for the target workflow.""" mark_feature_used(FeatureIndex.HOSTING_MCP) workflow = await self.state.get_target() - return [self._tool_for_workflow(workflow)] + return types.ListToolsResult(tools=[self._tool_for_workflow(workflow)]) def _tool_for_workflow(self, workflow: WorkflowT) -> types.Tool: tool_name = self._name @@ -77,7 +77,7 @@ def _tool_for_workflow(self, workflow: WorkflowT) -> types.Tool: return types.Tool( name=tool_name, description=self._description if self._description is not None else workflow.description or "", - inputSchema=input_schema, + input_schema=input_schema, ) def _input_adapter(self, workflow: WorkflowT) -> TypeAdapter[Any]: @@ -91,11 +91,22 @@ def _input_adapter(self, workflow: WorkflowT) -> TypeAdapter[Any]: def _workflow_input(self, workflow: WorkflowT, arguments: Mapping[str, Any] | None) -> Any: input_adapter = self._input_adapter(workflow) input_schema = input_adapter.json_schema() + if input_schema.get("type") == "object": - return input_adapter.validate_python(dict(arguments or {})) - if arguments is None or self.argument_name not in arguments: - raise ValueError(f"MCP tool arguments must include '{self.argument_name}'.") - return input_adapter.validate_python(arguments[self.argument_name]) + input_value = dict(arguments or {}) + else: + if arguments is None or self.argument_name not in arguments: + raise MCPError( + types.INVALID_PARAMS, + f"MCP tool arguments must include '{self.argument_name}'.", + ) + input_value = arguments[self.argument_name] + + try: + return input_adapter.validate_python(input_value) + except ValidationError as exc: + # We raise MCPError only for validation errors which pertain to the input + raise MCPError(types.INVALID_PARAMS, str(exc)) from exc def mcp_from_run(self, result: WorkflowRunResult) -> list[types.ContentBlock]: """Convert completed workflow outputs into native MCP content blocks.""" @@ -124,7 +135,7 @@ async def call_tool( self, name: str, arguments: Mapping[str, Any] | None, - ) -> list[types.ContentBlock]: + ) -> types.CallToolResult: """Run the target workflow for a native MCP ``call_tool`` handler. Args: @@ -141,7 +152,7 @@ async def call_tool( workflow = await self.state.get_target() tool = self._tool_for_workflow(workflow) if name != tool.name: - raise ValueError(f"Unknown MCP tool: {name}") + raise MCPError(types.INVALID_PARAMS, f"Unknown MCP tool: {name}") result = await workflow.run(self._workflow_input(workflow, arguments), stream=False) - return self.mcp_from_run(result) + return types.CallToolResult(content=self.mcp_from_run(result)) diff --git a/python/packages/hosting-mcp/pyproject.toml b/python/packages/hosting-mcp/pyproject.toml index 1aac85a096a..cbeb69db9c4 100644 --- a/python/packages/hosting-mcp/pyproject.toml +++ b/python/packages/hosting-mcp/pyproject.toml @@ -25,7 +25,7 @@ classifiers = [ dependencies = [ "agent-framework-core>=1.19.0,<2", "agent-framework-hosting==1.0.0a260730", - "mcp>=1.11.0,<2", + "mcp>=2.2.0,<3", "pydantic>=2,<3", ] diff --git a/python/packages/hosting-mcp/tests/hosting_mcp/test_agent_tool.py b/python/packages/hosting-mcp/tests/hosting_mcp/test_agent_tool.py index f3aec8a0126..23426f7e6bb 100644 --- a/python/packages/hosting-mcp/tests/hosting_mcp/test_agent_tool.py +++ b/python/packages/hosting-mcp/tests/hosting_mcp/test_agent_tool.py @@ -4,6 +4,7 @@ from collections.abc import Awaitable, Mapping, Sequence from typing import Any +from unittest.mock import AsyncMock, patch from agent_framework import ( Agent, @@ -16,7 +17,7 @@ ResponseStream, ) from agent_framework_hosting import AgentState -from mcp import types +from mcp import MCPError, types from pytest import raises from agent_framework_hosting_mcp import AgentMCPTool @@ -66,13 +67,15 @@ async def test_agent_tool_generates_schema_from_agent_with_overrides() -> None: }, ) - definitions = await tool.list_tools() + list_tools_result = await tool.list_tools() + result_tools = list_tools_result.tools - assert len(definitions) == 1 - definition = definitions[0] - assert definition.name == "research" - assert definition.description == "Tool description" - assert definition.inputSchema == { + assert list_tools_result.result_type == "complete" + assert len(result_tools) == 1 + result_tool = result_tools[0] + assert result_tool.name == "research" + assert result_tool.description == "Tool description" + assert result_tool.input_schema == { "type": "object", "properties": { "prompt": {"type": "string", "description": "Research request"}, @@ -102,10 +105,10 @@ async def test_agent_tool_uses_agent_metadata_by_default() -> None: agent = Agent(client=RecordingClient(), name="Research Agent", description="Agent description") tool: AgentMCPTool[Any] = AgentMCPTool(agent) - definition = (await tool.list_tools())[0] + result_tool = (await tool.list_tools()).tools[0] - assert definition.name == "Research_Agent" - assert definition.description == "Agent description" + assert result_tool.name == "Research_Agent" + assert result_tool.description == "Agent description" async def test_agent_tool_runs_with_agent_state_session() -> None: @@ -126,10 +129,16 @@ async def test_agent_tool_runs_with_agent_state_session() -> None: first = await tool.call_tool("session-agent", {"task": "first", "session_id": "session-1"}) second = await tool.call_tool("session-agent", {"task": "second", "session_id": "session-1"}) - assert isinstance(first[0], types.TextContent) - assert first[0].text == "response: first" - assert isinstance(second[0], types.TextContent) - assert second[0].text == "response: second" + assert first.result_type == "complete" + assert not first.is_error + assert isinstance(first.content[0], types.TextContent) + assert first.content[0].text == "response: first" + + assert second.result_type == "complete" + assert not second.is_error + assert isinstance(second.content[0], types.TextContent) + assert second.content[0].text == "response: second" + assert client.calls[0] == ["first"] assert client.calls[1] == ["first", "response: first", "second"] assert await state.session_store.get("session-1") is not None @@ -150,6 +159,43 @@ async def test_agent_tool_always_requires_session_parameter() -> None: session_id_parameter="session_id", ) - definition = (await tool.list_tools())[0] + definition = (await tool.list_tools()).tools[0] + + assert definition.input_schema["required"] == ["task", "session_id"] + + +async def test_agent_tool_unknown_tool_name() -> None: + agent = Agent(client=RecordingClient(), name="agent") + tool: AgentMCPTool[Any] = AgentMCPTool( + agent, parameters={"session_id": {"type": "string"}}, session_id_parameter="session_id", name="available_tool" + ) + + with raises(MCPError, match="Unknown MCP tool") as exc_info: + await tool.call_tool("made_up_tool", None) + + assert exc_info.value.code == types.INVALID_PARAMS + + +async def test_agent_tool_rejects_missing_runtime_session_id() -> None: + agent = Agent(client=RecordingClient(), name="agent") + tool: AgentMCPTool[Any] = AgentMCPTool( + agent, + parameters={"session_id": {"type": "string"}}, + session_id_parameter="session_id", + ) + + with raises(MCPError) as exc_info: + await tool.call_tool("agent", {"task": "hello"}) + + assert exc_info.value.code == types.INVALID_PARAMS + + +async def test_agent_tool_propagates_agent_execution_failure() -> None: + agent = Agent(client=RecordingClient(), name="agent") + tool: AgentMCPTool[Any] = AgentMCPTool(agent) - assert definition.inputSchema["required"] == ["task", "session_id"] + with ( + patch.object(agent, "run", AsyncMock(side_effect=RuntimeError("agent execution failed"))), + raises(RuntimeError, match="agent execution failed"), + ): + await tool.call_tool("agent", {"task": "hello"}) diff --git a/python/packages/hosting-mcp/tests/hosting_mcp/test_conversion.py b/python/packages/hosting-mcp/tests/hosting_mcp/test_conversion.py index 21de73f7afc..b21ab44f32e 100644 --- a/python/packages/hosting-mcp/tests/hosting_mcp/test_conversion.py +++ b/python/packages/hosting-mcp/tests/hosting_mcp/test_conversion.py @@ -71,11 +71,11 @@ def test_mcp_from_run_converts_final_response() -> None: assert blocks[1].name == "image.png" assert isinstance(blocks[2], types.ImageContent) assert blocks[2].data == "aW1hZ2U=" - assert blocks[2].mimeType == "image/png" + assert blocks[2].mime_type == "image/png" assert blocks[2].meta == {"source": "image"} assert isinstance(blocks[3], types.AudioContent) assert blocks[3].data == "YXVkaW8=" - assert blocks[3].mimeType == "audio/wav" + assert blocks[3].mime_type == "audio/wav" assert isinstance(blocks[4], types.EmbeddedResource) assert isinstance(blocks[4].resource, types.BlobResourceContents) assert str(blocks[4].resource.uri) == "af://binary" diff --git a/python/packages/hosting-mcp/tests/hosting_mcp/test_integration.py b/python/packages/hosting-mcp/tests/hosting_mcp/test_integration.py index 428a40548c3..d1ac00060ed 100644 --- a/python/packages/hosting-mcp/tests/hosting_mcp/test_integration.py +++ b/python/packages/hosting-mcp/tests/hosting_mcp/test_integration.py @@ -5,7 +5,7 @@ import asyncio import socket import time -from collections.abc import AsyncIterator, Awaitable, Mapping, Sequence +from collections.abc import AsyncGenerator, Awaitable, Mapping, Sequence from contextlib import asynccontextmanager from typing import Any @@ -22,6 +22,7 @@ ResponseStream, ) from mcp import types +from mcp.server import ServerRequestContext from mcp.server.lowlevel import Server from mcp.server.streamable_http_manager import StreamableHTTPSessionManager from starlette.applications import Starlette @@ -63,16 +64,21 @@ async def test_mcp_tool_calls_locally_hosted_agent() -> None: name="run_agent", chat_option_parameters={"reasoning_effort": {"type": "string"}}, ) - mcp_server = Server("hosting-mcp-integration") - @mcp_server.list_tools() - async def list_tools() -> list[types.Tool]: + async def list_tools( + _ctx: ServerRequestContext[dict[str, Any]], params: types.PaginatedRequestParams | None + ) -> types.ListToolsResult: return await agent_tool.list_tools() - @mcp_server.call_tool() - async def call_tool(name: str, arguments: dict[str, Any] | None) -> list[types.ContentBlock]: + async def call_tool( + _ctx: ServerRequestContext[dict[str, Any]], params: types.CallToolRequestParams + ) -> types.CallToolResult: + name = params.name + arguments = params.arguments or {} return await agent_tool.call_tool(name, arguments) + mcp_server = Server("hosting-mcp-integration", on_call_tool=call_tool, on_list_tools=list_tools) + session_manager = StreamableHTTPSessionManager( app=mcp_server, event_store=None, @@ -81,7 +87,7 @@ async def call_tool(name: str, arguments: dict[str, Any] | None) -> list[types.C ) @asynccontextmanager - async def lifespan(_app: Starlette) -> AsyncIterator[None]: + async def lifespan(_app: Starlette) -> AsyncGenerator[None]: async with session_manager.run(): yield diff --git a/python/packages/hosting-mcp/tests/hosting_mcp/test_workflow_tool.py b/python/packages/hosting-mcp/tests/hosting_mcp/test_workflow_tool.py index 20705f6677c..08737904cae 100644 --- a/python/packages/hosting-mcp/tests/hosting_mcp/test_workflow_tool.py +++ b/python/packages/hosting-mcp/tests/hosting_mcp/test_workflow_tool.py @@ -4,6 +4,7 @@ from dataclasses import dataclass from typing import Any +from unittest.mock import AsyncMock, patch from agent_framework import ( Executor, @@ -15,7 +16,7 @@ handler, ) from agent_framework_hosting import WorkflowState -from mcp import types +from mcp import MCPError, types from pytest import raises from agent_framework_hosting_mcp import WorkflowMCPTool @@ -50,15 +51,15 @@ async def test_workflow_tool_derives_object_schema_and_runs_workflow() -> None: name="repeat_text", ) - definition = (await tool.list_tools())[0] + definition = (await tool.list_tools()).tools[0] result = await tool.call_tool("repeat_text", {"text": "go", "repeat": 2}) assert definition.description == "Repeat text a requested number of times." - assert definition.inputSchema["type"] == "object" - assert definition.inputSchema["properties"]["text"]["type"] == "string" - assert definition.inputSchema["properties"]["repeat"]["type"] == "integer" - assert set(definition.inputSchema["required"]) == {"text", "repeat"} - assert result == [types.TextContent(type="text", text="gogo")] + assert definition.input_schema["type"] == "object" + assert definition.input_schema["properties"]["text"]["type"] == "string" + assert definition.input_schema["properties"]["repeat"]["type"] == "integer" + assert set(definition.input_schema["required"]) == {"text", "repeat"} + assert result.content == [types.TextContent(type="text", text="gogo")] async def test_workflow_tool_wraps_primitive_input() -> None: @@ -69,16 +70,16 @@ async def uppercase(value: str, ctx: WorkflowContext[object, str]) -> None: workflow = WorkflowBuilder(start_executor=uppercase, name="uppercase", output_from=[uppercase]).build() tool: WorkflowMCPTool[Any] = WorkflowMCPTool(workflow, argument_name="text") - definition = (await tool.list_tools())[0] + definition = (await tool.list_tools()).tools[0] result = await tool.call_tool("uppercase", {"text": "hello"}) - assert definition.inputSchema == { + assert definition.input_schema == { "type": "object", "properties": {"text": {"type": "string"}}, "required": ["text"], "additionalProperties": False, } - assert result == [types.TextContent(type="text", text="HELLO")] + assert result.content == [types.TextContent(type="text", text="HELLO")] async def test_workflow_tool_serializes_structured_output_as_json_text() -> None: @@ -91,7 +92,7 @@ async def structured(value: str, ctx: WorkflowContext[object, dict[str, str]]) - result = await tool.call_tool("structured", {"input": "hello"}) - assert result == [types.TextContent(type="text", text='{"value":"hello"}')] + assert result.content == [types.TextContent(type="text", text='{"value":"hello"}')] def test_workflow_tool_rejects_unhandled_external_input_requests() -> None: @@ -129,3 +130,32 @@ async def handle_number(self, value: int, ctx: WorkflowContext[object, str]) -> with raises(ValueError, match="exactly one"): tool._tool_for_workflow(workflow) # pyright: ignore[reportPrivateUsage] + + +async def test_workflow_tool_rejects_unknown_tool_name() -> None: + tool: WorkflowMCPTool[Any] = WorkflowMCPTool(create_workflow(), name="repeat_text") + + with raises(MCPError, match="Unknown MCP tool") as exc_info: + await tool.call_tool("unknown", {"text": "go", "repeat": 2}) + + assert exc_info.value.code == types.INVALID_PARAMS + + +async def test_workflow_tool_rejects_invalid_input() -> None: + tool: WorkflowMCPTool[Any] = WorkflowMCPTool(create_workflow(), name="repeat_text") + + with raises(MCPError) as exc_info: + await tool.call_tool("repeat_text", {"text": "go"}) + + assert exc_info.value.code == types.INVALID_PARAMS + + +async def test_workflow_tool_propagates_execution_failure() -> None: + workflow = create_workflow() + tool: WorkflowMCPTool[Any] = WorkflowMCPTool(workflow, name="repeat_text") + + with ( + patch.object(workflow, "run", AsyncMock(side_effect=RuntimeError("workflow execution failed"))), + raises(RuntimeError, match="workflow execution failed"), + ): + await tool.call_tool("repeat_text", {"text": "go", "repeat": 2}) diff --git a/python/packages/lab/pyproject.toml b/python/packages/lab/pyproject.toml index e9ac80a039d..b9ed34626be 100644 --- a/python/packages/lab/pyproject.toml +++ b/python/packages/lab/pyproject.toml @@ -83,7 +83,7 @@ prerelease = "if-necessary" # litellm>=1.100.0 drops the `fastapi.dependencies.utils.get_flat_dependant` import removed in FastAPI 0.141.0. constraint-dependencies = ["litellm>=1.100.0", "fastapi-sso>=0.19.0"] # python-multipart>=0.0.31 overrides litellm[proxy]'s exact pin of <=0.0.27 for security. -override-dependencies = ["mcp[ws]>=1.27.0,<2", "uvicorn[standard]>=0.34.0", "python-multipart>=0.0.31"] +override-dependencies = ["mcp>=2.2.0,<3", "uvicorn[standard]>=0.34.0", "python-multipart>=0.0.31"] environments = [ "sys_platform == 'darwin'", "sys_platform == 'linux'", diff --git a/python/pyproject.toml b/python/pyproject.toml index fd013d258be..43bb05ee897 100644 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -52,7 +52,7 @@ dev = [ ] test = [ "azure-monitor-opentelemetry>=1.8.10,<2", - "mcp[ws]", + "mcp", # Optional SDK used by the agent-hooks tests and isolated source checks. "agent-hooks-sdk>=0.1.0a4,<0.2", ] @@ -61,7 +61,7 @@ test = [ package = false prerelease = "if-necessary" # Security floors for transitive dependencies. -override-dependencies = ["mcp[ws]>=1.27.0,<2", "uvicorn[standard]>=0.34.0", "python-multipart>=0.0.31"] +override-dependencies = ["mcp>=2.2.0,<3", "uvicorn[standard]>=0.34.0", "python-multipart>=0.0.31"] environments = [ "sys_platform == 'darwin'", "sys_platform == 'linux'", diff --git a/python/samples/02-agents/mcp/README.md b/python/samples/02-agents/mcp/README.md index e3fce52794b..a409a8d6cf8 100644 --- a/python/samples/02-agents/mcp/README.md +++ b/python/samples/02-agents/mcp/README.md @@ -13,7 +13,8 @@ The Model Context Protocol (MCP) is an open standard for connecting AI agents to | **Agent as MCP Server** | [`agent_as_mcp_server.py`](agent_as_mcp_server.py) | Shows how to expose an Agent Framework agent as an MCP server that other AI applications can connect to | | **API Key Authentication** | [`mcp_api_key_auth.py`](mcp_api_key_auth.py) | Demonstrates API key authentication with MCP servers using `header_provider`, runtime invocation kwargs, and a command-line API key argument | | **GitHub Integration with PAT** | [`mcp_github_pat.py`](mcp_github_pat.py) | Demonstrates connecting to GitHub's MCP server using Personal Access Token (PAT) authentication | -| **Long-Running Task** | [`mcp_long_running_task.py`](mcp_long_running_task.py) | Demonstrates transparent SEP-2663 long-running task handling for MCP tools that advertise `taskSupport=required`. Self-spawns a stdio MCP child server | + + | **Progressive Disclosure** | [`mcp_progressive_disclosure.py`](mcp_progressive_disclosure.py) | Demonstrates `use_progressive_disclosure`, `always_load`, `allowed_tools`, and prefixed `list_mcp_tools` / `load_tool` / `unload_tool` names. `load_tool` and `unload_tool` can accept one tool name or multiple names. Self-spawns a stdio MCP child server | | **Sampling Approval** | [`mcp_sampling_approval.py`](mcp_sampling_approval.py) | Demonstrates gating server-initiated `sampling/createMessage` requests with a `sampling_approval_callback`, plus the `sampling_max_tokens` and `sampling_max_requests` guardrails. MCP sampling is denied by default | diff --git a/python/samples/02-agents/mcp/mcp_long_running_task.py b/python/samples/02-agents/mcp/mcp_long_running_task.py deleted file mode 100644 index 50ce078471b..00000000000 --- a/python/samples/02-agents/mcp/mcp_long_running_task.py +++ /dev/null @@ -1,181 +0,0 @@ -# Copyright (c) Microsoft. All rights reserved. - -""" -MCP Long-Running Task (SEP-2663) Example - -Demonstrates that ``MCPStdioTool`` transparently drives the MCP long-running -task lifecycle for tools that advertise ``execution.taskSupport == "required"``. -The agent observes a single function-call result; the framework handles the -``tools/call`` → ``tasks/get`` (polled) → ``tasks/result`` sequence in the -background. - -Run it as a single file. The script doubles as both the client and the stdio -MCP child server (the child branch is selected via ``--server``): - - python mcp_long_running_task.py - -Requirements: -- Azure CLI sign-in (``az login``) — used for Entra-ID auth against Azure OpenAI. -- ``AZURE_OPENAI_ENDPOINT`` — your Azure OpenAI resource endpoint, e.g. - ``https://.openai.azure.com/``. -- ``AZURE_OPENAI_CHAT_MODEL`` (or ``AZURE_OPENAI_MODEL``) — the deployment name, - e.g. ``gpt-4o-mini``. - -This sample uses the lower-level ``mcp.server.lowlevel.Server`` so it can: -1. Advertise a tool with ``execution=ToolExecution(taskSupport="required")``. -2. Enable the SDK's experimental task support for the ``tasks/*`` lifecycle. -""" - -import asyncio -import sys -from datetime import timedelta -from typing import Any - -from agent_framework import Agent, MCPStdioTool, MCPTaskOptions -from agent_framework.openai import OpenAIChatClient -from azure.identity import AzureCliCredential -from dotenv import load_dotenv - -load_dotenv() - - -# --------------------------------------------------------------------------- -# MCP stdio server (child-process branch) -# --------------------------------------------------------------------------- - - -async def _run_server() -> None: - """Run a minimal stdio MCP server exposing one long-running tool.""" - import mcp.types as types - from mcp.server.lowlevel import Server - from mcp.server.stdio import stdio_server - - server: Server[Any, Any] = Server("mcp-long-running-task-demo") - # Auto-registers handlers for tasks/get, tasks/result, tasks/cancel, tasks/list - # backed by an in-memory store. - server.experimental.enable_tasks() # ty: ignore[deprecated] - - @server.list_tools() - async def _list_tools() -> list[types.Tool]: # pyright: ignore[reportUnusedFunction] - return [ - types.Tool( - name="slow_summary", - description=( - "Produces a short summary of the supplied text after simulating several seconds of expensive work." - ), - inputSchema={ - "type": "object", - "properties": { - "text": { - "type": "string", - "description": "Text to summarize.", - } - }, - "required": ["text"], - }, - # Advertise that this tool MUST be invoked via the task lifecycle. - execution=types.ToolExecution(taskSupport="required"), - ) - ] - - @server.call_tool() - async def _call_tool(name: str, arguments: dict[str, Any]) -> Any: # pyright: ignore[reportUnusedFunction] - if name != "slow_summary": - raise ValueError(f"Unknown tool: {name}") - - ctx = server.request_context - - async def _work(task: Any) -> types.CallToolResult: - await task.update_status("Thinking...") - await asyncio.sleep(15.0) - text: str = (arguments.get("text") or "").strip() - words = text.split() - preview = " ".join(words[:6]) + ("..." if len(words) > 6 else "") - summary = ( - f"Summarized {len(words)} word(s). First few words: '{preview}'." - if words - else "No input text was provided." - ) - return types.CallToolResult( - content=[types.TextContent(type="text", text=summary)], - isError=False, - ) - - if not ctx.experimental.is_task: - # Client invoked the tool without task augmentation. Return a hard - # error so a misconfigured client surfaces the problem clearly. - return types.CallToolResult( - content=[ - types.TextContent( - type="text", - text="'slow_summary' must be invoked as a task.", - ) - ], - isError=True, - ) - - return await ctx.experimental.run_task(_work) - - async with stdio_server() as (read_stream, write_stream): - await server.run(read_stream, write_stream, server.create_initialization_options()) - - -# --------------------------------------------------------------------------- -# Agent client (default branch) -# --------------------------------------------------------------------------- - - -async def _run_client() -> None: - mcp_tool = MCPStdioTool( - name="LongRunningDemo", - description="Demo MCP server exposing a tool that advertises taskSupport=required.", - command=sys.executable, - args=[__file__, "--server"], - # Optional: cap individual tasks at two minutes. The server may apply its - # own default if this is omitted. - task_options=MCPTaskOptions(default_ttl=timedelta(minutes=2)), - ) - - async with Agent( - client=OpenAIChatClient(credential=AzureCliCredential()), - name="LROAgent", - instructions=( - "You are a helpful assistant. Use the slow_summary tool when the user " - "asks for a summary. Wait for the result and present it directly." - ), - tools=mcp_tool, - ) as agent: - prompt = ( - "Please summarize the following text using your slow_summary tool: " - "'The Model Context Protocol lets language models talk to external " - "tools and resources through a small JSON-RPC surface.'" - ) - - print("=== run() ===") - print(f"User: {prompt}") - response = await agent.run(prompt) - print(f"Agent: {response.text}\n") - - print("=== run(stream=True) ===") - print(f"User: {prompt}") - print("Agent: ", end="", flush=True) - async for update in agent.run(prompt, stream=True): - if update.text: - print(update.text, end="", flush=True) - print() - - -# --------------------------------------------------------------------------- -# Entry point -# --------------------------------------------------------------------------- - - -def main() -> None: - if len(sys.argv) > 1 and sys.argv[1] == "--server": - asyncio.run(_run_server()) - return - asyncio.run(_run_client()) - - -if __name__ == "__main__": - main() diff --git a/python/samples/02-agents/mcp/mcp_progressive_disclosure.py b/python/samples/02-agents/mcp/mcp_progressive_disclosure.py index 4f48e12910f..06024ef6d2f 100644 --- a/python/samples/02-agents/mcp/mcp_progressive_disclosure.py +++ b/python/samples/02-agents/mcp/mcp_progressive_disclosure.py @@ -7,6 +7,7 @@ from agent_framework import Agent, MCPStdioTool from agent_framework.openai import OpenAIChatClient from dotenv import load_dotenv +from mcp.server import ServerRequestContext __doc__ = """ MCP Progressive Disclosure Example @@ -51,20 +52,17 @@ async def _run_server() -> None: from mcp.server.lowlevel import Server from mcp.server.stdio import stdio_server - server: Server[Any, Any] = Server("mcp-progressive-disclosure-demo") - - @server.list_tools() - async def _list_tools() -> list[types.Tool]: # pyright: ignore[reportUnusedFunction] - return [ + async def _list_tools(_ctx: ServerRequestContext[dict[str, Any]], params: types.PaginatedRequestParams | None) -> types.ListToolsResult: + toolList = [ types.Tool( name="get_server_status", description="Return the health of the demo MCP server.", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), types.Tool( name="search_docs", description="Search short documentation snippets about MCP progressive disclosure.", - inputSchema={ + input_schema={ "type": "object", "properties": { "query": { @@ -78,12 +76,15 @@ async def _list_tools() -> list[types.Tool]: # pyright: ignore[reportUnusedFunc types.Tool( name="internal_admin_report", description="Internal server details that are intentionally filtered out by allowed_tools.", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), ] + return types.ListToolsResult(tools=toolList) + + async def _call_tool(_ctx: ServerRequestContext[dict[str, Any]], params: types.CallToolRequestParams) -> types.CallToolResult: + name = params.name + arguments = params.arguments or {} - @server.call_tool() - async def _call_tool(name: str, arguments: dict[str, Any]) -> types.CallToolResult: # pyright: ignore[reportUnusedFunction] if name == "get_server_status": text = "The demo MCP server is healthy. Use search_docs for progressive disclosure details." elif name == "search_docs": @@ -96,9 +97,12 @@ async def _call_tool(name: str, arguments: dict[str, Any]) -> types.CallToolResu elif name == "internal_admin_report": text = "This tool should not be discoverable because it is excluded by allowed_tools." else: - text = f"Unknown tool: {name}" + # text = f"Unknown tool: {name}" + return types.CallToolResult(content=[types.TextContent(type="text", text=f"Unknown tool: {name}")], is_error=True) return types.CallToolResult(content=[types.TextContent(type="text", text=text)]) + server: Server[Any] = Server("mcp-progressive-disclosure-demo", on_list_tools=_list_tools, on_call_tool=_call_tool) + async with stdio_server() as (read_stream, write_stream): await server.run(read_stream, write_stream, server.create_initialization_options()) diff --git a/python/samples/02-agents/mcp/mcp_sampling_approval.py b/python/samples/02-agents/mcp/mcp_sampling_approval.py index 0d359b7aecb..75b32c13af8 100644 --- a/python/samples/02-agents/mcp/mcp_sampling_approval.py +++ b/python/samples/02-agents/mcp/mcp_sampling_approval.py @@ -42,8 +42,8 @@ async def approve_sampling(params: types.CreateMessageRequestParams) -> bool: approve or deny. Returning ``False`` rejects the request. """ print("\n--- MCP server requested a sampling/createMessage ---") - if params.systemPrompt: - print(f"System prompt: {params.systemPrompt}") + if params.system_prompt: + print(f"System prompt: {params.system_prompt}") for message in params.messages: text = getattr(message.content, "text", message.content) print(f"{message.role}: {text}") diff --git a/python/samples/02-agents/providers/foundry/foundry_chat_client_with_toolbox_skills.py b/python/samples/02-agents/providers/foundry/foundry_chat_client_with_toolbox_skills.py index 8ac192289ef..9d669e5af82 100644 --- a/python/samples/02-agents/providers/foundry/foundry_chat_client_with_toolbox_skills.py +++ b/python/samples/02-agents/providers/foundry/foundry_chat_client_with_toolbox_skills.py @@ -4,7 +4,7 @@ import os from collections.abc import Generator -import httpx +import httpx2 from agent_framework import Agent, MCPSkillsSource, SkillsProvider, ToolApprovalMiddleware from agent_framework.foundry import FoundryChatClient from azure.core.credentials import TokenCredential @@ -34,13 +34,13 @@ """ -class _BearerAuth(httpx.Auth): +class _BearerAuth(httpx2.Auth): """Attach a fresh Foundry bearer token to every request.""" def __init__(self, credential: TokenCredential) -> None: self._get_token = get_bearer_token_provider(credential, "https://ai.azure.com/.default") - def auth_flow(self, request: httpx.Request) -> Generator[httpx.Request, httpx.Response, None]: + def auth_flow(self, request: httpx2.Request) -> Generator[httpx2.Request, httpx2.Response, None]: request.headers["Authorization"] = f"Bearer {self._get_token()}" yield request @@ -53,15 +53,15 @@ async def main() -> None: # and advertises the toolbox preview feature flag, plus the MCP streamable # HTTP transport that uses it. async with ( - httpx.AsyncClient( + httpx2.AsyncClient( auth=_BearerAuth(credential), - timeout=httpx.Timeout(30.0, read=300.0), + timeout=httpx2.Timeout(30.0, read=300.0), follow_redirects=True, ) as http_client, streamable_http_client( url=os.environ["FOUNDRY_TOOLBOX_MCP_SERVER_URL"], http_client=http_client, - ) as (read, write, _), + ) as (read, write), ClientSession(read, write) as session, ): await session.initialize() diff --git a/python/samples/02-agents/skills/mcp_based_skill/mcp_based_skill.py b/python/samples/02-agents/skills/mcp_based_skill/mcp_based_skill.py index 56379118f40..65f5a10317f 100644 --- a/python/samples/02-agents/skills/mcp_based_skill/mcp_based_skill.py +++ b/python/samples/02-agents/skills/mcp_based_skill/mcp_based_skill.py @@ -43,7 +43,7 @@ async def main() -> None: print("-" * 60) # 1. Connect to the MCP server over streamable HTTP. - async with streamable_http_client(url=mcp_url) as (read, write, _), ClientSession(read, write) as session: + async with streamable_http_client(url=mcp_url) as (read, write), ClientSession(read, write) as session: await session.initialize() # 2. Build a SkillsProvider that discovers skills over MCP. diff --git a/python/samples/04-hosting/mcp/agent_app.py b/python/samples/04-hosting/mcp/agent_app.py index c53bf9ec44c..d65f965ff9d 100644 --- a/python/samples/04-hosting/mcp/agent_app.py +++ b/python/samples/04-hosting/mcp/agent_app.py @@ -4,7 +4,7 @@ # "agent-framework-foundry", # "agent-framework-hosting-mcp", # "azure-identity", -# "mcp>=1.27.0,<2", +# "mcp>=2.2.0,<3", # "starlette>=0.40", # "uvicorn>=0.30", # ] @@ -30,8 +30,9 @@ from __future__ import annotations import os -from collections.abc import AsyncIterator +from collections.abc import AsyncGenerator from contextlib import asynccontextmanager +from typing import Any import uvicorn from agent_framework import Agent @@ -39,12 +40,24 @@ from agent_framework_hosting_mcp import AgentMCPTool from azure.identity.aio import DefaultAzureCredential from mcp import types +from mcp.server import ServerRequestContext from mcp.server.lowlevel import Server from mcp.server.streamable_http_manager import StreamableHTTPSessionManager from starlette.applications import Starlette from starlette.routing import Mount -server = Server("agent-framework-hosting-mcp-sample") + +async def list_tools(_ctx: ServerRequestContext[dict[str, Any]], params: types.PaginatedRequestParams | None) -> types.ListToolsResult: + """Describe the app-owned MCP tool schema.""" + return await agent_tool.list_tools() + + +async def call_tool(_ctx: ServerRequestContext[dict[str, Any]], params: types.CallToolRequestParams) -> types.CallToolResult: + """Run the app-owned tool with native MCP and Agent Framework values.""" + return await agent_tool.call_tool(params.name, params.arguments) + + +server = Server("agent-framework-hosting-mcp-sample", on_list_tools=list_tools, on_call_tool=call_tool) credential = DefaultAzureCredential() agent = Agent( client=FoundryChatClient( @@ -70,18 +83,6 @@ ) -@server.list_tools() -async def list_tools() -> list[types.Tool]: - """Describe the app-owned MCP tool schema.""" - return await agent_tool.list_tools() - - -@server.call_tool() -async def call_tool(name: str, arguments: dict[str, object] | None) -> list[types.ContentBlock]: - """Run the app-owned tool with native MCP and Agent Framework values.""" - return await agent_tool.call_tool(name, arguments) - - session_manager = StreamableHTTPSessionManager( app=server, event_store=None, @@ -91,7 +92,7 @@ async def call_tool(name: str, arguments: dict[str, object] | None) -> list[type @asynccontextmanager -async def lifespan(_app: Starlette) -> AsyncIterator[None]: +async def lifespan(_app: Starlette) -> AsyncGenerator[None]: """Start and stop native MCP and model-client resources.""" async with session_manager.run(), credential: yield diff --git a/python/samples/04-hosting/mcp/fastmcp_app.py b/python/samples/04-hosting/mcp/fastmcp_app.py index fb0901804e3..fc253e1b37a 100644 --- a/python/samples/04-hosting/mcp/fastmcp_app.py +++ b/python/samples/04-hosting/mcp/fastmcp_app.py @@ -4,7 +4,7 @@ # "agent-framework-foundry", # "agent-framework-hosting-mcp", # "azure-identity", -# "mcp>=1.27.0,<2", +# "mcp>=2.2.0,<3", # ] # /// # Run with: uv run fastmcp_app.py @@ -27,7 +27,7 @@ from __future__ import annotations import os -from collections.abc import AsyncIterator +from collections.abc import AsyncGenerator from contextlib import asynccontextmanager from typing import Literal @@ -36,7 +36,7 @@ from agent_framework_hosting_mcp import mcp_from_run, mcp_to_run from azure.identity.aio import DefaultAzureCredential from mcp import types -from mcp.server.fastmcp import FastMCP +from mcp.server.mcpserver import MCPServer credential = DefaultAzureCredential() agent = Agent( @@ -52,20 +52,15 @@ @asynccontextmanager -async def lifespan(_server: FastMCP[None]) -> AsyncIterator[None]: +async def lifespan(_server: MCPServer[None]) -> AsyncGenerator[None]: """Close the model credential when the FastMCP server stops.""" async with credential: yield -server = FastMCP( +server = MCPServer( name="agent-framework-hosting-fastmcp-sample", instructions="Expose an Agent Framework agent as an MCP tool.", - host="127.0.0.1", - port=8000, - streamable_http_path="/mcp", - json_response=True, - stateless_http=True, lifespan=lifespan, ) @@ -78,7 +73,7 @@ async def lifespan(_server: FastMCP[None]) -> AsyncIterator[None]: async def run_agent( task: str, reasoning_effort: Literal["low", "medium", "high"] | None = None, -) -> list[types.ContentBlock]: +) -> types.CallToolResult: """Run the agent with FastMCP-validated arguments.""" arguments: dict[str, object] = {"task": task} if reasoning_effort is not None: @@ -90,8 +85,8 @@ async def run_agent( options=run["options"], stream=False, ) - return mcp_from_run(result) + return types.CallToolResult(content=mcp_from_run(result)) if __name__ == "__main__": - server.run(transport="streamable-http") + server.run(transport="streamable-http", host="127.0.0.1", port=8000, streamable_http_path="/mcp", stateless_http=True, json_response=True) diff --git a/python/samples/04-hosting/mcp/manual_app.py b/python/samples/04-hosting/mcp/manual_app.py index 1faf24bc5a6..6f4d4e3a4b4 100644 --- a/python/samples/04-hosting/mcp/manual_app.py +++ b/python/samples/04-hosting/mcp/manual_app.py @@ -4,7 +4,7 @@ # "agent-framework-foundry", # "agent-framework-hosting-mcp", # "azure-identity", -# "mcp>=1.27.0,<2", +# "mcp>=2.2.0,<3", # "starlette>=0.40", # "uvicorn>=0.30", # ] @@ -23,8 +23,9 @@ from __future__ import annotations import os -from collections.abc import AsyncIterator +from collections.abc import AsyncGenerator from contextlib import asynccontextmanager +from typing import Any import uvicorn from agent_framework import Agent @@ -32,6 +33,7 @@ from agent_framework_hosting_mcp import mcp_from_run, mcp_to_run from azure.identity.aio import DefaultAzureCredential from mcp import types +from mcp.server import ServerRequestContext from mcp.server.lowlevel import Server from mcp.server.streamable_http_manager import StreamableHTTPSessionManager from starlette.applications import Starlette @@ -46,28 +48,14 @@ } } -server = Server("agent-framework-hosting-mcp-manual-sample") -credential = DefaultAzureCredential() -agent = Agent( - client=FoundryChatClient( - project_endpoint=os.environ["FOUNDRY_PROJECT_ENDPOINT"], - model=os.environ["FOUNDRY_MODEL"], - credential=credential, - ), - name="ManualMCPAgent", - description="Answer requests through a manually defined MCP tool.", - instructions="Answer the user's request clearly and concisely.", -) - -@server.list_tools() -async def list_tools() -> list[types.Tool]: +async def list_tools(_ctx: ServerRequestContext[dict[str, Any]], params: types.PaginatedRequestParams | None) -> types.ListToolsResult: """Return the app-owned native MCP tool definition.""" - return [ + return types.ListToolsResult(tools=[ types.Tool( name="run_agent_manually", description=agent.description or "", - inputSchema={ + input_schema={ "type": "object", "properties": { TASK_ARGUMENT: { @@ -80,11 +68,13 @@ async def list_tools() -> list[types.Tool]: "additionalProperties": False, }, ) - ] + ]) + +async def call_tool(_ctx: ServerRequestContext[dict[str, Any]], params: types.CallToolRequestParams) -> types.CallToolResult: + name = params.name + arguments = params.arguments or {} -@server.call_tool() -async def call_tool(name: str, arguments: dict[str, object] | None) -> list[types.ContentBlock]: """Convert, run, and render without the agent-backed adapter.""" if name != "run_agent_manually": raise ValueError(f"Unknown MCP tool: {name}") @@ -94,7 +84,21 @@ async def call_tool(name: str, arguments: dict[str, object] | None) -> list[type chat_option_arguments=CHAT_OPTION_ARGUMENTS, ) result = await agent.run(run["messages"], options=run["options"]) - return mcp_from_run(result) + return types.CallToolResult(content=mcp_from_run(result)) + + +server = Server("agent-framework-hosting-mcp-manual-sample", on_list_tools=list_tools, on_call_tool=call_tool) +credential = DefaultAzureCredential() +agent = Agent( + client=FoundryChatClient( + project_endpoint=os.environ["FOUNDRY_PROJECT_ENDPOINT"], + model=os.environ["FOUNDRY_MODEL"], + credential=credential, + ), + name="ManualMCPAgent", + description="Answer requests through a manually defined MCP tool.", + instructions="Answer the user's request clearly and concisely.", +) session_manager = StreamableHTTPSessionManager( @@ -106,7 +110,7 @@ async def call_tool(name: str, arguments: dict[str, object] | None) -> list[type @asynccontextmanager -async def lifespan(_app: Starlette) -> AsyncIterator[None]: +async def lifespan(_app: Starlette) -> AsyncGenerator[None]: """Start and stop native MCP and model-client resources.""" async with session_manager.run(), credential: yield diff --git a/python/samples/04-hosting/mcp/pyproject.toml b/python/samples/04-hosting/mcp/pyproject.toml index 2d1bb331dcb..d5b40a5d5fc 100644 --- a/python/samples/04-hosting/mcp/pyproject.toml +++ b/python/samples/04-hosting/mcp/pyproject.toml @@ -9,7 +9,7 @@ dependencies = [ "agent-framework-hosting", "agent-framework-hosting-mcp", "azure-identity", - "mcp>=1.27.0,<2", + "mcp>=2.2.0,<3", "starlette>=0.40", "uvicorn>=0.30", ] diff --git a/python/samples/04-hosting/mcp/session_app.py b/python/samples/04-hosting/mcp/session_app.py index c185c0ef130..5ee74513c8d 100644 --- a/python/samples/04-hosting/mcp/session_app.py +++ b/python/samples/04-hosting/mcp/session_app.py @@ -4,7 +4,7 @@ # "agent-framework-foundry", # "agent-framework-hosting-mcp", # "azure-identity", -# "mcp>=1.27.0,<2", +# "mcp>=2.2.0,<3", # "starlette>=0.40", # "uvicorn>=0.30", # ] @@ -34,8 +34,9 @@ import asyncio import os -from collections.abc import AsyncIterator +from collections.abc import AsyncGenerator from contextlib import asynccontextmanager +from typing import Any import uvicorn from agent_framework import Agent, InMemoryHistoryProvider @@ -44,12 +45,31 @@ from agent_framework_hosting_mcp import AgentMCPTool from azure.identity.aio import DefaultAzureCredential from mcp import types +from mcp.server import ServerRequestContext from mcp.server.lowlevel import Server from mcp.server.streamable_http_manager import StreamableHTTPSessionManager from starlette.applications import Starlette from starlette.routing import Mount -server = Server("agent-framework-hosting-mcp-session-sample") + +async def list_tools(_ctx: ServerRequestContext[dict[str, Any]], params: types.PaginatedRequestParams | None) -> types.ListToolsResult: + """Return the agent-derived MCP tool definition.""" + return await agent_tool.list_tools() + + +async def call_tool(_ctx: ServerRequestContext[dict[str, Any]], params: types.CallToolRequestParams) -> types.CallToolResult: + name = params.name + arguments = params.arguments or {} + """Serialize calls per app-owned session before using ``AgentState``.""" + session_id = arguments.get("session_id") if arguments else None + if not isinstance(session_id, str) or not session_id: + raise ValueError("MCP tool argument 'session_id' must be a non-empty string.") + lock = session_locks.setdefault(session_id, asyncio.Lock()) + async with lock: + return await agent_tool.call_tool(name, arguments) + + +server = Server("agent-framework-hosting-mcp-session-sample", on_list_tools=list_tools, on_call_tool=call_tool) credential = DefaultAzureCredential() agent = Agent( client=FoundryChatClient( @@ -90,23 +110,6 @@ session_locks: dict[str, asyncio.Lock] = {} -@server.list_tools() -async def list_tools() -> list[types.Tool]: - """Return the agent-derived MCP tool definition.""" - return await agent_tool.list_tools() - - -@server.call_tool() -async def call_tool(name: str, arguments: dict[str, object] | None) -> list[types.ContentBlock]: - """Serialize calls per app-owned session before using ``AgentState``.""" - session_id = arguments.get("session_id") if arguments else None - if not isinstance(session_id, str) or not session_id: - raise ValueError("MCP tool argument 'session_id' must be a non-empty string.") - lock = session_locks.setdefault(session_id, asyncio.Lock()) - async with lock: - return await agent_tool.call_tool(name, arguments) - - session_manager = StreamableHTTPSessionManager( app=server, event_store=None, @@ -116,7 +119,7 @@ async def call_tool(name: str, arguments: dict[str, object] | None) -> list[type @asynccontextmanager -async def lifespan(_app: Starlette) -> AsyncIterator[None]: +async def lifespan(_app: Starlette) -> AsyncGenerator[None]: """Start and stop native MCP and model-client resources.""" async with session_manager.run(), credential: yield diff --git a/python/samples/04-hosting/mcp/workflow_app.py b/python/samples/04-hosting/mcp/workflow_app.py index 8bc715a53ed..36fd766301e 100644 --- a/python/samples/04-hosting/mcp/workflow_app.py +++ b/python/samples/04-hosting/mcp/workflow_app.py @@ -2,7 +2,7 @@ # requires-python = ">=3.10" # dependencies = [ # "agent-framework-hosting-mcp", -# "mcp>=1.27.0,<2", +# "mcp>=2.2.0,<3", # "starlette>=0.40", # "uvicorn>=0.30", # ] @@ -20,15 +20,17 @@ from __future__ import annotations -from collections.abc import AsyncIterator +from collections.abc import AsyncGenerator from contextlib import asynccontextmanager from dataclasses import dataclass +from typing import Any import uvicorn from agent_framework import WorkflowBuilder, WorkflowContext, executor from agent_framework_hosting import WorkflowState from agent_framework_hosting_mcp import WorkflowMCPTool from mcp import types +from mcp.server import ServerRequestContext from mcp.server.lowlevel import Server from mcp.server.streamable_http_manager import StreamableHTTPSessionManager from starlette.applications import Starlette @@ -63,25 +65,24 @@ async def draft(request: DraftRequest, ctx: WorkflowContext[object, str]) -> Non ).build() -server = Server("agent-framework-hosting-mcp-workflow-sample") -workflow_tool = WorkflowMCPTool( - WorkflowState(create_workflow, cache_target=False), - name="draft_content", -) - - -@server.list_tools() -async def list_tools() -> list[types.Tool]: +async def list_tools(_ctx: ServerRequestContext[dict[str, Any]], params: types.PaginatedRequestParams | None) -> types.ListToolsResult: """Return the workflow-derived MCP tool definition.""" return await workflow_tool.list_tools() -@server.call_tool() -async def call_tool(name: str, arguments: dict[str, object] | None) -> list[types.ContentBlock]: +async def call_tool(_ctx: ServerRequestContext[dict[str, Any]], params: types.CallToolRequestParams) -> types.CallToolResult: + name = params.name + arguments = params.arguments or {} """Run a fresh workflow instance with validated MCP arguments.""" return await workflow_tool.call_tool(name, arguments) +server = Server("agent-framework-hosting-mcp-workflow-sample", on_list_tools=list_tools, on_call_tool=call_tool) +workflow_tool = WorkflowMCPTool( + WorkflowState(create_workflow, cache_target=False), + name="draft_content", +) + session_manager = StreamableHTTPSessionManager( app=server, event_store=None, @@ -91,7 +92,7 @@ async def call_tool(name: str, arguments: dict[str, object] | None) -> list[type @asynccontextmanager -async def lifespan(_app: Starlette) -> AsyncIterator[None]: +async def lifespan(_app: Starlette) -> AsyncGenerator[None]: """Start and stop the native MCP transport.""" async with session_manager.run(): yield diff --git a/python/scripts/dependencies/tests/test_dependency_bounds_runtime.py b/python/scripts/dependencies/tests/test_dependency_bounds_runtime.py index 12e184057c5..0c90de0bde4 100644 --- a/python/scripts/dependencies/tests/test_dependency_bounds_runtime.py +++ b/python/scripts/dependencies/tests/test_dependency_bounds_runtime.py @@ -180,7 +180,7 @@ def test_dependency_pyright_reuses_root_test_requirements(tmp_path: Path) -> Non (tmp_path / "pyproject.toml").write_text( """ [dependency-groups] -test = ["azure-monitor-opentelemetry", "mcp[ws]"] +test = ["azure-monitor-opentelemetry", "mcp"] """ ) command = ["uv", "run"] @@ -193,6 +193,6 @@ def test_dependency_pyright_reuses_root_test_requirements(tmp_path: Path) -> Non "--with", "azure-monitor-opentelemetry", "--with", - "mcp[ws]", + "mcp", ] assert command[-3:-1] == ["python", "-c"] diff --git a/python/uv.lock b/python/uv.lock index b52c7ff6ad8..869bcdfcd19 100644 --- a/python/uv.lock +++ b/python/uv.lock @@ -3,16 +3,16 @@ revision = 3 requires-python = ">=3.11, <3.15" resolution-markers = [ "python_full_version >= '3.14' and sys_platform == 'darwin'", - "python_full_version == '3.13.*' and sys_platform == 'darwin'", - "python_full_version == '3.12.*' and sys_platform == 'darwin'", - "python_full_version < '3.12' and sys_platform == 'darwin'", "python_full_version >= '3.14' and sys_platform == 'linux'", - "python_full_version == '3.13.*' and sys_platform == 'linux'", - "python_full_version == '3.12.*' and sys_platform == 'linux'", - "python_full_version < '3.12' and sys_platform == 'linux'", "python_full_version >= '3.14' and sys_platform == 'win32'", + "python_full_version == '3.13.*' and sys_platform == 'darwin'", + "python_full_version == '3.13.*' and sys_platform == 'linux'", "python_full_version == '3.13.*' and sys_platform == 'win32'", + "python_full_version == '3.12.*' and sys_platform == 'darwin'", + "python_full_version == '3.12.*' and sys_platform == 'linux'", "python_full_version == '3.12.*' and sys_platform == 'win32'", + "python_full_version < '3.12' and sys_platform == 'darwin'", + "python_full_version < '3.12' and sys_platform == 'linux'", "python_full_version < '3.12' and sys_platform == 'win32'", ] supported-markers = [ @@ -70,7 +70,7 @@ members = [ "agent-framework-typesafe", ] overrides = [ - { name = "mcp", extras = ["ws"], specifier = ">=1.27.0,<2" }, + { name = "mcp", specifier = ">=2.2.0,<3" }, { name = "python-multipart", specifier = ">=0.0.31" }, { name = "uvicorn", extras = ["standard"], specifier = ">=0.34.0" }, ] @@ -149,7 +149,7 @@ dev = [ test = [ { name = "agent-hooks-sdk", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "azure-monitor-opentelemetry", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "mcp", extra = ["ws"], marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "mcp", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, ] [package.metadata] @@ -181,7 +181,7 @@ dev = [ test = [ { name = "agent-hooks-sdk", specifier = ">=0.1.0a4,<0.2" }, { name = "azure-monitor-opentelemetry", specifier = ">=1.8.10,<2" }, - { name = "mcp", extras = ["ws"] }, + { name = "mcp" }, ] [[package]] @@ -483,7 +483,7 @@ all = [ { name = "agent-framework-purview", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "agent-framework-redis", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "agent-framework-tools", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "mcp", extra = ["ws"], marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "mcp", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, ] [package.dev-dependencies] @@ -525,7 +525,7 @@ requires-dist = [ { name = "agent-framework-purview", marker = "extra == 'all'", editable = "packages/purview" }, { name = "agent-framework-redis", marker = "python_full_version < '3.15' and extra == 'all'", editable = "packages/redis" }, { name = "agent-framework-tools", marker = "extra == 'all'", editable = "packages/tools" }, - { name = "mcp", marker = "extra == 'all'", specifier = ">=1.24.0,<2" }, + { name = "mcp", marker = "extra == 'all'", specifier = ">=2.2.0,<3" }, { name = "msgspec", specifier = ">=0.20.0,<0.23" }, { name = "opentelemetry-api", specifier = ">=1.39.0,<2" }, { name = "pydantic", specifier = ">=2,<3" }, @@ -674,7 +674,7 @@ dependencies = [ { name = "azure-ai-agentserver-invocations", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "azure-ai-agentserver-responses", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "httpx", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "mcp", extra = ["ws"], marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "mcp", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, ] [package.metadata] @@ -684,7 +684,7 @@ requires-dist = [ { name = "azure-ai-agentserver-invocations", specifier = ">=1.1.0,<2" }, { name = "azure-ai-agentserver-responses", specifier = ">=2.2.0b1,<3" }, { name = "httpx", specifier = ">=0.28,<1" }, - { name = "mcp", specifier = ">=1.24.0,<2" }, + { name = "mcp", specifier = ">=2.2.0,<3" }, ] [[package]] @@ -769,7 +769,7 @@ source = { editable = "packages/hosting-mcp" } dependencies = [ { name = "agent-framework-core", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "agent-framework-hosting", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "mcp", extra = ["ws"], marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "mcp", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "pydantic", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, ] @@ -777,7 +777,7 @@ dependencies = [ requires-dist = [ { name = "agent-framework-core", editable = "packages/core" }, { name = "agent-framework-hosting", editable = "packages/hosting" }, - { name = "mcp", specifier = ">=1.11.0,<2" }, + { name = "mcp", specifier = ">=2.2.0,<3" }, { name = "pydantic", specifier = ">=2,<3" }, ] @@ -1887,7 +1887,7 @@ source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "anyio", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "jsonschema", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "mcp", extra = ["ws"], marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "mcp", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "sniffio", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/9b/74/635d38e2ef813a2d1ba4ff7f3de736c35ee3bce9ebf4a5b7e89f3293604f/claude_agent_sdk-0.2.163.tar.gz", hash = "sha256:269821ad5acff5967522ff15f444b0598840fd9bdc45c645de569a5478061a53", size = 365180, upload-time = "2026-09-30T19:46:37.125Z" } @@ -2694,15 +2694,6 @@ http2 = [ { name = "h2", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, ] -[[package]] -name = "httpx-sse" -version = "0.4.3" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/0f/4c/751061ffa58615a32c31b2d82e8482be8dd4a89154f003147acee90f2be9/httpx_sse-0.4.3.tar.gz", hash = "sha256:9b1ed0127459a66014aec3c56bebd93da3c1bc8bb6618c8082039a44889a755d", size = 15943, upload-time = "2025-10-10T21:48:22.271Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/d2/fd/6668e5aec43ab844de6fc74927e155a3b37bf40d7c3790e49fc0406b6578/httpx_sse-0.4.3-py3-none-any.whl", hash = "sha256:0ac1c9fe3c0afad2e0ebb25a934a59f4c7823b60792691f779fad2c5568830fc", size = 8960, upload-time = "2025-10-10T21:48:21.158Z" }, -] - [[package]] name = "httpx2" version = "2.13.1" @@ -3137,15 +3128,15 @@ wheels = [ [[package]] name = "mcp" -version = "1.30.0" +version = "2.2.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "anyio", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "httpx", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "httpx-sse", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "httpx2", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "jsonschema", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "mcp-types", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "opentelemetry-api", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "pydantic", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "pydantic-settings", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "pyjwt", extra = ["crypto"], marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "python-multipart", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "pywin32", marker = "sys_platform == 'win32'" }, @@ -3155,14 +3146,22 @@ dependencies = [ { name = "typing-inspection", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "uvicorn", extra = ["standard"], marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/ba/93/0142dc84a666daf8ad51a34268f34c12fd6fda4f3810c4be2504eecc8212/mcp-1.30.0.tar.gz", hash = "sha256:445414625fce5c295faa505bb11bacece661ab6f4028d57c935db57820b7a3e4", size = 680511, upload-time = "2026-09-07T14:34:15.845Z" } +sdist = { url = "https://files.pythonhosted.org/packages/76/31/ac54fb0fdd5b37de704486e288bba4fbbb463f24cfcfedbede407b854513/mcp-2.2.0.tar.gz", hash = "sha256:2dc37ecb1974becdcebdbf7561e7c15a07dbbf20ba21ba16c3593b3038b3afbd", size = 4084129, upload-time = "2026-09-07T16:06:23.439Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/f5/f4/e58bc33317c92a0203664daaf00bf6f41166cc0149e5d6870a03f7cd004a/mcp-1.30.0-py3-none-any.whl", hash = "sha256:666edb5009503e1047c9d60346a756f94b261f05cc2625f23d41c728ffc484d0", size = 234581, upload-time = "2026-09-07T14:34:14.266Z" }, + { url = "https://files.pythonhosted.org/packages/1b/ff/8e7eade68b8a28f7da0ed1085544341b51f9c935dbf6b95c76b7edfea6a0/mcp-2.2.0-py3-none-any.whl", hash = "sha256:bde982589473a060ae145e3406e9a5333fe538c97229ba841f5a7f92be004f81", size = 365656, upload-time = "2026-09-07T16:06:19.711Z" }, ] -[package.optional-dependencies] -ws = [ - { name = "websockets", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, +[[package]] +name = "mcp-types" +version = "2.2.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pydantic", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "typing-extensions", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/ae/91/762d7755d971aff8a28d75f7961656148edf27875c8026e6385aaab08ae7/mcp_types-2.2.0.tar.gz", hash = "sha256:d3ed53703ddd10d9c6399f29d322bb66f3f67ab41348ac8556ba23e07fedefad", size = 65892, upload-time = "2026-09-07T16:06:25.187Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/8f/d7/6ffba5d8cd5dd9b8a19478875c50e04945314ba5074e84d749283f27f62d/mcp_types-2.2.0-py3-none-any.whl", hash = "sha256:ea476b73ee86709ab5abc9452385ed36cc05907e582355622e294595c9a04f13", size = 69106, upload-time = "2026-09-07T16:06:21.461Z" }, ] [[package]] @@ -3750,13 +3749,13 @@ version = "2.5.3" source = { registry = "https://pypi.org/simple" } resolution-markers = [ "python_full_version >= '3.14' and sys_platform == 'darwin'", - "python_full_version == '3.13.*' and sys_platform == 'darwin'", - "python_full_version == '3.12.*' and sys_platform == 'darwin'", "python_full_version >= '3.14' and sys_platform == 'linux'", - "python_full_version == '3.13.*' and sys_platform == 'linux'", - "python_full_version == '3.12.*' and sys_platform == 'linux'", "python_full_version >= '3.14' and sys_platform == 'win32'", + "python_full_version == '3.13.*' and sys_platform == 'darwin'", + "python_full_version == '3.13.*' and sys_platform == 'linux'", "python_full_version == '3.13.*' and sys_platform == 'win32'", + "python_full_version == '3.12.*' and sys_platform == 'darwin'", + "python_full_version == '3.12.*' and sys_platform == 'linux'", "python_full_version == '3.12.*' and sys_platform == 'win32'", ] sdist = { url = "https://files.pythonhosted.org/packages/13/01/11703282db468b85f6f7b8c7f22d058de5970d5c7e60a3a8aaa313c3de36/numpy-2.5.3.tar.gz", hash = "sha256:df2d5874ff183595a4ba404edd04f6bd9b5505c1d7708573f6a6c17489a67563", size = 20791231, upload-time = "2026-09-06T16:27:47.073Z" } @@ -3852,7 +3851,7 @@ source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "griffelib", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "httpx2", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "mcp", extra = ["ws"], marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "mcp", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "openai", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "pydantic", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "pyjwt", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, @@ -4890,20 +4889,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/e6/c0/3f7c604ac93ec47a36a3733d6e7717b7096f0989f669d4133bd3e85f3375/pydantic_monty_runtime-0.0.23-py3-none-win_amd64.whl", hash = "sha256:742375f494e298a4f96933ac3694a0fe34c9cddd9b666669bc193f348bbed0b8", size = 10404629, upload-time = "2026-09-05T19:26:01.143Z" }, ] -[[package]] -name = "pydantic-settings" -version = "2.15.0" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "pydantic", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "python-dotenv", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, - { name = "typing-inspection", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/68/ca/31c57507b13119d7d3cfa1576dad2911a4861e3be07b579395f4e9d393f9/pydantic_settings-2.15.0.tar.gz", hash = "sha256:694b793e84f766ba76a90ebdefc01d0a9a045dab0382bee70393da93712ad117", size = 261253, upload-time = "2026-08-07T09:24:57.419Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/30/a4/2bffa9f8e804325a09867f0e9d30795c80ea9f8d62560bd1b6ad6220eb2f/pydantic_settings-2.15.0-py3-none-any.whl", hash = "sha256:0ba092c291c94baceb5eff768aa0d56400a457585bc0175925a5a5510303da42", size = 69413, upload-time = "2026-08-07T09:24:55.839Z" }, -] - [[package]] name = "pygments" version = "2.21.0" From 6532629bb3b329f11e9cc0f74136dccaf832ac29 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 30 Sep 2026 14:29:55 +0200 Subject: [PATCH 02/42] Added client test for agents as mcp_servers --- .../packages/core/agent_framework/_agents.py | 5 ++++- .../packages/core/tests/core/test_agents.py | 20 +++++++++++++++++++ 2 files changed, 24 insertions(+), 1 deletion(-) diff --git a/python/packages/core/agent_framework/_agents.py b/python/packages/core/agent_framework/_agents.py index c4597013ad2..5f1def4c559 100644 --- a/python/packages/core/agent_framework/_agents.py +++ b/python/packages/core/agent_framework/_agents.py @@ -1892,9 +1892,12 @@ def as_mcp_server( server_args: dict[str, Any] = { "name": server_name, - "version": version, "instructions": instructions, } + + if version is not None: + server_args["version"] = version + if lifespan: server_args["lifespan"] = lifespan if kwargs: diff --git a/python/packages/core/tests/core/test_agents.py b/python/packages/core/tests/core/test_agents.py index 5bd4b4d8754..1b1f177a6d3 100644 --- a/python/packages/core/tests/core/test_agents.py +++ b/python/packages/core/tests/core/test_agents.py @@ -2925,6 +2925,26 @@ async def test_chat_agent_as_mcp_server_basic(client: SupportsChatGetResponse) - assert hasattr(server, "name") assert hasattr(server, "version") + from mcp import Client + + async with Client(server) as mcp_client: # auto -> 2026-07-28; aka V2 + tools_result = await mcp_client.list_tools() + call_result = await mcp_client.call_tool("TestAgent", {"task": "hello"}) + + assert [tool.name for tool in tools_result.tools] == ["TestAgent"] + assert call_result.result_type == "complete" + assert not call_result.is_error + assert mcp_client.protocol_version == "2026-07-28" + + async with Client(server, mode="legacy") as mcp_client: # 2025-11-25; aka legacy + tools_result = await mcp_client.list_tools() + call_result = await mcp_client.call_tool("TestAgent", {"task": "hello"}) + + assert [tool.name for tool in tools_result.tools] == ["TestAgent"] + assert call_result.result_type == "complete" + assert not call_result.is_error + assert mcp_client.protocol_version == "2025-11-25" + async def test_agent_prepares_mcp_run_before_copying_functions(chat_client_base: Any) -> None: captured_options: list[dict[str, Any]] = [] From 9cbf4ed7dc3dffd8cadce44a5981ace6f944c616 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 30 Sep 2026 15:02:22 +0200 Subject: [PATCH 03/42] More casing fixes --- python/packages/core/agent_framework/_mcp.py | 137 +++++++++---------- 1 file changed, 67 insertions(+), 70 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index a89f3f26cb9..aeb6d4ee94a 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -1099,7 +1099,7 @@ def _parse_prompt_result_from_mcp( { "type": "image" if isinstance(content, types.ImageContent) else "audio", "data": content.data, - "mimeType": content.mimeType, + "mimeType": content.mime_type, }, default=str, ) @@ -1114,7 +1114,7 @@ def _parse_prompt_result_from_mcp( { "type": "blob", "data": content.resource.blob, - "mimeType": content.resource.mimeType, + "mimeType": content.resource.mime_type, }, default=str, ) @@ -1177,7 +1177,7 @@ def _parse_tool_result_from_mcp( result.append( Content.from_data( data=decoded, - media_type=item.mimeType, + media_type=item.mime_type, **additional_kwargs, ) ) @@ -1185,7 +1185,7 @@ def _parse_tool_result_from_mcp( result.append( Content.from_uri( uri=str(item.uri), - media_type=item.mimeType, + media_type=item.mime_type, **additional_kwargs, ) ) @@ -1195,7 +1195,7 @@ def _parse_tool_result_from_mcp( result.append(Content.from_text(item.resource.text, **additional_kwargs)) case types.BlobResourceContents(): blob = item.resource.blob - mime = item.resource.mimeType or "application/octet-stream" + mime = item.resource.mime_type or "application/octet-stream" if not blob.startswith("data:"): blob = f"data:{mime};base64,{blob}" result.append( @@ -1209,9 +1209,9 @@ def _parse_tool_result_from_mcp( result.append(Content.from_text(str(item), **additional_kwargs)) structured_block: Content | None = None - if mcp_type.structuredContent is not None: + if mcp_type.structured_content is not None: structured_block = Content.from_text( - json.dumps(mcp_type.structuredContent, default=str), **additional_kwargs + json.dumps(mcp_type.structured_content, default=str), **additional_kwargs ) # Select model-visible content per explicit policy (#7866). Host payload still @@ -1272,7 +1272,7 @@ def _parse_content_from_mcp( return_types.append( Content.from_data( data=data_bytes, - media_type=mcp_type.mimeType, + media_type=mcp_type.mime_type, raw_representation=mcp_type, ) ) @@ -1280,7 +1280,7 @@ def _parse_content_from_mcp( return_types.append( Content.from_uri( uri=str(mcp_type.uri), - media_type=mcp_type.mimeType or "application/json", + media_type=mcp_type.mime_type or "application/json", raw_representation=mcp_type, ) ) @@ -1296,11 +1296,11 @@ def _parse_content_from_mcp( case types.ToolResultContent(): return_types.append( Content.from_function_result( - call_id=mcp_type.toolUseId, + call_id=mcp_type.tool_use_id, result=self._parse_content_from_mcp(mcp_type.content) if mcp_type.content - else mcp_type.structuredContent, - exception=str(Exception()) if mcp_type.isError else None, + else mcp_type.structured_content, + exception=str(Exception()) if mcp_type.is_error else None, raw_representation=mcp_type, ) ) @@ -1320,7 +1320,7 @@ def _parse_content_from_mcp( return_types.append( Content.from_uri( uri=mcp_type.resource.blob, - media_type=mcp_type.resource.mimeType, + media_type=mcp_type.resource.mime_type, raw_representation=mcp_type, additional_properties=( mcp_type.annotations.model_dump() if mcp_type.annotations else None @@ -1351,20 +1351,20 @@ def _prepare_content_for_mcp( ) if content.type == "data": if content.media_type and content.media_type.startswith("image/"): - return types.ImageContent(type="image", data=content.uri, mimeType=content.media_type) # type: ignore[attr-defined] + return types.ImageContent(type="image", data=content.uri, mime_type=content.media_type) # type: ignore[attr-defined] if content.media_type and content.media_type.startswith("audio/"): - return types.AudioContent(type="audio", data=content.uri, mimeType=content.media_type) # type: ignore[attr-defined] + return types.AudioContent(type="audio", data=content.uri, mime_type=content.media_type) # type: ignore[attr-defined] if content.media_type and content.media_type.startswith("application/"): return types.EmbeddedResource( type="resource", resource=types.BlobResourceContents( blob=content.uri, # type: ignore[attr-defined] - mimeType=content.media_type, + mime_type=content.media_type, uri=( content.additional_properties.get("uri", "af://binary") if content.additional_properties else "af://binary" - ), # type: ignore[arg-type] + ), ), ) return None @@ -1375,7 +1375,7 @@ def _prepare_content_for_mcp( return types.ResourceLink( type="resource_link", uri=content.uri, # type: ignore[arg-type,attr-defined] - mimeType=content.media_type, + mime_type=content.media_type, name=resource_name, ) return None @@ -2026,7 +2026,7 @@ async def _connect_on_owner( try: with create_mcp_client_span("initialize", attributes=self._mcp_base_span_attributes()) as init_span: initialize_result = await session.initialize() - init_span.set_attribute(OtelAttr.MCP_PROTOCOL_VERSION, initialize_result.protocolVersion) + init_span.set_attribute(OtelAttr.MCP_PROTOCOL_VERSION, initialize_result.protocol_version) self._set_server_capabilities(getattr(initialize_result, "capabilities", None)) except (Exception, asyncio.CancelledError) as ex: cancelled, cleanup_error = await self._close_and_check_cancelled(ex) @@ -2052,7 +2052,7 @@ async def _connect_on_owner( # If the session is not initialized, we need to reinitialize it with create_mcp_client_span("initialize", attributes=self._mcp_base_span_attributes()) as init_span: initialize_result = await self.session.initialize() - init_span.set_attribute(OtelAttr.MCP_PROTOCOL_VERSION, initialize_result.protocolVersion) + init_span.set_attribute(OtelAttr.MCP_PROTOCOL_VERSION, initialize_result.protocol_version) self._set_server_capabilities(getattr(initialize_result, "capabilities", None)) elif self._server_capabilities is None: self._set_server_capabilities(getattr(self.session, "_server_capabilities", None)) @@ -2212,7 +2212,7 @@ async def sampling_callback( "MCP server '%s' sent a sampling/createMessage request (%d message(s), maxTokens=%s).", self.name, len(params.messages), - params.maxTokens, + params.max_tokens, ) if self.sampling_max_requests is not None: @@ -2244,25 +2244,25 @@ async def sampling_callback( messages.append(self._parse_message_from_mcp(msg)) options: ChatOptions[None] = {} - if params.systemPrompt is not None: - options["instructions"] = params.systemPrompt + if params.system_prompt is not None: + options["instructions"] = params.system_prompt if params.tools is not None: options["tools"] = [ FunctionTool( name=tool.name, description=tool.description or "", - input_model=tool.inputSchema, + input_model=tool.input_schema, ) for tool in params.tools ] - if params.toolChoice is not None and params.toolChoice.mode is not None: - options["tool_choice"] = params.toolChoice.mode + if params.tool_choice is not None and params.tool_choice.mode is not None: + options["tool_choice"] = params.tool_choice.mode if params.temperature is not None: options["temperature"] = params.temperature - options["max_tokens"] = self._capped_sampling_max_tokens(params.maxTokens) - if params.stopSequences is not None: - options["stop"] = params.stopSequences + options["max_tokens"] = self._capped_sampling_max_tokens(params.max_tokens) + if params.stop_sequences is not None: + options["stop"] = params.stop_sequences try: chat_client: Any = self.client @@ -2293,7 +2293,7 @@ async def sampling_callback( role="assistant", content=tool_use_contents, model=response.model or "unknown", - stopReason="toolUse", + stop_reason="toolUse", ) # grab the first content that is of type TextContent or ImageContent @@ -2511,9 +2511,9 @@ async def _load_prompts_locked(self) -> None: existing_names.add(local_name) # Check if there are more pages - if not prompt_list.nextCursor: + if not prompt_list.next_cursor: break - params = types.PaginatedRequestParams(cursor=prompt_list.nextCursor) + params = types.PaginatedRequestParams(cursor=prompt_list.next_cursor) self._validate_config_names([*self._functions, *new_functions]) if self._function_load_callback is not None: @@ -2598,7 +2598,7 @@ async def _load_tools_locked(self) -> None: if tool.meta is not None: tool_call_meta_by_name[tool.name] = _validate_mcp_meta(tool.meta) or {} - task_support = getattr(getattr(tool, "execution", None), "taskSupport", None) + task_support = getattr(getattr(tool, "execution", None), "task_support", None) if task_support is not None: tool_task_support_by_name[tool.name] = task_support @@ -2607,7 +2607,7 @@ async def _load_tools_locked(self) -> None: # which causes OpenAI API to reject the schema with a 400 error. # Guard against non-conforming MCP servers that send inputSchema=None # despite the MCP spec typing it as dict[str, Any]. - input_schema = dict(tool.inputSchema or {}) + input_schema = dict(tool.input_schema or {}) if input_schema.get("type") == "object" and "properties" not in input_schema: input_schema["properties"] = {} @@ -2667,9 +2667,9 @@ async def _load_tools_locked(self) -> None: new_functions.append(func) # Check if there are more pages - if not tool_list.nextCursor: + if not tool_list.next_cursor: break - params = types.PaginatedRequestParams(cursor=tool_list.nextCursor) + params = types.PaginatedRequestParams(cursor=tool_list.next_cursor) current_functions = [ func @@ -2744,14 +2744,14 @@ async def _ensure_connected(self) -> None: Raises: ToolExecutionException: If reconnection fails. """ - from mcp.shared.exceptions import McpError + from mcp import MCPError if not self._ping_available: return try: await self.session.send_ping() # type: ignore[union-attr] - except McpError as mcp_exc: + except MCPError as mcp_exc: if mcp_exc.error.code == -32601: self._ping_available = False logger.debug("Skipping future MCP pings because the server does not support ping.") @@ -2870,13 +2870,13 @@ async def _call_tool_with_retries( ) -> str | list[Content]: """Execute the MCP tools/call RPC with retry logic.""" from anyio import ClosedResourceError - from mcp.shared.exceptions import McpError + from mcp import MCPError for attempt in range(2): try: result = await self.session.call_tool(tool_name, arguments=filtered_kwargs, meta=meta) # type: ignore _capture_mcp_tool_result(result) - if result.isError: + if result.is_error: parsed = parser(result) text = ( "\n".join(c.text for c in parsed if c.type == "text" and c.text) @@ -2890,13 +2890,13 @@ async def _call_tool_with_retries( return parser(result) except ToolExecutionException: raise - except (ClosedResourceError, McpError) as call_ex: + except (ClosedResourceError, MCPError) as call_ex: is_session_terminated = ( - isinstance(call_ex, McpError) and "session terminated" in call_ex.error.message.lower() + isinstance(call_ex, MCPError) and "session terminated" in call_ex.error.message.lower() ) is_connection_lost = isinstance(call_ex, ClosedResourceError) or is_session_terminated if not is_connection_lost: - error_message = call_ex.error.message if isinstance(call_ex, McpError) else str(call_ex) + error_message = call_ex.error.message if isinstance(call_ex, MCPError) else str(call_ex) if span.is_recording(): set_mcp_span_error(span, type(call_ex).__name__, error_message) raise ToolExecutionException(error_message, inner_exception=call_ex) from call_ex @@ -3005,7 +3005,7 @@ async def _call_tool_as_task( kwargs: dict[str, Any], ) -> str | list[Content]: from anyio import ClosedResourceError - from mcp.shared.exceptions import McpError + from mcp import MCPError if not self.load_tools_flag: raise ToolExecutionException( @@ -3021,9 +3021,9 @@ async def _call_tool_as_task( # Reconnect-and-retry is only safe after the task_id is known. try: task_id, fallback_result = await self._call_tool_as_task_create(tool_name, filtered_kwargs, meta) - except (ClosedResourceError, McpError) as ex: + except (ClosedResourceError, MCPError) as ex: if not self._is_connection_lost(ex): - error_message = ex.error.message if isinstance(ex, McpError) else str(ex) + error_message = ex.error.message if isinstance(ex, MCPError) else str(ex) raise ToolExecutionException(error_message, inner_exception=ex) from ex raise ToolExecutionException( f"Failed to call tool '{tool_name}' - connection lost; task state unknown.", @@ -3037,7 +3037,7 @@ async def _call_tool_as_task( # Server returned a CallToolResult (no task created) or fell back to plain tools/call. if fallback_result is not None: _capture_mcp_tool_result(fallback_result) - if fallback_result.isError: + if fallback_result.is_error: parsed = parser(fallback_result) text = ( "\n".join(c.text for c in parsed if c.type == "text" and c.text) @@ -3100,8 +3100,7 @@ async def _call_tool_as_task_create( ``(None, CallToolResult)`` when it returned a non-task result, falling back to plain ``tools/call`` if the server rejects the ``task`` field outright. """ - from mcp import types - from mcp.shared.exceptions import McpError + from mcp import MCPError, types from pydantic import ValidationError opts = self._effective_task_options() @@ -3128,7 +3127,7 @@ async def _call_tool_as_task_create( request, types.Result, ) - except McpError as ex: + except MCPError as ex: if ex.error.code not in (types.METHOD_NOT_FOUND, types.INVALID_PARAMS): raise logger.debug( @@ -3165,20 +3164,19 @@ async def _call_tool_as_task_create( async def _poll_task_until_terminal(self, task_id: str) -> types.GetTaskResult: """Poll ``tasks/get`` until the task reaches a terminal status.""" import httpx - from mcp import types - from mcp.shared.exceptions import McpError + from mcp import MCPError, types - # SDK raises McpError(code=httpx.REQUEST_TIMEOUT=408) on session read timeout. + # SDK raises MCPError(code=httpx.REQUEST_TIMEOUT=408) on session read timeout. transient_codes: frozenset[int] = frozenset({int(httpx.codes.REQUEST_TIMEOUT)}) while True: - request = types.ClientRequest(types.GetTaskRequest(params=types.GetTaskRequestParams(taskId=task_id))) + request = types.ClientRequest(types.GetTaskRequest(params=types.GetTaskRequestParams(task_id=task_id))) try: # GetTaskResult.ttl is required-but-Optional in the SDK; coerce below. lenient = await self._send_with_one_reconnect( request, types.Result, operation="tasks/get", task_id=task_id ) - except McpError as ex: + except MCPError as ex: if ex.error.code in transient_codes: logger.debug("Transient %s on tasks/get for '%s'; will retry.", ex.error.code, task_id) await asyncio.sleep(_MCP_TASK_MIN_POLL_INTERVAL.total_seconds()) @@ -3195,7 +3193,7 @@ async def _poll_task_until_terminal(self, task_id: str) -> types.GetTaskResult: if snapshot.status in _MCP_TASK_TERMINAL_STATUSES: return snapshot - await asyncio.sleep(self._compute_poll_delay(snapshot.pollInterval).total_seconds()) + await asyncio.sleep(self._compute_poll_delay(snapshot.poll_interval).total_seconds()) @staticmethod def _coerce_get_task_result(lenient: types.Result, task_id: str) -> types.GetTaskResult: @@ -3237,7 +3235,7 @@ async def _handle_terminal_task( if status == "completed": payload = await self._fetch_task_result(task_id) _capture_mcp_tool_result(payload) - if payload.isError: + if payload.is_error: parsed = parser(payload) text = ( "\n".join(c.text for c in parsed if c.type == "text" and c.text) @@ -3249,21 +3247,20 @@ async def _handle_terminal_task( # Non-completed terminal statuses surface as ToolExecutionException so the # function-calling loop sees a normal failure for tool_name. - message = snapshot.statusMessage or f"MCP task ended with status '{status}'." + message = snapshot.status_message or f"MCP task ended with status '{status}'." if status == "input_required": # Spec-non-terminal; treated as terminal here because the framework does # not implement the interactive input flow. - message = snapshot.statusMessage or "MCP task requires additional input and cannot continue." + message = snapshot.status_message or "MCP task requires additional input and cannot continue." raise ToolExecutionException(f"Tool '{tool_name}' task {status}: {message}") async def _fetch_task_result(self, task_id: str) -> types.CallToolResult: """Send ``tasks/result`` and reinterpret the open-typed payload as a CallToolResult.""" - from mcp import types - from mcp.shared.exceptions import McpError + from mcp import MCPError, types from pydantic import ValidationError request = types.ClientRequest( - types.GetTaskPayloadRequest(params=types.GetTaskPayloadRequestParams(taskId=task_id)) + types.GetTaskPayloadRequest(params=types.GetTaskPayloadRequestParams(task_id=task_id)) ) # Connection-loss retry only via the helper; no transient-code retry — server # has already completed the task, so a slow payload fetch is anomalous. @@ -3271,7 +3268,7 @@ async def _fetch_task_result(self, task_id: str) -> types.CallToolResult: payload = await self._send_with_one_reconnect( request, types.GetTaskPayloadResult, operation="tasks/result", task_id=task_id ) - except McpError as ex: + except MCPError as ex: # Server reported completed; a hard fetch error is a plain failure (no cancel). raise ToolExecutionException(ex.error.message, inner_exception=ex) from ex @@ -3300,12 +3297,12 @@ async def _send_with_one_reconnect( Non-connection errors propagate unchanged. """ from anyio import ClosedResourceError - from mcp.shared.exceptions import McpError + from mcp import MCPError for attempt in range(_MCP_RECONNECT_ATTEMPTS): try: return await self.session.send_request(request, result_type) # type: ignore[union-attr] - except (ClosedResourceError, McpError) as ex: + except (ClosedResourceError, MCPError) as ex: if not self._is_connection_lost(ex): raise if attempt < _MCP_RECONNECT_ATTEMPTS - 1: @@ -3368,7 +3365,7 @@ async def _try_cancel_task(self, task_id: str) -> None: """ from mcp import types - request = types.ClientRequest(types.CancelTaskRequest(params=types.CancelTaskRequestParams(taskId=task_id))) + request = types.ClientRequest(types.CancelTaskRequest(params=types.CancelTaskRequestParams(task_id=task_id))) try: await asyncio.wait_for( self.session.send_request(request, types.CancelTaskResult), # type: ignore[union-attr] @@ -3393,11 +3390,11 @@ async def _try_cancel_task(self, task_id: str) -> None: def _is_connection_lost(ex: BaseException) -> bool: """Return True if *ex* indicates the MCP transport was torn down.""" from anyio import ClosedResourceError - from mcp.shared.exceptions import McpError + from mcp import MCPError if isinstance(ex, ClosedResourceError): return True - if isinstance(ex, McpError): + if isinstance(ex, MCPError): return "session terminated" in ex.error.message.lower() return False @@ -3427,7 +3424,7 @@ async def get_prompt(self, prompt_name: str, **kwargs: Any) -> str: or the prompt call fails. """ from anyio import ClosedResourceError - from mcp.shared.exceptions import McpError + from mcp import MCPError if not self.load_prompts_flag: raise ToolExecutionException( @@ -3463,7 +3460,7 @@ async def get_prompt(self, prompt_name: str, **kwargs: Any) -> str: f"Failed to call prompt '{prompt_name}' - connection lost.", inner_exception=cl_ex, ) from cl_ex - except McpError as mcp_exc: + except MCPError as mcp_exc: error_message = mcp_exc.error.message set_mcp_span_error(span, type(mcp_exc).__name__, error_message) raise ToolExecutionException(error_message, inner_exception=mcp_exc) from mcp_exc From f42fe1f802993a377580b55249bbd936364afc8b Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 30 Sep 2026 15:25:00 +0200 Subject: [PATCH 04/42] Unit tests for legacy and v2 MCP server --- .../tests/hosting_mcp/test_agent_tool.py | 45 ++++++++++++++++++- .../tests/hosting_mcp/test_workflow_tool.py | 37 ++++++++++++++- 2 files changed, 80 insertions(+), 2 deletions(-) diff --git a/python/packages/hosting-mcp/tests/hosting_mcp/test_agent_tool.py b/python/packages/hosting-mcp/tests/hosting_mcp/test_agent_tool.py index 23426f7e6bb..5dee3253088 100644 --- a/python/packages/hosting-mcp/tests/hosting_mcp/test_agent_tool.py +++ b/python/packages/hosting-mcp/tests/hosting_mcp/test_agent_tool.py @@ -6,6 +6,7 @@ from typing import Any from unittest.mock import AsyncMock, patch +import pytest from agent_framework import ( Agent, BaseChatClient, @@ -17,7 +18,8 @@ ResponseStream, ) from agent_framework_hosting import AgentState -from mcp import MCPError, types +from mcp import Client, MCPError, types +from mcp.server import Server, ServerRequestContext from pytest import raises from agent_framework_hosting_mcp import AgentMCPTool @@ -199,3 +201,44 @@ async def test_agent_tool_propagates_agent_execution_failure() -> None: raises(RuntimeError, match="agent execution failed"), ): await tool.call_tool("agent", {"task": "hello"}) + + +@pytest.mark.parametrize( + ("mode", "expected_version"), + [ + ("auto", "2026-07-28"), + ("legacy", "2025-11-25"), + ], +) +async def test_agent_tool_serves_both_protocol_eras( + mode: str, + expected_version: str, +) -> None: + agent = Agent(client=RecordingClient(), name="agent") + agent_tool: AgentMCPTool[Any] = AgentMCPTool(agent) + + async def list_tools( + _ctx: ServerRequestContext[dict[str, Any]], _params: types.PaginatedRequestParams | None + ) -> types.ListToolsResult: + return await agent_tool.list_tools() + + async def call_tool( + _ctx: ServerRequestContext[dict[str, Any]], params: types.CallToolRequestParams + ) -> types.CallToolResult: + return await agent_tool.call_tool(params.name, params.arguments or {}) + + server = Server( + "test-server", + on_list_tools=list_tools, + on_call_tool=call_tool, + ) + + async with Client(server, mode=mode) as mcp_client: + tools = await mcp_client.list_tools() + result = await mcp_client.call_tool("agent", {"task": "hello"}) + + assert mcp_client.protocol_version == expected_version + assert [tool.name for tool in tools.tools] == ["agent"] + assert result.result_type == "complete" + assert not result.is_error + assert result.content diff --git a/python/packages/hosting-mcp/tests/hosting_mcp/test_workflow_tool.py b/python/packages/hosting-mcp/tests/hosting_mcp/test_workflow_tool.py index 08737904cae..b5c5ffbce05 100644 --- a/python/packages/hosting-mcp/tests/hosting_mcp/test_workflow_tool.py +++ b/python/packages/hosting-mcp/tests/hosting_mcp/test_workflow_tool.py @@ -6,6 +6,7 @@ from typing import Any from unittest.mock import AsyncMock, patch +import pytest from agent_framework import ( Executor, WorkflowBuilder, @@ -16,7 +17,8 @@ handler, ) from agent_framework_hosting import WorkflowState -from mcp import MCPError, types +from mcp import Client, MCPError, types +from mcp.server import Server, ServerRequestContext from pytest import raises from agent_framework_hosting_mcp import WorkflowMCPTool @@ -159,3 +161,36 @@ async def test_workflow_tool_propagates_execution_failure() -> None: raises(RuntimeError, match="workflow execution failed"), ): await tool.call_tool("repeat_text", {"text": "go", "repeat": 2}) + + +@pytest.mark.parametrize( + ("mode", "expected_version"), + [ + ("auto", "2026-07-28"), + ("legacy", "2025-11-25"), + ], +) +async def test_workflow_tool_serves_both_protocol_eras(mode: str, expected_version: str) -> None: + workflow_tool: WorkflowMCPTool[Any] = WorkflowMCPTool(create_workflow(), name="repeat_text") + + async def list_tools( + _ctx: ServerRequestContext[dict[str, Any]], _params: types.PaginatedRequestParams | None + ) -> types.ListToolsResult: + return await workflow_tool.list_tools() + + async def call_tool( + _ctx: ServerRequestContext[dict[str, Any]], params: types.CallToolRequestParams + ) -> types.CallToolResult: + return await workflow_tool.call_tool(params.name, params.arguments or {}) + + server = Server("test-server", on_list_tools=list_tools, on_call_tool=call_tool) + + async with Client(server, mode=mode) as mcp_client: + tools = await mcp_client.list_tools() + result = await mcp_client.call_tool("repeat_text", {"text": "go", "repeat": 2}) + + assert mcp_client.protocol_version == expected_version + assert [tool.name for tool in tools.tools] == ["repeat_text"] + assert result.result_type == "complete" + assert not result.is_error + assert result.content == [types.TextContent(type="text", text="gogo")] From 215d465e753dbba395b3bb5c575f547f727b5507 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 30 Sep 2026 15:37:50 +0200 Subject: [PATCH 05/42] Using MCPError in samples for consistency with internal behaviour --- python/samples/04-hosting/mcp/manual_app.py | 4 ++-- python/samples/04-hosting/mcp/session_app.py | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/python/samples/04-hosting/mcp/manual_app.py b/python/samples/04-hosting/mcp/manual_app.py index 6f4d4e3a4b4..5172ecff968 100644 --- a/python/samples/04-hosting/mcp/manual_app.py +++ b/python/samples/04-hosting/mcp/manual_app.py @@ -32,7 +32,7 @@ from agent_framework.foundry import FoundryChatClient from agent_framework_hosting_mcp import mcp_from_run, mcp_to_run from azure.identity.aio import DefaultAzureCredential -from mcp import types +from mcp import MCPError, types from mcp.server import ServerRequestContext from mcp.server.lowlevel import Server from mcp.server.streamable_http_manager import StreamableHTTPSessionManager @@ -77,7 +77,7 @@ async def call_tool(_ctx: ServerRequestContext[dict[str, Any]], params: types.Ca """Convert, run, and render without the agent-backed adapter.""" if name != "run_agent_manually": - raise ValueError(f"Unknown MCP tool: {name}") + raise MCPError(types.INVALID_PARAMS, f"Unknown MCP tool: {name}") run = mcp_to_run( arguments, argument_name=TASK_ARGUMENT, diff --git a/python/samples/04-hosting/mcp/session_app.py b/python/samples/04-hosting/mcp/session_app.py index 5ee74513c8d..76107ecd04d 100644 --- a/python/samples/04-hosting/mcp/session_app.py +++ b/python/samples/04-hosting/mcp/session_app.py @@ -44,7 +44,7 @@ from agent_framework_hosting import AgentState from agent_framework_hosting_mcp import AgentMCPTool from azure.identity.aio import DefaultAzureCredential -from mcp import types +from mcp import MCPError, types from mcp.server import ServerRequestContext from mcp.server.lowlevel import Server from mcp.server.streamable_http_manager import StreamableHTTPSessionManager @@ -63,7 +63,7 @@ async def call_tool(_ctx: ServerRequestContext[dict[str, Any]], params: types.Ca """Serialize calls per app-owned session before using ``AgentState``.""" session_id = arguments.get("session_id") if arguments else None if not isinstance(session_id, str) or not session_id: - raise ValueError("MCP tool argument 'session_id' must be a non-empty string.") + raise MCPError(types.INVALID_PARAMS, "MCP tool argument 'session_id' must be a non-empty string.") lock = session_locks.setdefault(session_id, asyncio.Lock()) async with lock: return await agent_tool.call_tool(name, arguments) From b21504569cfb9c3a3fed5f7024aee99fc9f88f9f Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 30 Sep 2026 16:33:16 +0200 Subject: [PATCH 06/42] Migrated test_mcp to use new types --- python/packages/core/tests/core/test_mcp.py | 563 ++++++++++---------- 1 file changed, 271 insertions(+), 292 deletions(-) diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 367bdcdc7f5..40971e99e36 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -15,9 +15,8 @@ from unittest.mock import AsyncMock, Mock, patch import pytest -from mcp import types +from mcp import MCPError, types from mcp.client.session import ClientSession -from mcp.shared.exceptions import McpError from pydantic import AnyUrl, BaseModel from agent_framework import ( @@ -191,10 +190,10 @@ async def test_load_tools_with_tool_name_prefix_preserves_matching_configuration types.Tool( name="search_docs", description="Search docs", - inputSchema={"type": "object", "properties": {"query": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"query": {"type": "string"}}}, ), ] - page.nextCursor = None + page.next_cursor = None mock_session.list_tools = AsyncMock(return_value=page) await tool.load_tools() @@ -233,14 +232,14 @@ async def test_load_tools_rejects_ambiguous_policy_names( use_progressive_disclosure=progressive, ) advertised_tools = [ - types.Tool(name=name, inputSchema={"type": "object", "properties": {}}) for name in remote_names + types.Tool(name=name, input_schema={"type": "object", "properties": {}}) for name in remote_names ] async def list_tools(params: types.PaginatedRequestParams | None = None) -> types.ListToolsResult: assert tool._functions == [] if paginated: if params is None: - return types.ListToolsResult(tools=advertised_tools[:1], nextCursor="second") + return types.ListToolsResult(tools=advertised_tools[:1], next_cursor="second") assert params.cursor == "second" return types.ListToolsResult(tools=advertised_tools[1:]) return types.ListToolsResult(tools=advertised_tools) @@ -280,7 +279,7 @@ async def test_ambiguous_policy_reload_preserves_previous_discovery( tool.session = AsyncMock() original = types.Tool( name=remote_names[0], - inputSchema={"type": "object", "properties": {"query": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"query": {"type": "string"}}}, _meta={"original": True}, ) tool.session.list_tools = AsyncMock(return_value=types.ListToolsResult(tools=[original])) @@ -294,14 +293,14 @@ async def list_tools(params: types.PaginatedRequestParams | None = None) -> type assert tool._functions == original_functions if params is None: return types.ListToolsResult( - tools=[original, types.Tool(name="extra", inputSchema={"type": "object"})], - nextCursor="second", + tools=[original, types.Tool(name="extra", input_schema={"type": "object"})], + next_cursor="second", ) - return types.ListToolsResult(tools=[types.Tool(name=remote_names[1], inputSchema={"type": "object"})]) + return types.ListToolsResult(tools=[types.Tool(name=remote_names[1], input_schema={"type": "object"})]) tool.session.list_tools = AsyncMock(side_effect=list_tools) with caplog.at_level(logging.WARNING, logger=logger.name): - await tool.message_handler(types.ServerNotification(types.ToolListChangedNotification())) + await tool.message_handler(types.ToolListChangedNotification()) await asyncio.gather(*tool._pending_reload_tasks) assert "configuration name 'docs_search' is ambiguous" in caplog.text @@ -334,12 +333,12 @@ async def test_tool_refresh_accepts_unambiguous_rename( tool.session = AsyncMock() tool.session.list_tools = AsyncMock( side_effect=[ - types.ListToolsResult(tools=[types.Tool(name=name, inputSchema={"type": "object"})]) + types.ListToolsResult(tools=[types.Tool(name=name, input_schema={"type": "object"})]) for name in remote_names ] ) await tool.load_tools() - await tool.message_handler(types.ServerNotification(types.ToolListChangedNotification())) + await tool.message_handler(types.ToolListChangedNotification()) await asyncio.gather(*tool._pending_reload_tasks) assert [function.name for function in tool.functions] == [f"docs_{remote_names[1]}"] @@ -357,14 +356,14 @@ async def test_tool_refresh_replaces_snapshot_and_preserves_other_functions(empt tool.session = AsyncMock() keep = types.Tool( name="keep", - inputSchema={"type": "object", "properties": {"query": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"query": {"type": "string"}}}, _meta={"version": 1}, ) removed = types.Tool( name="removed", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, _meta={"version": 1}, - execution=types.ToolExecution(taskSupport="required"), + execution=types.ToolExecution(task_support="required"), ) tool.session.list_tools = AsyncMock(return_value=types.ListToolsResult(tools=[keep, removed])) await tool.load_tools() @@ -385,7 +384,7 @@ async def list_tools(params: types.PaginatedRequestParams | None = None) -> type assert tool._functions == original_functions assert tool._tool_call_meta_by_name == original_meta if params is None: - return types.ListToolsResult(tools=[], nextCursor="second") + return types.ListToolsResult(tools=[], next_cursor="second") assert params.cursor == "second" return types.ListToolsResult(tools=[] if empty_snapshot else [keep]) @@ -414,7 +413,7 @@ async def test_tool_refresh_preserves_prompt_when_server_advertises_same_raw_nam prompt_function = tool._functions[0] tool.session.list_tools = AsyncMock( side_effect=[ - types.ListToolsResult(tools=[types.Tool(name="summary", inputSchema={"type": "object"})]), + types.ListToolsResult(tools=[types.Tool(name="summary", input_schema={"type": "object"})]), types.ListToolsResult(tools=[]), ] ) @@ -428,7 +427,7 @@ async def test_tool_refresh_preserves_prompt_when_server_advertises_same_raw_nam async def test_tool_refresh_forgets_removed_progressive_tools(replacement_name: str | None) -> None: tool = await _load_progressive_test_server( tool_name_prefix="docs", - tools=[types.Tool(name=name, inputSchema={"type": "object"}) for name in ("search/docs", "keep")], + tools=[types.Tool(name=name, input_schema={"type": "object"}) for name in ("search/docs", "keep")], ) loader = tool.functions[1] context = FunctionInvocationContext(function=loader, arguments={}, tools=list(tool.functions)) @@ -441,7 +440,7 @@ async def test_tool_refresh_forgets_removed_progressive_tools(replacement_name: "list_tools", return_value=types.ListToolsResult( tools=[ - types.Tool(name=name, inputSchema={"type": "object"}) + types.Tool(name=name, input_schema={"type": "object"}) for name in (["keep", replacement_name] if replacement_name else ["keep"]) ] ), @@ -486,7 +485,7 @@ async def test_overlapping_names_accept_unambiguous_policy( tool.session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ - types.Tool(name=name, inputSchema={"type": "object", "properties": {}}) + types.Tool(name=name, input_schema={"type": "object", "properties": {}}) for name in ("search", "docs_search") ] ) @@ -511,7 +510,7 @@ async def test_prompts_reject_ambiguous_policy_names(load_order: str) -> None: tool.session = AsyncMock() tool.session.list_prompts = AsyncMock( side_effect=[ - types.ListPromptsResult(prompts=[types.Prompt(name="search")], nextCursor="second"), + types.ListPromptsResult(prompts=[types.Prompt(name="search")], next_cursor="second"), types.ListPromptsResult(prompts=[types.Prompt(name="docs_search")]), ] ) @@ -523,7 +522,7 @@ async def test_prompts_reject_ambiguous_policy_names(load_order: str) -> None: tool.session.list_tools = AsyncMock( return_value=types.ListToolsResult( - tools=[types.Tool(name="search", inputSchema={"type": "object", "properties": {}})] + tools=[types.Tool(name="search", input_schema={"type": "object", "properties": {}})] ) ) tool.session.list_prompts = AsyncMock( @@ -554,10 +553,10 @@ async def test_allowed_tools_does_not_authorize_normalized_remote_name_collision types.Tool( name="delete/file", description="Delete a file", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), ] - page.nextCursor = None + page.next_cursor = None mock_session.list_tools = AsyncMock(return_value=page) await tool.load_tools() @@ -579,15 +578,15 @@ async def test_load_tools_rejects_colliding_normalized_tool_names() -> None: types.Tool( name="delete/file", description="Unauthorized tool", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), types.Tool( name="delete-file", description="Authorized tool", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), ] - page.nextCursor = None + page.next_cursor = None mock_session.list_tools = AsyncMock(return_value=page) with pytest.raises(ToolExecutionException, match="map to the same local function name"): @@ -607,10 +606,10 @@ async def test_allowed_tools_exact_raw_name_allows_normalized_function_name() -> types.Tool( name="delete/file", description="Delete a file", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), ] - page.nextCursor = None + page.next_cursor = None mock_session.list_tools = AsyncMock(return_value=page) await tool.load_tools() @@ -636,10 +635,10 @@ async def test_approval_mode_does_not_match_normalized_colliding_name() -> None: types.Tool( name="delete/file", description="Delete a file", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), ] - page.nextCursor = None + page.next_cursor = None mock_session.list_tools = AsyncMock(return_value=page) await tool.load_tools() @@ -664,7 +663,7 @@ async def test_load_prompts_with_tool_name_prefix() -> None: arguments=[types.PromptArgument(name="topic", description="Topic", required=True)], ), ] - page.nextCursor = None + page.next_cursor = None mock_session.list_prompts = AsyncMock(return_value=page) await tool.load_prompts() @@ -692,11 +691,11 @@ def test_mcp_tool_str_and_parse_prompt_result_rich_content() -> None: types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hello")), types.PromptMessage( role="assistant", - content=types.ImageContent(type="image", data="eHl6", mimeType="image/png"), + content=types.ImageContent(type="image", data="eHl6", mime_type="image/png"), ), types.PromptMessage( role="assistant", - content=types.AudioContent(type="audio", data="YXVkaW8=", mimeType="audio/wav"), + content=types.AudioContent(type="audio", data="YXVkaW8=", mime_type="audio/wav"), ), types.PromptMessage( role="assistant", @@ -704,7 +703,7 @@ def test_mcp_tool_str_and_parse_prompt_result_rich_content() -> None: type="resource", resource=types.TextResourceContents( uri=AnyUrl("file://prompt.txt"), - mimeType="text/plain", + mime_type="text/plain", text="Embedded prompt", ), ), @@ -715,7 +714,7 @@ def test_mcp_tool_str_and_parse_prompt_result_rich_content() -> None: type="resource", resource=types.BlobResourceContents( uri=AnyUrl("file://prompt.bin"), - mimeType="application/pdf", + mime_type="application/pdf", blob="ZGF0YQ==", ), ), @@ -739,9 +738,9 @@ def test_parse_tool_result_from_mcp(): mcp_result = types.CallToolResult( content=[ types.TextContent(type="text", text="Result text"), - types.ImageContent(type="image", data="eHl6", mimeType="image/png"), + types.ImageContent(type="image", data="eHl6", mime_type="image/png"), types.TextContent(type="text", text="After image"), - types.ImageContent(type="image", data="YWJj", mimeType="image/webp"), + types.ImageContent(type="image", data="YWJj", mime_type="image/webp"), ] ) result = _HELPER_MCP_TOOL._parse_tool_result_from_mcp(mcp_result) @@ -804,7 +803,7 @@ def test_parse_tool_result_from_mcp_audio_content(): """Test conversion from MCP tool result with audio returns rich content list.""" mcp_result = types.CallToolResult( content=[ - types.AudioContent(type="audio", data="YXVkaW8=", mimeType="audio/wav"), + types.AudioContent(type="audio", data="YXVkaW8=", mime_type="audio/wav"), ] ) result = _HELPER_MCP_TOOL._parse_tool_result_from_mcp(mcp_result) @@ -824,7 +823,7 @@ def test_parse_tool_result_from_mcp_blob_plain_base64(): type="resource", resource=types.BlobResourceContents( uri=AnyUrl("file://test.bin"), - mimeType="application/pdf", + mime_type="application/pdf", blob="dGVzdCBkYXRh", ), ), @@ -847,13 +846,13 @@ def test_parse_tool_result_from_mcp_resource_link_text_resource_and_unknown(): type="resource_link", uri=AnyUrl("https://example.com/resource"), name="resource", - mimeType="application/json", + mime_type="application/json", ), types.EmbeddedResource( type="resource", resource=types.TextResourceContents( uri=AnyUrl("file://prompt.txt"), - mimeType="text/plain", + mime_type="text/plain", text="Embedded result", ), ), @@ -872,7 +871,7 @@ def test_parse_tool_result_from_mcp_structured_content_only(): """Test that structuredContent is parsed when content list is empty.""" mcp_result = types.CallToolResult( content=[], - structuredContent={"Tables": [{"Name": "Sales", "Columns": ["Amount", "Date"]}]}, + structured_content={"Tables": [{"Name": "Sales", "Columns": ["Amount", "Date"]}]}, ) result = _HELPER_MCP_TOOL._parse_tool_result_from_mcp(mcp_result) @@ -888,7 +887,7 @@ def test_parse_tool_result_from_mcp_structured_content_with_text(): """Default structured_first prefers structuredContent when both are present (#7866).""" mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="Summary")], - structuredContent={"data": [1, 2, 3]}, + structured_content={"data": [1, 2, 3]}, ) result = _HELPER_MCP_TOOL._parse_tool_result_from_mcp(mcp_result) @@ -903,7 +902,7 @@ def test_parse_tool_result_content_modes_for_complementary_and_duplicate_payload """tool_result_content covers structured-only, content-only, either-first, and both.""" mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="Summary")], - structuredContent={"data": [1, 2, 3]}, + structured_content={"data": [1, 2, 3]}, ) content_first = MCPTool(name="helper", tool_result_content="content_first") # type: ignore[abstract] # ty: ignore[call-non-callable] @@ -923,12 +922,12 @@ def test_parse_tool_result_content_modes_for_complementary_and_duplicate_payload assert both_result[1].text is not None assert json.loads(both_result[1].text) == {"data": [1, 2, 3]} - structured_only_empty = types.CallToolResult(content=[], structuredContent={"x": 1}) + structured_only_empty = types.CallToolResult(content=[], structured_content={"x": 1}) empty_structured_text = content_first._parse_tool_result_from_mcp(structured_only_empty)[0].text assert empty_structured_text is not None assert json.loads(empty_structured_text) == {"x": 1} - empty = types.CallToolResult(content=[], structuredContent=None) + empty = types.CallToolResult(content=[], structured_content=None) assert content_only._parse_tool_result_from_mcp(empty)[0].text == "null" @@ -940,8 +939,8 @@ async def test_generated_mcp_tool_preserves_complete_host_payload_once() -> None } mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="Summary", _meta=file_meta)], - structuredContent={"image_url": "https://example.test/widget.png"}, - isError=False, + structured_content={"image_url": "https://example.test/widget.png"}, + is_error=False, _meta={"widget": "image"}, ) tool = MCPTool(name="helper") # type: ignore[abstract] # ty: ignore[call-non-callable] @@ -995,7 +994,7 @@ async def test_custom_mcp_result_parser_preserves_direct_shape_and_generated_hos """A custom parser controls model content while generated calls retain the Host payload.""" mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="Server summary")], - structuredContent={"image_url": "https://example.test/widget.png"}, + structured_content={"image_url": "https://example.test/widget.png"}, _meta={"source": "server"}, ) tool = MCPTool(name="helper", parse_tool_results=lambda _: "Custom model summary") # type: ignore[abstract] # ty: ignore[call-non-callable] @@ -1021,7 +1020,7 @@ async def test_oversized_mcp_host_payload_is_omitted_without_changing_model_resu """An oversized Host payload is omitted while bounded model content and metadata survive.""" mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="Server summary")], - structuredContent={"widget_data": "x" * 1024}, + structured_content={"widget_data": "x" * 1024}, _meta={"source": "oversized"}, ) tool = MCPTool( # type: ignore[abstract] # ty: ignore[call-non-callable] @@ -1056,7 +1055,7 @@ def test_mcp_host_payload_size_boundary_uri_serialization_and_early_abort( ) -> None: mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text='Escaped "\n\u2603" text')], - structuredContent={"widget_data": "x" * 1024}, + structured_content={"widget_data": "x" * 1024}, ) encoded_size = len(json.dumps(mcp_result.model_dump(by_alias=True, exclude_none=True)).encode("utf-8")) @@ -1090,8 +1089,8 @@ async def test_generated_mcp_error_preserves_complete_host_payload_on_function_r """An MCP error keeps its Host payload after generic function error conversion.""" mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="Widget failed")], - structuredContent={"reason": "invalid input"}, - isError=True, + structured_content={"reason": "invalid input"}, + is_error=True, _meta={"source": "server"}, ) tool = MCPTool(name="helper") # type: ignore[abstract] # ty: ignore[call-non-callable] @@ -1112,7 +1111,7 @@ async def test_generated_mcp_parser_failure_preserves_complete_host_payload_on_f """Capture raw MCP data before a custom parser can fail.""" mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="Server summary")], - structuredContent={"widget": "complete"}, + structured_content={"widget": "complete"}, _meta={"source": "server"}, ) tool = MCPTool(name="helper", parse_tool_results=_raise_result_parser) # type: ignore[abstract] # ty: ignore[call-non-callable] @@ -1131,7 +1130,7 @@ async def test_generated_mcp_parser_failure_preserves_complete_host_payload_on_f async def test_direct_mcp_calls_do_not_materialize_host_payload(monkeypatch: pytest.MonkeyPatch) -> None: """Public direct success and error calls retain their established behavior.""" success = types.CallToolResult(content=[types.TextContent(type="text", text="ok")]) - error = types.CallToolResult(content=[types.TextContent(type="text", text="failed")], isError=True) + error = types.CallToolResult(content=[types.TextContent(type="text", text="failed")], is_error=True) tool = MCPTool( # type: ignore[abstract] # ty: ignore[call-non-callable] name="helper", parse_tool_results=lambda result: cast(types.TextContent, result.content[0]).text, @@ -1173,7 +1172,7 @@ def fail_if_captured(*_args: Any, **_kwargs: Any) -> Any: async def test_function_tool_result_parser_cannot_discard_mcp_host_payload() -> None: mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="server projection")], - structuredContent={"widget": "complete"}, + structured_content={"widget": "complete"}, _meta={"source": "server"}, ) tool = MCPTool(name="helper", parse_tool_results=lambda _: "MCP parser projection") # type: ignore[abstract] # ty: ignore[call-non-callable] @@ -1199,7 +1198,7 @@ async def test_function_tool_result_parser_cannot_discard_mcp_host_payload() -> async def test_empty_custom_parser_projection_remains_empty(parser_layer: str) -> None: mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="server projection")], - structuredContent={"widget": "complete"}, + structured_content={"widget": "complete"}, ) tool = MCPTool( # type: ignore[abstract] # ty: ignore[call-non-callable] name="helper", @@ -1224,8 +1223,8 @@ async def test_empty_custom_parser_projection_remains_empty(parser_layer: str) - async def test_oversized_mcp_error_preserves_independently_bounded_meta() -> None: mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="failed")], - structuredContent={"large": "x" * 1024}, - isError=True, + structured_content={"large": "x" * 1024}, + is_error=True, _meta={"source": "small"}, ) tool = MCPTool(name="helper", max_host_payload_size_bytes=128) # type: ignore[abstract] # ty: ignore[call-non-callable] @@ -1316,7 +1315,7 @@ async def test_mcp_host_payload_survives_real_function_loop( ) -> None: mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="model projection")], - structuredContent={"widget": "complete"}, + structured_content={"widget": "complete"}, _meta={"source": "server"}, ) tool = MCPTool(name="helper") # type: ignore[abstract] # ty: ignore[call-non-callable] @@ -1420,7 +1419,7 @@ async def test_mcp_host_payload_has_aggregate_request_budget( async def call_tool(tool_name: str, **_kwargs: Any) -> types.CallToolResult: return types.CallToolResult( content=[types.TextContent(type="text", text=tool_name)], - structuredContent={"data": tool_name * 120}, + structured_content={"data": tool_name * 120}, _meta={"source": tool_name * 12}, ) @@ -1480,7 +1479,7 @@ async def test_secure_mcp_auto_hide_preserves_outer_host_payload() -> None: mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="untrusted payload")], - structuredContent={"widget": "complete"}, + structured_content={"widget": "complete"}, _meta={"ifc": {"integrity": "trusted", "confidentiality": "public"}}, ) tool = MCPTool( # type: ignore[abstract] # ty: ignore[call-non-callable] @@ -1540,7 +1539,7 @@ async def test_secure_mcp_builtin_parser_restricts_all_result_shapes(result_shap content=[types.TextContent(type="text", text="server trusted payload")] if result_shape in ("content", "both") else [], - structuredContent={"payload": "server trusted structured payload"} + structured_content={"payload": "server trusted structured payload"} if result_shape in ("structured", "both") else None, _meta={"ifc": {"integrity": "trusted", "confidentiality": "public"}}, @@ -1639,7 +1638,7 @@ def test_parse_tool_result_from_mcp_structured_content_none(): """Test that None structuredContent does not affect results.""" mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="Hello")], - structuredContent=None, + structured_content=None, ) result = _HELPER_MCP_TOOL._parse_tool_result_from_mcp(mcp_result) @@ -1653,7 +1652,7 @@ def test_parse_tool_result_from_mcp_structured_content_non_serializable(): """Test that non-JSON-serializable values in structuredContent degrade gracefully.""" mcp_result = types.CallToolResult( content=[], - structuredContent={"data": b"raw bytes", "count": 42}, + structured_content={"data": b"raw bytes", "count": 42}, ) result = _HELPER_MCP_TOOL._parse_tool_result_from_mcp(mcp_result) @@ -1680,7 +1679,7 @@ def test_mcp_content_types_to_ai_content_text(): def test_mcp_content_types_to_ai_content_image(): """Test conversion of MCP image content to AI content.""" # MCP can send data as base64 string or as bytes - mcp_content = types.ImageContent(type="image", data="YWJj", mimeType="image/jpeg") # base64 for b"abc" + mcp_content = types.ImageContent(type="image", data="YWJj", mime_type="image/jpeg") # base64 for b"abc" ai_content = _HELPER_MCP_TOOL._parse_content_from_mcp(mcp_content)[0] assert ai_content.type == "data" @@ -1692,7 +1691,7 @@ def test_mcp_content_types_to_ai_content_image(): def test_mcp_content_types_to_ai_content_audio(): """Test conversion of MCP audio content to AI content.""" # Use properly padded base64 - mcp_content = types.AudioContent(type="audio", data="ZGVm", mimeType="audio/wav") # base64 for b"def" + mcp_content = types.AudioContent(type="audio", data="ZGVm", mime_type="audio/wav") # base64 for b"def" ai_content = _HELPER_MCP_TOOL._parse_content_from_mcp(mcp_content)[0] assert ai_content.type == "data" @@ -1707,7 +1706,7 @@ def test_mcp_content_types_to_ai_content_resource_link(): type="resource_link", uri=AnyUrl("https://example.com/resource"), name="test_resource", - mimeType="application/json", + mime_type="application/json", ) ai_content = _HELPER_MCP_TOOL._parse_content_from_mcp(mcp_content)[0] @@ -1721,7 +1720,7 @@ def test_mcp_content_types_to_ai_content_embedded_resource_text(): """Test conversion of MCP embedded text resource to AI content.""" text_resource = types.TextResourceContents( uri=AnyUrl("file://test.txt"), - mimeType="text/plain", + mime_type="text/plain", text="Embedded text content", ) mcp_content = types.EmbeddedResource(type="resource", resource=text_resource) @@ -1737,7 +1736,7 @@ def test_mcp_content_types_to_ai_content_embedded_resource_blob(): # Use a proper data URI in the blob field since that's what the MCP implementation expects blob_resource = types.BlobResourceContents( uri=AnyUrl("file://test.bin"), - mimeType="application/octet-stream", + mime_type="application/octet-stream", blob="data:application/octet-stream;base64,dGVzdCBkYXRh", ) mcp_content = types.EmbeddedResource(type="resource", resource=blob_resource) @@ -1754,9 +1753,9 @@ def test_mcp_content_types_to_ai_content_tool_use_and_tool_result(): tool_use_content = types.ToolUseContent(type="tool_use", id="call-1", name="calculator", input={"x": 1}) tool_result_content = types.ToolResultContent( type="tool_result", - toolUseId="call-1", + tool_use_id="call-1", content=[types.TextContent(type="text", text="done")], - isError=True, + is_error=True, ) function_call = _HELPER_MCP_TOOL._parse_content_from_mcp(tool_use_content)[0] @@ -1790,7 +1789,7 @@ def test_ai_content_to_mcp_content_types_data_image(): assert isinstance(mcp_content, types.ImageContent) assert mcp_content.type == "image" assert mcp_content.data == "data:image/png;base64,xyz" - assert mcp_content.mimeType == "image/png" + assert mcp_content.mime_type == "image/png" def test_ai_content_to_mcp_content_types_data_audio(): @@ -1801,7 +1800,7 @@ def test_ai_content_to_mcp_content_types_data_audio(): assert isinstance(mcp_content, types.AudioContent) assert mcp_content.type == "audio" assert mcp_content.data == "data:audio/mpeg;base64,xyz" - assert mcp_content.mimeType == "audio/mpeg" + assert mcp_content.mime_type == "audio/mpeg" def test_ai_content_to_mcp_content_types_data_binary(): @@ -1815,7 +1814,7 @@ def test_ai_content_to_mcp_content_types_data_binary(): assert isinstance(mcp_content, types.EmbeddedResource) assert mcp_content.type == "resource" assert mcp_content.resource.blob == "data:application/octet-stream;base64,xyz" # type: ignore[union-attr] # ty: ignore[unresolved-attribute] - assert mcp_content.resource.mimeType == "application/octet-stream" + assert mcp_content.resource.mime_type == "application/octet-stream" def test_ai_content_to_mcp_content_types_uri(): @@ -1826,7 +1825,7 @@ def test_ai_content_to_mcp_content_types_uri(): assert isinstance(mcp_content, types.ResourceLink) assert mcp_content.type == "resource_link" assert str(mcp_content.uri) == "https://example.com/resource" - assert mcp_content.mimeType == "application/json" + assert mcp_content.mime_type == "application/json" def test_prepare_message_for_mcp(): @@ -2178,8 +2177,8 @@ def test_get_input_model_from_mcp_tool_parametrized(test_id: str, input_schema: - test_id: A descriptive name for the test case - input_schema: The JSON schema (inputSchema dict) """ - tool = types.Tool(name="test_tool", description="A test tool", inputSchema=input_schema) - schema = tool.inputSchema + tool = types.Tool(name="test_tool", description="A test tool", input_schema=input_schema) + schema = tool.input_schema # Verify schema is returned as-is (dict) assert isinstance(schema, dict), f"Expected dict, got {type(schema)}" @@ -2275,7 +2274,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="test_tool", description="Test tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"], @@ -2339,7 +2338,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="test_tool", description="Test tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"], @@ -2384,7 +2383,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="test_tool", description="Test tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"], @@ -2424,7 +2423,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="get_customer_detail", description="Get customer details", - inputSchema={ + input_schema={ "type": "object", "properties": { "params": { @@ -2477,7 +2476,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="test_tool", description="Test tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"], @@ -2487,9 +2486,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ) ) # Mock a tool call that raises an MCP error - self.session.call_tool = AsyncMock( - side_effect=McpError(types.ErrorData(code=-1, message="Tool execution failed")) - ) + self.session.call_tool = AsyncMock(side_effect=MCPError(-1, "Tool execution failed")) def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: return None # type: ignore[return-value] # pyrefly: ignore[bad-return] # ty: ignore[invalid-return-type] @@ -2517,9 +2514,7 @@ async def connect(self, *, reset: bool = False) -> None: self.session = Mock(spec=ClientSession) self.sessions.append(self.session) if self.connect_count == 1: - self.session.call_tool = AsyncMock( - side_effect=McpError(types.ErrorData(code=-32000, message="Session terminated")) - ) + self.session.call_tool = AsyncMock(side_effect=MCPError(-32000, "Session terminated")) else: self.session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="recovered")]) @@ -2541,7 +2536,7 @@ def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: async def test_mcp_tool_call_tool_raises_on_is_error(): - """Test that call_tool raises ToolExecutionException when MCP returns isError=True.""" + """Test that call_tool raises ToolExecutionException when MCP returns is_error=True.""" class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] @@ -2552,7 +2547,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="test_tool", description="Test tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"], @@ -2564,7 +2559,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri self.session.call_tool = AsyncMock( return_value=types.CallToolResult( content=[types.TextContent(type="text", text="Something went wrong")], - isError=True, + is_error=True, ) ) @@ -2581,7 +2576,7 @@ def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: async def test_mcp_tool_call_tool_succeeds_when_is_error_false(): - """Test that call_tool returns normally when MCP returns isError=False.""" + """Test that call_tool returns normally when MCP returns is_error=False.""" class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] @@ -2592,7 +2587,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="test_tool", description="Test tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"], @@ -2604,7 +2599,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri self.session.call_tool = AsyncMock( return_value=types.CallToolResult( content=[types.TextContent(type="text", text="Success")], - isError=False, + is_error=False, ) ) @@ -2621,7 +2616,7 @@ def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: async def test_mcp_tool_is_error_propagates_through_function_middleware(): - """Test that MCP isError=True propagates as ToolExecutionException through function middleware.""" + """Test that MCP is_error=True propagates as ToolExecutionException through function middleware.""" error_seen_in_middleware = False class ErrorCheckMiddleware(FunctionMiddleware): @@ -2642,7 +2637,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="test_tool", description="Test tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"], @@ -2654,7 +2649,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri self.session.call_tool = AsyncMock( return_value=types.CallToolResult( content=[types.TextContent(type="text", text="MCP error occurred")], - isError=True, + is_error=True, ) ) @@ -2757,7 +2752,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="tool_one", description="First tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, }, @@ -2765,7 +2760,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="tool_two", description="Second tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, }, @@ -2834,7 +2829,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="tool_one", description="First tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, }, @@ -2842,7 +2837,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="tool_two", description="Second tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, }, @@ -2850,7 +2845,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="tool_three", description="Third tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, }, @@ -2948,7 +2943,7 @@ def _progressive_tool_list_page(*, tools: list[types.Tool] | None = None) -> typ types.Tool( name="tool_one", description="First tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"], @@ -2957,7 +2952,7 @@ def _progressive_tool_list_page(*, tools: list[types.Tool] | None = None) -> typ types.Tool( name="tool_two", description="Second tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"value": {"type": "integer"}}, "required": ["value"], @@ -2966,7 +2961,7 @@ def _progressive_tool_list_page(*, tools: list[types.Tool] | None = None) -> typ types.Tool( name="secret_tool", description="Secret tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"secret": {"type": "string"}}, }, @@ -3015,12 +3010,12 @@ async def test_mcp_progressive_disclosure_filters_always_loaded_loader_name_coll types.Tool( name="list_mcp_tools", description="Remote tool whose local name collides with the list loader.", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), types.Tool( name="tool_one", description="First tool", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), ], ) @@ -3068,12 +3063,12 @@ async def test_mcp_progressive_list_mcp_tools_skips_loader_name_collisions() -> types.Tool( name="list_mcp_tools", description="Remote tool whose local name collides with the list loader.", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), types.Tool( name="tool_one", description="First tool", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), ], ) @@ -3382,7 +3377,7 @@ async def test_mcp_progressive_load_tool_rejects_loader_name_collision() -> None types.Tool( name="load_tool", description="Remote tool whose local name collides with the load loader.", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), ], ) @@ -3479,7 +3474,7 @@ def provider(kwargs: dict[str, Any]) -> dict[str, str]: types.Tool( name="greet", description="Says hello", - inputSchema={ + input_schema={ "type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"], @@ -3654,9 +3649,7 @@ async def test_mcp_tool_message_handler_notification(): tool.load_prompts = AsyncMock() # type: ignore[method-assign] # Test tools list changed notification - tools_notification = Mock(spec=types.ServerNotification) - tools_notification.root = Mock() - tools_notification.root.method = "notifications/tools/list_changed" + tools_notification = types.ToolListChangedNotification() result = await tool.message_handler(tools_notification) # type: ignore[func-returns-value] assert result is None @@ -3668,9 +3661,7 @@ async def test_mcp_tool_message_handler_notification(): tool.load_tools.reset_mock() # Test prompts list changed notification - prompts_notification = Mock(spec=types.ServerNotification) - prompts_notification.root = Mock() - prompts_notification.root.method = "notifications/prompts/list_changed" + prompts_notification = types.PromptListChangedNotification() result = await tool.message_handler(prompts_notification) # type: ignore[func-returns-value] assert result is None @@ -3678,9 +3669,7 @@ async def test_mcp_tool_message_handler_notification(): tool.load_prompts.assert_called_once() # Test unhandled notification - unknown_notification = Mock(spec=types.ServerNotification) - unknown_notification.root = Mock() - unknown_notification.root.method = "notifications/unknown" + unknown_notification = types.ResourceListChangedNotification() result = await tool.message_handler(unknown_notification) # type: ignore[func-returns-value] assert result is None @@ -3719,9 +3708,7 @@ async def slow_load_tools(): tool.load_tools = slow_load_tools # type: ignore[assignment] # ty: ignore[invalid-assignment] - tools_notification = Mock(spec=types.ServerNotification) - tools_notification.root = Mock() - tools_notification.root.method = "notifications/tools/list_changed" + tools_notification = types.ToolListChangedNotification() # message_handler must return immediately even though load_tools blocks. await tool.message_handler(tools_notification) @@ -3743,9 +3730,7 @@ async def test_mcp_tool_message_handler_reload_failure_is_logged(caplog: pytest. tool = MCPStdioTool(name="test_tool", command="python") tool.load_tools = AsyncMock(side_effect=RuntimeError("connection lost")) # type: ignore[method-assign] - tools_notification = Mock(spec=types.ServerNotification) - tools_notification.root = Mock() - tools_notification.root.method = "notifications/tools/list_changed" + tools_notification = types.ToolListChangedNotification() await tool.message_handler(tools_notification) # Let the background task run — it should not propagate the exception. @@ -3777,9 +3762,7 @@ async def blocking_load_tools(): tool.load_tools = blocking_load_tools # type: ignore[assignment] # ty: ignore[invalid-assignment] - notification = Mock(spec=types.ServerNotification) - notification.root = Mock() - notification.root.method = "notifications/tools/list_changed" + notification = types.ToolListChangedNotification() # First notification — starts a blocking reload task. await tool.message_handler(notification) @@ -3877,12 +3860,12 @@ async def test_mcp_tool_sampling_configuration_warns_once(): params = Mock() params.messages = [] - params.maxTokens = 128 - params.systemPrompt = None + params.max_tokens = 128 + params.system_prompt = None params.tools = None params.temperature = None - params.stopSequences = None - params.toolChoice = None + params.stop_sequences = None + params.tool_choice = None with warnings.catch_warnings(record=True) as caught: warnings.simplefilter("always", DeprecationWarning) @@ -3916,7 +3899,7 @@ async def test_mcp_tool_sampling_callback_denies_by_default(): params = Mock() params.messages = [] - params.maxTokens = 128 + params.max_tokens = 128 result = await _invoke_sampling_callback(tool, params) @@ -3935,7 +3918,7 @@ async def test_mcp_tool_sampling_callback_denied_by_callback(): params = Mock() params.messages = [] - params.maxTokens = 128 + params.max_tokens = 128 result = await _invoke_sampling_callback(tool, params) @@ -3957,7 +3940,7 @@ def boom(_params: object) -> bool: params = Mock() params.messages = [] - params.maxTokens = 128 + params.max_tokens = 128 result = await _invoke_sampling_callback(tool, params) @@ -3980,11 +3963,11 @@ async def approve(_params: object) -> bool: params = Mock() params.messages = [types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hi"))] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = None params.tools = None - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) @@ -4009,11 +3992,11 @@ async def test_mcp_tool_sampling_callback_clamps_max_tokens(): params = Mock() params.messages = [types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hi"))] params.temperature = None - params.maxTokens = 1_000_000 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 1_000_000 + params.stop_sequences = None + params.system_prompt = None params.tools = None - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) @@ -4037,11 +4020,11 @@ async def test_mcp_tool_sampling_callback_does_not_clamp_under_cap(): params = Mock() params.messages = [types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hi"))] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = None params.tools = None - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) @@ -4066,11 +4049,11 @@ def make_params() -> Mock: params = Mock() params.messages = [types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hi"))] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = None params.tools = None - params.toolChoice = None + params.tool_choice = None return params first = await _invoke_sampling_callback(tool, make_params()) @@ -4108,11 +4091,11 @@ async def test_mcp_tool_sampling_callback_chat_client_exception(): mock_message.content.text = "Test question" params.messages = [mock_message] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = None params.tools = None - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) @@ -4154,11 +4137,11 @@ async def test_mcp_tool_sampling_callback_no_valid_content(): mock_message.content.text = "Test question" params.messages = [mock_message] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = None params.tools = None - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) @@ -4178,11 +4161,11 @@ async def test_mcp_tool_sampling_callback_no_response_and_successful_message_cre params = Mock() params.messages = [types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hi"))] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = None params.tools = None - params.toolChoice = None + params.tool_choice = None tool.client.get_response.return_value = None no_response = await _invoke_sampling_callback(tool, params) @@ -4239,21 +4222,21 @@ async def test_mcp_tool_sampling_callback_returns_tool_use_results(): params = Mock() params.messages = [types.PromptMessage(role="user", content=types.TextContent(type="text", text="Answer"))] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = None params.tools = [ - types.Tool(name="Answer", description="Return an answer", inputSchema={"type": "object"}), - types.Tool(name="Citations", description="Return source citations", inputSchema={"type": "object"}), + types.Tool(name="Answer", description="Return an answer", input_schema={"type": "object"}), + types.Tool(name="Citations", description="Return source citations", input_schema={"type": "object"}), ] - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) assert isinstance(result, types.CreateMessageResultWithTools) assert result.role == "assistant" assert result.model == "test-model" - assert result.stopReason == "toolUse" + assert result.stop_reason == "toolUse" assert isinstance(result.content, list) tool_use_contents = [content for content in result.content if isinstance(content, types.ToolUseContent)] assert tool_use_contents == result.content @@ -4293,11 +4276,11 @@ async def test_mcp_tool_sampling_callback_forwards_system_prompt(): mock_message.content.text = "Test question" params.messages = [mock_message] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = "You are a helpful assistant" + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = "You are a helpful assistant" params.tools = None - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) @@ -4324,7 +4307,7 @@ async def test_mcp_tool_sampling_callback_forwards_tools(): mcp_tool = types.Tool( name="get_weather", description="Get weather", - inputSchema={"type": "object", "properties": {"city": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"city": {"type": "string"}}}, ) params = Mock() @@ -4334,11 +4317,11 @@ async def test_mcp_tool_sampling_callback_forwards_tools(): mock_message.content.text = "Test question" params.messages = [mock_message] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = None params.tools = [mcp_tool] - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) @@ -4374,11 +4357,11 @@ async def test_mcp_tool_sampling_callback_forwards_tool_choice(): mock_message.content.text = "Test question" params.messages = [mock_message] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = None params.tools = None - params.toolChoice = types.ToolChoice(mode="required") + params.tool_choice = types.ToolChoice(mode="required") result = await _invoke_sampling_callback(tool, params) @@ -4409,11 +4392,11 @@ async def test_mcp_tool_sampling_callback_forwards_empty_system_prompt(): mock_message.content.text = "Test question" params.messages = [mock_message] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = "" + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = "" params.tools = None - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) @@ -4444,11 +4427,11 @@ async def test_mcp_tool_sampling_callback_forwards_empty_tools_list(): mock_message.content.text = "Test question" params.messages = [mock_message] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = None params.tools = [] - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) @@ -4479,11 +4462,11 @@ async def test_mcp_tool_sampling_callback_forwards_generation_params_in_options( mock_message.content.text = "Test question" params.messages = [mock_message] params.temperature = 0.7 - params.maxTokens = 256 - params.stopSequences = ["STOP"] - params.systemPrompt = None + params.max_tokens = 256 + params.stop_sequences = ["STOP"] + params.system_prompt = None params.tools = None - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) @@ -4520,11 +4503,11 @@ async def test_mcp_tool_sampling_callback_omits_temperature_when_none(): mock_message.content.text = "Test question" params.messages = [mock_message] params.temperature = None - params.maxTokens = 100 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 100 + params.stop_sequences = None + params.system_prompt = None params.tools = None - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) @@ -4557,11 +4540,11 @@ async def test_mcp_tool_sampling_callback_always_passes_max_tokens(): mock_message.content.text = "Test question" params.messages = [mock_message] params.temperature = None - params.maxTokens = 200 - params.stopSequences = None - params.systemPrompt = None + params.max_tokens = 200 + params.stop_sequences = None + params.system_prompt = None params.tools = None - params.toolChoice = None + params.tool_choice = None result = await _invoke_sampling_callback(tool, params) @@ -5263,7 +5246,7 @@ async def test_load_tools_prevents_multiple_calls(): mock_session = AsyncMock() mock_tool_list = MagicMock() mock_tool_list.tools = [] - mock_tool_list.nextCursor = None # No pagination + mock_tool_list.next_cursor = None # No pagination mock_session.list_tools = AsyncMock(return_value=mock_tool_list) mock_session.initialize = AsyncMock() @@ -5302,7 +5285,7 @@ async def test_load_prompts_prevents_multiple_calls(): mock_session = AsyncMock() mock_prompt_list = MagicMock() mock_prompt_list.prompts = [] - mock_prompt_list.nextCursor = None # No pagination + mock_prompt_list.next_cursor = None # No pagination mock_session.list_prompts = AsyncMock(return_value=mock_prompt_list) tool.session = mock_session @@ -5397,35 +5380,35 @@ async def test_load_tools_with_pagination(): types.Tool( name="tool_1", description="First tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, ), types.Tool( name="tool_2", description="Second tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, ), ] - page1.nextCursor = "cursor_page2" + page1.next_cursor = "cursor_page2" page2 = MagicMock() page2.tools = [ types.Tool( name="tool_3", description="Third tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, ), ] - page2.nextCursor = "cursor_page3" + page2.next_cursor = "cursor_page3" page3 = MagicMock() page3.tools = [ types.Tool( name="tool_4", description="Fourth tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, ), ] - page3.nextCursor = None # No more pages + page3.next_cursor = None # No more pages # Mock list_tools to return different pages based on params async def mock_list_tools(params=None): @@ -5452,7 +5435,7 @@ async def test_load_tools_adds_properties_to_zero_arg_tool_schema(): """Test that load_tools normalizes inputSchema for zero-argument MCP tools. Some MCP servers (e.g. matlab-mcp-core-server) declare zero-argument tools - with inputSchema={"type": "object"} and no "properties" key. OpenAI's API + with input_schema={"type": "object"} and no "properties" key. OpenAI's API requires "properties" to be present on object schemas, so load_tools must inject an empty "properties" dict when it is missing. """ @@ -5475,22 +5458,22 @@ async def test_load_tools_adds_properties_to_zero_arg_tool_schema(): types.Tool( name="zero_arg_tool", description="A tool with no parameters", - inputSchema=original_zero_arg_schema, + input_schema=original_zero_arg_schema, ), types.Tool( name="normal_tool", description="A tool with parameters", - inputSchema={"type": "object", "properties": {"x": {"type": "string"}}, "required": ["x"]}, + input_schema={"type": "object", "properties": {"x": {"type": "string"}}, "required": ["x"]}, ), types.Tool( name="string_schema_tool", description="A tool with a non-object schema", - inputSchema=original_string_schema, + input_schema=original_string_schema, ), types.Tool( name="empty_schema_tool", description="A tool with an empty schema", - inputSchema=original_empty_schema, + input_schema=original_empty_schema, ), ] @@ -5499,10 +5482,10 @@ async def test_load_tools_adds_properties_to_zero_arg_tool_schema(): none_schema_tool = MagicMock() none_schema_tool.name = "none_schema_tool" none_schema_tool.description = "A tool with None inputSchema" - none_schema_tool.inputSchema = None + none_schema_tool.input_schema = None none_schema_tool.meta = None page.tools.append(none_schema_tool) - page.nextCursor = None + page.next_cursor = None mock_session.list_tools = AsyncMock(return_value=page) @@ -5570,7 +5553,7 @@ async def test_load_prompts_with_pagination(): arguments=[types.PromptArgument(name="arg2", description="Arg 2", required=True)], ), ] - page1.nextCursor = "cursor_page2" + page1.next_cursor = "cursor_page2" page2 = MagicMock() page2.prompts = [ @@ -5580,7 +5563,7 @@ async def test_load_prompts_with_pagination(): arguments=[types.PromptArgument(name="arg3", description="Arg 3", required=False)], ), ] - page2.nextCursor = None # No more pages + page2.next_cursor = None # No more pages # Mock list_prompts to return different pages based on params async def mock_list_prompts(params=None): @@ -5620,30 +5603,30 @@ async def test_load_tools_pagination_with_duplicates(): types.Tool( name="tool_1", description="First tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, ), types.Tool( name="tool_2", description="Second tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, ), ] - page1.nextCursor = "cursor_page2" + page1.next_cursor = "cursor_page2" page2 = MagicMock() page2.tools = [ types.Tool( name="tool_1", # Duplicate from page1 description="Duplicate tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, ), types.Tool( name="tool_3", description="Third tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, ), ] - page2.nextCursor = None + page2.next_cursor = None # Mock list_tools to return different pages async def mock_list_tools(params=None): @@ -5686,7 +5669,7 @@ async def test_load_prompts_pagination_with_duplicates(): arguments=[types.PromptArgument(name="arg1", description="Arg 1", required=True)], ), ] - page1.nextCursor = "cursor_page2" + page1.next_cursor = "cursor_page2" page2 = MagicMock() page2.prompts = [ @@ -5701,7 +5684,7 @@ async def test_load_prompts_pagination_with_duplicates(): arguments=[types.PromptArgument(name="arg3", description="Arg 3", required=True)], ), ] - page2.nextCursor = None + page2.next_cursor = None # Mock list_prompts to return different pages async def mock_list_prompts(params=None): @@ -5734,11 +5717,11 @@ async def test_load_tools_concurrent_reload_does_not_duplicate_tools_and_preserv types.Tool( name="tool_1", description="First tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, _meta={"echo": "tool_1"}, ), ] - page.nextCursor = None + page.next_cursor = None async def mock_list_tools(params: Any = None) -> Any: assert params is None @@ -5769,7 +5752,7 @@ async def test_load_prompts_concurrent_reload_does_not_duplicate_prompts(): arguments=[types.PromptArgument(name="arg1", description="Arg 1", required=True)], ), ] - page.nextCursor = None + page.next_cursor = None async def mock_list_prompts(params: Any = None) -> Any: assert params is None @@ -5850,7 +5833,7 @@ async def test_load_tools_empty_pagination(): # Create empty response page1 = MagicMock() page1.tools = [] - page1.nextCursor = None + page1.next_cursor = None mock_session.list_tools = AsyncMock(return_value=page1) @@ -5878,7 +5861,7 @@ async def test_load_prompts_empty_pagination(): # Create empty response page1 = MagicMock() page1.prompts = [] - page1.nextCursor = None + page1.next_cursor = None mock_session.list_prompts = AsyncMock(return_value=page1) @@ -6105,7 +6088,7 @@ async def test_generated_mcp_function_ignores_model_supplied_remote_tool_name() types.Tool( name="search_docs", description="Search docs.", - inputSchema={ + input_schema={ "type": "object", "properties": {"query": {"type": "string"}}, "required": ["query"], @@ -6114,7 +6097,7 @@ async def test_generated_mcp_function_ignores_model_supplied_remote_tool_name() types.Tool( name="delete_repo", description="Delete a repository.", - inputSchema={ + input_schema={ "type": "object", "properties": {"repo": {"type": "string"}}, "required": ["repo"], @@ -6633,9 +6616,9 @@ async def test_connect_skips_tools_and_prompts_when_server_does_not_advertise_ca tool.session._request_id = 0 tool.session.initialize = AsyncMock( return_value=types.InitializeResult( - protocolVersion=types.LATEST_PROTOCOL_VERSION, + protocol_version=types.LATEST_PROTOCOL_VERSION, capabilities=types.ServerCapabilities(), - serverInfo=types.Implementation(name="test", version="1.0"), + server_info=types.Implementation(name="test", version="1.0"), ) ) tool.session.list_tools = AsyncMock() @@ -6683,9 +6666,9 @@ async def test_connect_sets_logging_level_when_server_advertises_logging() -> No tool.session._request_id = 0 tool.session.initialize = AsyncMock( return_value=types.InitializeResult( - protocolVersion=types.LATEST_PROTOCOL_VERSION, + protocol_version=types.LATEST_PROTOCOL_VERSION, capabilities=types.ServerCapabilities(logging=types.LoggingCapability()), - serverInfo=types.Implementation(name="test", version="1.0"), + server_info=types.Implementation(name="test", version="1.0"), ) ) tool.session.set_logging_level = AsyncMock() @@ -6699,11 +6682,7 @@ async def test_connect_sets_logging_level_when_server_advertises_logging() -> No async def test_ensure_connected_skips_future_pings_when_ping_is_not_available() -> None: tool = MCPTool(name="test_tool") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock( - send_ping=AsyncMock( - side_effect=McpError(types.ErrorData(code=-32601, message="Method 'ping' is not available.")) - ) - ) + tool.session = Mock(send_ping=AsyncMock(side_effect=MCPError(-32601, "Method 'ping' is not available."))) with patch.object(tool, "_reconnect_without_loading", AsyncMock()) as mock_reconnect: await tool._ensure_connected() @@ -6747,7 +6726,7 @@ async def test_load_tools_reconnects_on_closed_resource_when_ping_is_unavailable page = Mock() page.tools = [] - page.nextCursor = None + page.next_cursor = None second_session = Mock() second_session.list_tools = AsyncMock(return_value=page) @@ -6776,7 +6755,7 @@ async def test_load_prompts_reconnects_on_closed_resource_when_ping_is_unavailab page = Mock() page.prompts = [] - page.nextCursor = None + page.next_cursor = None second_session = Mock() second_session.list_prompts = AsyncMock(return_value=page) @@ -6810,7 +6789,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="test_tool", description="Test tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"], @@ -6894,7 +6873,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="test_tool", description="Test tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"], @@ -6955,7 +6934,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="WorkIQSharePoint.readSmallBinaryFile", description="Read a binary file", - inputSchema={ + input_schema={ "type": "object", "properties": {"fileId": {"type": "string"}}, "required": ["fileId"], @@ -7000,7 +6979,7 @@ async def connect(self) -> None: # type: ignore[override] # pyrefly: ignore[ba types.Tool( name="test_tool", description="Test tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, _meta=tool_meta, ) ] @@ -7043,7 +7022,7 @@ async def connect(self) -> None: # type: ignore[override] # pyrefly: ignore[ba types.Tool( name="test_tool", description="Test tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, ) ] ) @@ -7087,7 +7066,7 @@ async def connect(self) -> None: # type: ignore[override] # pyrefly: ignore[ba types.Tool( name="test_tool", description="Test tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, ) ] ) @@ -7138,7 +7117,7 @@ async def connect(self) -> None: # type: ignore[override] # pyrefly: ignore[ba types.Tool( name="test_tool", description="Test tool", - inputSchema={"type": "object", "properties": {"param": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"param": {"type": "string"}}}, _meta=tool_meta, ) ] @@ -7216,7 +7195,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="greet", description="Says hello", - inputSchema={ + input_schema={ "type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"], @@ -7276,7 +7255,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="greet", description="Says hello", - inputSchema={"type": "object", "properties": {"name": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"name": {"type": "string"}}}, ) ] ) @@ -7318,7 +7297,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="greet", description="Says hello", - inputSchema={"type": "object", "properties": {"name": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"name": {"type": "string"}}}, ) ] ) @@ -7358,7 +7337,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="greet", description="Says hello", - inputSchema={"type": "object", "properties": {"name": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"name": {"type": "string"}}}, ) ] ) @@ -8018,7 +7997,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="greet", description="Says hello", - inputSchema={ + input_schema={ "type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"], @@ -8381,7 +8360,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="greet", description="Says hello", - inputSchema={"type": "object", "properties": {"name": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"name": {"type": "string"}}}, ) ] ) @@ -8433,7 +8412,7 @@ async def connect(self, *, reset: bool = False) -> None: types.Tool( name="greet", description="Says hello", - inputSchema={"type": "object", "properties": {"name": {"type": "string"}}}, + input_schema={"type": "object", "properties": {"name": {"type": "string"}}}, ) ] ) @@ -8494,13 +8473,13 @@ def _make_task_snapshot( ) -> types.GetTaskResult: now = _utc_now() return types.GetTaskResult( - taskId=task_id, + task_id=task_id, status=status, # type: ignore[arg-type] # ty: ignore[invalid-argument-type] - statusMessage=status_message, - createdAt=now, - lastUpdatedAt=now, + status_message=status_message, + created_at=now, + last_updated_at=now, ttl=None, - pollInterval=poll_interval_ms, + poll_interval=poll_interval_ms, ) @@ -8508,11 +8487,11 @@ def _make_create_task_result(task_id: str = "task-1") -> types.CreateTaskResult: now = _utc_now() return types.CreateTaskResult( task=types.Task( - taskId=task_id, + task_id=task_id, status="working", - statusMessage=None, - createdAt=now, - lastUpdatedAt=now, + status_message=None, + created_at=now, + last_updated_at=now, ttl=None, ) ) @@ -8610,16 +8589,16 @@ async def test_load_tools_captures_task_support() -> None: types.Tool( name="slow_op", description="slow", - inputSchema={"type": "object", "properties": {}}, - execution=types.ToolExecution(taskSupport="required"), + input_schema={"type": "object", "properties": {}}, + execution=types.ToolExecution(task_support="required"), ), types.Tool( name="fast_op", description="fast", - inputSchema={"type": "object", "properties": {}}, + input_schema={"type": "object", "properties": {}}, ), ] - page.nextCursor = None + page.next_cursor = None tool.session.list_tools = AsyncMock(return_value=page) await tool.load_tools() @@ -8697,7 +8676,7 @@ async def test_call_tool_as_task_fallback_preserves_custom_parser_host_payload() tool.parse_tool_results = lambda _: "custom fallback summary" fallback_result = types.CallToolResult( content=[types.TextContent(type="text", text="fallback")], - structuredContent={"widget": "fallback"}, + structured_content={"widget": "fallback"}, _meta={"source": "fallback"}, ) tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] @@ -8726,7 +8705,7 @@ async def test_secure_mcp_task_results_cannot_relax_local_label(result_path: str if result_path == "fallback": raw_result = types.CallToolResult( content=[types.TextContent(type="text", text="fallback")], - structuredContent=structured_content, + structured_content=structured_content, _meta=result_meta, ) tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] @@ -8777,7 +8756,7 @@ async def test_task_parser_failure_preserves_complete_host_payload(result_path: if result_path == "fallback": raw_result = types.CallToolResult( content=[types.TextContent(type="text", text="fallback")], - structuredContent={"widget": result_path}, + structured_content={"widget": result_path}, _meta=result_meta, ) tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] @@ -8962,7 +8941,7 @@ async def test_call_tool_as_task_malformed_payload_raises() -> None: async def test_call_tool_as_task_method_not_found_falls_back() -> None: tool = _make_task_tool() tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=McpError(types.ErrorData(code=types.METHOD_NOT_FOUND, message="no tasks here")) + side_effect=MCPError(types.METHOD_NOT_FOUND, "no tasks here") ) tool.session.call_tool = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] return_value=types.CallToolResult(content=[types.TextContent(type="text", text="fell back")]) @@ -8977,7 +8956,7 @@ async def test_call_tool_as_task_method_not_found_falls_back() -> None: async def test_call_tool_as_task_invalid_params_falls_back() -> None: tool = _make_task_tool() tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=McpError(types.ErrorData(code=types.INVALID_PARAMS, message="unknown field")) + side_effect=MCPError(types.INVALID_PARAMS, "unknown field") ) tool.session.call_tool = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] return_value=types.CallToolResult(content=[types.TextContent(type="text", text="plain ok")]) @@ -9149,7 +9128,7 @@ async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> An return _make_create_task_result(task_id="abc") if method == "tasks/get": poll_calls += 1 - assert request.root.params.taskId == "abc" + assert request.root.params.task_id == "abc" if poll_calls == 1: raise ClosedResourceError return _make_task_snapshot(task_id="abc", status="completed") @@ -9428,7 +9407,7 @@ async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> An if method == "tasks/get": poll_calls += 1 if poll_calls == 1: - raise McpError(types.ErrorData(code=int(httpx.codes.REQUEST_TIMEOUT), message="slow poll")) + raise MCPError(int(httpx.codes.REQUEST_TIMEOUT), "slow poll") return _make_task_snapshot(task_id="t1", status="completed") if method == "tasks/result": return _make_payload("recovered after transient") @@ -9466,7 +9445,7 @@ async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> An if method == "tools/call": return _make_create_task_result(task_id="h1") if method == "tasks/get": - raise McpError(types.ErrorData(code=types.INVALID_PARAMS, message="bad task id")) + raise MCPError(types.INVALID_PARAMS, "bad task id") if method == "tasks/cancel": cancel_called = True return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] @@ -9594,7 +9573,7 @@ async def test_mcp_task_options_max_task_wait_rejects_non_positive() -> None: async def test_fetch_task_result_hard_mcperror_raises_without_cancel() -> None: - """tasks/result hard McpError must wrap as ToolExecutionException without cancel (server done).""" + """tasks/result hard MCPError must wrap as ToolExecutionException without cancel (server done).""" tool = _make_task_tool() cancel_called = False @@ -9607,7 +9586,7 @@ async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> An if method == "tasks/get": return _make_task_snapshot(task_id="hf", status="completed") if method == "tasks/result": - raise McpError(types.ErrorData(code=types.INTERNAL_ERROR, message="payload vanished")) + raise MCPError(types.INTERNAL_ERROR, "payload vanished") if method == "tasks/cancel": cancel_called = True return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] @@ -9618,7 +9597,7 @@ async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> An with pytest.raises(ToolExecutionException, match="payload vanished"): await tool.call_tool("slow_op") - # No raw McpError leak and no cancel — server already reported the task as done. + # No raw MCPError leak and no cancel — server already reported the task as done. await asyncio.sleep(0.02) assert cancel_called is False @@ -9935,7 +9914,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="test_tool", description="Test tool", - inputSchema={ + input_schema={ "type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"], @@ -9989,7 +9968,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri types.Tool( name="test_tool", description="Test tool", - inputSchema={ + input_schema={ "type": "object", "properties": { "param": {"type": "string"}, @@ -10049,7 +10028,7 @@ async def connect(self, *, reset: bool = False) -> None: types.Tool( name="get_weather", description="Weather", - inputSchema={ + input_schema={ "type": "object", "properties": {"city": {"type": "string"}, "api_key": {"type": "string"}}, "required": ["city"], From 9f30e29272e987320d7cb51fc7a5695cb4acba4f Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 30 Sep 2026 16:41:58 +0200 Subject: [PATCH 07/42] collection for test_mcp --- python/packages/core/tests/core/test_mcp.py | 80 ++++++++++----------- 1 file changed, 40 insertions(+), 40 deletions(-) diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 40971e99e36..9ea1cb52a8a 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -17,7 +17,7 @@ import pytest from mcp import MCPError, types from mcp.client.session import ClientSession -from pydantic import AnyUrl, BaseModel +from pydantic import BaseModel from agent_framework import ( Agent, @@ -702,7 +702,7 @@ def test_mcp_tool_str_and_parse_prompt_result_rich_content() -> None: content=types.EmbeddedResource( type="resource", resource=types.TextResourceContents( - uri=AnyUrl("file://prompt.txt"), + uri="file://prompt.txt", mime_type="text/plain", text="Embedded prompt", ), @@ -713,7 +713,7 @@ def test_mcp_tool_str_and_parse_prompt_result_rich_content() -> None: content=types.EmbeddedResource( type="resource", resource=types.BlobResourceContents( - uri=AnyUrl("file://prompt.bin"), + uri="file://prompt.bin", mime_type="application/pdf", blob="ZGF0YQ==", ), @@ -822,7 +822,7 @@ def test_parse_tool_result_from_mcp_blob_plain_base64(): types.EmbeddedResource( type="resource", resource=types.BlobResourceContents( - uri=AnyUrl("file://test.bin"), + uri="file://test.bin", mime_type="application/pdf", blob="dGVzdCBkYXRh", ), @@ -844,14 +844,14 @@ def test_parse_tool_result_from_mcp_resource_link_text_resource_and_unknown(): content=[ types.ResourceLink( type="resource_link", - uri=AnyUrl("https://example.com/resource"), + uri="https://example.com/resource", name="resource", mime_type="application/json", ), types.EmbeddedResource( type="resource", resource=types.TextResourceContents( - uri=AnyUrl("file://prompt.txt"), + uri="file://prompt.txt", mime_type="text/plain", text="Embedded result", ), @@ -1066,7 +1066,7 @@ def test_mcp_host_payload_size_boundary_uri_serialization_and_early_abort( content=[ types.ResourceLink( type="resource_link", - uri=AnyUrl("file:///abc"), + uri="file:///abc", name="resource", ) ] @@ -1243,7 +1243,7 @@ async def test_oversized_mcp_error_preserves_independently_bounded_meta() -> Non def test_mcp_result_meta_has_independent_boundary_and_early_abort(monkeypatch: pytest.MonkeyPatch) -> None: mcp_result = types.CallToolResult( content=[types.TextContent(type="text", text="ok")], - _meta={"source": "server", "uri": AnyUrl("https://example.test/resource")}, + _meta={"source": "server", "uri": "https://example.test/resource"}, ) expected = {"source": "server", "uri": "https://example.test/resource"} encoded_size = len(json.dumps(expected).encode("utf-8")) @@ -1704,7 +1704,7 @@ def test_mcp_content_types_to_ai_content_resource_link(): """Test conversion of MCP resource link to AI content.""" mcp_content = types.ResourceLink( type="resource_link", - uri=AnyUrl("https://example.com/resource"), + uri="https://example.com/resource", name="test_resource", mime_type="application/json", ) @@ -1719,7 +1719,7 @@ def test_mcp_content_types_to_ai_content_resource_link(): def test_mcp_content_types_to_ai_content_embedded_resource_text(): """Test conversion of MCP embedded text resource to AI content.""" text_resource = types.TextResourceContents( - uri=AnyUrl("file://test.txt"), + uri="file://test.txt", mime_type="text/plain", text="Embedded text content", ) @@ -1735,7 +1735,7 @@ def test_mcp_content_types_to_ai_content_embedded_resource_blob(): """Test conversion of MCP embedded blob resource to AI content.""" # Use a proper data URI in the blob field since that's what the MCP implementation expects blob_resource = types.BlobResourceContents( - uri=AnyUrl("file://test.bin"), + uri="file://test.bin", mime_type="application/octet-stream", blob="data:application/octet-stream;base64,dGVzdCBkYXRh", ) @@ -8458,10 +8458,10 @@ def get_mcp_client(self): # pyrefly: ignore[bad-override] # region: MCP long-running task (SEP-2663) tests -def _utc_now() -> Any: +def _utc_now() -> str: from datetime import datetime, timezone - return datetime.now(timezone.utc) + return datetime.now(timezone.utc).isoformat() def _make_task_snapshot( @@ -8546,7 +8546,7 @@ def _send_request_dispatcher(*responses_by_method: tuple[str, Any]) -> Any: queues[method].append(response) async def _dispatch(request: Any, _result_type: Any, *_args: Any, **_kw: Any) -> Any: - method = getattr(request.root, "method", None) or getattr(request, "method", None) + method = getattr(request, "method", None) or getattr(request, "method", None) queue = queues.get(method) # type: ignore[arg-type, call-overload] # pyrefly: ignore[bad-argument-type] if not queue: raise AssertionError(f"No mocked send_request response for method '{method}'.") @@ -8798,7 +8798,7 @@ async def test_call_tool_as_task_default_ttl_propagates() -> None: async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: captured.append(request) - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result() if method == "tasks/get": @@ -8812,9 +8812,9 @@ async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> An await tool.call_tool("slow_op") create_req = captured[0] - assert create_req.root.method == "tools/call" - assert create_req.root.params.task is not None - assert create_req.root.params.task.ttl == 7 * 60 * 1000 + assert create_req.method == "tools/call" + assert create_req.params.task is not None + assert create_req.params.task.ttl == 7 * 60 * 1000 async def test_call_tool_as_task_sends_empty_task_metadata_when_ttl_none() -> None: @@ -8826,7 +8826,7 @@ async def test_call_tool_as_task_sends_empty_task_metadata_when_ttl_none() -> No async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: captured.append(request) - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result() if method == "tasks/get": @@ -8840,9 +8840,9 @@ async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> An await tool.call_tool("slow_op") create_req = captured[0] - assert create_req.root.method == "tools/call" - assert create_req.root.params.task is not None - assert create_req.root.params.task.ttl is None + assert create_req.method == "tools/call" + assert create_req.params.task is not None + assert create_req.params.task.ttl is None async def test_call_tool_skips_task_path_for_optional_and_forbidden() -> None: @@ -9037,7 +9037,7 @@ async def test_call_tool_as_task_local_cancellation_fires_remote_cancel( create_seen = asyncio.Event() async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - method = request.root.method + method = request.method if method == "tools/call": create_seen.set() return _make_create_task_result() @@ -9084,7 +9084,7 @@ async def test_call_tool_as_task_cancellation_suppressed_when_disabled( async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal cancel_called - method = request.root.method + method = request.method if method == "tools/call": create_seen.set() return _make_create_task_result() @@ -9123,12 +9123,12 @@ async def test_call_tool_as_task_reconnects_during_poll(monkeypatch: pytest.Monk async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal poll_calls - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="abc") if method == "tasks/get": poll_calls += 1 - assert request.root.params.task_id == "abc" + assert request.params.task_id == "abc" if poll_calls == 1: raise ClosedResourceError return _make_task_snapshot(task_id="abc", status="completed") @@ -9155,7 +9155,7 @@ async def fake_connect(reset: bool = False) -> None: sum( 1 # type: ignore[misc] for c in tool.session.send_request.await_args_list # type: ignore[union-attr] # ty: ignore[unresolved-attribute] - if c.args[0].root.method == "tools/call" + if c.args[0].method == "tools/call" ) == 1 ) @@ -9173,7 +9173,7 @@ async def test_call_tool_as_task_second_disconnect_raises_connection_lost( tool = _make_task_tool() async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="abc") if method == "tasks/get": @@ -9230,7 +9230,7 @@ async def test_fetch_task_result_reconnects_during_fetch() -> None: async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal fetch_calls - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="r1") if method == "tasks/get": @@ -9268,7 +9268,7 @@ async def test_fetch_task_result_second_disconnect_raises_task_state_unknown_and async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal cancel_called - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="r2") if method == "tasks/get": @@ -9323,7 +9323,7 @@ async def test_call_tool_as_task_max_wait_exceeded_raises_and_cancels(monkeypatc async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal cancel_called - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="mw") if method == "tasks/get": @@ -9364,7 +9364,7 @@ async def test_call_tool_as_task_max_wait_cancels_even_when_local_cancel_option_ async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal cancel_called - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="mw2") if method == "tasks/get": @@ -9401,7 +9401,7 @@ async def test_call_tool_as_task_poll_transient_request_timeout_keeps_polling( async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal poll_calls, cancel_called - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="t1") if method == "tasks/get": @@ -9441,7 +9441,7 @@ async def test_call_tool_as_task_poll_hard_mcperror_cancels_and_raises( async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal cancel_called - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="h1") if method == "tasks/get": @@ -9479,7 +9479,7 @@ async def test_call_tool_as_task_malformed_tasks_get_response_cancels_and_raises async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal cancel_called - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="m1") if method == "tasks/get": @@ -9512,7 +9512,7 @@ async def test_call_tool_as_task_failed_terminal_does_not_cancel(monkeypatch: py async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal cancel_called - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="f1") if method == "tasks/get": @@ -9580,7 +9580,7 @@ async def test_fetch_task_result_hard_mcperror_raises_without_cancel() -> None: async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal cancel_called - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="hf") if method == "tasks/get": @@ -9621,7 +9621,7 @@ def boom_parser(_: Any) -> list[Content]: async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal cancel_called - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="t2") if method == "tasks/get": @@ -9666,7 +9666,7 @@ def boom_parser(_: Any) -> list[Content]: async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: nonlocal cancel_called - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="t3") if method == "tasks/get": @@ -9695,7 +9695,7 @@ async def test_max_wait_interrupts_long_poll_sleep(monkeypatch: pytest.MonkeyPat tool = _make_task_tool(task_options=MCPTaskOptions(max_task_wait=timedelta(milliseconds=100))) async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - method = request.root.method + method = request.method if method == "tools/call": return _make_create_task_result(task_id="ds") if method == "tasks/get": From e00abb402ec75ef629955ff12b0ee77a2efd73bf Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 30 Sep 2026 16:49:22 +0200 Subject: [PATCH 08/42] collection passes for test_mcp_observability and test_mcp_skills --- .../core/tests/core/test_mcp_observability.py | 32 +++---- .../core/tests/core/test_mcp_skills.py | 96 +++++++++---------- 2 files changed, 59 insertions(+), 69 deletions(-) diff --git a/python/packages/core/tests/core/test_mcp_observability.py b/python/packages/core/tests/core/test_mcp_observability.py index 7afd1dd402d..9f0e03b1a5a 100644 --- a/python/packages/core/tests/core/test_mcp_observability.py +++ b/python/packages/core/tests/core/test_mcp_observability.py @@ -11,9 +11,7 @@ from unittest.mock import AsyncMock, Mock import pytest -from mcp import types -from mcp.shared.exceptions import McpError -from mcp.types import ErrorData +from mcp import MCPError, types from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter from opentelemetry.trace import SpanKind, StatusCode @@ -50,10 +48,10 @@ def _make_tool_list_result( tools = [{"name": "get-weather", "description": "Get weather", "inputSchema": {"type": "object"}}] result = Mock() result.tools = [ - types.Tool(name=t["name"], description=t.get("description", ""), inputSchema=t.get("inputSchema", {})) + types.Tool(name=t["name"], description=t.get("description", ""), input_schema=t.get("inputSchema", {})) for t in tools ] - result.nextCursor = None + result.next_cursor = None return result @@ -67,16 +65,16 @@ def _make_prompt_list_result( result.prompts = [ types.Prompt(name=p["name"], description=p.get("description", ""), arguments=None) for p in prompts ] - result.nextCursor = None + result.next_cursor = None return result def _make_call_tool_result(text: str = "result", is_error: bool = False) -> Mock: """Create a mock CallToolResult.""" result = Mock() - result.isError = is_error + result.is_error = is_error result.content = [types.TextContent(type="text", text=text)] - result.structuredContent = None + result.structured_content = None return result @@ -106,7 +104,7 @@ async def test_mcp_initialize_span(span_exporter: InMemorySpanExporter): mock_session_cls = AsyncMock() init_result = Mock() init_result.capabilities = None - init_result.protocolVersion = "2025-06-18" + init_result.protocol_version = "2025-06-18" mock_session_cls.initialize = AsyncMock(return_value=init_result) # Create a mock transport context manager @@ -133,7 +131,7 @@ async def patched_connect( with create_mcp_client_span("initialize", attributes=self_._mcp_base_span_attributes()) as init_span: result = await mock_session_cls.initialize() - protocol_version = getattr(result, "protocolVersion", None) + protocol_version = getattr(result, "protocol_version", None) if protocol_version: init_span.set_attribute(OtelAttr.MCP_PROTOCOL_VERSION, protocol_version) @@ -223,7 +221,7 @@ async def test_mcp_tools_call_creates_client_span_when_no_parent(span_exporter: async def test_mcp_tools_call_tool_error_sets_error_type(span_exporter: InMemorySpanExporter): - """When CallToolResult.isError is true, error.type should be 'tool_error' per MCP spec.""" + """When CallToolResult.is_error is true, error.type should be 'tool_error' per MCP spec.""" tool = _make_connected_mcp_tool() tool.session.call_tool = AsyncMock(return_value=_make_call_tool_result("bad input", is_error=True)) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] @@ -240,9 +238,9 @@ async def test_mcp_tools_call_tool_error_sets_error_type(span_exporter: InMemory async def test_mcp_tools_call_mcp_error_sets_error_type(span_exporter: InMemorySpanExporter): - """When session.call_tool() raises McpError, error.type should be the exception class name.""" + """When session.call_tool() raises MCPError, error.type should be the exception class name.""" tool = _make_connected_mcp_tool() - tool.session.call_tool = AsyncMock(side_effect=McpError(ErrorData(code=-32600, message="invalid request"))) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] + tool.session.call_tool = AsyncMock(side_effect=MCPError(-32600, "invalid request")) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] span_exporter.clear() with pytest.raises(ToolExecutionException): @@ -252,7 +250,7 @@ async def test_mcp_tools_call_mcp_error_sets_error_type(span_exporter: InMemoryS call_spans = [s for s in spans if "tools/call" in s.name] assert len(call_spans) == 1 span = call_spans[0] - assert span.attributes.get(OtelAttr.ERROR_TYPE) == "McpError" # type: ignore[union-attr] # ty: ignore[unresolved-attribute] + assert span.attributes.get(OtelAttr.ERROR_TYPE) == "MCPError" # type: ignore[union-attr] # ty: ignore[unresolved-attribute] assert span.status.status_code == StatusCode.ERROR @@ -282,10 +280,10 @@ async def test_mcp_prompts_get_creates_client_span(span_exporter: InMemorySpanEx async def test_mcp_prompts_get_mcp_error_sets_error_type(span_exporter: InMemorySpanExporter): - """When session.get_prompt() raises McpError, the span should have error.type and ERROR status.""" + """When session.get_prompt() raises MCPError, the span should have error.type and ERROR status.""" tool = _make_connected_mcp_tool() tool.session.get_prompt = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=McpError(ErrorData(code=-32602, message="prompt not found")) + side_effect=MCPError(-32602, "prompt not found") ) span_exporter.clear() @@ -296,7 +294,7 @@ async def test_mcp_prompts_get_mcp_error_sets_error_type(span_exporter: InMemory prompt_spans = [s for s in spans if "prompts/get" in s.name] assert len(prompt_spans) == 1 span = prompt_spans[0] - assert span.attributes.get(OtelAttr.ERROR_TYPE) == "McpError" # type: ignore[union-attr] # ty: ignore[unresolved-attribute] + assert span.attributes.get(OtelAttr.ERROR_TYPE) == "MCPError" # type: ignore[union-attr] # ty: ignore[unresolved-attribute] assert span.status.status_code == StatusCode.ERROR diff --git a/python/packages/core/tests/core/test_mcp_skills.py b/python/packages/core/tests/core/test_mcp_skills.py index 1584312628e..6259359a025 100644 --- a/python/packages/core/tests/core/test_mcp_skills.py +++ b/python/packages/core/tests/core/test_mcp_skills.py @@ -16,14 +16,12 @@ from urllib.parse import unquote import pytest -from mcp.shared.exceptions import McpError +from mcp import MCPError from mcp.types import ( BlobResourceContents, - ErrorData, ReadResourceResult, TextResourceContents, ) -from pydantic import AnyUrl from agent_framework import CachingSkillsSource, MCPSkill, MCPSkillResource, MCPSkillsSource, SkillsSourceContext from agent_framework._skills import _fully_unquote, _parse_mcp_skill_index @@ -63,7 +61,7 @@ def _make_text_result(text: str, uri: str = "skill://test") -> ReadResourceResult: """Create a ReadResourceResult with a single TextResourceContents.""" - return ReadResourceResult(contents=[TextResourceContents(uri=AnyUrl(uri), text=text, mimeType="text/markdown")]) + return ReadResourceResult(contents=[TextResourceContents(uri=uri, text=text, mime_type="text/markdown")]) def _make_blob_result( @@ -73,7 +71,7 @@ def _make_blob_result( ) -> ReadResourceResult: """Create a ReadResourceResult with a single BlobResourceContents.""" return ReadResourceResult( - contents=[BlobResourceContents(uri=AnyUrl(uri), blob=base64.b64encode(data).decode(), mimeType=mime_type)] + contents=[BlobResourceContents(uri=uri, blob=base64.b64encode(data).decode(), mime_type=mime_type)] ) @@ -87,16 +85,16 @@ def _make_client(**read_resource_responses: ReadResourceResult) -> AsyncMock: Args: **read_resource_responses: Mapping of URI string to ReadResourceResult. - Any URI not in this mapping raises McpError with the MCP-spec + Any URI not in this mapping raises MCPError with the MCP-spec "Resource not found" code (-32002). """ client = AsyncMock() - async def _read_resource(uri: AnyUrl) -> ReadResourceResult: + async def _read_resource(uri: str) -> ReadResourceResult: uri_str = str(uri) if uri_str in read_resource_responses: return read_resource_responses[uri_str] - raise McpError(error=ErrorData(code=-32002, message=f"Resource not found: {uri_str}")) + raise MCPError(-32002, f"Resource not found: {uri_str}") client.read_resource = AsyncMock(side_effect=_read_resource) return client @@ -225,8 +223,8 @@ async def test_read_empty_returns_none(self) -> None: async def test_read_multiple_text_contents_joined(self) -> None: result = ReadResourceResult( contents=[ - TextResourceContents(uri=AnyUrl("skill://a"), text="line1", mimeType="text/plain"), - TextResourceContents(uri=AnyUrl("skill://b"), text="line2", mimeType="text/plain"), + TextResourceContents(uri="skill://a", text="line1", mime_type="text/plain"), + TextResourceContents(uri="skill://b", text="line2", mime_type="text/plain"), ] ) resource = MCPSkillResource(name="multi", result=result) @@ -237,11 +235,11 @@ async def test_binary_takes_precedence_over_text(self) -> None: data = b"\xff\xfe" result = ReadResourceResult( contents=[ - TextResourceContents(uri=AnyUrl("skill://a"), text="text", mimeType="text/plain"), + TextResourceContents(uri="skill://a", text="text", mime_type="text/plain"), BlobResourceContents( - uri=AnyUrl("skill://b"), + uri="skill://b", blob=base64.b64encode(data).decode(), - mimeType="application/octet-stream", + mime_type="application/octet-stream", ), ] ) @@ -451,7 +449,7 @@ async def test_get_resource_preserves_safe_names_and_schemes(self, name: str, ro assert resource is not None assert resource.name == name assert await resource.read() == "safe content" - client.read_resource.assert_awaited_once_with(AnyUrl(root + name.replace("\\", "/"))) + client.read_resource.assert_awaited_once_with(root + name.replace("\\", "/")) @pytest.mark.parametrize("depth", [1, 31, 32, 33, 4096]) @pytest.mark.parametrize( @@ -476,7 +474,7 @@ async def test_get_resource_decoding_depth_is_bounded(self, depth: int, template if depth <= 32 and not name.startswith("../"): assert resource is not None assert resource.name == name - client.read_resource.assert_awaited_once_with(AnyUrl(root + name)) + client.read_resource.assert_awaited_once_with(root + name) else: assert resource is None client.read_resource.assert_not_called() @@ -752,23 +750,23 @@ def test_requires_exactly_one_of_client_or_session_provider(self) -> None: # --------------------------------------------------------------------------- -# McpError code branching tests +# MCPError code branching tests # --------------------------------------------------------------------------- class TestMCPSkillsSourceErrorCodeBranching: - """Tests that MCPSkillsSource and MCPSkill branch on McpError.error.code. + """Tests that MCPSkillsSource and MCPSkill branch on MCPError.error.code. Only "not found" codes (RESOURCE_NOT_FOUND -32002, METHOD_NOT_FOUND -32601) - should be silently swallowed as "no skills available." Other McpError codes - and non-McpError exceptions must propagate so that auth failures, server + should be silently swallowed as "no skills available." Other MCPError codes + and non-MCPError exceptions must propagate so that auth failures, server crashes, and connection drops are visible. """ async def test_index_method_not_found_returns_empty(self) -> None: """METHOD_NOT_FOUND (-32601) -> server doesn't support resources/read.""" client = AsyncMock() - client.read_resource = AsyncMock(side_effect=McpError(error=ErrorData(code=-32601, message="Method not found"))) + client.read_resource = AsyncMock(side_effect=MCPError(-32601, "Method not found")) source = MCPSkillsSource(client=client) skills = await source.get_skills(_SOURCE_CTX) assert skills == [] @@ -776,9 +774,7 @@ async def test_index_method_not_found_returns_empty(self) -> None: async def test_index_resource_not_found_returns_empty(self) -> None: """MCP-spec "Resource not found" (-32002) -> server has no index.""" client = AsyncMock() - client.read_resource = AsyncMock( - side_effect=McpError(error=ErrorData(code=-32002, message="Resource not found")) - ) + client.read_resource = AsyncMock(side_effect=MCPError(-32002, "Resource not found")) source = MCPSkillsSource(client=client) skills = await source.get_skills(_SOURCE_CTX) assert skills == [] @@ -786,39 +782,37 @@ async def test_index_resource_not_found_returns_empty(self) -> None: async def test_index_invalid_params_propagates(self) -> None: """INVALID_PARAMS (-32602) is a real bug, must propagate (not "not found").""" client = AsyncMock() - client.read_resource = AsyncMock(side_effect=McpError(error=ErrorData(code=-32602, message="Invalid params"))) + client.read_resource = AsyncMock(side_effect=MCPError(-32602, "Invalid params")) source = MCPSkillsSource(client=client) - with pytest.raises(McpError): + with pytest.raises(MCPError): await source.get_skills(_SOURCE_CTX) async def test_index_internal_error_propagates(self) -> None: """INTERNAL_ERROR (-32603) must propagate, not silently return empty.""" client = AsyncMock() - client.read_resource = AsyncMock(side_effect=McpError(error=ErrorData(code=-32603, message="Internal error"))) + client.read_resource = AsyncMock(side_effect=MCPError(-32603, "Internal error")) source = MCPSkillsSource(client=client) - with pytest.raises(McpError): + with pytest.raises(MCPError): await source.get_skills(_SOURCE_CTX) async def test_index_connection_closed_propagates(self) -> None: """CONNECTION_CLOSED (-32000) must propagate.""" client = AsyncMock() - client.read_resource = AsyncMock( - side_effect=McpError(error=ErrorData(code=-32000, message="Connection closed")) - ) + client.read_resource = AsyncMock(side_effect=MCPError(-32000, "Connection closed")) source = MCPSkillsSource(client=client) - with pytest.raises(McpError): + with pytest.raises(MCPError): await source.get_skills(_SOURCE_CTX) async def test_index_generic_error_code_propagates(self) -> None: """Generic handler error (code 0) must propagate.""" client = AsyncMock() - client.read_resource = AsyncMock(side_effect=McpError(error=ErrorData(code=0, message="Some handler error"))) + client.read_resource = AsyncMock(side_effect=MCPError(0, "Some handler error")) source = MCPSkillsSource(client=client) - with pytest.raises(McpError): + with pytest.raises(MCPError): await source.get_skills(_SOURCE_CTX) async def test_index_non_mcp_error_propagates(self) -> None: - """Non-McpError exceptions (connection drop, timeout) must propagate.""" + """Non-MCPError exceptions (connection drop, timeout) must propagate.""" client = AsyncMock() client.read_resource = AsyncMock(side_effect=ConnectionError("connection lost")) source = MCPSkillsSource(client=client) @@ -826,24 +820,22 @@ async def test_index_non_mcp_error_propagates(self) -> None: await source.get_skills(_SOURCE_CTX) async def test_get_resource_internal_error_propagates(self) -> None: - """McpError with INTERNAL_ERROR on get_resource must propagate.""" + """MCPError with INTERNAL_ERROR on get_resource must propagate.""" from agent_framework import SkillFrontmatter client = AsyncMock() - client.read_resource = AsyncMock(side_effect=McpError(error=ErrorData(code=-32603, message="Server crashed"))) + client.read_resource = AsyncMock(side_effect=MCPError(-32603, "Server crashed")) fm = SkillFrontmatter(name="test-skill", description="Test.") skill = MCPSkill(frontmatter=fm, skill_md_uri="skill://test/SKILL.md", client=client) - with pytest.raises(McpError): + with pytest.raises(MCPError): await skill.get_resource("references/file.md") async def test_get_resource_not_found_returns_none(self) -> None: - """McpError with RESOURCE_NOT_FOUND (-32002) on get_resource returns None.""" + """MCPError with RESOURCE_NOT_FOUND (-32002) on get_resource returns None.""" from agent_framework import SkillFrontmatter client = AsyncMock() - client.read_resource = AsyncMock( - side_effect=McpError(error=ErrorData(code=-32002, message="Resource not found")) - ) + client.read_resource = AsyncMock(side_effect=MCPError(-32002, "Resource not found")) fm = SkillFrontmatter(name="test-skill", description="Test.") skill = MCPSkill(frontmatter=fm, skill_md_uri="skill://test/SKILL.md", client=client) result = await skill.get_resource("references/file.md") @@ -872,14 +864,14 @@ async def test_get_resource_timeout_error_propagates(self) -> None: await skill.get_resource("references/file.md") async def test_get_resource_generic_mcp_error_propagates(self) -> None: - """McpError with a generic code (0) on get_resource must propagate.""" + """MCPError with a generic code (0) on get_resource must propagate.""" from agent_framework import SkillFrontmatter client = AsyncMock() - client.read_resource = AsyncMock(side_effect=McpError(error=ErrorData(code=0, message="Handler error"))) + client.read_resource = AsyncMock(side_effect=MCPError(0, "Handler error")) fm = SkillFrontmatter(name="test-skill", description="Test.") skill = MCPSkill(frontmatter=fm, skill_md_uri="skill://test/SKILL.md", client=client) - with pytest.raises(McpError): + with pytest.raises(MCPError): await skill.get_resource("references/file.md") async def test_index_timeout_error_propagates(self) -> None: @@ -939,7 +931,7 @@ def _archive_client(index_json: str, archive_url: str, archive_bytes: bytes, mim """Build a mock client that serves the index and a single archive blob resource.""" return _make_client(**{ "skill://index.json": _make_text_result(index_json, uri="skill://index.json"), - str(AnyUrl(archive_url)): _make_blob_result(archive_bytes, uri=archive_url, mime_type=mime_type), + archive_url: _make_blob_result(archive_bytes, uri=archive_url, mime_type=mime_type), }) @@ -1279,17 +1271,17 @@ async def test_archive_download_internal_error_propagates(self) -> None: url = "skill://archives/packaged-skill.zip" index = _make_archive_index("packaged-skill", url) - async def _read_resource(uri: AnyUrl) -> ReadResourceResult: + async def _read_resource(uri: str) -> ReadResourceResult: uri_str = str(uri) if uri_str == "skill://index.json": return _make_text_result(index, uri="skill://index.json") - raise McpError(error=ErrorData(code=-32603, message="Internal error")) + raise MCPError(-32603, "Internal error") client = AsyncMock() client.read_resource = AsyncMock(side_effect=_read_resource) source = MCPSkillsSource(client=client) - with pytest.raises(McpError): + with pytest.raises(MCPError): await source.get_skills(_SOURCE_CTX) async def test_archive_download_connection_error_propagates(self) -> None: @@ -1297,7 +1289,7 @@ async def test_archive_download_connection_error_propagates(self) -> None: url = "skill://archives/packaged-skill.zip" index = _make_archive_index("packaged-skill", url) - async def _read_resource(uri: AnyUrl) -> ReadResourceResult: + async def _read_resource(uri: str) -> ReadResourceResult: uri_str = str(uri) if uri_str == "skill://index.json": return _make_text_result(index, uri="skill://index.json") @@ -1333,7 +1325,7 @@ async def test_mixed_skill_md_and_archive_entries(self) -> None: client = _make_client(**{ "skill://index.json": _make_text_result(index, uri="skill://index.json"), "skill://unit-converter/SKILL.md": _make_text_result(SAMPLE_SKILL_MD), - str(AnyUrl(archive_url)): _make_blob_result(archive, uri=archive_url, mime_type="application/zip"), + archive_url: _make_blob_result(archive, uri=archive_url, mime_type="application/zip"), }) source = MCPSkillsSource(client=client) @@ -1635,7 +1627,7 @@ async def test_skill_md_digest_does_not_change_lazy_discovery(self, digest: obje assert len(skills) == 1 assert isinstance(skills[0], MCPSkill) - client.read_resource.assert_awaited_once_with(AnyUrl("skill://index.json")) + client.read_resource.assert_awaited_once_with("skill://index.json") @pytest.mark.parametrize("refresh_failure", ["digest", "transport"]) async def test_cache_refresh_handles_verification_and_transport_failures( @@ -1652,7 +1644,7 @@ async def test_cache_refresh_handles_verification_and_transport_failures( first = await source.get_skills(_SOURCE_CTX) assert len(first) == 1 assert await source.get_skills(_SOURCE_CTX) is first - client.read_resource.assert_any_await(AnyUrl(url)) + client.read_resource.assert_any_await(url) assert client.read_resource.await_count == 2 tampered = _make_zip({ From 33e61b989e9c0e9510ad0e06a970e95c7b90fb5c Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Thu, 1 Oct 2026 11:13:19 +0200 Subject: [PATCH 09/42] Migrated message_handler --- python/packages/core/agent_framework/_mcp.py | 22 +++++++++----------- python/packages/core/tests/core/test_mcp.py | 1 + 2 files changed, 11 insertions(+), 12 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index aeb6d4ee94a..9361d22c1de 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -26,6 +26,7 @@ else: from typing_extensions import deprecated # pragma: no cover +from mcp.client import IncomingMessage from opentelemetry import propagate from opentelemetry import trace as otel_trace from pydantic_core import to_jsonable_python @@ -67,7 +68,6 @@ from mcp import types from mcp.client.session import ClientSession from mcp.shared.context import RequestContext - from mcp.shared.session import RequestResponder from ._clients import SupportsChatGetResponse from ._middleware import FunctionInvocationContext @@ -2328,7 +2328,7 @@ async def logging_callback(self, params: types.LoggingMessageNotificationParams) async def message_handler( self, - message: (RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception), + message: IncomingMessage, ) -> None: """Handle messages from the MCP server. @@ -2344,19 +2344,17 @@ async def message_handler( Args: message: The message from the MCP server (request responder, notification, or exception). """ - from mcp import types - if isinstance(message, Exception): logger.error("Error from MCP server: %s", message, exc_info=message) return - if isinstance(message, types.ServerNotification): - match message.root.method: - case "notifications/tools/list_changed": - self._schedule_reload(self.load_tools()) - case "notifications/prompts/list_changed": - self._schedule_reload(self.load_prompts()) - case _: - logger.debug("Unhandled notification: %s", message.root.method) + + match message.method: + case "notifications/tools/list_changed": + self._schedule_reload(self.load_tools()) + case "notifications/prompts/list_changed": + self._schedule_reload(self.load_prompts()) + case _: + logger.debug("Unhandled notification: %s", message.method) def _schedule_reload(self, coro: Coroutine[Any, Any, None]) -> None: """Schedule a reload coroutine as a background task. diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 9ea1cb52a8a..4ef9a51fd58 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -954,6 +954,7 @@ async def test_generated_mcp_tool_preserves_complete_host_payload_once() -> None "content": [{"type": "text", "text": "Summary", "_meta": file_meta}], "structuredContent": {"image_url": "https://example.test/widget.png"}, "isError": False, + "resultType": "complete", } assert [item.additional_properties["_meta"] for item in function_result.items] == [{"widget": "image"}] From c6f177cfd0d563835ed8955d73c0aff39dd83e3f Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Thu, 1 Oct 2026 11:24:51 +0200 Subject: [PATCH 10/42] deprecated MCPWebsocketTool --- python/packages/core/agent_framework/_mcp.py | 30 ++------------------ 1 file changed, 2 insertions(+), 28 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 9361d22c1de..15ba67ff4d6 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -4328,6 +4328,7 @@ async def call_tool(self, tool_name: str, **kwargs: Any) -> str | list[Content]: return await super().call_tool(tool_name, **kwargs) +@deprecated("Websocket transport was removed in MCP v2. Use MCPStreamableHTTPTool instead.") class MCPWebsocketTool(MCPTool): """MCP tool for connecting to WebSocket-based MCP servers. @@ -4517,31 +4518,4 @@ def _mcp_base_span_attributes(self) -> dict[str, Any]: return attrs def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: - """Get an MCP WebSocket client. - - Returns: - An async context manager for the WebSocket client transport. - """ - try: - websocket_module = __import__("mcp.client.websocket", fromlist=["websocket_client"]) - except ModuleNotFoundError as ex: - missing_name = ex.name or "mcp/websocket dependencies" - if missing_name == "mcp" or missing_name.startswith("mcp."): - reason = "The `mcp` package is not installed." - elif missing_name == "websockets" or missing_name.startswith("websockets."): - reason = "WebSocket transport support is not installed." - else: - reason = f"The optional dependency `{missing_name}` is not installed." - raise ModuleNotFoundError( - f"`MCPWebsocketTool` requires websocket transport support. {reason} " - "Please install `mcp` and update your dependencies." - ) from ex - - # Support MCP releases from before and after the transport gained its deprecation marker. - websocket_client = websocket_module.websocket_client - args: dict[str, Any] = { - "url": self.url, - } - if self._client_kwargs: - args.update(self._client_kwargs) - return websocket_client(**args) + raise RuntimeError("MCP WebSocket transport was removed in MCP v2. Use MCPStreamableHTTPTool instead.") From db0af2e8f7e53f41db8d4ae90e4c8fcdc00fd9c8 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Thu, 1 Oct 2026 11:56:46 +0200 Subject: [PATCH 11/42] deprecated MCPWebsocketTool, part 2 --- python/packages/core/agent_framework/_mcp.py | 5 +- python/packages/core/tests/core/test_mcp.py | 56 +++---------------- .../core/tests/core/test_mcp_observability.py | 19 +------ .../tests/core/test_optional_dependencies.py | 47 ---------------- 4 files changed, 12 insertions(+), 115 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 15ba67ff4d6..4e5fd3ed805 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -877,14 +877,13 @@ class MCPTool: Note: MCPTool cannot be instantiated directly. Use one of the subclasses: - MCPStdioTool, MCPStreamableHTTPTool, or MCPWebsocketTool. + MCPStdioTool or MCPStreamableHTTPTool. Examples: See the subclass documentation for usage examples: - :class:`MCPStdioTool` for stdio-based MCP servers - :class:`MCPStreamableHTTPTool` for HTTP-based MCP servers - - :class:`MCPWebsocketTool` for WebSocket-based MCP servers """ def __init__( @@ -4328,7 +4327,7 @@ async def call_tool(self, tool_name: str, **kwargs: Any) -> str | list[Content]: return await super().call_tool(tool_name, **kwargs) -@deprecated("Websocket transport was removed in MCP v2. Use MCPStreamableHTTPTool instead.") +@deprecated("MCP WebSocket transport was removed in MCP v2. Use MCPStreamableHTTPTool instead.") class MCPWebsocketTool(MCPTool): """MCP tool for connecting to WebSocket-based MCP servers. diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 4ef9a51fd58..a513913303c 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -161,14 +161,6 @@ def test_mcp_transport_subclasses_accept_tool_name_prefix() -> None: ).tool_name_prefix == "http" ) - assert ( - MCPWebsocketTool( - name="ws", - url="wss://example.com/mcp", - tool_name_prefix="ws", - ).tool_name_prefix - == "ws" - ) @pytest.mark.parametrize("configured_name", ["search_docs", "docs_search_docs"]) @@ -2885,29 +2877,20 @@ def test_mcp_transport_subclasses_accept_progressive_disclosure_options() -> Non use_progressive_disclosure=True, always_load=["search"], ) - websocket = MCPWebsocketTool( - name="ws", - url="wss://example.com/mcp", - use_progressive_disclosure=True, - always_load=["search"], - ) assert stdio.use_progressive_disclosure is True assert http.use_progressive_disclosure is True - assert websocket.use_progressive_disclosure is True assert stdio.always_load == ["search"] assert http.always_load == ["search"] - assert websocket.always_load == ["search"] def test_mcp_transport_subclasses_forward_host_payload_limit() -> None: tools = [ MCPStdioTool(name="stdio", command="python", max_host_payload_size_bytes=101), MCPStreamableHTTPTool(name="http", url="https://example.com/mcp", max_host_payload_size_bytes=102), - MCPWebsocketTool(name="ws", url="wss://example.com/mcp", max_host_payload_size_bytes=103), ] - assert [tool.max_host_payload_size_bytes for tool in tools] == [101, 102, 103] + assert [tool.max_host_payload_size_bytes for tool in tools] == [101, 102] def test_mcp_progressive_disclosure_requires_loading_tools() -> None: @@ -3529,13 +3512,6 @@ def test_local_mcp_stdio_tool_init(): assert tool.args == ["hello"] -def test_local_mcp_websocket_tool_init(): - """Test MCPWebsocketTool initialization.""" - tool = MCPWebsocketTool(name="test", url="ws://localhost:8080") - assert tool.name == "test" - assert tool.url == "ws://localhost:8080" - - def test_local_mcp_streamable_http_tool_init(): """Test MCPStreamableHTTPTool initialization.""" tool = MCPStreamableHTTPTool(name="test", url="http://localhost:8080") @@ -3543,6 +3519,14 @@ def test_local_mcp_streamable_http_tool_init(): assert tool.url == "http://localhost:8080" +def test_mcp_websocket_tool_is_deprecated() -> None: + with pytest.warns(DeprecationWarning, match="MCP WebSocket transport was removed in MCP v2"): + tool = MCPWebsocketTool(name="test", url="ws://localhost:8080") # pyright: ignore[reportDeprecated] + + with pytest.raises(RuntimeError, match="Use MCPStreamableHTTPTool instead"): + tool.get_mcp_client() + + # Integration test @pytest.mark.flaky @pytest.mark.integration @@ -5149,28 +5133,6 @@ async def test_mcp_streamable_http_tool_get_mcp_client_all_params(): assert http_client.is_closed -def test_mcp_websocket_tool_get_mcp_client_with_kwargs(): - """Test MCPWebsocketTool.get_mcp_client() with client kwargs.""" - tool = MCPWebsocketTool( - name="test", - url="wss://example.com", - max_size=1024, - ping_interval=30, - compression="deflate", - ) - - with patch("mcp.client.websocket.websocket_client") as mock_ws_client: - tool.get_mcp_client() - - # Verify all kwargs were passed - mock_ws_client.assert_called_once_with( - url="wss://example.com", - max_size=1024, - ping_interval=30, - compression="deflate", - ) - - async def test_mcp_tool_deduplication(): """Test that MCP tools are not duplicated in MCPTool""" from agent_framework._mcp import MCPTool diff --git a/python/packages/core/tests/core/test_mcp_observability.py b/python/packages/core/tests/core/test_mcp_observability.py index 9f0e03b1a5a..d2353af90e4 100644 --- a/python/packages/core/tests/core/test_mcp_observability.py +++ b/python/packages/core/tests/core/test_mcp_observability.py @@ -15,7 +15,7 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter from opentelemetry.trace import SpanKind, StatusCode -from agent_framework import MCPStdioTool, MCPStreamableHTTPTool, MCPWebsocketTool +from agent_framework import MCPStdioTool, MCPStreamableHTTPTool from agent_framework._mcp import MCPTool from agent_framework.exceptions import ToolExecutionException from agent_framework.observability import OtelAttr @@ -336,23 +336,6 @@ def test_mcp_http_tool_http_default_port(): assert attrs[OtelAttr.PORT] == 80 -def test_mcp_websocket_tool_transport_attributes(): - """MCPWebsocketTool should have tcp transport and URL-based server address/port.""" - tool = MCPWebsocketTool(name="test", url="wss://ws.example.com:9090/mcp") - attrs = tool._mcp_base_span_attributes() - assert attrs[OtelAttr.NETWORK_TRANSPORT] == "tcp" - assert attrs[OtelAttr.NETWORK_PROTOCOL_NAME] == "websocket" - assert attrs[OtelAttr.ADDRESS] == "ws.example.com" - assert attrs[OtelAttr.PORT] == 9090 - - -def test_mcp_websocket_tool_default_port(): - """MCPWebsocketTool should default to 443 for wss.""" - tool = MCPWebsocketTool(name="test", url="wss://ws.example.com/mcp") - attrs = tool._mcp_base_span_attributes() - assert attrs[OtelAttr.PORT] == 443 - - # endregion diff --git a/python/packages/core/tests/core/test_optional_dependencies.py b/python/packages/core/tests/core/test_optional_dependencies.py index 0e6e8cc3429..262fb674a91 100644 --- a/python/packages/core/tests/core/test_optional_dependencies.py +++ b/python/packages/core/tests/core/test_optional_dependencies.py @@ -132,50 +132,3 @@ def _import_without_mcp( with pytest.raises(ModuleNotFoundError, match=r"Please install `mcp`\.$"): agent.as_mcp_server() - - -def test_mcp_websocket_tool_requires_ws_support(monkeypatch: pytest.MonkeyPatch) -> None: - import builtins - - real_import = builtins.__import__ - - sys.modules.pop("mcp.client.websocket", None) - - def _import_without_websocket_support( - name: str, - globals_: dict[str, object] | None = None, - locals_: dict[str, object] | None = None, - fromlist: tuple[str, ...] = (), - level: int = 0, - ) -> object: - if name == "mcp.client.websocket": - raise ModuleNotFoundError("No module named 'websockets'", name="websockets") - return real_import(name, globals_, locals_, fromlist, level) - - monkeypatch.setattr(builtins, "__import__", _import_without_websocket_support) - - with pytest.raises(ModuleNotFoundError, match=r"mcp\[ws\]"): - agent_framework.MCPWebsocketTool(name="test", url="wss://example.com").get_mcp_client() - - -def test_mcp_websocket_tool_requires_mcp(monkeypatch: pytest.MonkeyPatch) -> None: - import builtins - - real_import = builtins.__import__ - sys.modules.pop("mcp.client.websocket", None) - - def _import_without_mcp( - name: str, - globals_: dict[str, object] | None = None, - locals_: dict[str, object] | None = None, - fromlist: tuple[str, ...] = (), - level: int = 0, - ) -> object: - if name == "mcp.client.websocket": - raise ModuleNotFoundError("No module named 'mcp.client.websocket'", name="mcp.client.websocket") - return real_import(name, globals_, locals_, fromlist, level) - - monkeypatch.setattr(builtins, "__import__", _import_without_mcp) - - with pytest.raises(ModuleNotFoundError, match=r"agent-framework-core\[mcp\]|mcp\[ws\]"): - agent_framework.MCPWebsocketTool(name="test", url="wss://example.com").get_mcp_client() From 29f3b1b8aac8583e96cca70c18bb79cb6febd0c3 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Thu, 1 Oct 2026 13:47:14 +0200 Subject: [PATCH 12/42] documentation updates --- .../0033-feature-usage-bitmask-user-agent.md | 7 ++++ python/packages/core/AGENTS.md | 4 +- python/packages/core/agent_framework/_mcp.py | 42 +++++++------------ .../packages/core/agent_framework/security.py | 4 +- .../samples/02-agents/observability/README.md | 2 +- 5 files changed, 29 insertions(+), 30 deletions(-) diff --git a/docs/decisions/0033-feature-usage-bitmask-user-agent.md b/docs/decisions/0033-feature-usage-bitmask-user-agent.md index 9a0b0b14e8b..2089b883d02 100644 --- a/docs/decisions/0033-feature-usage-bitmask-user-agent.md +++ b/docs/decisions/0033-feature-usage-bitmask-user-agent.md @@ -22,6 +22,13 @@ The detailed mechanism is in [SPEC-004](../specs/004-feature-usage-telemetry.md) the per-language bit tables are in [feature-usage-bit-registry.md](../specs/feature-usage-bit-registry.md). +### Amendment: MCP WebSocket transport + +As of the MCP Python SDK v2 migration, `MCPWebsocketTool` is a deprecated +compatibility symbol because upstream removed WebSocket transport. The examples +below record the public surface considered when this ADR was accepted; the +currently supported MCP transports are stdio and Streamable HTTP. + ## Decision Drivers - **Transparency** — openly documented, human-decodable, user-controllable. No diff --git a/python/packages/core/AGENTS.md b/python/packages/core/AGENTS.md index 6f6913b71d0..0efc1f19bb8 100644 --- a/python/packages/core/AGENTS.md +++ b/python/packages/core/AGENTS.md @@ -204,7 +204,9 @@ The vector store API is experimental under the shared `VECTOR_STORES` feature ID ### Model Context Protocol (`_mcp.py`) - **`MCPTool`** - Base wrapper that owns the MCP `ClientSession` and exposes the remote server's tools as `FunctionTool`s. -- **`MCPStdioTool`** / **`MCPStreamableHTTPTool`** / **`MCPWebsocketTool`** - Transport-specific subclasses. +- **`MCPStdioTool`** / **`MCPStreamableHTTPTool`** - Supported transport-specific subclasses. + **`MCPWebsocketTool`** remains only as a deprecated compatibility symbol because MCP v2 removed WebSocket + transport; it cannot create a connection. - **Argument allowlist (`_prepare_call_kwargs`)** - Before each `tools/call`, kwargs are filtered to an **allowlist** built from the tool's declared parameters (`inputSchema.properties`) plus any user-configured extras. **The declared half comes from the server's advertised schema**, and runtime kwargs (`FunctionInvocationContext.kwargs`, seeded from `function_invocation_kwargs`) are merged with the model-supplied arguments upstream in `_call_tool_with_runtime_kwargs`, so provenance is gone by the time the filter runs. A runtime kwarg is therefore forwarded whenever the server declares a property of that name, without the model mentioning it — the server, not the caller, decides which runtime kwarg names it receives. A tool that declares no usable `properties` (including schemas with `additionalProperties: true`) forwards only the configured extras. `_MCP_FRAMEWORK_DENYLIST` is a narrow safety net covering only non-serializable framework objects a server *declares* in its schema (those are dropped); it does not generalize to arbitrary caller-chosen names, and explicit extras always win. The reserved `_meta` key is never forwarded as an argument; trusted caller/runtime `_meta` is validated as MCP request metadata, model-supplied `_meta` is discarded in generated MCP functions, and metadata precedence is caller/runtime < OpenTelemetry < tools/list metadata. - **`allowed_tools`** (constructor arg on all `MCPTool` subclasses) - Restricts exposed MCP tools by raw remote MCP tool identity. Prefixed local names remain accepted only when the raw remote name already matches its normalized form; normalized/local aliases do not authorize a different raw remote name. Configured allow/approval names must identify at most one raw remote name across loaded tools and prompts: a raw name that overlaps another tool's prefixed alias raises `ToolExecutionException`. Discovery validates all pages before publishing new functions; a failed reload retains the previous functions and metadata. The allowlist is also revalidated when exposing functions so runtime changes cannot select an ambiguous name. If multiple raw remote tool names map to the same local function name, tool loading raises `ToolExecutionException` instead of first-one-wins shadowing. - **Progressive MCP disclosure** (`use_progressive_disclosure`, `always_load`) - When enabled on any `MCPTool` subclass, the initial model-facing surface is loader tools (`list_mcp_tools` / `load_tool` / `unload_tool`, prefixed by `tool_name_prefix` when configured) plus allowed tools selected by `always_load` and tools loaded earlier on the same `MCPTool` instance. `list_mcp_tools` only reports tools that pass `allowed_tools`; filtered tools are not listed or loadable. Loader tool names are reserved in progressive mode: remote MCP tools whose local generated name collides with a loader name are omitted from the initial/listed surface, and explicit `load_tool` calls return a model-visible message pointing callers to `tool_name_prefix` or excluding the colliding tool. `load_tool` accepts one tool name or a list of tool names and uses `FunctionInvocationContext.add_tools(...)` so the selected generated MCP `FunctionTool`s become available on the next function-calling iteration while keeping existing approval mode, argument filtering, header-provider runtime kwargs, result parsing, OTel, and task behavior. `unload_tool` accepts one dynamically loaded tool name or a list of names and removes them from the live tool list and persisted progressive surface, but it does not remove tools configured in `always_load`. Invalid `always_load` entries are ignored like unmatched `allowed_tools` entries. diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 4e5fd3ed805..172630d7857 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -915,8 +915,7 @@ def __init__( """Initialize the MCP Tool base. Note: - Do not use this method, use one of the subclasses: MCPStreamableHTTPTool, MCPWebsocketTool - or MCPStdioTool. + Do not use this method directly. Use MCPStreamableHTTPTool or MCPStdioTool. Args: name: The name of the MCP tool. @@ -4329,24 +4328,11 @@ async def call_tool(self, tool_name: str, **kwargs: Any) -> str | list[Content]: @deprecated("MCP WebSocket transport was removed in MCP v2. Use MCPStreamableHTTPTool instead.") class MCPWebsocketTool(MCPTool): - """MCP tool for connecting to WebSocket-based MCP servers. + """Deprecated compatibility symbol for the removed MCP WebSocket transport. - This class connects to MCP servers that communicate via WebSocket. - - Examples: - .. code-block:: python - - from agent_framework import MCPWebsocketTool, Agent - - # Create an MCP WebSocket tool - mcp_tool = MCPWebsocketTool( - name="realtime-service", url="wss://service.example.com/mcp", description="Real-time service operations" - ) - - # Use with a chat agent - async with mcp_tool: - agent = Agent(client=client, name="assistant", tools=mcp_tool) - response = await agent.run("Connect to the real-time service") + MCP v2 removed WebSocket transport because it was never part of the MCP + specification. Use :class:`MCPStreamableHTTPTool` instead. This class remains + importable during the deprecation window but cannot create a connection. """ def __init__( @@ -4377,17 +4363,16 @@ def __init__( tool_result_content: MCPToolResultContentMode = "structured_first", **kwargs: Any, ) -> None: - """Initialize the MCP WebSocket tool. + """Initialize the deprecated MCP WebSocket compatibility wrapper. Note: - The arguments are used to create a WebSocket client. - See ``mcp.client.websocket.websocket_client`` for more details. - Any extra arguments passed to the constructor will be passed to the - WebSocket client constructor. + The constructor signature is retained for source compatibility. + MCP v2 cannot create a WebSocket transport, and extra ``kwargs`` are + accepted but unused. Args: name: The name of the tool. - url: The URL of the MCP server. + url: The former WebSocket URL, retained for source compatibility. Keyword Args: tool_name_prefix: Optional prefix to prepend to exposed MCP function names. @@ -4469,7 +4454,7 @@ def __init__( transports. ``None`` disables the limit. tool_result_content: How to choose model-visible text when both ``content`` and ``structuredContent`` are present. See :data:`MCPToolResultContentMode`. - kwargs: Any extra arguments to pass to the WebSocket client. + kwargs: Deprecated compatibility arguments. They are not used. """ super().__init__( name=name, @@ -4517,4 +4502,9 @@ def _mcp_base_span_attributes(self) -> dict[str, Any]: return attrs def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: + """Raise because MCP v2 removed WebSocket transport. + + Raises: + RuntimeError: Always. Use :class:`MCPStreamableHTTPTool` instead. + """ raise RuntimeError("MCP WebSocket transport was removed in MCP v2. Use MCPStreamableHTTPTool instead.") diff --git a/python/packages/core/agent_framework/security.py b/python/packages/core/agent_framework/security.py index 7433d3f940b..1eabc0c2b16 100644 --- a/python/packages/core/agent_framework/security.py +++ b/python/packages/core/agent_framework/security.py @@ -4431,7 +4431,7 @@ async def apply_mcp_security_labels( Args: mcp_tool: A connected ``MCPTool`` instance (``MCPStdioTool``, - ``MCPStreamableHTTPTool``, ``MCPWebsocketTool``). + ``MCPStreamableHTTPTool``). default_integrity: Integrity label to assign when the server provides no annotations. Defaults to ``UNTRUSTED`` (conservative). annotation_overrides: Optional per-tool-name overrides. Keys are @@ -4637,7 +4637,7 @@ class SecureMCPToolProxy: There are two ways to create a proxy: - 1. **Wrap an existing MCPTool** (local binary, WebSocket, or HTTP):: + 1. **Wrap an existing MCPTool** (local binary or Streamable HTTP):: async with SecureMCPToolProxy( MCPStdioTool(name="github", command="gh-mcp", args=["stdio"]) diff --git a/python/samples/02-agents/observability/README.md b/python/samples/02-agents/observability/README.md index f456f93dd7d..84f5db883bc 100644 --- a/python/samples/02-agents/observability/README.md +++ b/python/samples/02-agents/observability/README.md @@ -220,7 +220,7 @@ Because Agent Framework is **natively instrumented** with OpenTelemetry, you do Whenever there is an active OpenTelemetry span context, Agent Framework automatically propagates trace context to MCP servers via the `params._meta` field of `tools/call` requests. It uses the globally configured OpenTelemetry propagator(s) — W3C Trace Context by default (producing `traceparent` and `tracestate`) — so custom propagators (B3, Jaeger, etc.) are also supported. This enables distributed tracing across agent-to-MCP-server boundaries, compliant with the [MCP `_meta` specification](https://modelcontextprotocol.io/specification/2025-11-25/basic#_meta). -**Scope:** automatic `_meta` injection applies only to MCP sessions that the agent process itself opens — `MCPStreamableHTTPTool`, `MCPStdioTool`, and `MCPWebsocketTool` (or any other client-opened `MCPTool` subclass). It does **not** apply to hosted or provider-managed MCP tool configurations such as `FoundryChatClient.get_mcp_tool(...)`, `OpenAIChatClient.get_mcp_tool(...)`, `AnthropicClient.get_mcp_tool(...)`, `GeminiChatClient.get_mcp_tool(...)`, or toolbox-fetched tools (e.g. `toolbox = await client.get_toolbox(...)` then `Agent(tools=toolbox.tools)`). In those cases the `tools/call` message is issued by the provider service runtime rather than by the agent process, so propagating `traceparent`/`tracestate` across that boundary is the service runtime's responsibility. If you need end-to-end distributed tracing to the downstream MCP server, use a client-opened MCP transport instead of a hosted connector. +**Scope:** automatic `_meta` injection applies only to MCP sessions that the agent process itself opens — `MCPStreamableHTTPTool`, `MCPStdioTool`, or another supported client-opened `MCPTool` subclass. It does **not** apply to hosted or provider-managed MCP tool configurations such as `FoundryChatClient.get_mcp_tool(...)`, `OpenAIChatClient.get_mcp_tool(...)`, `AnthropicClient.get_mcp_tool(...)`, `GeminiChatClient.get_mcp_tool(...)`, or toolbox-fetched tools (e.g. `toolbox = await client.get_toolbox(...)` then `Agent(tools=toolbox.tools)`). In those cases the `tools/call` message is issued by the provider service runtime rather than by the agent process, so propagating `traceparent`/`tracestate` across that boundary is the service runtime's responsibility. If you need end-to-end distributed tracing to the downstream MCP server, use a client-opened MCP transport instead of a hosted connector. ## Configuration From dc86593a3c237a587ea765e56e3339c38b96ae46 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Fri, 2 Oct 2026 15:46:30 +0200 Subject: [PATCH 13/42] Moved transport setup to mcp.Client, added test for v2. Needs legacy check --- python/packages/core/agent_framework/_mcp.py | 63 ++++++----------- python/packages/core/tests/core/test_mcp.py | 72 ++++++++++++++++++++ 2 files changed, 94 insertions(+), 41 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 172630d7857..f3ad33ca437 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -64,10 +64,12 @@ from typing_extensions import Self # pragma: no cover if TYPE_CHECKING: - from httpx import AsyncClient, Request, Response + # TODO(jpalvarezl): clean up and consolidate under httpx2 + import httpx2 + from httpx import Request from mcp import types + from mcp.client.context import ClientRequestContext from mcp.client.session import ClientSession - from mcp.shared.context import RequestContext from ._clients import SupportsChatGetResponse from ._middleware import FunctionInvocationContext @@ -448,7 +450,7 @@ def _capture_mcp_tool_result(mcp_type: Any) -> None: class _MCPHeaderScopedClient: """Attach private tool context to MCP transport requests.""" - def __init__(self, client: AsyncClient, owner: object) -> None: + def __init__(self, client: httpx2.AsyncClient, owner: object) -> None: self._client = client self._owner = owner @@ -472,7 +474,7 @@ def _tagged_kwargs(self, kwargs: dict[str, Any]) -> dict[str, Any]: def stream(self, *args: Any, **kwargs: Any) -> Any: return self._client.stream(*args, **self._tagged_kwargs(kwargs)) - async def send(self, request: Request, **kwargs: Any) -> Response: + async def send(self, request: httpx2.Request, **kwargs: Any) -> httpx2.Response: request.extensions[_MCP_HEADER_OWNER_EXTENSION] = self._owner return await self._client.send(request, **kwargs) @@ -1961,30 +1963,9 @@ async def _connect_on_owner( self._reset_session_discovery_state() self._exit_stack = AsyncExitStack() if not self.session: - try: - transport = await self._exit_stack.enter_async_context(self.get_mcp_client()) - except (Exception, asyncio.CancelledError) as ex: - # On Python >= 3.11, re-raise genuine task cancellation (task.cancelling() > 0) - # instead of wrapping it in ToolException. On Python < 3.11, task.cancelling() - # is unavailable so MCP-internal CancelledErrors cannot be distinguished from - # caller-driven cancellation; they are wrapped as ToolException in that case. - cancelled, cleanup_error = await self._close_and_check_cancelled(ex) - if cancelled: - raise - command = getattr(self, "command", None) - if command: - error_msg = f"Failed to start MCP server '{command}': {_describe_with_cleanup(ex, cleanup_error)}" - else: - error_msg = f"Failed to connect to MCP server: {_describe_with_cleanup(ex, cleanup_error)}" - # CancelledError is a BaseException (not Exception) on Python >= 3.8, so - # inner_exception=None and ToolException.__init__ won't log exc_info. - if isinstance(ex, asyncio.CancelledError): - logger.debug(error_msg, exc_info=True) - raise ToolException(error_msg, inner_exception=ex if isinstance(ex, Exception) else None) from ex try: try: - from mcp import types - from mcp.client.session import ClientSession as runtime_client_session + from mcp import Client, types except ModuleNotFoundError as ex: await self._safe_close_exit_stack() raise ToolException( @@ -1997,19 +1978,19 @@ async def _connect_on_owner( sampling_capabilities = types.SamplingCapability( tools=types.SamplingToolsCapability(), ) - session = await self._exit_stack.enter_async_context( - runtime_client_session( - read_stream=transport[0], - write_stream=transport[1], + mcp_client = await self._exit_stack.enter_async_context( + Client( + server=self.get_mcp_client(), read_timeout_seconds=( - timedelta(seconds=self.request_timeout) if self.request_timeout else None + timedelta(seconds=self.request_timeout).seconds if self.request_timeout else None ), message_handler=self.message_handler, logging_callback=self.logging_callback, - sampling_callback=self.sampling_callback, # pyright: ignore[reportDeprecated] sampling_capabilities=sampling_capabilities, + sampling_callback=self.sampling_callback, # pyright: ignore[reportDeprecated] ) ) + session = mcp_client.session except (Exception, asyncio.CancelledError) as ex: cancelled, cleanup_error = await self._close_and_check_cancelled(ex) if cancelled: @@ -2023,9 +2004,9 @@ async def _connect_on_owner( ) from ex try: with create_mcp_client_span("initialize", attributes=self._mcp_base_span_attributes()) as init_span: - initialize_result = await session.initialize() - init_span.set_attribute(OtelAttr.MCP_PROTOCOL_VERSION, initialize_result.protocol_version) - self._set_server_capabilities(getattr(initialize_result, "capabilities", None)) + init_span.set_attribute(OtelAttr.MCP_PROTOCOL_VERSION, mcp_client.protocol_version) + self._set_server_capabilities(mcp_client.server_capabilities) + self._ping_available = session.initialize_result is not None except (Exception, asyncio.CancelledError) as ex: cancelled, cleanup_error = await self._close_and_check_cancelled(ex) if cancelled: @@ -2160,7 +2141,7 @@ def _warn_sampling_deprecated(self, *, stacklevel: int) -> None: @deprecated(_MCP_SAMPLING_DEPRECATION_MESSAGE, category=None) async def sampling_callback( self, - context: RequestContext[ClientSession, Any], + context: ClientRequestContext, params: types.CreateMessageRequestParams, ) -> types.CreateMessageResult | types.CreateMessageResultWithTools | types.ErrorData: """Callback function for sampling. @@ -3763,7 +3744,7 @@ def __init__( sampling_max_tokens: int | None = _DEFAULT_SAMPLING_MAX_TOKENS, sampling_max_requests: int | None = _DEFAULT_SAMPLING_MAX_REQUESTS, additional_properties: dict[str, Any] | None = None, - http_client: AsyncClient | None = None, + http_client: httpx2.AsyncClient | None = None, static_headers: Mapping[str, str] | None = None, header_provider: Callable[[dict[str, Any]], dict[str, str]] | None = None, task_options: MCPTaskOptions | None = None, @@ -3964,7 +3945,7 @@ def __init__( ) self.url = url self.terminate_on_close = terminate_on_close - self._httpx_client: AsyncClient | None = http_client + self._httpx_client: httpx2.AsyncClient | None = http_client self._static_headers = dict(static_headers or {}) self._header_provider = header_provider # Headers for the in-flight call_tool invocation. The streamable HTTP transport @@ -3986,7 +3967,7 @@ def __init__( self._pending_connection_kwargs: dict[str, Any] | None = None self._call_headers_lock = asyncio.Lock() self._header_request_owner = object() - self._header_hook_client: AsyncClient | None = None + self._header_hook_client: httpx2.AsyncClient | None = None def _mcp_base_span_attributes(self) -> dict[str, Any]: attrs = super()._mcp_base_span_attributes() @@ -4012,7 +3993,7 @@ def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: Returns: An async context manager for the streamable HTTP client transport. """ - from httpx import URL, AsyncClient, Timeout + from httpx2 import URL, AsyncClient, Timeout self._promote_pending_session_headers() @@ -4109,7 +4090,7 @@ async def _inject_headers(request: Request) -> None: # ruff:ignore[unused-async terminate_on_close=self.terminate_on_close if self.terminate_on_close is not None else True, ) - async def _close_owned_http_client(self, http_client: AsyncClient) -> None: + async def _close_owned_http_client(self, http_client: httpx2.AsyncClient) -> None: """Release a framework-created client without retaining it for reconnect.""" try: await http_client.aclose() diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index a513913303c..43e6740fe72 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -8096,6 +8096,78 @@ async def handler(request: httpx.Request) -> httpx.Response: assert initialize_headers[0].get("x-api-key") == "connect-token" +async def test_mcp_streamble_http_tool_connects_to_v2_server() -> None: + + from httpx2 import AsyncClient, MockTransport, Request, Response + + captured_methods: list[str] = [] + + async def mcp_v2_server_mock_handler(request: Request) -> Response: + if request.method == "DELETE": + return Response(200) + + body = json.loads(request.content) + method = body["method"] + captured_methods.append(method) + + if method == "initialize": + return Response( + 200, + json={ + "jsonrpc": "2.0", + "id": body["id"], + "error": {"code": -32601, "message": "Method not found"}, + }, + ) + + if method == "server/discover": + return Response( + 200, + json={ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "supportedVersions": ["2026-07-28"], + "capabilities": {"tools": {}}, + }, + }, + ) + + if method == "tools/list": + return Response( + 200, + json={ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "cacheScope": "private", + "resultType": "complete", + "ttlMs": 0, + "tools": [{"name": "greet", "inputSchema": {"type": "object", "properties": {}}}], + }, + }, + ) + + raise AssertionError(f"Unexpected MCP method: {method}") + + user_client = AsyncClient(transport=MockTransport(mcp_v2_server_mock_handler)) + + tool_a = MCPStreamableHTTPTool( + name="a", + url="http://example.com/mcp", + http_client=user_client, + header_provider=lambda _kw: {"Authorization": "Bearer token-a"}, + ) + + async with tool_a: + assert tool_a.session is not None + assert tool_a.session.protocol_version == "2026-07-28" + assert [function.name for function in tool_a.functions] == ["greet"] + + assert "server/discover" in captured_methods + assert "initialize" not in captured_methods + + async def test_agent_context_manager_authenticates_connect_with_closure_provider( client: SupportsChatGetResponse, ) -> None: From 2697e80de713f2639707ca3388895d78faab0fc5 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Mon, 5 Oct 2026 10:28:54 +0200 Subject: [PATCH 14/42] Added legacy handshake unit test --- python/packages/core/agent_framework/_mcp.py | 1 + python/packages/core/tests/core/test_mcp.py | 81 ++++++++++++++++++++ 2 files changed, 82 insertions(+) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index f3ad33ca437..a90e3c94bdf 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -1979,6 +1979,7 @@ async def _connect_on_owner( tools=types.SamplingToolsCapability(), ) mcp_client = await self._exit_stack.enter_async_context( + # default "mode" is `auto` which automatically negotiates with the right protocol version Client( server=self.get_mcp_client(), read_timeout_seconds=( diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 43e6740fe72..91e37629be6 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -8168,6 +8168,87 @@ async def mcp_v2_server_mock_handler(request: Request) -> Response: assert "initialize" not in captured_methods +async def test_mcp_streamable_http_tool_connects_to_legacy_server() -> None: + + from httpx2 import AsyncClient, MockTransport, Request, Response + + captured_methods: list[str] = [] + + async def mcp_v2_server_mock_handler(request: Request) -> Response: + if request.method == "DELETE": + return Response(200) + + body = json.loads(request.content) + method = body["method"] + captured_methods.append(method) + + if method == "notifications/initialized": + return Response(202) + + if method == "ping": + return Response( + 200, + json={"jsonrpc": "2.0", "id": body["id"], "result": {}}, + ) + if method == "server/discover": + return Response( + 200, + json={ + "jsonrpc": "2.0", + "id": body["id"], + "error": {"code": -32601, "message": "Method not found"}, + }, + ) + + if method == "initialize": + return Response( + 200, + json={ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "protocolVersion": "2025-11-25", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "legacy-server", "version": "1.0.0"}, + }, + }, + ) + + if method == "tools/list": + return Response( + 200, + json={ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "cacheScope": "private", + "resultType": "complete", + "ttlMs": 0, + "tools": [{"name": "greet", "inputSchema": {"type": "object", "properties": {}}}], + }, + }, + ) + + raise AssertionError(f"Unexpected MCP method: {method}") + + user_client = AsyncClient(transport=MockTransport(mcp_v2_server_mock_handler)) + + tool_a = MCPStreamableHTTPTool( + name="a", + url="http://example.com/mcp", + http_client=user_client, + header_provider=lambda _kw: {"Authorization": "Bearer token-a"}, + ) + + async with tool_a: + assert tool_a.session is not None + assert tool_a.session.protocol_version == "2025-11-25" + assert [function.name for function in tool_a.functions] == ["greet"] + + assert "server/discover" in captured_methods + assert "initialize" in captured_methods + + async def test_agent_context_manager_authenticates_connect_with_closure_provider( client: SupportsChatGetResponse, ) -> None: From d70cb898a1cf53c8e6f7221c49587f72705b1318 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Mon, 5 Oct 2026 10:30:26 +0200 Subject: [PATCH 15/42] Added ADR for mcp v2 --- .../0045-python-mcp-v2-client-lifecycle.md | 134 ++++++++++++++++++ 1 file changed, 134 insertions(+) create mode 100644 docs/decisions/0045-python-mcp-v2-client-lifecycle.md diff --git a/docs/decisions/0045-python-mcp-v2-client-lifecycle.md b/docs/decisions/0045-python-mcp-v2-client-lifecycle.md new file mode 100644 index 00000000000..90d1440f937 --- /dev/null +++ b/docs/decisions/0045-python-mcp-v2-client-lifecycle.md @@ -0,0 +1,134 @@ +--- +status: proposed +date: 2026-10-01 +deciders: eavanvalkenburg, westey-m +--- + +# Adopt the MCP v2 negotiating client behind Python MCP tools + +## Context and Problem Statement + +The Python `MCPTool` implementation owns an MCP transport, constructs a low-level `ClientSession`, and always calls +`initialize()`. That lifecycle supports handshake-era protocol versions through 2025-11-25, but it cannot connect to +a 2026-07-28 server, where `server/discover` replaces the initialization handshake. + +The MCP Python SDK v2 provides a high-level `Client(mode="auto")` that probes `server/discover` and falls back to +`initialize` when the probe does not provide positive evidence of a compatible modern peer. Agent Framework needs to +adopt that negotiation without replacing its public `MCPTool`, `MCPStdioTool`, or `MCPStreamableHTTPTool` classes, or +regressing async context management, reconnect behavior, caller-owned sessions, and request-scoped authentication +headers. + +## Decision Drivers + +- One Agent Framework code path must support 2026-07-28 and 2025-era MCP peers. +- Protocol negotiation should remain owned by the MCP SDK rather than being independently reimplemented. +- Existing public tool classes, `async with` ergonomics, and tool/prompt discovery behavior must remain compatible. +- Framework-created resources must still be closed by the task that opened them, including cancellation and failed + connection attempts. +- Caller-supplied sessions and HTTP clients must remain caller-owned. +- Dynamic HTTP header identity must continue to bind discovery and later requests to the same authenticated identity. +- Stdio and Streamable HTTP must share the same protocol lifecycle. + +## Considered Options + +- Enter the MCP SDK v2 `Client(mode="auto")` behind the existing Agent Framework tool classes. +- Reimplement `server/discover` and legacy fallback directly around `ClientSession`. +- Add separate modern and legacy Agent Framework tool classes or a public protocol-mode switch. + +## Decision Outcome + +Chosen option: "Enter the MCP SDK v2 `Client(mode="auto")` behind the existing Agent Framework tool classes", +because it uses the SDK's supported negotiation path while preserving the Agent Framework API and transport-specific +behavior. + +For framework-created connections: + +- `MCPStdioTool` and `MCPStreamableHTTPTool` continue to create their existing transport context managers. +- `MCPTool` passes that transport to `mcp.Client(mode="auto")` and enters the client on its existing `AsyncExitStack`. +- The SDK probes `server/discover`. A compatible modern peer is adopted without sending `initialize`; a probe that + does not establish a compatible modern peer falls back to the initialization handshake on the same code path. + `METHOD_NOT_FOUND` is the representative legacy response covered by Agent Framework's contract test, not the SDK's + only fallback condition. +- Agent Framework retains a private reference to the high-level client and exposes its underlying `ClientSession` + through the existing `session` attribute for compatibility with current integrations. +- Negotiated protocol version and server capabilities are read from the client's era-neutral properties rather than + captured only from an initialize result. +- Tools configured for the existing server-initiated sampling callback use `mode="legacy"`. The MCP migration guide + requires legacy mode for that back-channel behavior because 2026-07-28 refuses server-initiated sampling on every + transport. Auto mode remains the default when no legacy-only callback behavior is requested. + +Agent Framework retains its lifecycle owner task and locks. They protect framework state, preserve AnyIO task +ownership during teardown, serialize identity-changing reconnects, and roll back partially loaded discovery state. +They no longer implement protocol negotiation. A reset or reconnect closes the whole SDK client and transport, then +constructs a new auto-negotiating client. `ping` is not used as a universal liveness preflight because it does not +exist in 2026-07-28; operation failures drive reconnect, while any retained legacy ping behavior is gated by the +negotiated protocol. + +Caller-supplied `ClientSession` remains a compatibility path: + +- Agent Framework does not enter, close, or replace the session. +- An already negotiated session is reused through its `protocol_version` and `server_capabilities` properties. +- An unnegotiated session retains the existing compatibility behavior and uses `initialize()`. A caller that supplies + a modern low-level session negotiates it with `discover()` before passing it to Agent Framework. This avoids + duplicating the SDK's broader auto-negotiation policy, which is implemented only by the high-level `Client`. + +Dynamic and static HTTP headers remain below the negotiating client at the Streamable HTTP transport boundary. The +effective header set continues to define connection identity; changing it closes and rebuilds the whole SDK client +before discovery or tool calls proceed. This preserves the trust and approval boundaries documented in +[ADR 0043](0043-python-mcp-runtime-context.md). + +The MCP Streamable HTTP implementation uses `httpx2` types, as required by MCP SDK v2. `httpx` and `httpx2` may +coexist elsewhere in the repository; this decision does not require unrelated HTTP features to migrate. + +### Consequences + +- Good, because modern and legacy servers use one public Agent Framework API and one SDK-supported negotiation path. +- Good, because stdio, Streamable HTTP, reconnect, cancellation, and header identity remain Agent Framework concerns + without duplicating protocol-version selection. +- Good, because later work can use the high-level client for per-request metadata, subscriptions, MRTR, and caching. +- Neutral, because Agent Framework keeps both a private high-level client and the public low-level `session` view. +- Neutral, because caller-supplied unnegotiated sessions remain handshake-era unless the caller negotiates them first. +- Bad, because tests that mocked `ClientSession` construction must move toward client/transport contract tests. + +## Validation + +Contract tests cover both branches through the same Agent Framework tool: + +- A modern-only Streamable HTTP peer accepts `server/discover`, rejects `initialize`, serves `tools/list`, and records + that the negotiated protocol is 2026-07-28. +- A 2025-era peer rejects `server/discover`, accepts `initialize`, and serves the same tool operations. +- Equivalent stdio coverage verifies that negotiation is transport-independent. +- Existing lifecycle, reconnect, cancellation, caller-ownership, and dynamic-header identity tests remain green. + +Follow-on tests cover per-request logging metadata, `subscriptions/listen` with legacy notification fallback, prompts, +skills, MRTR, and caching. Tasks remain deferred until the Python MCP SDK exposes the 2026 Tasks extension runtime. + +## Pros and Cons of the Options + +### Enter the MCP SDK v2 client behind existing tool classes + +- Good, because the SDK owns current and future negotiation rules. +- Good, because the existing transport subclasses and public API remain intact. +- Good, because it matches the .NET direction of delegating connection creation and negotiation to its MCP SDK. +- Bad, because the existing connection code and mocks must be reshaped around a higher-level lifecycle owner. + +### Reimplement negotiation around ClientSession + +- Good, because it minimizes the first code diff and preserves direct session construction. +- Bad, because Agent Framework would duplicate the SDK's discover, fallback, adoption, and future-version behavior. +- Bad, because later high-level SDK features would still require a second migration. + +### Add era-specific tool classes or a public mode switch + +- Good, because callers could force a known protocol era. +- Bad, because callers should not need to know a server's era before connecting. +- Bad, because it duplicates public classes, tests, documentation, and lifecycle behavior. +- Bad, because a mode switch can accidentally disable fallback and fragment compatibility. + +## More Information + +- [Python MCP 2026-07-28 umbrella issue](https://github.com/microsoft/agent-framework/issues/8245) +- [MCP Python SDK v1-to-v2 migration guide](https://py.sdk.modelcontextprotocol.io/migration/#clients) +- [.NET MCP Tasks migration](https://github.com/microsoft/agent-framework/pull/7774) +- The .NET declarative MCP handler delegates connection creation to `McpClient.CreateAsync`; its protocol stub rejects + `server/discover` with `METHOD_NOT_FOUND` before accepting `initialize`, demonstrating SDK-owned legacy fallback. From 2527fe3e651a15c7a5cbb0b2aeadc2d552a9a4f0 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Mon, 5 Oct 2026 11:17:49 +0200 Subject: [PATCH 16/42] Added tool/call mock and assertions --- python/packages/core/tests/core/test_mcp.py | 41 ++++++++++++++++++++- 1 file changed, 39 insertions(+), 2 deletions(-) diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 91e37629be6..bab29a4659a 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -8148,6 +8148,20 @@ async def mcp_v2_server_mock_handler(request: Request) -> Response: }, ) + if method == "tools/call": + return Response( + 200, + json={ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "resultType": "complete", + "content": [{"type": "text", "text": "Hello!"}], + "isError": False, + }, + }, + ) + raise AssertionError(f"Unexpected MCP method: {method}") user_client = AsyncClient(transport=MockTransport(mcp_v2_server_mock_handler)) @@ -8164,8 +8178,13 @@ async def mcp_v2_server_mock_handler(request: Request) -> Response: assert tool_a.session.protocol_version == "2026-07-28" assert [function.name for function in tool_a.functions] == ["greet"] + result = await tool_a.call_tool("greet") + assert isinstance(result, list) + assert [item.text for item in result if item.type == "text"] == ["Hello!"] + assert "server/discover" in captured_methods assert "initialize" not in captured_methods + assert "tools/list" in captured_methods async def test_mcp_streamable_http_tool_connects_to_legacy_server() -> None: @@ -8174,7 +8193,7 @@ async def test_mcp_streamable_http_tool_connects_to_legacy_server() -> None: captured_methods: list[str] = [] - async def mcp_v2_server_mock_handler(request: Request) -> Response: + async def mcp_legacy_server_mock_handler(request: Request) -> Response: if request.method == "DELETE": return Response(200) @@ -8229,9 +8248,22 @@ async def mcp_v2_server_mock_handler(request: Request) -> Response: }, ) + if method == "tools/call": + return Response( + 200, + json={ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "content": [{"type": "text", "text": "Hello!"}], + "isError": False, + }, + }, + ) + raise AssertionError(f"Unexpected MCP method: {method}") - user_client = AsyncClient(transport=MockTransport(mcp_v2_server_mock_handler)) + user_client = AsyncClient(transport=MockTransport(mcp_legacy_server_mock_handler)) tool_a = MCPStreamableHTTPTool( name="a", @@ -8245,8 +8277,13 @@ async def mcp_v2_server_mock_handler(request: Request) -> Response: assert tool_a.session.protocol_version == "2025-11-25" assert [function.name for function in tool_a.functions] == ["greet"] + result = await tool_a.call_tool("greet") + assert isinstance(result, list) + assert [item.text for item in result if item.type == "text"] == ["Hello!"] + assert "server/discover" in captured_methods assert "initialize" in captured_methods + assert "tools/list" in captured_methods async def test_agent_context_manager_authenticates_connect_with_closure_provider( From a1cd03869d3cc7e87603b0f844ec1fa75cbf47c2 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Mon, 5 Oct 2026 13:01:20 +0200 Subject: [PATCH 17/42] Adjusted sampling capabilities test --- python/packages/core/tests/core/test_mcp.py | 70 +++++++++------------ 1 file changed, 28 insertions(+), 42 deletions(-) diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index bab29a4659a..8f01dbabe6a 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -4540,59 +4540,45 @@ async def test_mcp_tool_sampling_callback_always_passes_max_tokens(): async def test_connect_sampling_capabilities_with_client(): - """Test connect() passes sampling_capabilities to ClientSession when client is set.""" + """Test connect() passes sampling_capabilities to mcp.Client""" tool = MCPStdioTool(name="test", command="test-command", load_tools=False, load_prompts=False) tool.client = Mock() - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] - - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session = AsyncMock() - mock_session._request_id = 1 + with patch("mcp.Client") as mock_client_class: + sdk_client = AsyncMock() + sdk_client.__aenter__.return_value = sdk_client + sdk_client.session = Mock(initialize_result=None) + sdk_client.protocol_version = "2026-07-28" + sdk_client.server_capabilities = types.ServerCapabilities() + mock_client_class.return_value = sdk_client - session_cm = AsyncMock() - session_cm.__aenter__ = AsyncMock(return_value=mock_session) - session_cm.__aexit__ = AsyncMock(return_value=None) - mock_session_class.return_value = session_cm - - await tool.connect() - - call_kwargs = mock_session_class.call_args.kwargs - sampling_caps = call_kwargs.get("sampling_capabilities") - assert sampling_caps is not None - assert isinstance(sampling_caps, types.SamplingCapability) - assert sampling_caps.tools is not None - assert isinstance(sampling_caps.tools, types.SamplingToolsCapability) + async with tool: + call_kwargs = mock_client_class.call_args.kwargs + sampling_caps = call_kwargs.get("sampling_capabilities") + assert sampling_caps is not None + assert isinstance(sampling_caps, types.SamplingCapability) + assert sampling_caps.tools is not None + assert isinstance(sampling_caps.tools, types.SamplingToolsCapability) async def test_connect_no_sampling_capabilities_without_client(): """Test connect() does not pass sampling_capabilities when no client is set.""" tool = MCPStdioTool(name="test", command="test-command", load_tools=False, load_prompts=False) - # No client set - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] - - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session = AsyncMock() - mock_session._request_id = 1 - - session_cm = AsyncMock() - session_cm.__aenter__ = AsyncMock(return_value=mock_session) - session_cm.__aexit__ = AsyncMock(return_value=None) - mock_session_class.return_value = session_cm + with patch("mcp.Client") as mock_client_class: + sdk_client = AsyncMock() + sdk_client.__aenter__.return_value = sdk_client + sdk_client.session = Mock(initialize_result=None) + sdk_client.protocol_version = "2026-07-28" + sdk_client.server_capabilities = types.ServerCapabilities() + mock_client_class.return_value = sdk_client - await tool.connect() - - call_kwargs = mock_session_class.call_args.kwargs - assert call_kwargs.get("sampling_capabilities") is None + try: + await tool.connect() + call_kwargs = mock_client_class.call_args.kwargs + assert call_kwargs.get("sampling_capabilities") is None + finally: + await tool.close() # Test error handling in connect() method From 3f827472e4ddc695b37b08ec99ceb9a7d3adc8fb Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Mon, 5 Oct 2026 14:19:53 +0200 Subject: [PATCH 18/42] removing ClientSession references in favour of using mcp.Client --- python/packages/core/agent_framework/_mcp.py | 9 +- python/packages/core/tests/core/test_mcp.py | 497 ++++++++----------- 2 files changed, 222 insertions(+), 284 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index a90e3c94bdf..0906b49ab47 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -1996,7 +1996,14 @@ async def _connect_on_owner( cancelled, cleanup_error = await self._close_and_check_cancelled(ex) if cancelled: raise - session_error_msg = f"Failed to create MCP session: {_describe_with_cleanup(ex, cleanup_error)}" + described = _describe_with_cleanup(ex, cleanup_error) + command = getattr(self, "command", None) + if command: + args_str = " ".join(getattr(self, "args", [])) + full_command = f"{command} {args_str}".strip() + session_error_msg = f"Failed to create MCP session for server '{full_command}': {described}" + else: + session_error_msg = f"Failed to create MCP session: {described}" if isinstance(ex, asyncio.CancelledError): logger.debug(session_error_msg, exc_info=True) raise ToolException( diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 8f01dbabe6a..9e9864dd100 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -8,9 +8,10 @@ import os import sys import warnings -from contextlib import _AsyncGeneratorContextManager # type: ignore +from contextlib import AbstractAsyncContextManager, _AsyncGeneratorContextManager # type: ignore from contextvars import ContextVar from datetime import timedelta +from types import TracebackType from typing import Any, cast from unittest.mock import AsyncMock, Mock, patch @@ -74,6 +75,57 @@ def _mcp_result_to_text(result: str | list[Content]) -> str: return text or str(result) +def _mock_sdk_client( + *, + session: Mock | None = None, + capabilities: types.ServerCapabilities | None = None, + protocol_version: str = "2025-11-25", +) -> AsyncMock: + """Model the SDK client's connected session and negotiated metadata.""" + capabilities = capabilities if capabilities is not None else types.ServerCapabilities() + session = session if session is not None else Mock(spec=ClientSession) + session.protocol_version = protocol_version + session.server_capabilities = capabilities + session.initialize_result = ( + types.InitializeResult( + protocol_version=protocol_version, + capabilities=capabilities, + server_info=types.Implementation(name="mock-server", version="1.0"), + ) + if protocol_version == "2025-11-25" + else None + ) + + client = AsyncMock() + client.session = session + client.protocol_version = protocol_version + client.server_capabilities = capabilities + client.__aenter__.return_value = client + + return client + + +class _TransportBoundClientContext: + """Model SDK-owned transport entry and exit without running a protocol dispatcher.""" + + def __init__(self, transport: AbstractAsyncContextManager[Any], client: AsyncMock) -> None: + self.transport = transport + self.client = client + self.exit_stack = contextlib.AsyncExitStack() + + async def __aenter__(self) -> AsyncMock: + async with contextlib.AsyncExitStack() as stack: + await stack.enter_async_context(self.transport) + client = await stack.enter_async_context(self.client) + self.exit_stack = stack.pop_all() + return client + + async def __aexit__( + self, exc_type: type[BaseException] | None, exc: BaseException | None, tb: TracebackType | None + ) -> bool | None: + return await self.exit_stack.__aexit__(exc_type, exc, tb) + + _HELPER_MCP_TOOL = MCPTool(name="helper") # type: ignore[abstract] # ty: ignore[call-non-callable] @@ -4585,53 +4637,37 @@ async def test_connect_no_sampling_capabilities_without_client(): async def test_connect_session_creation_failure(): - """Test connect() raises ToolException when ClientSession creation fails.""" + """Test connect() preserves the cause when SDK client construction fails.""" tool = MCPStdioTool(name="test", command="test-command") - # Mock successful transport creation - mock_transport = (Mock(), Mock()) # (read_stream, write_stream) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] - - # Mock ClientSession to raise an exception - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session_class.side_effect = RuntimeError("Session creation failed") - + with patch("mcp.Client", side_effect=RuntimeError("Client creation failed")): with pytest.raises(ToolException) as exc_info: await tool.connect() assert "Failed to create MCP session" in str(exc_info.value) - assert "Session creation failed" in str(exc_info.value) # exception text is now part of the message - assert "Session creation failed" in str(exc_info.value.__cause__) + assert "Client creation failed" in str(exc_info.value) + assert "Client creation failed" in str(exc_info.value.__cause__) + await tool.close() async def test_connect_initialization_failure_http_no_command(): - """Test connect() when session.initialize() fails for HTTP tool (no command attribute).""" + """SDK negotiation fails during client entry, before a connected session is available.""" tool = MCPStreamableHTTPTool(name="test", url="http://example.com") + tool.get_mcp_client = Mock(return_value=Mock()) # type: ignore[method-assign] + sdk_client = _mock_sdk_client() + failure = ConnectionError("Server not ready") + sdk_client.__aenter__.side_effect = failure - # Mock successful transport creation - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] - - # Mock successful session creation but failed initialization - mock_session = Mock() - mock_session.initialize = AsyncMock(side_effect=ConnectionError("Server not ready")) - - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session_class.return_value.__aenter__ = AsyncMock(return_value=mock_session) - mock_session_class.return_value.__aexit__ = AsyncMock(return_value=None) - + with patch("mcp.Client", return_value=sdk_client): with pytest.raises(ToolException) as exc_info: await tool.connect() - # Should use generic error message since HTTP tool doesn't have command - assert "MCP server failed to initialize" in str(exc_info.value) + assert "Failed to create MCP session" in str(exc_info.value) assert "Server not ready" in str(exc_info.value) + assert exc_info.value.__cause__ is failure + assert tool.session is None + assert tool.is_connected is False + await tool.close() async def test_connect_cleanup_on_transport_failure(): @@ -4649,6 +4685,7 @@ async def test_connect_cleanup_on_transport_failure(): # Verify cleanup was called tool._exit_stack.aclose.assert_called_once() + await tool.close() async def test_connect_cleanup_on_transport_failure_http_uses_generic_message(): @@ -4657,39 +4694,30 @@ async def test_connect_cleanup_on_transport_failure_http_uses_generic_message(): tool._exit_stack.aclose = AsyncMock() # type: ignore[method-assign] tool.get_mcp_client = Mock(side_effect=RuntimeError("Transport failed")) # type: ignore[method-assign] - with pytest.raises(ToolException, match="Failed to connect to MCP server: Transport failed"): + with pytest.raises(ToolException, match="Failed to create MCP session: Transport failed"): await tool.connect() tool._exit_stack.aclose.assert_called_once() + await tool.close() async def test_connect_cleanup_on_initialization_failure(): - """Test that _exit_stack.aclose() is called when initialization fails.""" + """Test that framework cleanup runs when SDK negotiation fails during entry.""" tool = MCPStdioTool(name="test", command="test-command") # Mock _exit_stack.aclose to verify it's called tool._exit_stack.aclose = AsyncMock() # type: ignore[method-assign] - # Mock successful transport creation - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] - - # Mock successful session creation but failed initialization - mock_session = Mock() - mock_session.initialize = AsyncMock(side_effect=RuntimeError("Init failed")) - - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session_class.return_value.__aenter__ = AsyncMock(return_value=mock_session) - mock_session_class.return_value.__aexit__ = AsyncMock(return_value=None) + sdk_client = _mock_sdk_client() + sdk_client.__aenter__.side_effect = RuntimeError("Init failed") + with patch("mcp.Client", return_value=sdk_client): with pytest.raises(ToolException): await tool.connect() # Verify cleanup was called tool._exit_stack.aclose.assert_called_once() + await tool.close() async def test_connect_cancelled_error_during_transport_creation_raises_tool_exception(): @@ -4698,10 +4726,11 @@ async def test_connect_cancelled_error_during_transport_creation_raises_tool_exc tool._exit_stack.aclose = AsyncMock() # type: ignore[method-assign] tool.get_mcp_client = Mock(side_effect=asyncio.CancelledError("cancel scope")) # type: ignore[method-assign] - with pytest.raises(ToolException, match="Failed to connect to MCP server"): + with pytest.raises(ToolException, match="Failed to create MCP session"): await tool.connect() tool._exit_stack.aclose.assert_called_once() + await tool.close() async def test_connect_cancelled_error_during_transport_creation_stdio_raises_tool_exception(): @@ -4710,32 +4739,30 @@ async def test_connect_cancelled_error_during_transport_creation_stdio_raises_to tool._exit_stack.aclose = AsyncMock() # type: ignore[method-assign] tool.get_mcp_client = Mock(side_effect=asyncio.CancelledError("cancel scope")) # type: ignore[method-assign] - with pytest.raises(ToolException, match="Failed to start MCP server 'my-server'"): + with pytest.raises(ToolException, match="Failed to create MCP session for server 'my-server'"): await tool.connect() tool._exit_stack.aclose.assert_called_once() + await tool.close() async def test_connect_cancelled_error_during_session_creation_raises_tool_exception(): - """Test that CancelledError from session creation is wrapped in ToolException.""" + """Test that an SDK client-entry CancelledError is wrapped in ToolException.""" tool = MCPStreamableHTTPTool(name="test", url="http://example.com") + tool.get_mcp_client = Mock(return_value=Mock()) # type: ignore[method-assign] + sdk_client = _mock_sdk_client() + sdk_client.__aenter__.side_effect = asyncio.CancelledError("cancel scope") - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] - - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session_class.return_value.__aenter__ = AsyncMock(side_effect=asyncio.CancelledError("cancel scope")) - mock_session_class.return_value.__aexit__ = AsyncMock(return_value=None) - - with pytest.raises(ToolException, match="Failed to create MCP session"): - await tool.connect() + with ( + patch("mcp.Client", return_value=sdk_client), + pytest.raises(ToolException, match="Failed to create MCP session"), + ): + await tool.connect() + await tool.close() async def test_connect_cancelled_error_during_initialize_raises_tool_exception(): - """Test that CancelledError from session.initialize() is wrapped in ToolException. + """Test that CancelledError from SDK negotiation is wrapped in ToolException. This is the primary regression test for the bug: when an MCP server is unreachable, the MCP library raises asyncio.CancelledError internally, which previously escaped @@ -4743,42 +4770,31 @@ async def test_connect_cancelled_error_during_initialize_raises_tool_exception() """ tool = MCPStreamableHTTPTool(name="test", url="http://example.com") - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] - - mock_session = Mock() - mock_session.initialize = AsyncMock(side_effect=asyncio.CancelledError("Cancelled via cancel scope")) + tool.get_mcp_client = Mock(return_value=Mock()) # type: ignore[method-assign] + sdk_client = _mock_sdk_client() + sdk_client.__aenter__.side_effect = asyncio.CancelledError("Cancelled via cancel scope") - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session_class.return_value.__aenter__ = AsyncMock(return_value=mock_session) - mock_session_class.return_value.__aexit__ = AsyncMock(return_value=None) - - with pytest.raises(ToolException, match="MCP server failed to initialize"): - await tool.connect() + with ( + patch("mcp.Client", return_value=sdk_client), + pytest.raises(ToolException, match="Failed to create MCP session"), + ): + await tool.connect() + await tool.close() async def test_connect_cancelled_error_during_initialize_stdio_raises_tool_exception(): - """Test that CancelledError from session.initialize() uses the command-specific message for MCPStdioTool.""" + """SDK negotiation failures retain the full stdio command in the diagnostic.""" tool = MCPStdioTool(name="test", command="my-server", args=["--port", "8080"]) - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] - - mock_session = Mock() - mock_session.initialize = AsyncMock(side_effect=asyncio.CancelledError("Cancelled via cancel scope")) + sdk_client = _mock_sdk_client() + sdk_client.__aenter__.side_effect = asyncio.CancelledError("Cancelled via cancel scope") - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session_class.return_value.__aenter__ = AsyncMock(return_value=mock_session) - mock_session_class.return_value.__aexit__ = AsyncMock(return_value=None) - - with pytest.raises(ToolException, match="MCP server 'my-server --port 8080' failed to initialize"): - await tool.connect() + with ( + patch("mcp.Client", return_value=sdk_client), + pytest.raises(ToolException, match="Failed to create MCP session for server 'my-server --port 8080'"), + ): + await tool.connect() + await tool.close() @pytest.mark.skipif(sys.version_info < (3, 11), reason="task.cancelling() requires Python >= 3.11") @@ -4796,65 +4812,55 @@ async def test_connect_genuine_cancellation_during_transport_creation_propagates await tool.connect() tool._exit_stack.aclose.assert_called_once() + await tool.close() @pytest.mark.skipif(sys.version_info < (3, 11), reason="task.cancelling() requires Python >= 3.11") async def test_connect_genuine_cancellation_during_initialize_propagates(): - """Test that genuine task cancellation during initialize() propagates as CancelledError.""" + """Test that genuine task cancellation during SDK negotiation propagates.""" tool = MCPStreamableHTTPTool(name="test", url="http://example.com") tool._exit_stack.aclose = AsyncMock() # type: ignore[method-assign] - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] - - mock_session = Mock() - mock_session.initialize = AsyncMock(side_effect=asyncio.CancelledError("task cancelled")) + tool.get_mcp_client = Mock(return_value=Mock()) # type: ignore[method-assign] + sdk_client = _mock_sdk_client() + sdk_client.__aenter__.side_effect = asyncio.CancelledError("task cancelled") mock_cancelled_task = Mock() mock_cancelled_task.cancelling.return_value = 1 with ( patch("asyncio.current_task", return_value=mock_cancelled_task), - patch("mcp.client.session.ClientSession") as mock_session_class, + patch("mcp.Client", return_value=sdk_client), + pytest.raises(asyncio.CancelledError), ): - mock_session_class.return_value.__aenter__ = AsyncMock(return_value=mock_session) - mock_session_class.return_value.__aexit__ = AsyncMock(return_value=None) - - with pytest.raises(asyncio.CancelledError): - await tool.connect() + await tool.connect() tool._exit_stack.aclose.assert_called_once() + await tool.close() @pytest.mark.skipif(sys.version_info < (3, 11), reason="task.cancelling() requires Python >= 3.11") async def test_connect_genuine_cancellation_during_session_creation_propagates(): - """Test that genuine task cancellation during session creation propagates as CancelledError.""" + """Test that genuine task cancellation during SDK client entry propagates.""" tool = MCPStreamableHTTPTool(name="test", url="http://example.com") tool._exit_stack.aclose = AsyncMock() # type: ignore[method-assign] - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] + tool.get_mcp_client = Mock(return_value=Mock()) # type: ignore[method-assign] + sdk_client = _mock_sdk_client() + sdk_client.__aenter__.side_effect = asyncio.CancelledError("task cancelled") mock_cancelled_task = Mock() mock_cancelled_task.cancelling.return_value = 1 with ( patch("asyncio.current_task", return_value=mock_cancelled_task), - patch("mcp.client.session.ClientSession") as mock_session_class, + patch("mcp.Client", return_value=sdk_client), + pytest.raises(asyncio.CancelledError), ): - mock_session_class.return_value.__aenter__ = AsyncMock(side_effect=asyncio.CancelledError("task cancelled")) - mock_session_class.return_value.__aexit__ = AsyncMock(return_value=None) - - with pytest.raises(asyncio.CancelledError): - await tool.connect() + await tool.connect() tool._exit_stack.aclose.assert_called_once() + await tool.close() async def test_aenter_cancelled_error_during_connect_is_catchable_as_exception(): @@ -4865,19 +4871,11 @@ async def test_aenter_cancelled_error_during_connect_is_catchable_as_exception() """ tool = MCPStreamableHTTPTool(name="test", url="http://example.com") - mock_session = Mock() - mock_session.initialize = AsyncMock(side_effect=asyncio.CancelledError("Cancelled via cancel scope")) - - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] - - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session_class.return_value.__aenter__ = AsyncMock(return_value=mock_session) - mock_session_class.return_value.__aexit__ = AsyncMock(return_value=None) + tool.get_mcp_client = Mock(return_value=Mock()) # type: ignore[method-assign] + sdk_client = _mock_sdk_client() + sdk_client.__aenter__.side_effect = asyncio.CancelledError("Cancelled via cancel scope") + with patch("mcp.Client", return_value=sdk_client): caught = None try: async with tool: @@ -4887,6 +4885,7 @@ async def test_aenter_cancelled_error_during_connect_is_catchable_as_exception() assert caught is not None, "Expected an exception to be caught by except Exception" assert isinstance(caught, ToolException) + await tool.close() # Tests for _should_propagate_cancelled_error helper @@ -4918,26 +4917,19 @@ def test_should_propagate_cancelled_error_returns_false_when_task_not_cancelling async def test_connect_cancelled_error_during_session_creation_includes_exception_in_message(): - """Test that CancelledError from session creation includes exception details in ToolException message.""" + """Test that an SDK client-entry CancelledError retains its exception details.""" tool = MCPStreamableHTTPTool(name="test", url="http://example.com") + tool.get_mcp_client = Mock(return_value=Mock()) # type: ignore[method-assign] + sdk_client = _mock_sdk_client() + sdk_client.__aenter__.side_effect = asyncio.CancelledError("cancel scope detail") - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] - - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session_class.return_value.__aenter__ = AsyncMock( - side_effect=asyncio.CancelledError("cancel scope detail") - ) - mock_session_class.return_value.__aexit__ = AsyncMock(return_value=None) - + with patch("mcp.Client", return_value=sdk_client): with pytest.raises(ToolException) as exc_info: await tool.connect() assert "Failed to create MCP session" in str(exc_info.value) assert "cancel scope detail" in str(exc_info.value) + await tool.close() # Tests for _describe_error helper (cancel-scope / exception-group unmasking) @@ -4975,20 +4967,15 @@ async def test_connect_cancelled_error_unmasks_inner_auth_failure(): """A 401 swallowed by the MCP client's cancel scope must be named in the ToolException.""" tool = MCPStreamableHTTPTool(name="test", url="http://example.com") - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] + tool.get_mcp_client = Mock(return_value=Mock()) # type: ignore[method-assign] real = RuntimeError("401 Client Error: Unauthorized") masked = asyncio.CancelledError("Cancelled via cancel scope") masked.__context__ = real - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session_class.return_value.__aenter__ = AsyncMock(side_effect=masked) - mock_session_class.return_value.__aexit__ = AsyncMock(return_value=None) - + sdk_client = _mock_sdk_client() + sdk_client.__aenter__.side_effect = masked + with patch("mcp.Client", return_value=sdk_client): with pytest.raises(ToolException) as exc_info: await tool.connect() @@ -4996,13 +4983,12 @@ async def test_connect_cancelled_error_unmasks_inner_auth_failure(): assert "Failed to create MCP session" in message assert "401 Client Error: Unauthorized" in message assert "Cancelled via cancel scope" not in message + await tool.close() @pytest.mark.skipif(sys.version_info < (3, 11), reason="ExceptionGroup is Python >= 3.11") async def test_connect_bare_cancel_names_cleanup_error_from_exit_stack(): - """The reported 401 path: initialize() raises a bare CancelledError and the - real HTTP failure only surfaces from the exit-stack close. The ToolException - must name that close-time error, not the cancellation.""" + """SDK entry raises a bare cancellation and cleanup reveals the actual HTTP failure.""" import builtins exception_group_type = getattr(builtins, "ExceptionGroup", None) @@ -5011,46 +4997,37 @@ async def test_connect_bare_cancel_names_cleanup_error_from_exit_stack(): tool = MCPStreamableHTTPTool(name="test", url="http://example.com") - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] + tool.get_mcp_client = Mock(return_value=Mock()) # type: ignore[method-assign] cleanup_group = exception_group_type( "unhandled errors in a TaskGroup", [RuntimeError("401 Client Error: Unauthorized")] ) - mock_session = Mock() - mock_session.initialize = AsyncMock(side_effect=asyncio.CancelledError("Cancelled via cancel scope")) - - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session_class.return_value.__aenter__ = AsyncMock(return_value=mock_session) - mock_session_class.return_value.__aexit__ = AsyncMock(side_effect=cleanup_group) + sdk_client = _mock_sdk_client() + sdk_client.__aenter__.side_effect = asyncio.CancelledError("Cancelled via cancel scope") + tool._exit_stack.aclose = AsyncMock(side_effect=cleanup_group) # type: ignore[method-assign] + with patch("mcp.Client", return_value=sdk_client): with pytest.raises(ToolException) as exc_info: await tool.connect() message = str(exc_info.value) - assert "MCP server failed to initialize" in message + assert "Failed to create MCP session" in message assert "401 Client Error: Unauthorized" in message assert "Cancelled via cancel scope" not in message + tool._exit_stack.aclose.assert_awaited_once() + tool._exit_stack.aclose.side_effect = None + await tool.close() async def test_connect_cancelled_error_during_session_creation_logs_with_exc_info(): - """Test that CancelledError from session creation is logged with exc_info=True.""" + """Test that an SDK client-entry cancellation is logged with exc_info=True.""" tool = MCPStreamableHTTPTool(name="test", url="http://example.com") + tool.get_mcp_client = Mock(return_value=Mock()) # type: ignore[method-assign] + sdk_client = _mock_sdk_client() + sdk_client.__aenter__.side_effect = asyncio.CancelledError("cancel scope") - mock_transport = (Mock(), Mock()) - mock_context_manager = Mock() - mock_context_manager.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context_manager.__aexit__ = AsyncMock(return_value=None) - tool.get_mcp_client = Mock(return_value=mock_context_manager) # type: ignore[method-assign] - - with patch("mcp.client.session.ClientSession") as mock_session_class: - mock_session_class.return_value.__aenter__ = AsyncMock(side_effect=asyncio.CancelledError("cancel scope")) - mock_session_class.return_value.__aexit__ = AsyncMock(return_value=None) - + with patch("mcp.Client", return_value=sdk_client): from agent_framework._mcp import logger as mcp_logger with patch.object(mcp_logger, "debug") as mock_debug: @@ -5063,6 +5040,7 @@ async def test_connect_cancelled_error_during_session_creation_logs_with_exc_inf assert cancel_calls, "Expected a debug log for the cancelled session creation" _, kwargs = cancel_calls[0] assert kwargs.get("exc_info") is True + await tool.close() def test_mcp_stdio_tool_get_mcp_client_with_env_and_kwargs(): @@ -5265,8 +5243,8 @@ async def test_mcp_streamable_http_tool_httpx_client_cleanup(): # Mock the streamable_http_client to avoid actual connections with ( - patch("mcp.client.streamable_http.streamable_http_client") as mock_client, - patch("mcp.client.session.ClientSession") as mock_session_class, + patch("agent_framework._mcp.streamable_http_client") as mock_client, + patch("mcp.Client", side_effect=lambda **_: _mock_sdk_client()), ): # Setup mock context manager for streamable_http_client mock_transport = (Mock(), Mock()) @@ -5275,12 +5253,6 @@ async def test_mcp_streamable_http_tool_httpx_client_cleanup(): mock_context_manager.__aexit__ = AsyncMock(return_value=None) mock_client.return_value = mock_context_manager - # Setup mock session - mock_session = Mock() - mock_session.initialize = AsyncMock() - mock_session_class.return_value.__aenter__ = AsyncMock(return_value=mock_session) - mock_session_class.return_value.__aexit__ = AsyncMock(return_value=None) - tool1 = MCPStreamableHTTPTool( name="test", url="http://localhost:8081/mcp", @@ -6169,17 +6141,11 @@ async def __aexit__(self, exc_type, exc, tb): ) transport_context = TaskBoundTransportContext() - mock_session = Mock() - mock_session._request_id = 1 - mock_session.initialize = AsyncMock() - - mock_session_context = AsyncMock() - mock_session_context.__aenter__ = AsyncMock(return_value=mock_session) - mock_session_context.__aexit__ = AsyncMock(return_value=None) + sdk_client = _mock_sdk_client() with ( patch.object(tool, "get_mcp_client", return_value=transport_context), - patch("mcp.client.session.ClientSession", return_value=mock_session_context), + patch("mcp.Client", return_value=_TransportBoundClientContext(transport_context, sdk_client)), ): await asyncio.create_task(tool.connect()) @@ -6189,6 +6155,7 @@ async def __aexit__(self, exc_type, exc, tb): assert transport_context.closed_cleanly is True assert transport_context.exit_task is transport_context.enter_task + sdk_client.__aexit__.assert_awaited_once() assert not any("cancel scope" in record.getMessage().lower() for record in caplog.records) @@ -6222,37 +6189,39 @@ async def __aexit__(self, exc_type, exc, tb): transport_contexts = [TaskBoundTransportContext(), TaskBoundTransportContext()] sessions = [] - session_contexts = [] - for _ in range(2): - session = Mock() - session._request_id = 1 - session.initialize = AsyncMock() - session.set_logging_level = AsyncMock() + clients = [] + client_contexts = [] + for transport_context in transport_contexts: + session = Mock(spec=ClientSession) sessions.append(session) - - session_context = AsyncMock() - session_context.__aenter__ = AsyncMock(return_value=session) - session_context.__aexit__ = AsyncMock(return_value=None) - session_contexts.append(session_context) + client = _mock_sdk_client(session=session) + clients.append(client) + client_contexts.append(_TransportBoundClientContext(transport_context, client)) with ( patch.object(tool, "get_mcp_client", side_effect=transport_contexts), - patch("mcp.client.session.ClientSession", side_effect=session_contexts), + patch("mcp.Client", side_effect=client_contexts), ): - await tool.connect() - - caplog.clear() - with caplog.at_level(logging.WARNING, logger=logger.name): - await asyncio.create_task(tool.connect(reset=True)) + try: + await tool.connect() - assert transport_contexts[0].closed_cleanly is True - assert transport_contexts[0].exit_task is transport_contexts[0].enter_task - assert transport_contexts[1].enter_task is transport_contexts[0].enter_task - assert tool.session is sessions[1] - assert tool.is_connected is True - assert not any("cancel scope" in record.getMessage().lower() for record in caplog.records) + caplog.clear() + with caplog.at_level(logging.WARNING, logger=logger.name): + await asyncio.create_task(tool.connect(reset=True)) + + assert transport_contexts[0].closed_cleanly is True + assert transport_contexts[0].exit_task is transport_contexts[0].enter_task + assert transport_contexts[1].enter_task is transport_contexts[0].enter_task + assert tool.session is sessions[1] + assert tool.is_connected is True + clients[0].__aexit__.assert_awaited_once() + clients[1].__aenter__.assert_awaited_once() + assert not any("cancel scope" in record.getMessage().lower() for record in caplog.records) + finally: + await tool.close() - await tool.close() + assert transport_contexts[1].closed_cleanly is True + assert transport_contexts[1].exit_task is transport_contexts[1].enter_task async def test_mcp_tool_connect_from_lifecycle_owner_bypasses_request_lock() -> None: @@ -6429,30 +6398,18 @@ async def test_connect_sets_logging_level_when_logger_level_is_set(): load_prompts=False, ) - # Mock the transport and session - mock_transport = (Mock(), Mock()) - mock_context = AsyncMock() - mock_context.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context.__aexit__ = AsyncMock() - - mock_session = Mock() - mock_session._request_id = 1 - mock_session.initialize = AsyncMock() + mock_session = Mock(spec=ClientSession) mock_session.set_logging_level = AsyncMock() - - mock_session_context = AsyncMock() - mock_session_context.__aenter__ = AsyncMock(return_value=mock_session) - mock_session_context.__aexit__ = AsyncMock() + sdk_client = _mock_sdk_client( + session=mock_session, capabilities=types.ServerCapabilities(logging=types.LoggingCapability()) + ) with ( - patch.object(tool, "get_mcp_client", return_value=mock_context), - patch("mcp.client.session.ClientSession", return_value=mock_session_context), + patch("mcp.Client", return_value=sdk_client), patch.object(logger, "level", logging.DEBUG), # Set logger level to DEBUG ): - await tool.connect() - - # Verify set_logging_level was called with "debug" - mock_session.set_logging_level.assert_called_once_with("debug") + async with tool: + mock_session.set_logging_level.assert_awaited_once_with("debug") async def test_connect_does_not_set_logging_level_when_logger_level_is_notset(): @@ -6466,30 +6423,18 @@ async def test_connect_does_not_set_logging_level_when_logger_level_is_notset(): load_prompts=False, ) - # Mock the transport and session - mock_transport = (Mock(), Mock()) - mock_context = AsyncMock() - mock_context.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context.__aexit__ = AsyncMock() - - mock_session = Mock() - mock_session._request_id = 1 - mock_session.initialize = AsyncMock() + mock_session = Mock(spec=ClientSession) mock_session.set_logging_level = AsyncMock() - - mock_session_context = AsyncMock() - mock_session_context.__aenter__ = AsyncMock(return_value=mock_session) - mock_session_context.__aexit__ = AsyncMock() + sdk_client = _mock_sdk_client( + session=mock_session, capabilities=types.ServerCapabilities(logging=types.LoggingCapability()) + ) with ( - patch.object(tool, "get_mcp_client", return_value=mock_context), - patch("mcp.client.session.ClientSession", return_value=mock_session_context), + patch("mcp.Client", return_value=sdk_client), patch.object(logger, "level", logging.NOTSET), # Set logger level to NOTSET ): - await tool.connect() - - # Verify set_logging_level was NOT called - mock_session.set_logging_level.assert_not_called() + async with tool: + mock_session.set_logging_level.assert_not_called() async def test_connect_handles_set_logging_level_exception(): @@ -6503,38 +6448,24 @@ async def test_connect_handles_set_logging_level_exception(): load_prompts=False, ) - # Mock the transport and session - mock_transport = (Mock(), Mock()) - mock_context = AsyncMock() - mock_context.__aenter__ = AsyncMock(return_value=mock_transport) - mock_context.__aexit__ = AsyncMock() - - mock_session = Mock() - mock_session._request_id = 1 - mock_session.initialize = AsyncMock() + mock_session = Mock(spec=ClientSession) # Make set_logging_level raise an exception mock_session.set_logging_level = AsyncMock(side_effect=RuntimeError("Server doesn't support logging level")) - mock_session_context = AsyncMock() - mock_session_context.__aenter__ = AsyncMock(return_value=mock_session) - mock_session_context.__aexit__ = AsyncMock() + sdk_client = _mock_sdk_client( + session=mock_session, capabilities=types.ServerCapabilities(logging=types.LoggingCapability()) + ) with ( - patch.object(tool, "get_mcp_client", return_value=mock_context), - patch("mcp.client.session.ClientSession", return_value=mock_session_context), + patch("mcp.Client", return_value=sdk_client), patch.object(logger, "level", logging.INFO), # Set logger level to INFO patch.object(logger, "warning") as mock_warning, ): - # Should NOT raise - the exception should be caught and logged - await tool.connect() - - # Verify set_logging_level was called - mock_session.set_logging_level.assert_called_once_with("info") - - # Verify warning was logged - mock_warning.assert_called_once() - call_args = mock_warning.call_args - assert "Failed to set log level" in call_args[0][0] + async with tool: + mock_session.set_logging_level.assert_awaited_once_with("info") + mock_warning.assert_called_once() + call_args = mock_warning.call_args + assert "Failed to set log level" in call_args[0][0] async def test_connect_reinitializes_existing_session_and_loads_tools_and_prompts() -> None: From 1266aa1ff856ad795942eb3b7197475b6e4f9bb5 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Mon, 5 Oct 2026 14:25:56 +0200 Subject: [PATCH 19/42] Assert preserve headers --- python/packages/core/tests/core/test_mcp.py | 17 +++++++++++++++-- 1 file changed, 15 insertions(+), 2 deletions(-) diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 9e9864dd100..2e7c3e7fe9e 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -8017,7 +8017,7 @@ async def test_mcp_streamble_http_tool_connects_to_v2_server() -> None: from httpx2 import AsyncClient, MockTransport, Request, Response - captured_methods: list[str] = [] + captured_requests: list[tuple[dict[str, Any], dict[str, str]]] = [] async def mcp_v2_server_mock_handler(request: Request) -> Response: if request.method == "DELETE": @@ -8025,7 +8025,8 @@ async def mcp_v2_server_mock_handler(request: Request) -> Response: body = json.loads(request.content) method = body["method"] - captured_methods.append(method) + headers = {name.lower(): value for name, value in request.headers.items()} + captured_requests.append((body, headers)) if method == "initialize": return Response( @@ -8099,9 +8100,21 @@ async def mcp_v2_server_mock_handler(request: Request) -> Response: assert isinstance(result, list) assert [item.text for item in result if item.type == "text"] == ["Hello!"] + captured_methods = [body["method"] for body, _ in captured_requests] assert "server/discover" in captured_methods assert "initialize" not in captured_methods assert "tools/list" in captured_methods + assert "tools/call" in captured_methods + + tool_calls = [(body, headers) for body, headers in captured_requests if body["method"] == "tools/call"] + assert len(tool_calls) == 1 + body, headers = tool_calls[0] + params = body["params"] + meta = params["_meta"] + assert headers["mcp-protocol-version"] == meta["io.modelcontextprotocol/protocolVersion"] == "2026-07-28" + assert headers["mcp-method"] == body["method"] == "tools/call" + assert headers["mcp-name"] == params["name"] == "greet" + assert isinstance(meta["io.modelcontextprotocol/clientCapabilities"], dict) async def test_mcp_streamable_http_tool_connects_to_legacy_server() -> None: From c5355aa2b9a11d06ec8c116579b7da185eaea9ce Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Mon, 5 Oct 2026 17:04:11 +0200 Subject: [PATCH 20/42] Added session negotiation tests --- python/packages/core/agent_framework/_mcp.py | 11 +- python/packages/core/tests/core/test_mcp.py | 192 ++++++++++++++----- 2 files changed, 154 insertions(+), 49 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 0906b49ab47..2556cdd1e9d 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -2035,14 +2035,15 @@ async def _connect_on_owner( self._owns_session = True else: try: - if self.session._request_id == 0: # type: ignore[attr-defined] - # If the session is not initialized, we need to reinitialize it + if self.session.protocol_version is None: + # Preserve old compatibility behavior for an + # unnegotiated caller-supplied session: initialize it as legacy. with create_mcp_client_span("initialize", attributes=self._mcp_base_span_attributes()) as init_span: initialize_result = await self.session.initialize() init_span.set_attribute(OtelAttr.MCP_PROTOCOL_VERSION, initialize_result.protocol_version) - self._set_server_capabilities(getattr(initialize_result, "capabilities", None)) - elif self._server_capabilities is None: - self._set_server_capabilities(getattr(self.session, "_server_capabilities", None)) + + self._set_server_capabilities(self.session.server_capabilities) + self._ping_available = self.session.initialize_result is not None except (Exception, asyncio.CancelledError): await self._close_on_owner() raise diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 2e7c3e7fe9e..d633195e800 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -105,6 +105,32 @@ def _mock_sdk_client( return client +def _mock_unnegotiated_session(capabilities: types.ServerCapabilities | None) -> Mock: + """Model a caller-owned session that becomes legacy-negotiated on initialize.""" + initialize_result = ( + types.InitializeResult( + protocol_version="2025-11-25", + capabilities=capabilities, + server_info=types.Implementation(name="mock-server", version="1.0"), + ) + if capabilities is not None + else Mock(protocol_version="2025-11-25", capabilities=None) + ) + session = Mock(spec=ClientSession) + session.protocol_version = None + session.server_capabilities = None + session.initialize_result = None + + async def initialize() -> Any: + session.protocol_version = "2025-11-25" + session.server_capabilities = capabilities + session.initialize_result = initialize_result + return initialize_result + + session.initialize = AsyncMock(side_effect=initialize) + return session + + class _TransportBoundClientContext: """Model SDK-owned transport entry and exit without running a protocol dispatcher.""" @@ -5817,7 +5843,6 @@ async def test_mcp_tool_connection_properly_invalidated_after_closed_resource_er # Mock the session mock_session = MagicMock() - mock_session._request_id = 1 mock_session.call_tool = AsyncMock() # Mock _exit_stack.aclose to track cleanup calls @@ -5916,7 +5941,6 @@ async def test_mcp_tool_get_prompt_reconnection_on_closed_resource_error(): # Mock the session mock_session = MagicMock() - mock_session._request_id = 1 mock_session.get_prompt = AsyncMock() # Mock _exit_stack.aclose to track cleanup calls @@ -6468,12 +6492,95 @@ async def test_connect_handles_set_logging_level_exception(): assert "Failed to set log level" in call_args[0][0] +@pytest.mark.parametrize( + ("mode", "expected_version"), + [ + ("auto", "2026-07-28"), + ("legacy", "2025-11-25"), + ], +) +async def test_mcp_tool_reuses_supplied_session(mode: str, expected_version: str) -> None: + from mcp import Client + from mcp.server import Server, ServerRequestContext + + async def list_tools( + _ctx: ServerRequestContext[Any], + _params: types.PaginatedRequestParams | None, + ) -> types.ListToolsResult: + return types.ListToolsResult( + result_type="complete", + tools=[ + types.Tool( + name="greet", + input_schema={"type": "object", "properties": {}}, + ) + ], + ) + + async def call_tool( + _ctx: ServerRequestContext[Any], + params: types.CallToolRequestParams, + ) -> types.CallToolResult: + assert params.name == "greet" + return types.CallToolResult( + result_type="complete", + content=[types.TextContent(type="text", text="Hello!")], + is_error=False, + ) + + server = Server( + "test-server", + on_list_tools=list_tools, + on_call_tool=call_tool, + ) + + async with Client(server, mode=mode) as client: + session = client.session + assert session.protocol_version == expected_version + + initialize = AsyncMock(wraps=session.initialize) + discover = AsyncMock(wraps=session.discover) + send_ping = AsyncMock(wraps=session.send_ping) + + with ( + patch.object(session, "initialize", initialize), + patch.object(session, "discover", discover), + patch.object(session, "send_ping", send_ping), + ): + wrapper = MCPStdioTool( + name="test", + command="unused", + session=session, + load_prompts=False, + ) + + async with wrapper: + assert wrapper.session is session + assert [function.name for function in wrapper.functions] == ["greet"] + assert _mcp_result_to_text(await wrapper.call_tool("greet")) == "Hello!" + + initialize.assert_not_awaited() + discover.assert_not_awaited() + if mode == "auto": + send_ping.assert_not_awaited() + + # The wrapper has closed, but the caller-owned session must still work. + result = await session.call_tool("greet") + assert isinstance(result.content[0], types.TextContent) + assert result.content[0].text == "Hello!" + + async def test_connect_reinitializes_existing_session_and_loads_tools_and_prompts() -> None: - tool = MCPTool(name="test_tool", load_tools=True, load_prompts=True) # type: ignore[abstract] # ty: ignore[call-non-callable] + session = _mock_unnegotiated_session( + types.ServerCapabilities(tools=types.ToolsCapability(), prompts=types.PromptsCapability()) + ) + tool = MCPTool( # type: ignore[abstract] # ty: ignore[call-non-callable] + name="test_tool", + load_tools=True, + load_prompts=True, + session=session, + ) tool.is_connected = True - tool.session = Mock() - tool.session._request_id = 0 - tool.session.initialize = AsyncMock() with ( patch.object(tool, "load_tools", AsyncMock()) as mock_load_tools, @@ -6482,7 +6589,7 @@ async def test_connect_reinitializes_existing_session_and_loads_tools_and_prompt ): await tool._connect_on_owner() - tool.session.initialize.assert_awaited_once() + session.initialize.assert_awaited_once() mock_load_tools.assert_awaited_once() mock_load_prompts.assert_awaited_once() assert tool._tools_loaded is True @@ -6490,28 +6597,25 @@ async def test_connect_reinitializes_existing_session_and_loads_tools_and_prompt async def test_connect_skips_tools_and_prompts_when_server_does_not_advertise_capabilities() -> None: - tool = MCPTool(name="test_tool", load_tools=True, load_prompts=True) # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.is_connected = True - tool.session = Mock() - tool.session._request_id = 0 - tool.session.initialize = AsyncMock( - return_value=types.InitializeResult( - protocol_version=types.LATEST_PROTOCOL_VERSION, - capabilities=types.ServerCapabilities(), - server_info=types.Implementation(name="test", version="1.0"), - ) + session = _mock_unnegotiated_session(types.ServerCapabilities()) + session.list_tools = AsyncMock() + session.list_prompts = AsyncMock() + session.set_logging_level = AsyncMock() + tool = MCPTool( # type: ignore[abstract] # ty: ignore[call-non-callable] + name="test_tool", + load_tools=True, + load_prompts=True, + session=session, ) - tool.session.list_tools = AsyncMock() - tool.session.list_prompts = AsyncMock() - tool.session.set_logging_level = AsyncMock() + tool.is_connected = True with patch.object(logger, "level", logging.INFO): await tool._connect_on_owner() - tool.session.initialize.assert_awaited_once() - tool.session.list_tools.assert_not_called() - tool.session.list_prompts.assert_not_called() - tool.session.set_logging_level.assert_not_called() + session.initialize.assert_awaited_once() + session.list_tools.assert_not_called() + session.list_prompts.assert_not_called() + session.set_logging_level.assert_not_called() assert tool.is_connected is True assert tool._supports_tools is False assert tool._supports_prompts is False @@ -6521,42 +6625,42 @@ async def test_connect_skips_tools_and_prompts_when_server_does_not_advertise_ca async def test_connect_treats_missing_capabilities_as_unsupported() -> None: - tool = MCPTool(name="test_tool", load_tools=True, load_prompts=True) # type: ignore[abstract] # ty: ignore[call-non-callable] + session = _mock_unnegotiated_session(None) + session.list_tools = AsyncMock() + session.list_prompts = AsyncMock() + tool = MCPTool( # type: ignore[abstract] # ty: ignore[call-non-callable] + name="test_tool", + load_tools=True, + load_prompts=True, + session=session, + ) tool.is_connected = True - tool.session = Mock() - tool.session._request_id = 0 - tool.session.initialize = AsyncMock(return_value=Mock(capabilities=None)) - tool.session.list_tools = AsyncMock() - tool.session.list_prompts = AsyncMock() with patch.object(logger, "level", logging.NOTSET): await tool._connect_on_owner() - tool.session.list_tools.assert_not_called() - tool.session.list_prompts.assert_not_called() + session.list_tools.assert_not_called() + session.list_prompts.assert_not_called() assert tool._supports_tools is False assert tool._supports_prompts is False assert tool._supports_logging is False async def test_connect_sets_logging_level_when_server_advertises_logging() -> None: - tool = MCPTool(name="test_tool", load_tools=False, load_prompts=False) # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.is_connected = True - tool.session = Mock() - tool.session._request_id = 0 - tool.session.initialize = AsyncMock( - return_value=types.InitializeResult( - protocol_version=types.LATEST_PROTOCOL_VERSION, - capabilities=types.ServerCapabilities(logging=types.LoggingCapability()), - server_info=types.Implementation(name="test", version="1.0"), - ) + session = _mock_unnegotiated_session(types.ServerCapabilities(logging=types.LoggingCapability())) + session.set_logging_level = AsyncMock() + tool = MCPTool( # type: ignore[abstract] # ty: ignore[call-non-callable] + name="test_tool", + load_tools=False, + load_prompts=False, + session=session, ) - tool.session.set_logging_level = AsyncMock() + tool.is_connected = True with patch.object(logger, "level", logging.INFO): await tool._connect_on_owner() - tool.session.set_logging_level.assert_awaited_once_with("info") + session.set_logging_level.assert_awaited_once_with("info") assert tool._supports_logging is True From 2cc24d30a254524d9809c991b275a70198a785dc Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Mon, 5 Oct 2026 17:27:17 +0200 Subject: [PATCH 21/42] Sampling capabilities --- python/packages/core/agent_framework/_mcp.py | 3 ++- python/packages/core/tests/core/test_mcp.py | 18 ++++++------------ 2 files changed, 8 insertions(+), 13 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 2556cdd1e9d..8a1329008fe 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -1978,10 +1978,11 @@ async def _connect_on_owner( sampling_capabilities = types.SamplingCapability( tools=types.SamplingToolsCapability(), ) + client_mode = "legacy" if self.client is not None else "auto" mcp_client = await self._exit_stack.enter_async_context( - # default "mode" is `auto` which automatically negotiates with the right protocol version Client( server=self.get_mcp_client(), + mode=client_mode, read_timeout_seconds=( timedelta(seconds=self.request_timeout).seconds if self.request_timeout else None ), diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index d633195e800..97dbf1fee6d 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -4618,20 +4618,17 @@ async def test_mcp_tool_sampling_callback_always_passes_max_tokens(): async def test_connect_sampling_capabilities_with_client(): - """Test connect() passes sampling_capabilities to mcp.Client""" + """Test connect() uses legacy mode and advertises sampling when a chat client is configured.""" tool = MCPStdioTool(name="test", command="test-command", load_tools=False, load_prompts=False) tool.client = Mock() with patch("mcp.Client") as mock_client_class: - sdk_client = AsyncMock() - sdk_client.__aenter__.return_value = sdk_client - sdk_client.session = Mock(initialize_result=None) - sdk_client.protocol_version = "2026-07-28" - sdk_client.server_capabilities = types.ServerCapabilities() + sdk_client = _mock_sdk_client() mock_client_class.return_value = sdk_client async with tool: call_kwargs = mock_client_class.call_args.kwargs + assert call_kwargs["mode"] == "legacy" sampling_caps = call_kwargs.get("sampling_capabilities") assert sampling_caps is not None assert isinstance(sampling_caps, types.SamplingCapability) @@ -4640,20 +4637,17 @@ async def test_connect_sampling_capabilities_with_client(): async def test_connect_no_sampling_capabilities_without_client(): - """Test connect() does not pass sampling_capabilities when no client is set.""" + """Test connect() keeps auto mode and omits sampling capabilities without a chat client.""" tool = MCPStdioTool(name="test", command="test-command", load_tools=False, load_prompts=False) with patch("mcp.Client") as mock_client_class: - sdk_client = AsyncMock() - sdk_client.__aenter__.return_value = sdk_client - sdk_client.session = Mock(initialize_result=None) - sdk_client.protocol_version = "2026-07-28" - sdk_client.server_capabilities = types.ServerCapabilities() + sdk_client = _mock_sdk_client(protocol_version="2026-07-28") mock_client_class.return_value = sdk_client try: await tool.connect() call_kwargs = mock_client_class.call_args.kwargs + assert call_kwargs["mode"] == "auto" assert call_kwargs.get("sampling_capabilities") is None finally: await tool.close() From f043c00dcd1383d040c7a2245ce3180dad50ad07 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Mon, 5 Oct 2026 17:28:19 +0200 Subject: [PATCH 22/42] ADR update --- docs/decisions/0045-python-mcp-v2-client-lifecycle.md | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/docs/decisions/0045-python-mcp-v2-client-lifecycle.md b/docs/decisions/0045-python-mcp-v2-client-lifecycle.md index 90d1440f937..4ebf5a4976b 100644 --- a/docs/decisions/0045-python-mcp-v2-client-lifecycle.md +++ b/docs/decisions/0045-python-mcp-v2-client-lifecycle.md @@ -54,8 +54,9 @@ For framework-created connections: - Negotiated protocol version and server capabilities are read from the client's era-neutral properties rather than captured only from an initialize result. - Tools configured for the existing server-initiated sampling callback use `mode="legacy"`. The MCP migration guide - requires legacy mode for that back-channel behavior because 2026-07-28 refuses server-initiated sampling on every - transport. Auto mode remains the default when no legacy-only callback behavior is requested. + advises legacy mode for workflows that rely on that back-channel behavior because 2026-07-28 refuses + server-initiated sampling on every transport. Auto mode remains the default when no legacy-only callback behavior + is requested. Agent Framework retains its lifecycle owner task and locks. They protect framework state, preserve AnyIO task ownership during teardown, serialize identity-changing reconnects, and roll back partially loaded discovery state. From e3775e0ffae6fb9e4bd5d1ae195356b037eefd82 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Mon, 5 Oct 2026 17:35:55 +0200 Subject: [PATCH 23/42] MCPStdio test --- python/packages/core/tests/core/test_mcp.py | 79 +++++++++++++++++++++ 1 file changed, 79 insertions(+) diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 97dbf1fee6d..7d2297cef5b 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -11,6 +11,7 @@ from contextlib import AbstractAsyncContextManager, _AsyncGeneratorContextManager # type: ignore from contextvars import ContextVar from datetime import timedelta +from textwrap import dedent from types import TracebackType from typing import Any, cast from unittest.mock import AsyncMock, Mock, patch @@ -8314,6 +8315,84 @@ async def mcp_legacy_server_mock_handler(request: Request) -> Response: assert "tools/list" in captured_methods +@pytest.mark.parametrize( + ("era", "expected_version"), + [ + ("modern", "2026-07-28"), + ("legacy", "2025-11-25"), + ], +) +async def test_mcp_stdio_tool_connects_to_both_protocol_eras(era: str, expected_version: str) -> None: + server_script = dedent( + """ + import json + import sys + + era = sys.argv[1] + for line in sys.stdin: + request = json.loads(line) + if "id" not in request: + continue + + request_id = request["id"] + method = request["method"] + result = None + error = None + + if method == "server/discover": + if era == "modern": + result = { + "supportedVersions": ["2026-07-28"], + "capabilities": {"tools": {}}, + } + else: + error = {"code": -32601, "message": "Method not found"} + elif method == "initialize": + if era == "legacy": + result = { + "protocolVersion": "2025-11-25", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "legacy-test-server", "version": "1.0"}, + } + else: + error = {"code": -32601, "message": "Method not found"} + elif method == "ping": + result = {} + elif method == "tools/list": + result = { + "tools": [{"name": "greet", "inputSchema": {"type": "object", "properties": {}}}], + } + if era == "modern": + result.update({"resultType": "complete", "cacheScope": "private", "ttlMs": 0}) + elif method == "tools/call": + result = { + "content": [{"type": "text", "text": "Hello!"}], + "isError": False, + } + if era == "modern": + result["resultType"] = "complete" + else: + error = {"code": -32601, "message": f"Unexpected method: {method}"} + + response = {"jsonrpc": "2.0", "id": request_id} + response["error" if error is not None else "result"] = error if error is not None else result + print(json.dumps(response), flush=True) + """ + ) + tool = MCPStdioTool( + name=f"{era}-stdio", + command=sys.executable, + args=["-c", server_script, era], + load_prompts=False, + ) + + async with tool: + assert tool.session is not None + assert tool.session.protocol_version == expected_version + assert [function.name for function in tool.functions] == ["greet"] + assert _mcp_result_to_text(await tool.call_tool("greet")) == "Hello!" + + async def test_agent_context_manager_authenticates_connect_with_closure_provider( client: SupportsChatGetResponse, ) -> None: From bdb7badcb7cd55ac812e9d08ca8644191e68e838 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Tue, 6 Oct 2026 13:07:56 +0200 Subject: [PATCH 24/42] basic tools support completed --- python/packages/core/agent_framework/_mcp.py | 85 +++++- python/packages/core/tests/core/test_mcp.py | 270 +++++++++++++++++++ 2 files changed, 351 insertions(+), 4 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 8a1329008fe..14c5c263c0a 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -67,7 +67,7 @@ # TODO(jpalvarezl): clean up and consolidate under httpx2 import httpx2 from httpx import Request - from mcp import types + from mcp import Client, types from mcp.client.context import ClientRequestContext from mcp.client.session import ClientSession @@ -1070,6 +1070,7 @@ def __init__( self._supports_logging: bool | None = None self._ping_available: bool = True self._pending_reload_tasks: set[asyncio.Task[None]] = set() + self._capability_list_subscription_task: asyncio.Task[None] | None = None def __str__(self) -> str: return f"MCPTool(name={self.name}, description={self.description})" @@ -1605,6 +1606,64 @@ def _list_progressive_mcp_tools(self, ctx: FunctionInvocationContext) -> list[di }) return tools + async def _listen_capability_list_changes(self, mcp_client: Client) -> None: + """Tool and Prompt update handler for MCP Clients.""" + from mcp.client.subscriptions import ListenNotSupportedError, PromptsListChanged, ToolsListChanged + + if self._capability_list_subscription_task is not None: + return + + capabilities = self._server_capabilities + if capabilities is None: + return + + tools_changed = self.load_tools_flag and bool(capabilities.tools and capabilities.tools.list_changed) + prompts_changed = self.load_prompts_flag and bool(capabilities.prompts and capabilities.prompts.list_changed) + + if not tools_changed and not prompts_changed: + return + + listen_context = mcp_client.listen( + tools_list_changed=tools_changed, + prompts_list_changed=prompts_changed, + ) + + # v2 introduces `mcp_client.listen`, but messages are teed via the message_handler passed in + # in the constructor anyway, so we check for the feature availability an fallback to legacy behaviour + try: + subscription = await self._exit_stack.enter_async_context(listen_context) + except ListenNotSupportedError: + logger.debug("Listen not supported, falling back to legacy behaviour.") + return + + async def consume() -> None: + async for event in subscription: + match event: + case ToolsListChanged(): + self._schedule_reload(self.load_tools()) + case PromptsListChanged(): + self._schedule_reload(self.load_prompts()) + case _: + logger.debug("Unhandled event: %s", event) + + self._capability_list_subscription_task = asyncio.create_task( + consume(), + name=f"mcp-capability-list-subscription:{self.name}", + ) + + async def _cancel_capability_list_subscription(self) -> None: + """Cancel the capability list subscription task, if it exists.""" + task = self._capability_list_subscription_task + self._capability_list_subscription_task = None + + if task is None: + return + + if not task.done(): + task.cancel() + + await asyncio.gather(task, return_exceptions=True) + async def _load_progressive_mcp_tool(self, ctx: FunctionInvocationContext, tool: str | Sequence[str]) -> str: """Load an allowed MCP tool into the live function-calling tool list.""" if ctx.tools is None: @@ -1952,6 +2011,7 @@ async def _connect_on_owner( ToolException: If connection or session initialization fails. """ if reset: + await self._cancel_capability_list_subscription() if reset_discovery: await self._cancel_pending_reload_tasks() await self._safe_close_exit_stack() @@ -2034,6 +2094,11 @@ async def _connect_on_owner( raise ToolException(error_msg, inner_exception=ex if isinstance(ex, Exception) else None) from ex self.session = session self._owns_session = True + try: + await self._listen_capability_list_changes(mcp_client) + except (Exception, asyncio.CancelledError): + await self._close_on_owner() + raise else: try: if self.session.protocol_version is None: @@ -2319,24 +2384,36 @@ async def message_handler( self, message: IncomingMessage, ) -> None: - """Handle messages from the MCP server. + """Handle messages from the MCP server ("legacy"). Kept for backward compatibility. By default this function will handle exceptions on the server by logging them, and it will trigger a reload of the tools and prompts when the list changed notification is received. Note: - If you want to extend this behavior, you can subclass MCPTool and override + If you want to extend the legacy behavior, you can subclass MCPTool and override this function. If you want to keep the default behavior, make sure to call ``super().message_handler(message)``. + Alternatively for newer server versions, please see the `_listen_capability_list_changes` method. + Args: message: The message from the MCP server (request responder, notification, or exception). """ + from mcp.shared.subscriptions import SUBSCRIPTION_ID_META_KEY + if isinstance(message, Exception): logger.error("Error from MCP server: %s", message, exc_info=message) return + params = getattr(message, "params", None) + meta = getattr(params, "meta", None) + + # MCP v2 uses subscription directly with Client.listen and attaches the subscription ID to the message meta. + # To avoid double handling we skip at the message_handler level. + if isinstance(meta, Mapping) and SUBSCRIPTION_ID_META_KEY in meta: + return + match message.method: case "notifications/tools/list_changed": self._schedule_reload(self.load_tools()) @@ -2692,8 +2769,8 @@ async def _cancel_pending_reload_tasks(self) -> None: await asyncio.gather(*tasks, return_exceptions=True) async def _close_on_owner(self) -> None: + await self._cancel_capability_list_subscription() await self._cancel_pending_reload_tasks() - await self._safe_close_exit_stack() self._exit_stack = AsyncExitStack() if self._owns_session: diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 7d2297cef5b..57b3d82b4c7 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -8,6 +8,7 @@ import os import sys import warnings +from collections.abc import AsyncIterator from contextlib import AbstractAsyncContextManager, _AsyncGeneratorContextManager # type: ignore from contextvars import ContextVar from datetime import timedelta @@ -3739,6 +3740,275 @@ async def test_mcp_tool_message_handler_notification(): assert result is None +async def test_mcp_tool_refreshes_catalogs_from_modern_subscription() -> None: + from mcp.client.subscriptions import PromptsListChanged, ToolsListChanged + from mcp.shared.subscriptions import SUBSCRIPTION_ID_META_KEY + + tools_refreshed = asyncio.Event() + prompts_refreshed = asyncio.Event() + listen_entered = asyncio.Event() + listen_exited = asyncio.Event() + tool_load_count = 0 + prompt_load_count = 0 + + async def load_tools() -> None: + nonlocal tool_load_count + tool_load_count += 1 + if tool_load_count == 2: + tools_refreshed.set() + + async def load_prompts() -> None: + nonlocal prompt_load_count + prompt_load_count += 1 + if prompt_load_count == 2: + prompts_refreshed.set() + + async def events() -> AsyncIterator[ToolsListChanged | PromptsListChanged]: + yield ToolsListChanged() + yield PromptsListChanged() + + @contextlib.asynccontextmanager + async def listen_context() -> AsyncIterator[AsyncIterator[ToolsListChanged | PromptsListChanged]]: + listen_entered.set() + try: + yield events() + finally: + listen_exited.set() + + capabilities = types.ServerCapabilities( + tools=types.ToolsCapability(list_changed=True), + prompts=types.PromptsCapability(list_changed=True), + ) + sdk_client = _mock_sdk_client(capabilities=capabilities, protocol_version="2026-07-28") + sdk_client.listen = Mock(return_value=listen_context()) + tool = MCPStdioTool(name="test_tool", command="unused") + tool.load_tools = load_tools # type: ignore[method-assign] + tool.load_prompts = load_prompts # type: ignore[method-assign] + + with patch("mcp.Client", return_value=sdk_client): + async with tool: + await asyncio.wait_for(listen_entered.wait(), timeout=1) + sdk_client.listen.assert_called_once_with( + tools_list_changed=True, + prompts_list_changed=True, + ) + await asyncio.wait_for(tools_refreshed.wait(), timeout=1) + await asyncio.wait_for(prompts_refreshed.wait(), timeout=1) + + subscription_meta = {SUBSCRIPTION_ID_META_KEY: "listen-1"} + await tool.message_handler( + types.ToolListChangedNotification(params=types.NotificationParams(_meta=subscription_meta)) + ) + await tool.message_handler( + types.PromptListChangedNotification(params=types.NotificationParams(_meta=subscription_meta)) + ) + await asyncio.sleep(0) + assert tool_load_count == 2 + assert prompt_load_count == 2 + + assert listen_exited.is_set() + + +async def test_mcp_tool_reuses_and_closes_modern_catalog_subscription() -> None: + from mcp.client.subscriptions import PromptsListChanged, ToolsListChanged + + listen_entered = asyncio.Event() + never_set = asyncio.Event() + enter_count = 0 + exit_count = 0 + + async def events() -> AsyncIterator[ToolsListChanged | PromptsListChanged]: + await never_set.wait() + yield ToolsListChanged() + + @contextlib.asynccontextmanager + async def listen_context() -> AsyncIterator[AsyncIterator[ToolsListChanged | PromptsListChanged]]: + nonlocal enter_count, exit_count + enter_count += 1 + listen_entered.set() + try: + yield events() + finally: + exit_count += 1 + + capabilities = types.ServerCapabilities( + tools=types.ToolsCapability(list_changed=True), + prompts=types.PromptsCapability(list_changed=True), + ) + sdk_client = _mock_sdk_client(capabilities=capabilities, protocol_version="2026-07-28") + sdk_client.listen = Mock(side_effect=lambda **_: listen_context()) + tool = MCPStdioTool(name="test_tool", command="unused") + tool.load_tools = AsyncMock() # type: ignore[method-assign] + tool.load_prompts = AsyncMock() # type: ignore[method-assign] + + with patch("mcp.Client", return_value=sdk_client): + async with tool: + await asyncio.wait_for(listen_entered.wait(), timeout=1) + subscription_task = tool._capability_list_subscription_task + assert subscription_task is not None + + await tool.connect() + + sdk_client.listen.assert_called_once_with( + tools_list_changed=True, + prompts_list_changed=True, + ) + assert tool._capability_list_subscription_task is subscription_task + assert not subscription_task.done() + + assert subscription_task.done() + assert tool._capability_list_subscription_task is None + assert enter_count == 1 + assert exit_count == 1 + + +async def test_mcp_tool_replaces_modern_catalog_subscription_on_reset() -> None: + from mcp.client.subscriptions import PromptsListChanged, ToolsListChanged + + second_listen_entered = asyncio.Event() + never_set = asyncio.Event() + enter_count = 0 + exit_count = 0 + + async def events() -> AsyncIterator[ToolsListChanged | PromptsListChanged]: + await never_set.wait() + yield ToolsListChanged() + + @contextlib.asynccontextmanager + async def listen_context() -> AsyncIterator[AsyncIterator[ToolsListChanged | PromptsListChanged]]: + nonlocal enter_count, exit_count + enter_count += 1 + if enter_count == 2: + second_listen_entered.set() + try: + yield events() + finally: + exit_count += 1 + + capabilities = types.ServerCapabilities( + tools=types.ToolsCapability(list_changed=True), + prompts=types.PromptsCapability(list_changed=True), + ) + sdk_client = _mock_sdk_client(capabilities=capabilities, protocol_version="2026-07-28") + sdk_client.listen = Mock(side_effect=lambda **_: listen_context()) + tool = MCPStdioTool(name="test_tool", command="unused") + tool.load_tools = AsyncMock() # type: ignore[method-assign] + tool.load_prompts = AsyncMock() # type: ignore[method-assign] + + with patch("mcp.Client", return_value=sdk_client): + async with tool: + first_task = tool._capability_list_subscription_task + assert first_task is not None + + await tool.connect(reset=True) + await asyncio.wait_for(second_listen_entered.wait(), timeout=1) + + second_task = tool._capability_list_subscription_task + assert second_task is not None + assert second_task is not first_task + assert first_task.done() + assert enter_count == 2 + assert exit_count == 1 + + assert second_task.done() + assert tool._capability_list_subscription_task is None + assert exit_count == 2 + + +async def test_mcp_tool_uses_legacy_catalog_notifications_when_subscription_is_unsupported() -> None: + from mcp.client.subscriptions import ListenNotSupportedError + + listen_attempted = asyncio.Event() + tools_refreshed = asyncio.Event() + prompts_refreshed = asyncio.Event() + tool_load_count = 0 + prompt_load_count = 0 + + async def load_tools() -> None: + nonlocal tool_load_count + tool_load_count += 1 + if tool_load_count == 2: + tools_refreshed.set() + + async def load_prompts() -> None: + nonlocal prompt_load_count + prompt_load_count += 1 + if prompt_load_count == 2: + prompts_refreshed.set() + + @contextlib.asynccontextmanager + async def unsupported_listen() -> AsyncIterator[Any]: + listen_attempted.set() + raise ListenNotSupportedError("2025-11-25") + yield None # pragma: no cover + + capabilities = types.ServerCapabilities( + tools=types.ToolsCapability(list_changed=True), + prompts=types.PromptsCapability(list_changed=True), + ) + sdk_client = _mock_sdk_client(capabilities=capabilities) + sdk_client.listen = Mock(return_value=unsupported_listen()) + tool = MCPStdioTool(name="test_tool", command="unused") + tool.load_tools = load_tools # type: ignore[method-assign] + tool.load_prompts = load_prompts # type: ignore[method-assign] + + with patch("mcp.Client", return_value=sdk_client): + async with tool: + await asyncio.wait_for(listen_attempted.wait(), timeout=1) + await tool.message_handler(types.ToolListChangedNotification()) + await tool.message_handler(types.PromptListChangedNotification()) + await asyncio.wait_for(tools_refreshed.wait(), timeout=1) + await asyncio.wait_for(prompts_refreshed.wait(), timeout=1) + + assert tool_load_count == 2 + assert prompt_load_count == 2 + + +async def test_mcp_tool_cleans_up_when_catalog_subscription_setup_fails() -> None: + @contextlib.asynccontextmanager + async def rejected_listen() -> AsyncIterator[Any]: + raise RuntimeError("subscription rejected") + yield None # pragma: no cover + + capabilities = types.ServerCapabilities(tools=types.ToolsCapability(list_changed=True)) + sdk_client = _mock_sdk_client(capabilities=capabilities, protocol_version="2026-07-28") + sdk_client.listen = Mock(return_value=rejected_listen()) + tool = MCPStdioTool(name="test_tool", command="unused", load_prompts=False) + + with ( + patch("mcp.Client", return_value=sdk_client), + pytest.raises(RuntimeError, match="subscription rejected"), + ): + await tool._connect_on_owner() + + sdk_client.__aexit__.assert_awaited_once() + assert tool.session is None + assert tool.is_connected is False + + +async def test_mcp_tool_closes_subscription_before_reload_tasks_and_exit_stack() -> None: + tool = MCPStdioTool(name="test_tool", command="unused") + cleanup_order: list[str] = [] + + async def cancel_subscription() -> None: + cleanup_order.append("subscription") + + async def cancel_reloads() -> None: + cleanup_order.append("reloads") + + async def close_stack() -> None: + cleanup_order.append("exit_stack") + + with ( + patch.object(tool, "_cancel_capability_list_subscription", side_effect=cancel_subscription), + patch.object(tool, "_cancel_pending_reload_tasks", side_effect=cancel_reloads), + patch.object(tool, "_safe_close_exit_stack", side_effect=close_stack), + ): + await tool._close_on_owner() + + assert cleanup_order == ["subscription", "reloads", "exit_stack"] + + async def test_mcp_tool_message_handler_error(): """Test that message_handler gracefully handles exceptions by logging and returning None.""" tool = MCPStdioTool(name="test_tool", command="python") From 8522ca054b2cb80f5b9ac42c0435821d59ae6fb9 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Tue, 6 Oct 2026 14:01:37 +0200 Subject: [PATCH 25/42] test cleanup --- python/packages/core/tests/core/test_mcp.py | 299 +++++++++++--------- 1 file changed, 169 insertions(+), 130 deletions(-) diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 57b3d82b4c7..b404db87b8f 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -8,7 +8,7 @@ import os import sys import warnings -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Callable, Mapping from contextlib import AbstractAsyncContextManager, _AsyncGeneratorContextManager # type: ignore from contextvars import ContextVar from datetime import timedelta @@ -18,6 +18,7 @@ from unittest.mock import AsyncMock, Mock, patch import pytest +from httpx2 import AsyncClient, MockTransport, Request, Response from mcp import MCPError, types from mcp.client.session import ClientSession from pydantic import BaseModel @@ -8382,13 +8383,22 @@ async def handler(request: httpx.Request) -> httpx.Response: assert initialize_headers[0].get("x-api-key") == "connect-token" -async def test_mcp_streamble_http_tool_connects_to_v2_server() -> None: +_MCPProtocolEndpoint = dict[str, Any] | Callable[[dict[str, Any]], dict[str, Any]] + - from httpx2 import AsyncClient, MockTransport, Request, Response +def _make_mcp_protocol_server_mock( + *, + era: str, + capabilities: Mapping[str, Any], + endpoints: Mapping[str, _MCPProtocolEndpoint], +) -> tuple[MockTransport, list[tuple[dict[str, Any], dict[str, str]]]]: + """Create a protocol-era mock while keeping feature endpoints explicit in each test.""" + if era not in {"modern", "legacy"}: + raise ValueError(f"Unsupported MCP protocol era: {era}") captured_requests: list[tuple[dict[str, Any], dict[str, str]]] = [] - async def mcp_v2_server_mock_handler(request: Request) -> Response: + async def handler(request: Request) -> Response: if request.method == "DELETE": return Response(200) @@ -8397,61 +8407,62 @@ async def mcp_v2_server_mock_handler(request: Request) -> Response: headers = {name.lower(): value for name, value in request.headers.items()} captured_requests.append((body, headers)) - if method == "initialize": - return Response( - 200, - json={ - "jsonrpc": "2.0", - "id": body["id"], - "error": {"code": -32601, "message": "Method not found"}, - }, - ) - + error: dict[str, Any] | None = None if method == "server/discover": - return Response( - 200, - json={ - "jsonrpc": "2.0", - "id": body["id"], - "result": { - "supportedVersions": ["2026-07-28"], - "capabilities": {"tools": {}}, - }, - }, - ) + if era == "modern": + result = { + "supportedVersions": ["2026-07-28"], + "capabilities": dict(capabilities), + } + else: + error = {"code": -32601, "message": "Method not found"} + result = None + elif method == "initialize": + if era == "legacy": + result = { + "protocolVersion": "2025-11-25", + "capabilities": dict(capabilities), + "serverInfo": {"name": "legacy-server", "version": "1.0.0"}, + } + else: + error = {"code": -32601, "message": "Method not found"} + result = None + elif era == "legacy" and method == "notifications/initialized": + return Response(202) + elif era == "legacy" and method == "ping": + result = {} + elif method in endpoints: + endpoint = endpoints[method] + result = endpoint(body) if callable(endpoint) else endpoint + else: + raise AssertionError(f"Unexpected MCP method: {method}") - if method == "tools/list": - return Response( - 200, - json={ - "jsonrpc": "2.0", - "id": body["id"], - "result": { - "cacheScope": "private", - "resultType": "complete", - "ttlMs": 0, - "tools": [{"name": "greet", "inputSchema": {"type": "object", "properties": {}}}], - }, - }, - ) + response: dict[str, Any] = {"jsonrpc": "2.0", "id": body["id"]} + response["error" if error is not None else "result"] = error if error is not None else result + return Response(200, json=response) - if method == "tools/call": - return Response( - 200, - json={ - "jsonrpc": "2.0", - "id": body["id"], - "result": { - "resultType": "complete", - "content": [{"type": "text", "text": "Hello!"}], - "isError": False, - }, - }, - ) + return MockTransport(handler), captured_requests - raise AssertionError(f"Unexpected MCP method: {method}") - user_client = AsyncClient(transport=MockTransport(mcp_v2_server_mock_handler)) +async def test_mcp_streamble_http_tool_connects_to_v2_server() -> None: + transport, captured_requests = _make_mcp_protocol_server_mock( + era="modern", + capabilities={"tools": {}}, + endpoints={ + "tools/list": { + "cacheScope": "private", + "resultType": "complete", + "ttlMs": 0, + "tools": [{"name": "greet", "inputSchema": {"type": "object", "properties": {}}}], + }, + "tools/call": { + "resultType": "complete", + "content": [{"type": "text", "text": "Hello!"}], + "isError": False, + }, + }, + ) + user_client = AsyncClient(transport=transport) tool_a = MCPStreamableHTTPTool( name="a", @@ -8487,82 +8498,23 @@ async def mcp_v2_server_mock_handler(request: Request) -> Response: async def test_mcp_streamable_http_tool_connects_to_legacy_server() -> None: - - from httpx2 import AsyncClient, MockTransport, Request, Response - - captured_methods: list[str] = [] - - async def mcp_legacy_server_mock_handler(request: Request) -> Response: - if request.method == "DELETE": - return Response(200) - - body = json.loads(request.content) - method = body["method"] - captured_methods.append(method) - - if method == "notifications/initialized": - return Response(202) - - if method == "ping": - return Response( - 200, - json={"jsonrpc": "2.0", "id": body["id"], "result": {}}, - ) - if method == "server/discover": - return Response( - 200, - json={ - "jsonrpc": "2.0", - "id": body["id"], - "error": {"code": -32601, "message": "Method not found"}, - }, - ) - - if method == "initialize": - return Response( - 200, - json={ - "jsonrpc": "2.0", - "id": body["id"], - "result": { - "protocolVersion": "2025-11-25", - "capabilities": {"tools": {}}, - "serverInfo": {"name": "legacy-server", "version": "1.0.0"}, - }, - }, - ) - - if method == "tools/list": - return Response( - 200, - json={ - "jsonrpc": "2.0", - "id": body["id"], - "result": { - "cacheScope": "private", - "resultType": "complete", - "ttlMs": 0, - "tools": [{"name": "greet", "inputSchema": {"type": "object", "properties": {}}}], - }, - }, - ) - - if method == "tools/call": - return Response( - 200, - json={ - "jsonrpc": "2.0", - "id": body["id"], - "result": { - "content": [{"type": "text", "text": "Hello!"}], - "isError": False, - }, - }, - ) - - raise AssertionError(f"Unexpected MCP method: {method}") - - user_client = AsyncClient(transport=MockTransport(mcp_legacy_server_mock_handler)) + transport, captured_requests = _make_mcp_protocol_server_mock( + era="legacy", + capabilities={"tools": {}}, + endpoints={ + "tools/list": { + "cacheScope": "private", + "resultType": "complete", + "ttlMs": 0, + "tools": [{"name": "greet", "inputSchema": {"type": "object", "properties": {}}}], + }, + "tools/call": { + "content": [{"type": "text", "text": "Hello!"}], + "isError": False, + }, + }, + ) + user_client = AsyncClient(transport=transport) tool_a = MCPStreamableHTTPTool( name="a", @@ -8580,11 +8532,98 @@ async def mcp_legacy_server_mock_handler(request: Request) -> Response: assert isinstance(result, list) assert [item.text for item in result if item.type == "text"] == ["Hello!"] + captured_methods = [body["method"] for body, _ in captured_requests] assert "server/discover" in captured_methods assert "initialize" in captured_methods assert "tools/list" in captured_methods +@pytest.mark.parametrize( + ("era", "expected_version"), + [ + ("modern", "2026-07-28"), + ("legacy", "2025-11-25"), + ], +) +async def test_mcp_streamable_http_prompt_contract_supports_both_protocol_eras( + era: str, + expected_version: str, +) -> None: + prompt_list_result: dict[str, Any] = { + "prompts": [ + { + "name": "summarize", + "description": "Summarize a topic", + "arguments": [ + { + "name": "topic", + "description": "Topic to summarize", + "required": True, + } + ], + } + ] + } + if era == "modern": + prompt_list_result.update({"resultType": "complete", "cacheScope": "private", "ttlMs": 0}) + + def get_prompt(body: dict[str, Any]) -> dict[str, Any]: + topic = body["params"]["arguments"]["topic"] + result = { + "description": "Rendered summary prompt", + "messages": [ + { + "role": "user", + "content": {"type": "text", "text": f"Summarize {topic}"}, + } + ], + } + if era == "modern": + result["resultType"] = "complete" + return result + + transport, captured_requests = _make_mcp_protocol_server_mock( + era=era, + capabilities={"prompts": {}}, + endpoints={ + "prompts/list": prompt_list_result, + "prompts/get": get_prompt, + }, + ) + + async with AsyncClient(transport=transport) as user_client: + tool = MCPStreamableHTTPTool( + name=f"{era}-prompts", + url="http://example.com/mcp", + load_tools=False, + http_client=user_client, + ) + + async with tool: + assert tool.session is not None + assert tool.session.protocol_version == expected_version + assert [function.name for function in tool.functions] == ["summarize"] + + prompt = tool.functions[0] + context = FunctionInvocationContext( + function=prompt, + arguments={"topic": "Python"}, + kwargs={"runtime_only": "do-not-forward"}, + ) + result = await prompt.invoke(arguments={"topic": "Python"}, context=context) + assert _mcp_result_to_text(result) == "Summarize Python" + + captured_methods = [body["method"] for body, _ in captured_requests] + assert "server/discover" in captured_methods + assert ("initialize" in captured_methods) is (era == "legacy") + assert "prompts/list" in captured_methods + assert "prompts/get" in captured_methods + + prompt_get = next(body for body, _ in captured_requests if body["method"] == "prompts/get") + assert prompt_get["params"]["name"] == "summarize" + assert prompt_get["params"]["arguments"] == {"topic": "Python"} + + @pytest.mark.parametrize( ("era", "expected_version"), [ From a840613b7937bbd602be8b8bf04b9099358a735d Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Tue, 6 Oct 2026 14:29:30 +0200 Subject: [PATCH 26/42] Added support for ResourceLink for prompts --- python/packages/core/agent_framework/_mcp.py | 2 ++ python/packages/core/tests/core/test_mcp.py | 27 ++++++++++++++++++++ 2 files changed, 29 insertions(+) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 14c5c263c0a..2a298b0e1c4 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -1120,6 +1120,8 @@ def _parse_prompt_result_from_mcp( default=str, ) ) + elif isinstance(content, types.ResourceLink): + parts.append(json.dumps(content.model_dump(by_alias=True, exclude_none=True), default=str)) else: parts.append(str(content)) if not parts: diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index b404db87b8f..6b332bf1469 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -792,6 +792,15 @@ def test_mcp_tool_str_and_parse_prompt_result_rich_content() -> None: ), ), ), + types.PromptMessage( + role="assistant", + content=types.ResourceLink( + name="prompt-guide", + uri="https://example.test/prompt-guide", + description="Prompt guide", + mime_type="text/markdown", + ), + ), ] ) @@ -804,6 +813,13 @@ def test_mcp_tool_str_and_parse_prompt_result_rich_content() -> None: assert json.loads(parsed[2]) == {"type": "audio", "data": "YXVkaW8=", "mimeType": "audio/wav"} assert parsed[3] == "Embedded prompt" assert json.loads(parsed[4]) == {"type": "blob", "data": "ZGF0YQ==", "mimeType": "application/pdf"} + assert json.loads(parsed[5]) == { + "type": "resource_link", + "name": "prompt-guide", + "uri": "https://example.test/prompt-guide", + "description": "Prompt guide", + "mimeType": "text/markdown", + } def test_parse_tool_result_from_mcp(): @@ -8622,6 +8638,17 @@ def get_prompt(body: dict[str, Any]) -> dict[str, Any]: prompt_get = next(body for body, _ in captured_requests if body["method"] == "prompts/get") assert prompt_get["params"]["name"] == "summarize" assert prompt_get["params"]["arguments"] == {"topic": "Python"} + if era == "modern": + for method, name in (("prompts/list", None), ("prompts/get", "summarize")): + body, headers = next(item for item in captured_requests if item[0]["method"] == method) + meta = body["params"]["_meta"] + assert ( + headers["mcp-protocol-version"] == meta["io.modelcontextprotocol/protocolVersion"] == expected_version + ) + assert headers["mcp-method"] == method + assert isinstance(meta["io.modelcontextprotocol/clientCapabilities"], dict) + if name is not None: + assert headers["mcp-name"] == body["params"]["name"] == name @pytest.mark.parametrize( From 3a4b24d04334b9f6724f9b16bbbd2603cca9a6f1 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Tue, 6 Oct 2026 15:17:47 +0200 Subject: [PATCH 27/42] fixed and tested skill support --- .../packages/core/agent_framework/_skills.py | 16 +++- .../core/tests/core/test_mcp_skills.py | 95 ++++++++++++++----- 2 files changed, 85 insertions(+), 26 deletions(-) diff --git a/python/packages/core/agent_framework/_skills.py b/python/packages/core/agent_framework/_skills.py index 27073038b04..090768868e3 100644 --- a/python/packages/core/agent_framework/_skills.py +++ b/python/packages/core/agent_framework/_skills.py @@ -4450,6 +4450,7 @@ def _is_mcp_resource_not_found(ex: Exception) -> bool: * ``METHOD_NOT_FOUND`` (``-32601``) — the server does not implement ``resources/read`` at all, which for the skills source is functionally equivalent to "no skills available." + * ``INVALID_PARAMS`` from an :class:`MCPError` indicating that the ``uri`` parameter is invalid. All other codes — ``INVALID_PARAMS``, ``INTERNAL_ERROR``, ``PARSE_ERROR``, ``CONNECTION_CLOSED``, auth rejections, and generic handler errors @@ -4457,13 +4458,20 @@ def _is_mcp_resource_not_found(ex: Exception) -> bool: token or crashing server is not silently mistaken for "the server has no skills." """ - from mcp.shared.exceptions import MCPError as _McpError + from mcp.shared.exceptions import MCPError as _MCPError + from mcp.types import INVALID_PARAMS as _INVALID_PARAMS + from mcp.types import METHOD_NOT_FOUND as _METHOD_NOT_FOUND - if not isinstance(ex, _McpError): + if not isinstance(ex, _MCPError): return False - from mcp.types import METHOD_NOT_FOUND as _METHOD_NOT_FOUND - return ex.error.code in {-32002, _METHOD_NOT_FOUND} + data = ex.error.data + return ex.error.code in {-32002, _METHOD_NOT_FOUND} or ( + ex.error.code == _INVALID_PARAMS + and isinstance(data, dict) + and cast(dict[str, object], data).keys() == {"uri"} + and isinstance(data["uri"], str) + ) def _mcp_join_text(result: ReadResourceResult) -> str: diff --git a/python/packages/core/tests/core/test_mcp_skills.py b/python/packages/core/tests/core/test_mcp_skills.py index 6259359a025..9d172b5eff5 100644 --- a/python/packages/core/tests/core/test_mcp_skills.py +++ b/python/packages/core/tests/core/test_mcp_skills.py @@ -10,7 +10,7 @@ import json import warnings import zipfile -from collections.abc import Iterator, Mapping, Sequence +from collections.abc import Callable, Iterator, Mapping, Sequence from datetime import timedelta from unittest.mock import AsyncMock, patch from urllib.parse import unquote @@ -18,6 +18,7 @@ import pytest from mcp import MCPError from mcp.types import ( + INVALID_PARAMS, BlobResourceContents, ReadResourceResult, TextResourceContents, @@ -80,13 +81,28 @@ def _make_empty_result() -> ReadResourceResult: return ReadResourceResult(contents=[]) -def _make_client(**read_resource_responses: ReadResourceResult) -> AsyncMock: +def _legacy_resource_not_found(uri: str) -> MCPError: + return MCPError(-32002, f"Resource not found: {uri}") + + +def _modern_resource_not_found(uri: str) -> MCPError: + return MCPError(INVALID_PARAMS, "Resource not found", {"uri": uri}) + + +_RESOURCE_NOT_FOUND_ERRORS = [_legacy_resource_not_found, _modern_resource_not_found] + + +def _make_client( + *, + resource_not_found_error: Callable[[str], MCPError] = _legacy_resource_not_found, + **read_resource_responses: ReadResourceResult, +) -> AsyncMock: """Create a mock ClientSession whose read_resource returns different results per URI. Args: + resource_not_found_error: Creates the missing-resource error for an unknown URI. **read_resource_responses: Mapping of URI string to ReadResourceResult. - Any URI not in this mapping raises MCPError with the MCP-spec - "Resource not found" code (-32002). + Any URI not in this mapping raises the configured missing-resource error. """ client = AsyncMock() @@ -94,7 +110,7 @@ async def _read_resource(uri: str) -> ReadResourceResult: uri_str = str(uri) if uri_str in read_resource_responses: return read_resource_responses[uri_str] - raise MCPError(-32002, f"Resource not found: {uri_str}") + raise resource_not_found_error(uri_str) client.read_resource = AsyncMock(side_effect=_read_resource) return client @@ -644,7 +660,15 @@ async def test_missing_required_fields_is_skipped(self) -> None: skills = await source.get_skills(_SOURCE_CTX) assert skills == [] - async def test_archive_missing_resource_is_skipped(self) -> None: + @pytest.mark.parametrize( + "resource_not_found_error", + _RESOURCE_NOT_FOUND_ERRORS, + ids=["legacy-shape", "modern-shape"], + ) + async def test_archive_missing_resource_is_skipped( + self, + resource_not_found_error: Callable[[str], MCPError], + ) -> None: # An archive entry whose archive resource is not available on the server # is skipped (the index is read, but the archive download fails). index_json = json.dumps({ @@ -658,7 +682,10 @@ async def test_archive_missing_resource_is_skipped(self) -> None: } ], }) - client = _make_client(**{"skill://index.json": _make_text_result(index_json, uri="skill://index.json")}) + client = _make_client( + resource_not_found_error=resource_not_found_error, + **{"skill://index.json": _make_text_result(index_json, uri="skill://index.json")}, + ) source = MCPSkillsSource(client=client) skills = await source.get_skills(_SOURCE_CTX) assert skills == [] @@ -757,10 +784,9 @@ def test_requires_exactly_one_of_client_or_session_provider(self) -> None: class TestMCPSkillsSourceErrorCodeBranching: """Tests that MCPSkillsSource and MCPSkill branch on MCPError.error.code. - Only "not found" codes (RESOURCE_NOT_FOUND -32002, METHOD_NOT_FOUND -32601) - should be silently swallowed as "no skills available." Other MCPError codes - and non-MCPError exceptions must propagate so that auth failures, server - crashes, and connection drops are visible. + The legacy -32002 code, METHOD_NOT_FOUND, and INVALID_PARAMS carrying the + exact requested URI mean the resource is absent. Other MCPError shapes and + non-MCPError exceptions must propagate so failures remain visible. """ async def test_index_method_not_found_returns_empty(self) -> None: @@ -771,18 +797,36 @@ async def test_index_method_not_found_returns_empty(self) -> None: skills = await source.get_skills(_SOURCE_CTX) assert skills == [] - async def test_index_resource_not_found_returns_empty(self) -> None: - """MCP-spec "Resource not found" (-32002) -> server has no index.""" - client = AsyncMock() - client.read_resource = AsyncMock(side_effect=MCPError(-32002, "Resource not found")) + @pytest.mark.parametrize( + "resource_not_found_error", + _RESOURCE_NOT_FOUND_ERRORS, + ids=["legacy-shape", "modern-shape"], + ) + async def test_index_resource_not_found_returns_empty( + self, + resource_not_found_error: Callable[[str], MCPError], + ) -> None: + """Either valid missing-resource error shape means the server has no skill index.""" + client = _make_client(resource_not_found_error=resource_not_found_error) source = MCPSkillsSource(client=client) skills = await source.get_skills(_SOURCE_CTX) assert skills == [] - async def test_index_invalid_params_propagates(self) -> None: - """INVALID_PARAMS (-32602) is a real bug, must propagate (not "not found").""" + @pytest.mark.parametrize( + "data", + [ + None, + {"uri": "skill://index.json", "reason": "malformed request"}, + ], + ids=["no-uri", "extra-error-data"], + ) + async def test_index_invalid_params_propagates( + self, + data: dict[str, str] | None, + ) -> None: + """INVALID_PARAMS propagates unless its data identifies the requested URI.""" client = AsyncMock() - client.read_resource = AsyncMock(side_effect=MCPError(-32602, "Invalid params")) + client.read_resource = AsyncMock(side_effect=MCPError(INVALID_PARAMS, "Invalid params", data)) source = MCPSkillsSource(client=client) with pytest.raises(MCPError): await source.get_skills(_SOURCE_CTX) @@ -830,12 +874,19 @@ async def test_get_resource_internal_error_propagates(self) -> None: with pytest.raises(MCPError): await skill.get_resource("references/file.md") - async def test_get_resource_not_found_returns_none(self) -> None: - """MCPError with RESOURCE_NOT_FOUND (-32002) on get_resource returns None.""" + @pytest.mark.parametrize( + "resource_not_found_error", + _RESOURCE_NOT_FOUND_ERRORS, + ids=["legacy-shape", "modern-shape"], + ) + async def test_get_resource_not_found_returns_none( + self, + resource_not_found_error: Callable[[str], MCPError], + ) -> None: + """Either valid missing-resource error shape on get_resource returns None.""" from agent_framework import SkillFrontmatter - client = AsyncMock() - client.read_resource = AsyncMock(side_effect=MCPError(-32002, "Resource not found")) + client = _make_client(resource_not_found_error=resource_not_found_error) fm = SkillFrontmatter(name="test-skill", description="Test.") skill = MCPSkill(frontmatter=fm, skill_md_uri="skill://test/SKILL.md", client=client) result = await skill.get_resource("references/file.md") From e57fc3b04697ba927a9dcbeba2f36716f4ed8dcd Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Tue, 6 Oct 2026 18:12:39 +0200 Subject: [PATCH 28/42] ADR and agents update --- .../0045-python-mcp-v2-client-lifecycle.md | 336 +++++++++++++----- python/packages/core/AGENTS.md | 7 +- 2 files changed, 250 insertions(+), 93 deletions(-) diff --git a/docs/decisions/0045-python-mcp-v2-client-lifecycle.md b/docs/decisions/0045-python-mcp-v2-client-lifecycle.md index 4ebf5a4976b..a7f6199b03e 100644 --- a/docs/decisions/0045-python-mcp-v2-client-lifecycle.md +++ b/docs/decisions/0045-python-mcp-v2-client-lifecycle.md @@ -1,135 +1,287 @@ --- status: proposed -date: 2026-10-01 -deciders: eavanvalkenburg, westey-m +date: 2026-10-06 +deciders: jpalvarezl, eavanvalkenburg --- -# Adopt the MCP v2 negotiating client behind Python MCP tools +# Use the MCP SDK Client as the primary Python MCP client ## Context and Problem Statement -The Python `MCPTool` implementation owns an MCP transport, constructs a low-level `ClientSession`, and always calls -`initialize()`. That lifecycle supports handshake-era protocol versions through 2025-11-25, but it cannot connect to -a 2026-07-28 server, where `server/discover` replaces the initialization handshake. +The Python `MCPTool` implementation originally owned an MCP transport, constructed a low-level `ClientSession`, and +used that session for every operation. That design matched handshake-era protocol versions through 2025-11-25. -The MCP Python SDK v2 provides a high-level `Client(mode="auto")` that probes `server/discover` and falls back to -`initialize` when the probe does not provide positive evidence of a compatible modern peer. Agent Framework needs to -adopt that negotiation without replacing its public `MCPTool`, `MCPStdioTool`, or `MCPStreamableHTTPTool` classes, or -regressing async context management, reconnect behavior, caller-owned sessions, and request-scoped authentication -headers. +The MCP Python SDK v2 adds a high-level `Client` that owns protocol negotiation and standard client behavior. In +`mode="auto"` it probes `server/discover` for a 2026-07-28 peer and falls back to the legacy `initialize` handshake. +Its standard operation methods add automatic multi-round-trip request (MRTR) handling, response caching and cache +invalidation coordination, and claimed extension-result resolution that session-tier calls do not provide. +`ClientSession` can still declare callback capabilities from its configured callbacks, and the SDK exposes a typed +subscription helper over an entered session; the high-level `Client` is the normal API that configures and coordinates +those features with its own operations and cache. + +The first migration pass entered `Client` for framework-created connections but continued to route most operations +through `client.session`. That partial adoption supports both protocol eras, but it bypasses SDK-owned behavior and +leaves later migration work with two possible implementations. + +This decision defines the lasting ownership boundary: which work belongs to the SDK `Client`, which work remains +Agent Framework policy, and when direct `ClientSession` access is justified. + +This ADR covers the outbound Python MCP client. Hosting/server code uses the SDK server APIs and is not converted to +`Client`. ## Decision Drivers -- One Agent Framework code path must support 2026-07-28 and 2025-era MCP peers. -- Protocol negotiation should remain owned by the MCP SDK rather than being independently reimplemented. -- Existing public tool classes, `async with` ergonomics, and tool/prompt discovery behavior must remain compatible. +- One Agent Framework path must support 2026-07-28 and 2025-era MCP peers. +- Standard protocol behavior should remain owned by the MCP SDK rather than being copied into Agent Framework. +- Existing `MCPTool`, `MCPStdioTool`, and `MCPStreamableHTTPTool` APIs and `async with` ergonomics must remain + compatible. - Framework-created resources must still be closed by the task that opened them, including cancellation and failed connection attempts. - Caller-supplied sessions and HTTP clients must remain caller-owned. -- Dynamic HTTP header identity must continue to bind discovery and later requests to the same authenticated identity. -- Stdio and Streamable HTTP must share the same protocol lifecycle. +- Existing tool argument filtering, parsing, tracing, reconnect, header identity, and security behavior must remain + Agent Framework policy around the SDK operation. +- Moving from `ClientSession` to `Client` must not silently change cache or refresh behavior. +- Tools, prompts, and resources must share one client-ownership rule rather than acquire separate MRTR loops. +- Tasks and delayed/durable human-input continuation must remain separate from base MRTR support. ## Considered Options -- Enter the MCP SDK v2 `Client(mode="auto")` behind the existing Agent Framework tool classes. -- Reimplement `server/discover` and legacy fallback directly around `ClientSession`. +- Retain the high-level SDK `Client` and use it for standard operations, with `ClientSession` as an explicit escape + hatch. +- Use `Client` only for negotiation while continuing all operations through `client.session`. +- Reimplement negotiation, MRTR, caching, and subscriptions around `ClientSession`. - Add separate modern and legacy Agent Framework tool classes or a public protocol-mode switch. ## Decision Outcome -Chosen option: "Enter the MCP SDK v2 `Client(mode="auto")` behind the existing Agent Framework tool classes", -because it uses the SDK's supported negotiation path while preserving the Agent Framework API and transport-specific -behavior. - -For framework-created connections: - -- `MCPStdioTool` and `MCPStreamableHTTPTool` continue to create their existing transport context managers. -- `MCPTool` passes that transport to `mcp.Client(mode="auto")` and enters the client on its existing `AsyncExitStack`. -- The SDK probes `server/discover`. A compatible modern peer is adopted without sending `initialize`; a probe that - does not establish a compatible modern peer falls back to the initialization handshake on the same code path. - `METHOD_NOT_FOUND` is the representative legacy response covered by Agent Framework's contract test, not the SDK's - only fallback condition. -- Agent Framework retains a private reference to the high-level client and exposes its underlying `ClientSession` - through the existing `session` attribute for compatibility with current integrations. -- Negotiated protocol version and server capabilities are read from the client's era-neutral properties rather than - captured only from an initialize result. -- Tools configured for the existing server-initiated sampling callback use `mode="legacy"`. The MCP migration guide - advises legacy mode for workflows that rely on that back-channel behavior because 2026-07-28 refuses - server-initiated sampling on every transport. Auto mode remains the default when no legacy-only callback behavior - is requested. - -Agent Framework retains its lifecycle owner task and locks. They protect framework state, preserve AnyIO task -ownership during teardown, serialize identity-changing reconnects, and roll back partially loaded discovery state. -They no longer implement protocol negotiation. A reset or reconnect closes the whole SDK client and transport, then -constructs a new auto-negotiating client. `ping` is not used as a universal liveness preflight because it does not -exist in 2026-07-28; operation failures drive reconnect, while any retained legacy ping behavior is gated by the -negotiated protocol. - -Caller-supplied `ClientSession` remains a compatibility path: - -- Agent Framework does not enter, close, or replace the session. -- An already negotiated session is reused through its `protocol_version` and `server_capabilities` properties. -- An unnegotiated session retains the existing compatibility behavior and uses `initialize()`. A caller that supplies - a modern low-level session negotiates it with `discover()` before passing it to Agent Framework. This avoids - duplicating the SDK's broader auto-negotiation policy, which is implemented only by the high-level `Client`. - -Dynamic and static HTTP headers remain below the negotiating client at the Streamable HTTP transport boundary. The -effective header set continues to define connection identity; changing it closes and rebuilds the whole SDK client -before discovery or tool calls proceed. This preserves the trust and approval boundaries documented in -[ADR 0043](0043-python-mcp-runtime-context.md). - -The MCP Streamable HTTP implementation uses `httpx2` types, as required by MCP SDK v2. `httpx` and `httpx2` may -coexist elsewhere in the repository; this decision does not require unrelated HTTP features to migrate. +Chosen option: "Retain the high-level SDK `Client` and use it for standard operations, with `ClientSession` as an +explicit escape hatch", because it delegates protocol mechanics to the supported SDK surface while preserving Agent +Framework's lifecycle and application policy. + +### Decided target ownership model + +The following is the target state established by this decision; the branch has not implemented every item yet. + +For a framework-created connection, `MCPTool` will own and retain one entered `mcp.Client`. It will continue exposing +that client's underlying `ClientSession` through the existing `session` attribute for compatibility and advanced use. +The client and session form one connection unit: they are published only after successful negotiation and are cleared +and replaced together on close, reset, reconnect, cancellation cleanup, or authenticated-header identity change. + +Agent Framework continues to own: + +- the lifecycle-owner task and locks; +- transport construction and request-scoped HTTP header identity; +- reconnect policy around complete high-level operations; +- tool argument filtering and trusted request metadata merging; +- result parsing, Agent Framework content conversion, tracing, and error translation; +- catalog publication and rollback after a failed complete fetch; +- security and approval policy. + +The SDK `Client` owns: + +- `server/discover` negotiation and legacy `initialize` fallback; +- standard tools, prompts, and resources operations; +- MRTR callback dispatch, retries, round limits, state echo, and state-only backoff; +- response-cache storage and protocol notification invalidation; +- `subscriptions/listen` acknowledgment and event delivery; +- protocol/client/capability metadata stamps; +- extension result-claim resolution registered on the client. + +### Operation ownership + +| Operation | Default path | Direct `ClientSession` exception | +|---|---|---| +| Connection negotiation | `Client(mode="auto")` | A caller-supplied session is already negotiated, or an unnegotiated supplied session retains legacy `initialize()` compatibility. | +| `tools/list`, `tools/call` | `Client` | Manual MRTR or raw claimed-extension results. | +| `prompts/list`, `prompts/get` | `Client` | Manual MRTR. | +| `resources/list`, `resources/templates/list`, `resources/read` | `Client` | Manual MRTR or a deliberately uncached low-level caller-owned session. | +| `subscriptions/listen` | `Client.listen()` | The SDK session helper may be used for a caller-supplied session, without a framework-owned client cache. | +| Standard request metadata | High-level operation `meta=` | Raw extension requests use `ClientSession.send_request()`. | +| Modern log-level opt-in | `Client(log_level=...)` | Legacy `logging/setLevel` remains negotiated-version compatibility only. | +| Legacy ping | None for modern connections | A protocol-gated legacy session call may remain temporarily. | +| Tasks or arbitrary extension RPCs | Future extension/client API | Raw `ClientSession.send_request()` at the extension boundary only. | + +Code should use one private, connection-bound operation provider rather than repeat `Client`/`ClientSession` +selection independently across tools, prompts, skills, security helpers, and Foundry integrations. + +### Cache policy is explicit + +High-level and low-level calls are not interchangeable. `ClientSession` always reaches the server. In SDK 2.2.0, +`Client` makes page-one results for tools, prompts, resources, and resource templates, plus `resources/read`, cache +aware; `tools/call` and `prompts/get` are not cached. A result is reusable only when it has a positive effective TTL. +Absent server hints use `CacheConfig.default_ttl_ms`, whose default is `0`. Under `cache_mode="use"`, non-`None` +request metadata forces a wire refresh. Although `server/discover` carries protocol cache hints, SDK 2.2.0 deliberately +excludes it from the response cache; persisting or reusing `prior_discover` is caller-managed. + +The Client-first cleanup and the separate Caching checklist item are staged deliberately: + +1. While preserving pre-caching Agent Framework behavior, explicit catalog and resource refresh paths use + `cache_mode="bypass"`. +2. The Caching migration later assigns an intentional policy per operation: + - `"refresh"` for an explicit authoritative refetch that must update or evict the SDK cache; + - `"use"` only where Agent Framework intentionally accepts server `ttlMs` / `cacheScope` freshness; + - `"bypass"` only where neither reading nor updating the SDK cache is desired. +3. Resource and catalog refresh tests must cover positive TTLs, pagination, changed metadata, empty snapshots, + reconnect, and authenticated identity changes. +4. A reconnect or effective header-identity change replaces the whole Client and its default per-client cache. A + shared cache store must be partitioned by a verified authorization identity. Because Agent Framework constructs + `Client` from a transport rather than a URL, a future shared store also requires an explicit stable + `CacheConfig.target_id`. + +MRTR-seeded and MRTR-resolved resource reads are not cached by the SDK. + +### MRTR + +For framework-owned connections, Agent Framework delegates MRTR for `tools/call`, `prompts/get`, and +`resources/read` to the high-level `Client`. The SDK dispatches embedded elicitation, sampling, and roots requests to +the callbacks configured on that client, retries with `inputResponses`, echoes `requestState` unchanged, assigns a +new JSON-RPC request ID, and applies its round limit and state-only backoff. + +Agent Framework does not import the private `mcp.client._input_required` driver and does not implement separate +tool, prompt, or resource loops. + +Direct session handling with `allow_input_required=True` is reserved for a future requirement to persist, inspect, or +resume MRTR rounds outside the process that began the operation. `UserInputRequiredException` is currently a way to +surface a pause, not a complete persisted MCP continuation model. Delayed or durable human-input continuation +therefore requires a separate design and function-calling-loop review; it is not part of the base MRTR migration. + +### Caller-supplied clients and sessions + +A caller-supplied `ClientSession` remains a low-level compatibility path: + +- Agent Framework does not enter, close, replace, wrap, or reconnect it. +- An already negotiated session is reused through its era-neutral properties. +- An unnegotiated session retains legacy `initialize()` compatibility. +- Standard calls remain direct and therefore do not gain automatic MRTR or SDK response caching. +- Callbacks configured on the session do not by themselves drive MRTR. + +SDK 2.2.0 has no public API that adopts an existing `ClientSession` into `Client`, and it does not export the MRTR +driver. Agent Framework must not open a second connection to simulate adoption. + +Experiments proved that an already-entered, caller-owned high-level `Client` can be reused without transferring +ownership and can provide automatic MRTR. A public API for that path is preferred over a framework-owned MRTR loop, +but it requires a separate API decision. It should not force callers to provide dummy transport arguments to +transport-specific wrappers. + +### Legacy callbacks, logging, and liveness + +The existing sampling-enabled path uses `mode="legacy"` because it relies on the legacy server-to-client +back-channel. Modern 2026-07-28 sampling, elicitation, and roots requests travel inside MRTR instead. Moving the +existing sampling option to auto mode requires a separate compatibility decision; it is not an incidental part of +routing standard calls through `Client`. + +Modern protocol logging is per-request metadata. Agent Framework should pass the selected log level to `Client` so +the SDK stamps `io.modelcontextprotocol/logLevel`; `logging/setLevel` is retained only for a negotiated legacy peer. +New logging behavior should prefer normal Python logging and OpenTelemetry because MCP protocol logging is +deprecated. + +`ping` is removed in 2026-07-28. It is not a universal liveness preflight. Operation failures drive reconnect; any +remaining ping call is explicitly gated to a legacy session. + +### Checklist boundaries + +The Client-first cleanup strengthens the foundation of completed migration rows without reopening their accepted +scope: + +- Hosting/server remains server-side and checked. +- Tools, Tool refresh, Prompts, and Skills remain checked for their stated dual-era behavior. +- MRTR, Caching, Logging, Samples/docs, local validation, and live dual-era validation remain separate unchecked + work. +- The protocol-independent prompt snapshot bug remains in + [microsoft/agent-framework#9115](https://github.com/microsoft/agent-framework/issues/9115). +- Optional subscription stream recovery remains in + [microsoft/agent-framework#9109](https://github.com/microsoft/agent-framework/issues/9109). +- Tasks and the removed task sample remain blocked until the Python SDK exposes the current + `io.modelcontextprotocol/tasks` runtime. A task status of `input_required` is not MRTR `InputRequiredResult`. ### Consequences -- Good, because modern and legacy servers use one public Agent Framework API and one SDK-supported negotiation path. -- Good, because stdio, Streamable HTTP, reconnect, cancellation, and header identity remain Agent Framework concerns - without duplicating protocol-version selection. -- Good, because later work can use the high-level client for per-request metadata, subscriptions, MRTR, and caching. +- Good, because modern and legacy servers use one public Agent Framework API and SDK-supported negotiation. +- Good, because standard protocol behavior, MRTR, subscriptions, and caching stay with the SDK. +- Good, because stdio, Streamable HTTP, reconnect, cancellation, header identity, parsing, and policy remain Agent + Framework concerns. +- Good, because tools, prompts, and resources follow one ownership rule. - Neutral, because Agent Framework keeps both a private high-level client and the public low-level `session` view. -- Neutral, because caller-supplied unnegotiated sessions remain handshake-era unless the caller negotiates them first. -- Bad, because tests that mocked `ClientSession` construction must move toward client/transport contract tests. +- Neutral, because caller-supplied `ClientSession` objects intentionally remain lower level. +- Bad, because tests and integrations that mock `tool.session` may need to move toward operation or transport + contract tests. +- Bad, because adopting `Client` exposes cache behavior that must be selected and tested explicitly. ## Validation -Contract tests cover both branches through the same Agent Framework tool: +Existing committed contract tests cover: + +- modern Streamable HTTP discovery without `initialize`; +- legacy fallback on the same public tool path; +- modern and legacy stdio list/call behavior; +- required headers, `_meta`, parsing, tracing, reconnect, lifecycle ownership, and caller-owned sessions; +- modern `Client.listen()` with legacy notification fallback; +- dual-era prompts and skills resource-not-found behavior. -- A modern-only Streamable HTTP peer accepts `server/discover`, rejects `initialize`, serves `tools/list`, and records - that the negotiated protocol is 2026-07-28. -- A 2025-era peer rejects `server/discover`, accepts `initialize`, and serves the same tool operations. -- Equivalent stdio coverage verifies that negotiation is transport-independent. -- Existing lifecycle, reconnect, cancellation, caller-ownership, and dynamic-header identity tests remain green. +Separate detached-worktree experiments based on commit `a02f9ff8c487d04c719d4ce5e3a3a8326b856673` established: -Follow-on tests cover per-request logging metadata, `subscriptions/listen` with legacy notification fallback, prompts, -skills, MRTR, and caching. Tasks remain deferred until the Python MCP SDK exposes the 2026 Tasks extension runtime. +- the existing direct-session tool call raises on the first state-only `InputRequiredResult`; +- retaining and calling `Client` causes SDK 2.2.0 to retry with the same operation arguments, exact + `requestState`, omitted empty `inputResponses`, and a new JSON-RPC request ID; +- high-level tools, prompts, and resources MRTR work without function-calling-loop changes; +- metadata, parsing, tracing, reconnect, subscriptions, and legacy fallback remain compatible; +- a caller-owned, already-entered `Client` remains usable after an Agent Framework wrapper closes; +- a caller-supplied `ClientSession` does not gain automatic MRTR merely by having callbacks. + +Implementation validation for this decision must cover: + +- retained Client publication, failure cleanup, close, reset, reconnect, cancellation, and header identity change; +- standard operations routing through the current owned Client, never a stale prior Client; +- tool, prompt, and resource MRTR, including state-only and embedded callback rounds; +- explicit cache modes, positive TTLs, pagination, invalidation, and refresh; +- caller-owned `ClientSession` behavior remaining direct, uncached, and caller-owned; +- legacy sampling, logging, and ping compatibility; +- full core source and test typing plus the complete affected unit suites. ## Pros and Cons of the Options -### Enter the MCP SDK v2 client behind existing tool classes +### Retain Client and use it for standard operations + +- Good, because it uses the SDK's supported, documented client surface. +- Good, because it avoids duplicate negotiation, MRTR, caching, and subscription code. +- Good, because the low-level session remains available where genuinely required. +- Bad, because operation mocks and cache-sensitive tests need deliberate migration. + +### Use Client only for negotiation and session for operations -- Good, because the SDK owns current and future negotiation rules. -- Good, because the existing transport subclasses and public API remain intact. -- Good, because it matches the .NET direction of delegating connection creation and negotiation to its MCP SDK. -- Bad, because the existing connection code and mocks must be reshaped around a higher-level lifecycle owner. +- Good, because it minimizes immediate code changes. +- Bad, because it bypasses SDK-owned MRTR, caching, extension claims, and subscription cache coordination. +- Bad, because every new protocol feature reopens the same ownership question. +- Bad, because the difference between framework-owned and caller-owned sessions becomes implicit. -### Reimplement negotiation around ClientSession +### Reimplement behavior around ClientSession -- Good, because it minimizes the first code diff and preserves direct session construction. -- Bad, because Agent Framework would duplicate the SDK's discover, fallback, adoption, and future-version behavior. -- Bad, because later high-level SDK features would still require a second migration. +- Good, because it could expose every intermediate round to Agent Framework. +- Bad, because Agent Framework would duplicate SDK protocol policy and future changes. +- Bad, because exactly-once execution, cancellation, concurrent callback dispatch, round limits, and cache + invalidation become framework responsibilities. +- Bad, because the function-calling loop would be changed before a durable continuation requirement exists. -### Add era-specific tool classes or a public mode switch +### Add era-specific tool classes or a public protocol-mode switch -- Good, because callers could force a known protocol era. -- Bad, because callers should not need to know a server's era before connecting. -- Bad, because it duplicates public classes, tests, documentation, and lifecycle behavior. -- Bad, because a mode switch can accidentally disable fallback and fragment compatibility. +- Good, because callers could force a known era. +- Bad, because callers should not need to know the peer's era before connecting. +- Bad, because it fragments public classes, tests, and documentation. +- Bad, because it can disable SDK fallback accidentally. ## More Information - [Python MCP 2026-07-28 umbrella issue](https://github.com/microsoft/agent-framework/issues/8245) - [MCP Python SDK v1-to-v2 migration guide](https://py.sdk.modelcontextprotocol.io/migration/#clients) -- [.NET MCP Tasks migration](https://github.com/microsoft/agent-framework/pull/7774) -- The .NET declarative MCP handler delegates connection creation to `McpClient.CreateAsync`; its protocol stub rejects - `server/discover` with `METHOD_NOT_FOUND` before accepting `initialize`, demonstrating SDK-owned legacy fallback. +- [MCP Python SDK Client guide](https://py.sdk.modelcontextprotocol.io/client/) +- [MCP Python SDK MRTR guide](https://py.sdk.modelcontextprotocol.io/handlers/multi-round-trip/) +- [MCP Python SDK caching guide](https://py.sdk.modelcontextprotocol.io/client/caching/) +- [MCP Python SDK subscriptions guide](https://py.sdk.modelcontextprotocol.io/client/subscriptions/) +- [MCP 2026 MRTR specification](https://modelcontextprotocol.io/specification/2026-07-28/basic/patterns/mrtr) +- [MCP 2026 subscriptions specification](https://modelcontextprotocol.io/specification/2026-07-28/basic/patterns/subscriptions) +- [MCP 2026 caching specification](https://modelcontextprotocol.io/specification/2026-07-28/server/utilities/caching) +- [ADR 0043: MCP runtime context](0043-python-mcp-runtime-context.md) + +While this ADR remains `proposed`, implementation sessions should treat it as the current evidence-backed direction +for microsoft/agent-framework#8245 and revise it if new SDK or code evidence changes the decision. It becomes +`accepted` only through the repository ADR review process. diff --git a/python/packages/core/AGENTS.md b/python/packages/core/AGENTS.md index 0efc1f19bb8..70b2f8fa7e8 100644 --- a/python/packages/core/AGENTS.md +++ b/python/packages/core/AGENTS.md @@ -203,7 +203,12 @@ The vector store API is experimental under the shared `VECTOR_STORES` feature ID ### Model Context Protocol (`_mcp.py`) -- **`MCPTool`** - Base wrapper that owns the MCP `ClientSession` and exposes the remote server's tools as `FunctionTool`s. +- **Client lifecycle design** - [ADR 0045](../../../docs/decisions/0045-python-mcp-v2-client-lifecycle.md) is the + current evidence-backed target for high-level `mcp.Client` versus low-level `ClientSession` ownership during the + Python MCP v2 migration; not every target-state item is implemented yet. Keep the ADR aligned when implementation + evidence changes this boundary. +- **`MCPTool`** - Base wrapper that currently exposes the MCP `ClientSession` and the remote server's tools as + `FunctionTool`s. Follow ADR 0045 when migrating standard operations to the retained high-level client. - **`MCPStdioTool`** / **`MCPStreamableHTTPTool`** - Supported transport-specific subclasses. **`MCPWebsocketTool`** remains only as a deprecated compatibility symbol because MCP v2 removed WebSocket transport; it cannot create a connection. From 45d7611b6b6fbfa1e12e32501642e59b4edec361 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 7 Oct 2026 10:15:48 +0200 Subject: [PATCH 29/42] Added MCP client as a private field, set usage precedence 1st mcp_client, 2nd raw session --- python/packages/core/agent_framework/_mcp.py | 18 +++++++- python/packages/core/tests/core/test_mcp.py | 48 ++++++++++++++++++++ 2 files changed, 65 insertions(+), 1 deletion(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 2a298b0e1c4..04e410c7592 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -1034,6 +1034,7 @@ def __init__( self._lifecycle_owner_task: asyncio.Task[None] | None = None self.session = session self._owns_session = session is None + self._mcp_client: Client | None = None self.request_timeout = request_timeout self.client = client self.sampling_approval_callback = sampling_approval_callback @@ -2018,6 +2019,7 @@ async def _connect_on_owner( await self._cancel_pending_reload_tasks() await self._safe_close_exit_stack() if self._owns_session: + self._mcp_client = None self.session = None self.is_connected = False self._reset_session_state() @@ -2095,6 +2097,7 @@ async def _connect_on_owner( logger.debug(error_msg, exc_info=True) raise ToolException(error_msg, inner_exception=ex if isinstance(ex, Exception) else None) from ex self.session = session + self._mcp_client = mcp_client self._owns_session = True try: await self._listen_capability_list_changes(mcp_client) @@ -2776,6 +2779,7 @@ async def _close_on_owner(self) -> None: await self._safe_close_exit_stack() self._exit_stack = AsyncExitStack() if self._owns_session: + self._mcp_client = None self.session = None self.is_connected = False self._reset_session_state() @@ -2792,6 +2796,14 @@ async def close(self) -> None: async with self._lifecycle_request_lock: await self._run_on_lifecycle_owner("close") + def _operation_client(self) -> Client | ClientSession: + """Return the highest-level MCP client available for standard operations.""" + if self._mcp_client is not None: + return self._mcp_client + if self.session is None: + raise RuntimeError("MCPTool is not connected.") + return self.session + @abstractmethod def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: """Get an MCP client. @@ -2940,7 +2952,11 @@ async def _call_tool_with_retries( for attempt in range(2): try: - result = await self.session.call_tool(tool_name, arguments=filtered_kwargs, meta=meta) # type: ignore + result = await self._operation_client().call_tool( + tool_name, + arguments=filtered_kwargs, + meta=cast("types.RequestParamsMeta | None", meta), + ) _capture_mcp_tool_result(result) if result.is_error: parsed = parser(result) diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 6b332bf1469..d0cb0e41332 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -4941,6 +4941,52 @@ async def test_connect_no_sampling_capabilities_without_client(): await tool.close() +async def test_connect_retains_and_close_clears_sdk_client() -> None: + """Test a framework-owned Client and its session share one lifecycle.""" + tool = MCPStdioTool(name="test", command="test-command", load_tools=False, load_prompts=False) + sdk_client = _mock_sdk_client(protocol_version="2026-07-28") + + with patch("mcp.Client", return_value=sdk_client): + await tool.connect() + try: + assert tool._mcp_client is sdk_client + assert tool.session is sdk_client.session + finally: + await tool.close() + + assert tool._mcp_client is None + assert tool.session is None + + +async def test_call_tool_uses_sdk_client_for_framework_owned_connection() -> None: + """Test a standard tool call uses the retained high-level Client.""" + session = Mock(spec=ClientSession) + session.list_tools = AsyncMock( + return_value=types.ListToolsResult( + result_type="complete", + tools=[types.Tool(name="greet", input_schema={"type": "object", "properties": {}})], + ) + ) + session.call_tool = AsyncMock() + sdk_client = _mock_sdk_client(session=session, protocol_version="2026-07-28") + sdk_client.call_tool = AsyncMock( + return_value=types.CallToolResult( + result_type="complete", + content=[types.TextContent(type="text", text="Hello!")], + is_error=False, + ) + ) + tool = MCPStdioTool(name="test", command="test-command", load_prompts=False) + + with patch("mcp.Client", return_value=sdk_client): + async with tool: + result = await tool.call_tool("greet") + + assert _mcp_result_to_text(result) == "Hello!" + sdk_client.call_tool.assert_awaited_once_with("greet", arguments={}, meta=None) + session.call_tool.assert_not_awaited() + + # Test error handling in connect() method @@ -6837,10 +6883,12 @@ async def call_tool( ) async with wrapper: + assert wrapper._mcp_client is None assert wrapper.session is session assert [function.name for function in wrapper.functions] == ["greet"] assert _mcp_result_to_text(await wrapper.call_tool("greet")) == "Hello!" + assert wrapper._mcp_client is None initialize.assert_not_awaited() discover.assert_not_awaited() if mode == "auto": From 2041a837e4e2b8550fd34681cf55a4545d0774e2 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 7 Oct 2026 10:21:20 +0200 Subject: [PATCH 30/42] Prompt migrated to use mcp_client -> session --- python/packages/core/agent_framework/_mcp.py | 2 +- python/packages/core/tests/core/test_mcp.py | 32 ++++++++++++++++++++ 2 files changed, 33 insertions(+), 1 deletion(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 04e410c7592..279c052f5be 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -3520,7 +3520,7 @@ async def get_prompt(self, prompt_name: str, **kwargs: Any) -> str: with create_mcp_client_span("prompts/get", target=prompt_name, attributes=mcp_span_attrs) as span: for attempt in range(2): try: - prompt_result = await self.session.get_prompt(prompt_name, arguments=kwargs) # type: ignore + prompt_result = await self._operation_client().get_prompt(prompt_name, arguments=kwargs) return parser(prompt_result) except ClosedResourceError as cl_ex: if attempt == 0: diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index d0cb0e41332..6f4506517b0 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -4987,6 +4987,38 @@ async def test_call_tool_uses_sdk_client_for_framework_owned_connection() -> Non session.call_tool.assert_not_awaited() +async def test_get_prompt_uses_sdk_client_for_framework_owned_connection() -> None: + """Test a standard prompt get uses the retained high-level Client.""" + session = Mock(spec=ClientSession) + session.list_prompts = AsyncMock( + return_value=types.ListPromptsResult( + result_type="complete", + prompts=[types.Prompt(name="summarize", arguments=[])], + ) + ) + session.get_prompt = AsyncMock() + sdk_client = _mock_sdk_client(session=session, protocol_version="2026-07-28") + sdk_client.get_prompt = AsyncMock( + return_value=types.GetPromptResult( + messages=[ + types.PromptMessage( + role="user", + content=types.TextContent(type="text", text="Summarize this."), + ) + ] + ) + ) + tool = MCPStdioTool(name="test", command="test-command", load_tools=False) + + with patch("mcp.Client", return_value=sdk_client): + async with tool: + result = await tool.get_prompt("summarize") + + assert "Summarize this." in result + sdk_client.get_prompt.assert_awaited_once_with("summarize", arguments={}) + session.get_prompt.assert_not_awaited() + + # Test error handling in connect() method From 202be25fb84d744281ce47ed4544b25c512f9828 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 7 Oct 2026 10:53:22 +0200 Subject: [PATCH 31/42] Migrated list/tools and list/prompts --- python/packages/core/agent_framework/_mcp.py | 26 ++++++- python/packages/core/tests/core/test_mcp.py | 80 ++++++++++++++++++++ 2 files changed, 104 insertions(+), 2 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 279c052f5be..fddeaa4aecf 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -2527,7 +2527,7 @@ async def _load_prompts_locked(self) -> None: ) return with create_mcp_client_span("prompts/list", attributes=self._mcp_base_span_attributes()): - prompt_list = await self.session.list_prompts(params=params) # type: ignore[union-attr] + prompt_list = await self._list_prompts_page(params) break except ClosedResourceError as cl_ex: if attempt == 0: @@ -2640,7 +2640,7 @@ async def _load_tools_locked(self) -> None: logger.debug("Skipping MCP tool loading because the server did not advertise tools support.") return with create_mcp_client_span("tools/list", attributes=self._mcp_base_span_attributes()): - tool_list = await self.session.list_tools(params=params) # type: ignore[union-attr] + tool_list = await self._list_tools_page(params) break except ClosedResourceError as cl_ex: if attempt == 0: @@ -2804,6 +2804,28 @@ def _operation_client(self) -> Client | ClientSession: raise RuntimeError("MCPTool is not connected.") return self.session + async def _list_tools_page(self, params: types.PaginatedRequestParams | None) -> types.ListToolsResult: + """List one tools page without changing existing cache behavior.""" + if self._mcp_client is not None: + return await self._mcp_client.list_tools( + cursor=params.cursor if params is not None else None, + cache_mode="bypass", + ) + if self.session is None: + raise RuntimeError("MCPTool is not connected.") + return await self.session.list_tools(params=params) + + async def _list_prompts_page(self, params: types.PaginatedRequestParams | None) -> types.ListPromptsResult: + """List one prompts page without changing existing cache behavior.""" + if self._mcp_client is not None: + return await self._mcp_client.list_prompts( + cursor=params.cursor if params is not None else None, + cache_mode="bypass", + ) + if self.session is None: + raise RuntimeError("MCPTool is not connected.") + return await self.session.list_prompts(params=params) + @abstractmethod def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: """Get an MCP client. diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 6f4506517b0..4e99d3f2329 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -105,6 +105,29 @@ def _mock_sdk_client( client.server_capabilities = capabilities client.__aenter__.return_value = client + async def list_tools( + *, + cursor: str | None = None, + meta: types.RequestParamsMeta | None = None, + cache_mode: str = "use", + ) -> types.ListToolsResult: + del meta, cache_mode + params = types.PaginatedRequestParams(cursor=cursor) if cursor is not None else None + return await session.list_tools(params=params) + + async def list_prompts( + *, + cursor: str | None = None, + meta: types.RequestParamsMeta | None = None, + cache_mode: str = "use", + ) -> types.ListPromptsResult: + del meta, cache_mode + params = types.PaginatedRequestParams(cursor=cursor) if cursor is not None else None + return await session.list_prompts(params=params) + + client.list_tools = AsyncMock(side_effect=list_tools) + client.list_prompts = AsyncMock(side_effect=list_prompts) + return client @@ -5019,6 +5042,63 @@ async def test_get_prompt_uses_sdk_client_for_framework_owned_connection() -> No session.get_prompt.assert_not_awaited() +async def test_catalog_loading_uses_sdk_client_without_cache() -> None: + """Test framework-owned catalog pagination uses the Client without caching.""" + capabilities = types.ServerCapabilities( + tools=types.ToolsCapability(), + prompts=types.PromptsCapability(), + ) + session = Mock(spec=ClientSession) + session.list_tools = AsyncMock() + session.list_prompts = AsyncMock() + sdk_client = _mock_sdk_client( + session=session, + capabilities=capabilities, + protocol_version="2026-07-28", + ) + sdk_client.list_tools = AsyncMock( + side_effect=[ + types.ListToolsResult( + tools=[types.Tool(name="first_tool", input_schema={"type": "object", "properties": {}})], + next_cursor="tools-next", + ), + types.ListToolsResult( + tools=[types.Tool(name="second_tool", input_schema={"type": "object", "properties": {}})] + ), + ] + ) + sdk_client.list_prompts = AsyncMock( + side_effect=[ + types.ListPromptsResult( + prompts=[types.Prompt(name="first_prompt", arguments=[])], + next_cursor="prompts-next", + ), + types.ListPromptsResult(prompts=[types.Prompt(name="second_prompt", arguments=[])]), + ] + ) + tool = MCPStdioTool(name="test", command="test-command") + + with patch("mcp.Client", return_value=sdk_client): + async with tool: + assert [function.name for function in tool.functions] == [ + "first_tool", + "second_tool", + "first_prompt", + "second_prompt", + ] + + assert [awaited.kwargs for awaited in sdk_client.list_tools.await_args_list] == [ + {"cursor": None, "cache_mode": "bypass"}, + {"cursor": "tools-next", "cache_mode": "bypass"}, + ] + assert [awaited.kwargs for awaited in sdk_client.list_prompts.await_args_list] == [ + {"cursor": None, "cache_mode": "bypass"}, + {"cursor": "prompts-next", "cache_mode": "bypass"}, + ] + session.list_tools.assert_not_awaited() + session.list_prompts.assert_not_awaited() + + # Test error handling in connect() method From b527812f2cc5c7a83a6558867f1fe93ac990e2aa Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 7 Oct 2026 12:32:07 +0200 Subject: [PATCH 32/42] Made protocol seam explicit with _MCPConnection --- python/packages/core/AGENTS.md | 7 +- python/packages/core/agent_framework/_mcp.py | 235 ++++++++++++++---- python/packages/core/tests/core/test_mcp.py | 150 ++++++----- .../core/tests/core/test_mcp_http_auth.py | 19 ++ 4 files changed, 302 insertions(+), 109 deletions(-) diff --git a/python/packages/core/AGENTS.md b/python/packages/core/AGENTS.md index 70b2f8fa7e8..d41b8b1f763 100644 --- a/python/packages/core/AGENTS.md +++ b/python/packages/core/AGENTS.md @@ -207,8 +207,11 @@ The vector store API is experimental under the shared `VECTOR_STORES` feature ID current evidence-backed target for high-level `mcp.Client` versus low-level `ClientSession` ownership during the Python MCP v2 migration; not every target-state item is implemented yet. Keep the ADR aligned when implementation evidence changes this boundary. -- **`MCPTool`** - Base wrapper that currently exposes the MCP `ClientSession` and the remote server's tools as - `FunctionTool`s. Follow ADR 0045 when migrating standard operations to the retained high-level client. +- **`MCPTool` connection state** - The private `_MCPConnection` Protocol is the normalized tools/prompts/catalog + surface. `_ClientMCPConnection` and `_SessionMCPConnection` each own one lifecycle responsibility; `None` means + disconnected, and the Client implementation is the single update point for SDK cache policy. `MCPTool.session` + remains a compatibility view/setter; constructor or direct assignment of a session selects the caller-owned + low-level path. - **`MCPStdioTool`** / **`MCPStreamableHTTPTool`** - Supported transport-specific subclasses. **`MCPWebsocketTool`** remains only as a deprecated compatibility symbol because MCP v2 removed WebSocket transport; it cannot create a connection. diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index fddeaa4aecf..f177e5abcdf 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -19,7 +19,7 @@ from datetime import date, datetime, timedelta from http.cookiejar import CookieJar, DefaultCookiePolicy from inspect import isawaitable -from typing import TYPE_CHECKING, Any, Literal, TypeAlias, TypedDict, cast +from typing import TYPE_CHECKING, Any, Literal, Protocol, TypeAlias, TypedDict, cast if sys.version_info >= (3, 13): from warnings import deprecated # pragma: no cover @@ -482,6 +482,147 @@ async def delete(self, *args: Any, **kwargs: Any) -> Any: return await self._client.delete(*args, **self._tagged_kwargs(kwargs)) +class _MCPConnection(Protocol): + """Normalized MCP connection surface used by MCPTool.""" + + @property + def session(self) -> ClientSession: + """Return the low-level compatibility session.""" + ... + + @property + def client(self) -> Client | None: + """Return the high-level Client when this connection has one.""" + ... + + @property + def is_framework_owned(self) -> bool: + """Return whether Agent Framework owns this connection.""" + ... + + async def call_tool( + self, + name: str, + arguments: dict[str, Any] | None, + *, + meta: dict[str, Any] | None, + ) -> types.CallToolResult: + """Call a tool through this connection.""" + ... + + async def get_prompt(self, name: str, arguments: dict[str, Any] | None) -> types.GetPromptResult: + """Get a prompt through this connection.""" + ... + + async def list_tools_page(self, params: types.PaginatedRequestParams | None) -> types.ListToolsResult: + """List one tools page through this connection.""" + ... + + async def list_prompts_page(self, params: types.PaginatedRequestParams | None) -> types.ListPromptsResult: + """List one prompts page through this connection.""" + ... + + async def set_logging_level(self, level: Any) -> None: + """Set the legacy logging level through this connection.""" + ... + + +@dataclass(frozen=True) +class _ClientMCPConnection: + """MCP connection backed by a high-level Client.""" + + client: Client + + @property + def session(self) -> ClientSession: + """Return the Client's low-level compatibility session.""" + return self.client.session + + @property + def is_framework_owned(self) -> bool: + """Return True because Agent Framework owns the Client.""" + return True + + async def call_tool( + self, + name: str, + arguments: dict[str, Any] | None, + *, + meta: dict[str, Any] | None, + ) -> types.CallToolResult: + """Call a tool through the high-level Client.""" + return await self.client.call_tool( + name, + arguments=arguments, + meta=cast("types.RequestParamsMeta | None", meta), + ) + + async def get_prompt(self, name: str, arguments: dict[str, Any] | None) -> types.GetPromptResult: + """Get a prompt through the high-level Client.""" + return await self.client.get_prompt(name, arguments=cast("dict[str, str] | None", arguments)) + + async def list_tools_page(self, params: types.PaginatedRequestParams | None) -> types.ListToolsResult: + """List one tools page without changing existing cache behavior.""" + return await self.client.list_tools( + cursor=params.cursor if params is not None else None, + cache_mode="bypass", + ) + + async def list_prompts_page(self, params: types.PaginatedRequestParams | None) -> types.ListPromptsResult: + """List one prompts page without changing existing cache behavior.""" + return await self.client.list_prompts( + cursor=params.cursor if params is not None else None, + cache_mode="bypass", + ) + + async def set_logging_level(self, level: Any) -> None: + """Set the legacy logging level through the underlying session.""" + await self.session.set_logging_level(level) # pyright: ignore[reportDeprecated] + + +@dataclass(frozen=True) +class _SessionMCPConnection: + """MCP connection backed by a low-level ClientSession.""" + + session: ClientSession + client: None = None + + @property + def is_framework_owned(self) -> bool: + """Return False because the caller owns the session.""" + return False + + async def call_tool( + self, + name: str, + arguments: dict[str, Any] | None, + *, + meta: dict[str, Any] | None, + ) -> types.CallToolResult: + """Call a tool directly through the session.""" + return await self.session.call_tool( + name, + arguments=arguments, + meta=cast("types.RequestParamsMeta | None", meta), + ) + + async def get_prompt(self, name: str, arguments: dict[str, Any] | None) -> types.GetPromptResult: + """Get a prompt directly through the session.""" + return await self.session.get_prompt(name, arguments=cast("dict[str, str] | None", arguments)) + + async def list_tools_page(self, params: types.PaginatedRequestParams | None) -> types.ListToolsResult: + """List one tools page directly through the session.""" + return await self.session.list_tools(params=params) + + async def list_prompts_page(self, params: types.PaginatedRequestParams | None) -> types.ListPromptsResult: + """List one prompts page directly through the session.""" + return await self.session.list_prompts(params=params) + + async def set_logging_level(self, level: Any) -> None: + """Set the legacy logging level directly through the session.""" + await self.session.set_logging_level(level) # pyright: ignore[reportDeprecated] + + # Default safety limits applied to server-initiated MCP sampling requests # (``sampling/createMessage``). MCP servers are untrusted third parties, so the # default ``sampling_callback`` denies requests unless an approval callback is @@ -1032,9 +1173,7 @@ def __init__( asyncio.Queue[tuple[str, bool, bool, bool, asyncio.Future[None], asyncio.Future[bool]]] | None ) = None self._lifecycle_owner_task: asyncio.Task[None] | None = None - self.session = session - self._owns_session = session is None - self._mcp_client: Client | None = None + self._connection: _MCPConnection | None = _SessionMCPConnection(session) if session is not None else None self.request_timeout = request_timeout self.client = client self.sampling_approval_callback = sampling_approval_callback @@ -1076,6 +1215,21 @@ def __init__( def __str__(self) -> str: return f"MCPTool(name={self.name}, description={self.description})" + @property + def session(self) -> ClientSession | None: + """Return the low-level session for compatibility and advanced use.""" + return self._connection.session if self._connection is not None else None + + @session.setter + def session(self, value: ClientSession | None) -> None: + """Replace the connection with a caller-owned session compatibility path.""" + connection = self._connection + if connection is not None and connection.is_framework_owned: + if value is connection.session: + return + raise RuntimeError("Cannot replace the session while its framework-owned MCP Client is connected.") + self._connection = _SessionMCPConnection(value) if value is not None else None + def _mcp_base_span_attributes(self) -> dict[str, Any]: """Return base MCP span attributes shared across all operations. @@ -2018,9 +2172,9 @@ async def _connect_on_owner( if reset_discovery: await self._cancel_pending_reload_tasks() await self._safe_close_exit_stack() - if self._owns_session: - self._mcp_client = None - self.session = None + connection = self._connection + if connection is not None and connection.is_framework_owned: + self._connection = None self.is_connected = False self._reset_session_state() if reset_discovery: @@ -2096,9 +2250,7 @@ async def _connect_on_owner( if isinstance(ex, asyncio.CancelledError): logger.debug(error_msg, exc_info=True) raise ToolException(error_msg, inner_exception=ex if isinstance(ex, Exception) else None) from ex - self.session = session - self._mcp_client = mcp_client - self._owns_session = True + self._connection = _ClientMCPConnection(mcp_client) try: await self._listen_capability_list_changes(mcp_client) except (Exception, asyncio.CancelledError): @@ -2139,7 +2291,7 @@ async def _connect_on_owner( level_name = cast( Any, next(level for level, value in LOG_LEVEL_MAPPING.items() if value == logger.level) ) - await self.session.set_logging_level(level_name) + await self._require_connection().set_logging_level(level_name) except Exception as exc: logger.warning("Failed to set log level to %s", logger.level, exc_info=exc) except (Exception, asyncio.CancelledError): @@ -2527,7 +2679,7 @@ async def _load_prompts_locked(self) -> None: ) return with create_mcp_client_span("prompts/list", attributes=self._mcp_base_span_attributes()): - prompt_list = await self._list_prompts_page(params) + prompt_list = await self._require_connection().list_prompts_page(params) break except ClosedResourceError as cl_ex: if attempt == 0: @@ -2640,7 +2792,7 @@ async def _load_tools_locked(self) -> None: logger.debug("Skipping MCP tool loading because the server did not advertise tools support.") return with create_mcp_client_span("tools/list", attributes=self._mcp_base_span_attributes()): - tool_list = await self._list_tools_page(params) + tool_list = await self._require_connection().list_tools_page(params) break except ClosedResourceError as cl_ex: if attempt == 0: @@ -2778,9 +2930,9 @@ async def _close_on_owner(self) -> None: await self._cancel_pending_reload_tasks() await self._safe_close_exit_stack() self._exit_stack = AsyncExitStack() - if self._owns_session: - self._mcp_client = None - self.session = None + connection = self._connection + if connection is not None and connection.is_framework_owned: + self._connection = None self.is_connected = False self._reset_session_state() @@ -2796,35 +2948,11 @@ async def close(self) -> None: async with self._lifecycle_request_lock: await self._run_on_lifecycle_owner("close") - def _operation_client(self) -> Client | ClientSession: - """Return the highest-level MCP client available for standard operations.""" - if self._mcp_client is not None: - return self._mcp_client - if self.session is None: - raise RuntimeError("MCPTool is not connected.") - return self.session - - async def _list_tools_page(self, params: types.PaginatedRequestParams | None) -> types.ListToolsResult: - """List one tools page without changing existing cache behavior.""" - if self._mcp_client is not None: - return await self._mcp_client.list_tools( - cursor=params.cursor if params is not None else None, - cache_mode="bypass", - ) - if self.session is None: + def _require_connection(self) -> _MCPConnection: + """Return the active normalized MCP connection.""" + if self._connection is None: raise RuntimeError("MCPTool is not connected.") - return await self.session.list_tools(params=params) - - async def _list_prompts_page(self, params: types.PaginatedRequestParams | None) -> types.ListPromptsResult: - """List one prompts page without changing existing cache behavior.""" - if self._mcp_client is not None: - return await self._mcp_client.list_prompts( - cursor=params.cursor if params is not None else None, - cache_mode="bypass", - ) - if self.session is None: - raise RuntimeError("MCPTool is not connected.") - return await self.session.list_prompts(params=params) + return self._connection @abstractmethod def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: @@ -2974,10 +3102,10 @@ async def _call_tool_with_retries( for attempt in range(2): try: - result = await self._operation_client().call_tool( + result = await self._require_connection().call_tool( tool_name, - arguments=filtered_kwargs, - meta=cast("types.RequestParamsMeta | None", meta), + filtered_kwargs, + meta=meta, ) _capture_mcp_tool_result(result) if result.is_error: @@ -3542,7 +3670,7 @@ async def get_prompt(self, prompt_name: str, **kwargs: Any) -> str: with create_mcp_client_span("prompts/get", target=prompt_name, attributes=mcp_span_attrs) as span: for attempt in range(2): try: - prompt_result = await self._operation_client().get_prompt(prompt_name, arguments=kwargs) + prompt_result = await self._require_connection().get_prompt(prompt_name, kwargs) return parser(prompt_result) except ClosedResourceError as cl_ex: if attempt == 0: @@ -4090,6 +4218,7 @@ def __init__( # the replacement transport initializes. self._session_headers: dict[str, str] | None = None self._session_header_identity: _MCPHeaderIdentity | None = None + self._session_header_session: ClientSession | None = None self._pending_session_headers: dict[str, str] | None = None self._pending_connection_kwargs: dict[str, Any] | None = None self._call_headers_lock = asyncio.Lock() @@ -4283,6 +4412,7 @@ async def _prepare_for_run(self, kwargs: Mapping[str, Any]) -> None: def _bind_session_headers(self, headers: Mapping[str, str]) -> None: self._session_headers = dict(headers) self._session_header_identity = _mcp_header_identity(headers) + self._session_header_session = self.session def _stage_session_headers(self, headers: Mapping[str, str], kwargs: Mapping[str, Any]) -> None: self._pending_session_headers = dict(headers) @@ -4330,8 +4460,9 @@ async def _ensure_session_identity( kwargs: Mapping[str, Any], ) -> None: identity = _mcp_header_identity(headers) - if not self._owns_session: - if self._session_header_identity is None: + connection = self._connection + if connection is not None and not connection.is_framework_owned: + if self._session_header_identity is None or self._session_header_session is not self.session: raise ToolExecutionException( "MCP header identity is unknown for a caller-supplied session; " "use a separate framework-managed tool instance." @@ -4389,9 +4520,11 @@ def _release_connection_kwargs(self) -> None: self._connection_kwargs = None self._pending_session_headers = None self._pending_connection_kwargs = None - if self._owns_session: + connection = self._connection + if connection is None or connection.is_framework_owned: self._session_headers = None self._session_header_identity = None + self._session_header_session = None async def call_tool(self, tool_name: str, **kwargs: Any) -> str | list[Content]: """Call a tool, injecting headers from the header_provider if configured. diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 4e99d3f2329..53e2baf145d 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -44,6 +44,7 @@ MCPSpecificApproval, MCPTool, _build_prefixed_mcp_name, + _ClientMCPConnection, _describe_error, _get_input_model_from_mcp_prompt, _json_size_exceeds, @@ -2356,24 +2357,6 @@ async def test_local_mcp_server_initialization(): assert server.functions == [] -async def test_local_mcp_server_context_manager(): - """Test MCPTool as context manager.""" - - class TestServer(MCPTool): - async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - # Mock connection - self.session = Mock(spec=ClientSession) - - def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: - return None # type: ignore[return-value] # pyrefly: ignore[bad-return] # ty: ignore[invalid-return-type] - - server = TestServer(name="test_server") - async with server: - assert server.session is not None - - assert server.session is None - - async def test_local_mcp_server_load_functions(): """Test loading functions from MCP server.""" @@ -3571,8 +3554,9 @@ def provider(kwargs: dict[str, Any]) -> dict[str, str]: header_provider=provider, use_progressive_disclosure=True, ) - server.session = AsyncMock() - server.session.list_tools = AsyncMock( + sdk_client = AsyncMock() + sdk_client.session = AsyncMock() + sdk_client.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -3587,9 +3571,10 @@ def provider(kwargs: dict[str, Any]) -> dict[str, str]: ] ) ) - server.session.call_tool = AsyncMock( + sdk_client.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="Hello!")]) ) + server._connection = _ClientMCPConnection(sdk_client) await server.load_tools() load_tool = server.functions[1] load_context = FunctionInvocationContext( @@ -3618,7 +3603,7 @@ def provider(kwargs: dict[str, Any]) -> dict[str, str]: assert result[0].text == "Hello!" assert provider_received[0]["some_token"] == "my-secret" - call_args = server.session.call_tool.call_args # type: ignore[union-attr] + call_args = sdk_client.call_tool.call_args assert call_args.kwargs.get("arguments", {}).get("name") == "Alice" assert "some_token" not in call_args.kwargs.get("arguments", {}) @@ -4972,12 +4957,18 @@ async def test_connect_retains_and_close_clears_sdk_client() -> None: with patch("mcp.Client", return_value=sdk_client): await tool.connect() try: - assert tool._mcp_client is sdk_client + connection = tool._connection + tool.session = tool.session + assert tool._connection is connection + with pytest.raises(RuntimeError, match="framework-owned MCP Client"): + tool.session = None + await tool.connect(reset=True) + assert tool._connection.client is sdk_client assert tool.session is sdk_client.session finally: await tool.close() - assert tool._mcp_client is None + assert tool._connection is None assert tool.session is None @@ -6995,12 +6986,16 @@ async def call_tool( ) async with wrapper: - assert wrapper._mcp_client is None + assert wrapper._connection is not None + assert wrapper._connection.client is None + assert wrapper._connection.is_framework_owned is False assert wrapper.session is session assert [function.name for function in wrapper.functions] == ["greet"] assert _mcp_result_to_text(await wrapper.call_tool("greet")) == "Hello!" - assert wrapper._mcp_client is None + assert wrapper._connection is not None + assert wrapper._connection.client is None + assert wrapper._connection.is_framework_owned is False initialize.assert_not_awaited() discover.assert_not_awaited() if mode == "auto": @@ -7012,6 +7007,36 @@ async def call_tool( assert result.content[0].text == "Hello!" +async def test_replacing_supplied_session_preserves_caller_ownership() -> None: + """Test compatibility assignment does not transfer session ownership.""" + original_session = Mock(spec=ClientSession) + replacement_session = Mock(spec=ClientSession) + replacement_session.protocol_version = "2026-07-28" + replacement_session.server_capabilities = types.ServerCapabilities() + replacement_session.initialize_result = None + tool = MCPStdioTool( + name="test", + command="unused", + session=original_session, + load_tools=False, + load_prompts=False, + ) + + tool.session = None + tool.session = replacement_session + + with patch.object(tool, "get_mcp_client", side_effect=AssertionError("Unexpected transport")): + await tool.connect(reset=True) + assert tool._connection is not None + assert tool._connection.is_framework_owned is False + assert tool.session is replacement_session + await tool.close() + + assert tool._connection is not None + assert tool._connection.is_framework_owned is False + assert tool.session is replacement_session + + async def test_connect_reinitializes_existing_session_and_loads_tools_and_prompts() -> None: session = _mock_unnegotiated_session( types.ServerCapabilities(tools=types.ToolsCapability(), prompts=types.PromptsCapability()) @@ -7614,8 +7639,9 @@ async def test_mcp_streamable_http_tool_header_provider_injects_headers(): class _TestServer(MCPStreamableHTTPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + session = Mock(spec=ClientSession) + sdk_client = _mock_sdk_client(session=session, protocol_version="2026-07-28") + sdk_client.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -7630,10 +7656,10 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + sdk_client.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="Hello!")]) ) - self.session.send_ping = AsyncMock() + self._connection = _ClientMCPConnection(sdk_client) self.is_connected = True def get_mcp_client(self): # pyrefly: ignore[bad-override] @@ -7653,8 +7679,10 @@ def provider(kwargs): # Simulate the runtime kwargs that flow from FunctionInvocationContext.kwargs await server.call_tool("greet", name="Alice", some_token="my-secret") - # Verify the MCP session.call_tool was called - server.session.call_tool.assert_called_once() # type: ignore[union-attr] # ty: ignore[unresolved-attribute] + # Verify the high-level Client.call_tool was called. + sdk_client = server._connection.client + assert sdk_client is not None + cast(AsyncMock, sdk_client.call_tool).assert_awaited_once() async def test_mcp_streamable_http_tool_header_provider_sets_contextvar(): @@ -7674,8 +7702,9 @@ async def spy_call_tool(self, tool_name, **kwargs): class _TestServer(MCPStreamableHTTPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + session = Mock(spec=ClientSession) + sdk_client = _mock_sdk_client(session=session, protocol_version="2026-07-28") + sdk_client.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -7686,10 +7715,10 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + sdk_client.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="Hello!")]) ) - self.session.send_ping = AsyncMock() + self._connection = _ClientMCPConnection(sdk_client) self.is_connected = True def get_mcp_client(self): # pyrefly: ignore[bad-override] @@ -7716,8 +7745,9 @@ async def test_mcp_streamable_http_tool_header_provider_contextvar_reset_after_c class _TestServer(MCPStreamableHTTPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + session = Mock(spec=ClientSession) + sdk_client = _mock_sdk_client(session=session, protocol_version="2026-07-28") + sdk_client.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -7728,10 +7758,10 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + sdk_client.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="Hello!")]) ) - self.session.send_ping = AsyncMock() + self._connection = _ClientMCPConnection(sdk_client) self.is_connected = True def get_mcp_client(self): # pyrefly: ignore[bad-override] @@ -7756,8 +7786,9 @@ async def test_mcp_streamable_http_tool_without_header_provider(): class _TestServer(MCPStreamableHTTPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + session = Mock(spec=ClientSession) + sdk_client = _mock_sdk_client(session=session, protocol_version="2026-07-28") + sdk_client.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -7768,10 +7799,10 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + sdk_client.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="Hello!")]) ) - self.session.send_ping = AsyncMock() + self._connection = _ClientMCPConnection(sdk_client) self.is_connected = True def get_mcp_client(self): # pyrefly: ignore[bad-override] @@ -7784,7 +7815,9 @@ def get_mcp_client(self): # pyrefly: ignore[bad-override] async with server: await server.load_tools() await server.call_tool("greet", name="Alice") - server.session.call_tool.assert_called_once() # type: ignore[union-attr] # ty: ignore[unresolved-attribute] + sdk_client = server._connection.client + assert sdk_client is not None + cast(AsyncMock, sdk_client.call_tool).assert_awaited_once() # Without header_provider, call_tool should delegate directly to MCPTool assert server._header_provider is None @@ -8416,8 +8449,9 @@ async def spy_call_tool(self, tool_name, **kwargs): class _TestServer(MCPStreamableHTTPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + session = Mock(spec=ClientSession) + sdk_client = _mock_sdk_client(session=session, protocol_version="2026-07-28") + sdk_client.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -8432,10 +8466,10 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + sdk_client.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="Hello!")]) ) - self.session.send_ping = AsyncMock() + self._connection = _ClientMCPConnection(sdk_client) self.is_connected = True def get_mcp_client(self): # pyrefly: ignore[bad-override] @@ -8479,9 +8513,12 @@ def provider(kwargs): assert len(provider_received) == 1 assert provider_received[0]["some_token"] == "my-secret" - # Verify session.call_tool was called with the tool arguments (not the runtime kwargs) - server.session.call_tool.assert_called_once() # type: ignore[union-attr] # ty: ignore[unresolved-attribute] - call_args = server.session.call_tool.call_args # type: ignore[union-attr] # ty: ignore[unresolved-attribute] + # Verify Client.call_tool was called with the tool arguments (not the runtime kwargs). + sdk_client = server._connection.client + assert sdk_client is not None + call_tool = cast(AsyncMock, sdk_client.call_tool) + call_tool.assert_awaited_once() + call_args = call_tool.call_args assert call_args.kwargs.get("arguments", {}).get("name") == "Alice" @@ -9109,8 +9146,9 @@ async def spy_call_tool(self, tool_name, **kwargs): class _TestServer(MCPStreamableHTTPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + session = Mock(spec=ClientSession) + sdk_client = _mock_sdk_client(session=session, protocol_version="2026-07-28") + sdk_client.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -9121,10 +9159,10 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + sdk_client.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="Hello!")]) ) - self.session.send_ping = AsyncMock() + self._connection = _ClientMCPConnection(sdk_client) self.is_connected = True def get_mcp_client(self): # pyrefly: ignore[bad-override] diff --git a/python/packages/core/tests/core/test_mcp_http_auth.py b/python/packages/core/tests/core/test_mcp_http_auth.py index b0bc05ce70a..f0d9059e6ed 100644 --- a/python/packages/core/tests/core/test_mcp_http_auth.py +++ b/python/packages/core/tests/core/test_mcp_http_auth.py @@ -13,6 +13,7 @@ import httpx import pytest +from mcp.client.session import ClientSession from agent_framework import FunctionInvocationContext, MCPStreamableHTTPTool from agent_framework.exceptions import ToolException, ToolExecutionException @@ -1204,6 +1205,24 @@ async def test_caller_supplied_session_rejects_header_identity_changes(mcp_http_ await supplied_session.send_ping() +async def test_replacing_supplied_session_invalidates_bound_header_identity() -> None: + """A recorded header identity is valid only for the session that established it.""" + original_session = Mock(spec=ClientSession) + replacement_session = Mock(spec=ClientSession) + tool = MCPStreamableHTTPTool( + name="borrowed", + url="https://must-not-connect.example/mcp", + session=original_session, + header_provider=lambda _kwargs: {"Authorization": "A"}, + ) + tool._bind_session_headers({"Authorization": "A"}) + + tool.session = replacement_session + + with pytest.raises(ToolExecutionException, match="identity is unknown"): + await tool._ensure_session_identity({"Authorization": "A"}, {}) + + async def test_cancelled_redundant_connect_keeps_existing_session(mcp_http_server: MCPHTTPServer) -> None: client, _, _ = mcp_http_server async with _tool(client, "token-a") as tool: From aad4ace9fdec73e76e5de2a5b34b96a9d6cc9355 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 7 Oct 2026 13:53:48 +0200 Subject: [PATCH 33/42] Migrated skills --- python/packages/core/agent_framework/_mcp.py | 22 +++++++++ .../packages/core/agent_framework/_skills.py | 48 +++++++++++-------- 2 files changed, 50 insertions(+), 20 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index f177e5abcdf..74b2d2aad92 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -526,6 +526,10 @@ async def set_logging_level(self, level: Any) -> None: """Set the legacy logging level through this connection.""" ... + async def read_resource(self, uri: str) -> types.ReadResourceResult: + """Read a resource through this connection.""" + ... + @dataclass(frozen=True) class _ClientMCPConnection: @@ -579,6 +583,10 @@ async def set_logging_level(self, level: Any) -> None: """Set the legacy logging level through the underlying session.""" await self.session.set_logging_level(level) # pyright: ignore[reportDeprecated] + async def read_resource(self, uri: str) -> types.ReadResourceResult: + """Read a resource through the high-level Client.""" + return await self.client.read_resource(uri, cache_mode="bypass") + @dataclass(frozen=True) class _SessionMCPConnection: @@ -622,6 +630,20 @@ async def set_logging_level(self, level: Any) -> None: """Set the legacy logging level directly through the session.""" await self.session.set_logging_level(level) # pyright: ignore[reportDeprecated] + async def read_resource(self, uri: str) -> types.ReadResourceResult: + """Read a resource through this connection.""" + return await self.session.read_resource(uri) + + +def _as_mcp_connection( # pyright: ignore[reportUnusedFunction] + client: Client | ClientSession, +) -> _MCPConnection: + from mcp import Client as MCPClient + + if isinstance(client, MCPClient): + return _ClientMCPConnection(client) + return _SessionMCPConnection(client) + # Default safety limits applied to server-initiated MCP sampling requests # (``sampling/createMessage``). MCP servers are untrusted third parties, so the diff --git a/python/packages/core/agent_framework/_skills.py b/python/packages/core/agent_framework/_skills.py index 090768868e3..aeb20a4ab1e 100644 --- a/python/packages/core/agent_framework/_skills.py +++ b/python/packages/core/agent_framework/_skills.py @@ -88,12 +88,14 @@ from ._feature_stage import ExperimentalFeature, experimental from ._filesystem import _is_link_or_reparse_point # pyright: ignore[reportPrivateUsage] +from ._mcp import _as_mcp_connection, _MCPConnection # pyright: ignore[reportPrivateUsage] from ._middleware import FunctionInvocationContext from ._sessions import ContextProvider from ._telemetry import FeatureIndex, mark_feature_used from ._tools import ApprovalMode, FunctionTool if TYPE_CHECKING: + from mcp import Client from mcp.client.session import ClientSession from mcp.types import ReadResourceResult @@ -4589,11 +4591,11 @@ def _parse_mcp_skill_index(text: str) -> _McpSkillIndex: return _McpSkillIndex(schema=raw.get("$schema"), skills=entries) -def _resolve_mcp_session_provider( - client: ClientSession | None, - session_provider: Callable[[], ClientSession] | None, -) -> Callable[[], ClientSession]: - """Normalize the two MCP session inputs into a single session resolver. +def _resolve_mcp_connection_provider( + client: Client | ClientSession | None, + session_provider: Callable[[], Client | ClientSession] | None, +) -> Callable[[], _MCPConnection]: + """Normalize the two MCP connection inputs into a single session resolver. Callers supply **exactly one** of a fixed ``client`` or a ``session_provider`` callable. A fixed client is wrapped in a provider that @@ -4602,24 +4604,26 @@ def _resolve_mcp_session_provider( over time, e.g. a reconnecting :class:`~agent_framework.MCPTool`). Args: - client: A fixed MCP client session, or ``None``. + client: A fixed MCP client or session, or ``None``. session_provider: A callable returning the current MCP client session, or ``None``. Returns: - A callable that returns the MCP client session to use. + A callable that returns the MCP client or session to use. Raises: ValueError: If both or neither of *client* and *session_provider* are provided. """ - if client is not None and session_provider is not None: - raise ValueError("Provide exactly one of 'client' or 'session_provider', not both.") if session_provider is not None: - return session_provider + if client is not None: + raise ValueError("Provide exactly one of 'client' or 'session_provider'.") + return lambda: _as_mcp_connection(session_provider()) + if client is None: raise ValueError("Provide exactly one of 'client' or 'session_provider'.") - fixed: ClientSession = client + + fixed = _as_mcp_connection(client) return lambda: fixed @@ -4712,9 +4716,9 @@ def __init__( self, frontmatter: SkillFrontmatter, skill_md_uri: str, - client: ClientSession | None = None, + client: Client | ClientSession | None = None, *, - session_provider: Callable[[], ClientSession] | None = None, + session_provider: Callable[[], Client | ClientSession] | None = None, ) -> None: """Initialize an MCPSkill. @@ -4744,7 +4748,7 @@ def __init__( self._frontmatter = frontmatter self._skill_md_uri = skill_md_uri self._skill_root_uri = self._compute_skill_root_uri(skill_md_uri) - self._session_provider = _resolve_mcp_session_provider(client, session_provider) + self._session_provider = _resolve_mcp_connection_provider(client, session_provider) self._content: str | None = None @property @@ -5082,7 +5086,7 @@ class _ArchiveEntryLoader: def __init__( self, - session_provider: Callable[[], ClientSession], + session_provider: Callable[[], _MCPConnection], *, resource_extensions: tuple[str, ...] | None, resource_search_depth: int, @@ -5415,9 +5419,9 @@ class MCPSkillsSource(SkillsSource): def __init__( self, - client: ClientSession | None = None, + client: Client | ClientSession | None = None, *, - session_provider: Callable[[], ClientSession] | None = None, + session_provider: Callable[[], Client | ClientSession] | None = None, archive_resource_extensions: tuple[str, ...] | None = None, archive_resource_search_depth: int = DEFAULT_SEARCH_DEPTH, archive_max_file_count: int = _DEFAULT_ARCHIVE_MAX_FILE_COUNT, @@ -5429,7 +5433,7 @@ def __init__( Provide **exactly one** of *client* or *session_provider*. Args: - client: A fixed MCP client session connected to a server that exposes + client: A fixed MCP client or clientsession connected to a server that exposes Agent Skills resources. Use this when the session outlives the source (e.g. a caller-owned long-lived session). @@ -5463,7 +5467,7 @@ def __init__( ValueError: If both or neither of *client* and *session_provider* are provided. """ - self._session_provider = _resolve_mcp_session_provider(client, session_provider) + self._session_provider = _resolve_mcp_connection_provider(client, session_provider) self._archive_loader = _ArchiveEntryLoader( self._session_provider, resource_extensions=archive_resource_extensions, @@ -5554,6 +5558,10 @@ async def _try_read_index(self) -> _McpSkillIndex | None: logger.warning("Failed to parse skill://index.json JSON document.", exc_info=True) return None + def _current_client(self) -> Client | ClientSession: + connection = self._session_provider() + return connection.client if connection.client is not None else connection.session + def _try_create_skill(self, entry: _McpSkillIndexEntry) -> MCPSkill | None: """Attempt to create an :class:`MCPSkill` from a ``skill-md`` index entry. @@ -5585,7 +5593,7 @@ def _try_create_skill(self, entry: _McpSkillIndexEntry) -> MCPSkill | None: logger.debug("Skipping entry '%s': invalid metadata: %s", entry.name, ex) return None - return MCPSkill(frontmatter=fm, skill_md_uri=entry.url, session_provider=self._session_provider) + return MCPSkill(frontmatter=fm, skill_md_uri=entry.url, session_provider=self._current_client) # endregion From 236c4a72803884b9b67bed32c5ec9c3d219b90a0 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 7 Oct 2026 14:25:50 +0200 Subject: [PATCH 34/42] Toolbox migration --- .../packages/core/agent_framework/_skills.py | 5 ++- .../core/tests/core/test_mcp_skills.py | 34 ++++++++++++++++++- .../foundry_hosting/tests/test_toolbox.py | 32 +++++++++++++++++ 3 files changed, 67 insertions(+), 4 deletions(-) diff --git a/python/packages/core/agent_framework/_skills.py b/python/packages/core/agent_framework/_skills.py index aeb20a4ab1e..d8560b2735e 100644 --- a/python/packages/core/agent_framework/_skills.py +++ b/python/packages/core/agent_framework/_skills.py @@ -95,8 +95,7 @@ from ._tools import ApprovalMode, FunctionTool if TYPE_CHECKING: - from mcp import Client - from mcp.client.session import ClientSession + from mcp import Client, ClientSession from mcp.types import ReadResourceResult from ._agents import SupportsAgentRun @@ -5405,7 +5404,7 @@ class MCPSkillsSource(SkillsSource): Examples: .. code-block:: python - from mcp.client.session import ClientSession + from mcp import Client, ClientSession source = MCPSkillsSource(client=session) # `context` is normally supplied by SkillsProvider at runtime. diff --git a/python/packages/core/tests/core/test_mcp_skills.py b/python/packages/core/tests/core/test_mcp_skills.py index 9d172b5eff5..e2e35618392 100644 --- a/python/packages/core/tests/core/test_mcp_skills.py +++ b/python/packages/core/tests/core/test_mcp_skills.py @@ -12,7 +12,8 @@ import zipfile from collections.abc import Callable, Iterator, Mapping, Sequence from datetime import timedelta -from unittest.mock import AsyncMock, patch +from typing import Any +from unittest.mock import AsyncMock, call, patch from urllib.parse import unquote import pytest @@ -572,6 +573,37 @@ def test_requires_exactly_one_of_client_or_session_provider(self) -> None: class TestMCPSkillsSource: """Tests for MCPSkillsSource.""" + async def test_high_level_client_reads_resources_with_cache_bypass(self) -> None: + from mcp import Client + from mcp.server import Server, ServerRequestContext + from mcp.types import ReadResourceRequestParams + + responses = { + "skill://index.json": _make_text_result(SAMPLE_SKILL_INDEX, uri="skill://index.json"), + "skill://unit-converter/SKILL.md": _make_text_result(SAMPLE_SKILL_MD), + } + + async def read_resource( + _ctx: ServerRequestContext[Any], + params: ReadResourceRequestParams, + ) -> ReadResourceResult: + return responses[str(params.uri)] + + server = Server("skills-server", on_read_resource=read_resource) + + async with Client(server) as client: + read_resource_mock = AsyncMock(wraps=client.read_resource) + with patch.object(client, "read_resource", read_resource_mock): + source = MCPSkillsSource(client=client) + skill = (await source.get_skills(_SOURCE_CTX))[0] + content = await skill.get_content() + + assert content == SAMPLE_SKILL_MD + assert read_resource_mock.await_args_list == [ + call("skill://index.json", cache_mode="bypass"), + call("skill://unit-converter/SKILL.md", cache_mode="bypass"), + ] + @pytest.mark.parametrize( "uri", [ diff --git a/python/packages/foundry_hosting/tests/test_toolbox.py b/python/packages/foundry_hosting/tests/test_toolbox.py index 1a4b586ce0b..90430c6e0a1 100644 --- a/python/packages/foundry_hosting/tests/test_toolbox.py +++ b/python/packages/foundry_hosting/tests/test_toolbox.py @@ -669,6 +669,38 @@ async def get_skills(self, context: SkillsSourceContext) -> list[str]: assert captured_kwargs == {} +async def test_skills_source_uses_current_high_level_client(monkeypatch: pytest.MonkeyPatch) -> None: + toolbox = FoundryToolbox( + _FakeCredential(), # type: ignore + url="https://h/toolboxes/tb/mcp", + ) + sentinel_client = object() + sentinel_session = object() + current = {"connection": SimpleNamespace(client=sentinel_client, session=sentinel_session)} + monkeypatch.setattr(toolbox, "_require_connection", lambda: current["connection"]) + + captured: dict[str, Callable[[], object]] = {} + + class _StubSkillsSource: + def __init__(self, *, session_provider: Callable[[], object]) -> None: + captured["session_provider"] = session_provider + + async def get_skills(self, context: SkillsSourceContext) -> list[str]: + return ["skill-a"] + + monkeypatch.setattr("agent_framework_foundry_hosting._toolbox.MCPSkillsSource", _StubSkillsSource) + + result = await _FoundryToolboxSkillsSource(toolbox).get_skills(_source_context()) + + assert result == ["skill-a"] + provider = captured["session_provider"] + assert provider() is sentinel_client + + replacement_client = object() + current["connection"] = SimpleNamespace(client=replacement_client, session=object()) + assert provider() is replacement_client + + async def test_skills_source_forwards_archive_options(monkeypatch: pytest.MonkeyPatch) -> None: toolbox = FoundryToolbox( _FakeCredential(), # type: ignore From 9fd7bc762d766fc7ee6e47393be45b069424c753 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 7 Oct 2026 14:33:56 +0200 Subject: [PATCH 35/42] using right method for security labels --- .../packages/core/agent_framework/security.py | 12 ++-- python/packages/core/tests/test_security.py | 55 +++++++++++++------ 2 files changed, 44 insertions(+), 23 deletions(-) diff --git a/python/packages/core/agent_framework/security.py b/python/packages/core/agent_framework/security.py index 1eabc0c2b16..50c981e1e68 100644 --- a/python/packages/core/agent_framework/security.py +++ b/python/packages/core/agent_framework/security.py @@ -4347,7 +4347,7 @@ def _map_mcp_annotations_to_labels( if annotations is None: return (default_integrity, ConfidentialityLabel.PUBLIC, False) - open_world: bool | None = getattr(annotations, "openWorldHint", None) + open_world: bool | None = getattr(annotations, "open_world_hint", None) integrity = default_integrity if open_world is True: integrity = IntegrityLabel.UNTRUSTED @@ -4467,21 +4467,19 @@ async def apply_mcp_security_labels( "MCPTool is not connected. Call connect() or use 'async with' before applying security labels." ) - session = getattr(mcp_tool, "session", None) - if session is None: - raise RuntimeError("MCPTool has no active session.") + connection = mcp_tool._require_connection() from mcp import types as mcp_types annotation_map: dict[str, Any] = {} params: mcp_types.PaginatedRequestParams | None = None while True: - tool_list = await session.list_tools(params=params) + tool_list = await connection.list_tools_page(params) for remote_tool in tool_list.tools: annotation_map[remote_tool.name] = remote_tool.annotations - if not tool_list.nextCursor: + if not tool_list.next_cursor: break - params = mcp_types.PaginatedRequestParams(cursor=tool_list.nextCursor) + params = mcp_types.PaginatedRequestParams(cursor=tool_list.next_cursor) loaded_functions = getattr(mcp_tool, "_functions", None) if not isinstance(loaded_functions, list): diff --git a/python/packages/core/tests/test_security.py b/python/packages/core/tests/test_security.py index 41c8656799c..8aaf253a36b 100644 --- a/python/packages/core/tests/test_security.py +++ b/python/packages/core/tests/test_security.py @@ -10,7 +10,7 @@ from datetime import timedelta from types import MappingProxyType, SimpleNamespace from typing import Annotated, Any, cast -from unittest.mock import AsyncMock, Mock +from unittest.mock import AsyncMock, Mock, call import pytest from pydantic import AfterValidator, BaseModel, field_validator @@ -6303,7 +6303,7 @@ def test_map_mcp_annotations_to_labels( annotations = None if read_only is not None or open_world is not None: - annotations = SimpleNamespace(readOnlyHint=read_only, openWorldHint=open_world) + annotations = SimpleNamespace(read_only_hint=read_only, open_world_hint=open_world) integrity, max_conf, accepts_untrusted = _map_mcp_annotations_to_labels( annotations, @@ -6349,7 +6349,7 @@ async def fake_call(**kwargs: Any) -> list[Content]: mcp_tool.session.list_tools = AsyncMock( # type: ignore[method-assign] return_value=SimpleNamespace( tools=[SimpleNamespace(name="remote_tool", annotations=annotations)], - nextCursor=None, + next_cursor=None, ) ) mcp_tool.functions.append(function) @@ -6362,8 +6362,8 @@ def _make_mcp_tool_definition(name: str, *, open_world: bool = False) -> Any: return mcp_types.Tool( name=name, description=f"{name} description", - inputSchema={"type": "object", "properties": {}}, - annotations=mcp_types.ToolAnnotations(readOnlyHint=False, openWorldHint=open_world), + input_schema={"type": "object", "properties": {}}, + annotations=mcp_types.ToolAnnotations(read_only_hint=False, open_world_hint=open_world), ) @@ -6775,8 +6775,8 @@ async def fake_call(**kwargs: Any) -> list[Content]: async def test_wrap_mcp_function_reads_refreshed_local_policy_at_invocation(self): from agent_framework.security import SecureMCPToolProxy - trusted_annotations = SimpleNamespace(readOnlyHint=True, openWorldHint=False) - untrusted_annotations = SimpleNamespace(readOnlyHint=True, openWorldHint=True) + trusted_annotations = SimpleNamespace(read_only_hint=True, open_world_hint=False) + untrusted_annotations = SimpleNamespace(read_only_hint=True, open_world_hint=True) server_meta = {"ifc": {"integrity": "trusted", "confidentiality": "public"}} mcp_tool, function = _make_connected_mcp_tool_for_ifc( annotations=trusted_annotations, @@ -6785,10 +6785,10 @@ async def test_wrap_mcp_function_reads_refreshed_local_policy_at_invocation(self ) mcp_tool.session.list_tools.side_effect = [ SimpleNamespace( - tools=[SimpleNamespace(name="remote_tool", annotations=trusted_annotations)], nextCursor=None + tools=[SimpleNamespace(name="remote_tool", annotations=trusted_annotations)], next_cursor=None ), SimpleNamespace( - tools=[SimpleNamespace(name="remote_tool", annotations=untrusted_annotations)], nextCursor=None + tools=[SimpleNamespace(name="remote_tool", annotations=untrusted_annotations)], next_cursor=None ), ] proxy = SecureMCPToolProxy(mcp_tool, default_integrity=IntegrityLabel.TRUSTED) @@ -6888,7 +6888,7 @@ async def fake_call(**kwargs: Any) -> list[Content]: async def test_apply_mcp_security_labels_configures_result_authority(self, trust_server_ifc: bool): from agent_framework.security import apply_mcp_security_labels - annotations = SimpleNamespace(readOnlyHint=True, openWorldHint=False) + annotations = SimpleNamespace(read_only_hint=True, open_world_hint=False) server_meta = { "ifc": {"integrity": "trusted", "confidentiality": "public"}, "_mcp_trust_server_ifc": True, @@ -6915,10 +6915,35 @@ async def test_apply_mcp_security_labels_configures_result_authority(self, trust ) assert result[0].additional_properties["security_label"] == expected_label + async def test_apply_mcp_security_labels_uses_high_level_client_connection(self) -> None: + from agent_framework._mcp import _ClientMCPConnection + from agent_framework.security import apply_mcp_security_labels + + annotations = SimpleNamespace(read_only_hint=True, open_world_hint=False) + mcp_tool, _ = _make_connected_mcp_tool_for_ifc(annotations=annotations, server_meta={}) + sdk_client = AsyncMock() + sdk_client.session = AsyncMock() + sdk_client.list_tools.side_effect = [ + SimpleNamespace( + tools=[SimpleNamespace(name="remote_tool", annotations=annotations)], + next_cursor="next-page", + ), + SimpleNamespace(tools=[], next_cursor=None), + ] + mcp_tool._connection = _ClientMCPConnection(sdk_client) + + await apply_mcp_security_labels(mcp_tool) + + assert sdk_client.list_tools.await_args_list == [ + call(cursor=None, cache_mode="bypass"), + call(cursor="next-page", cache_mode="bypass"), + ] + sdk_client.session.list_tools.assert_not_awaited() + async def test_framework_stamped_mcp_label_remains_authoritative_through_tracking(self) -> None: from agent_framework.security import apply_mcp_security_labels - annotations = SimpleNamespace(readOnlyHint=True, openWorldHint=False) + annotations = SimpleNamespace(read_only_hint=True, open_world_hint=False) server_meta = {"ifc": {"integrity": "trusted", "confidentiality": "public"}} mcp_tool, function = _make_connected_mcp_tool_for_ifc( annotations=annotations, @@ -6943,7 +6968,7 @@ async def next_fn() -> None: async def test_apply_mcp_security_labels_reconfigures_existing_wrapper_authority(self): from agent_framework.security import apply_mcp_security_labels - annotations = SimpleNamespace(readOnlyHint=True, openWorldHint=False) + annotations = SimpleNamespace(read_only_hint=True, open_world_hint=False) server_meta = {"ifc": {"integrity": "trusted", "confidentiality": "public"}} mcp_tool, function = _make_connected_mcp_tool_for_ifc(annotations=annotations, server_meta=server_meta) await apply_mcp_security_labels(mcp_tool) @@ -6967,7 +6992,7 @@ async def test_apply_mcp_security_labels_reconfigures_existing_wrapper_authority async def test_secure_mcp_proxy_configures_result_authority(self, trust_server_ifc: bool): from agent_framework.security import SecureMCPToolProxy - annotations = SimpleNamespace(readOnlyHint=True, openWorldHint=False) + annotations = SimpleNamespace(read_only_hint=True, open_world_hint=False) server_meta = {"ifc": {"integrity": "trusted", "confidentiality": "public"}} mcp_tool, function = _make_connected_mcp_tool_for_ifc(annotations=annotations, server_meta=server_meta) proxy = ( @@ -7010,9 +7035,7 @@ async def test_secure_mcp_proxy_labels_notification_reload_before_publication(se reloaded_initial_tool = _make_mcp_tool_definition("initial_sink", open_world=True) mcp_tool.session.list_tools.return_value = mcp_types.ListToolsResult(tools=[reloaded_initial_tool, late_tool]) - notification = Mock(spec=mcp_types.ServerNotification) - notification.root = Mock() - notification.root.method = "notifications/tools/list_changed" + notification = mcp_types.ToolListChangedNotification() await mcp_tool.message_handler(notification) pending_reloads = list(mcp_tool._pending_reload_tasks) From 3e98b01b42d392d6087746cb8e4ccb4f6da49dd6 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 7 Oct 2026 15:15:48 +0200 Subject: [PATCH 36/42] legacy and modern MRTR validated --- python/packages/core/agent_framework/_mcp.py | 2 - python/packages/core/tests/core/test_mcp.py | 365 +++++++++++++++++- .../core/tests/core/test_mcp_skills.py | 35 ++ 3 files changed, 395 insertions(+), 7 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 74b2d2aad92..48df868caf3 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -2218,11 +2218,9 @@ async def _connect_on_owner( sampling_capabilities = types.SamplingCapability( tools=types.SamplingToolsCapability(), ) - client_mode = "legacy" if self.client is not None else "auto" mcp_client = await self._exit_stack.enter_async_context( Client( server=self.get_mcp_client(), - mode=client_mode, read_timeout_seconds=( timedelta(seconds=self.request_timeout).seconds if self.request_timeout else None ), diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 53e2baf145d..875eccd44ef 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -4914,17 +4914,17 @@ async def test_mcp_tool_sampling_callback_always_passes_max_tokens(): async def test_connect_sampling_capabilities_with_client(): - """Test connect() uses legacy mode and advertises sampling when a chat client is configured.""" + """Test connect() uses the SDK's auto mode and advertises sampling when a chat client is configured.""" tool = MCPStdioTool(name="test", command="test-command", load_tools=False, load_prompts=False) tool.client = Mock() with patch("mcp.Client") as mock_client_class: - sdk_client = _mock_sdk_client() + sdk_client = _mock_sdk_client(protocol_version="2026-07-28") mock_client_class.return_value = sdk_client async with tool: call_kwargs = mock_client_class.call_args.kwargs - assert call_kwargs["mode"] == "legacy" + assert "mode" not in call_kwargs sampling_caps = call_kwargs.get("sampling_capabilities") assert sampling_caps is not None assert isinstance(sampling_caps, types.SamplingCapability) @@ -4933,7 +4933,7 @@ async def test_connect_sampling_capabilities_with_client(): async def test_connect_no_sampling_capabilities_without_client(): - """Test connect() keeps auto mode and omits sampling capabilities without a chat client.""" + """Test connect() uses the SDK's auto mode and omits sampling capabilities without a chat client.""" tool = MCPStdioTool(name="test", command="test-command", load_tools=False, load_prompts=False) with patch("mcp.Client") as mock_client_class: @@ -4943,7 +4943,7 @@ async def test_connect_no_sampling_capabilities_without_client(): try: await tool.connect() call_kwargs = mock_client_class.call_args.kwargs - assert call_kwargs["mode"] == "auto" + assert "mode" not in call_kwargs assert call_kwargs.get("sampling_capabilities") is None finally: await tool.close() @@ -6923,6 +6923,361 @@ async def test_connect_handles_set_logging_level_exception(): assert "Failed to set log level" in call_args[0][0] +def _mcp_tool_for_in_process_server( + server: Any, + *, + load_tools: bool, + load_prompts: bool, + client: SupportsChatGetResponse | None = None, + sampling_approval_callback: Callable[[types.CreateMessageRequestParams], bool] | None = None, +) -> MCPTool: + class _InProcessMCPTool(MCPTool): + def get_mcp_client(self) -> Any: + return server + + return _InProcessMCPTool( + name="mrtr-test", + load_tools=load_tools, + load_prompts=load_prompts, + client=client, + sampling_approval_callback=sampling_approval_callback, + ) + + +async def test_tool_call_drives_state_only_mrtr_through_high_level_client() -> None: + from mcp.server import Server, ServerRequestContext + + calls: list[tuple[int | str | None, str, dict[str, Any] | None, str | None]] = [] + + async def list_tools( + _ctx: ServerRequestContext[Any], + _params: types.PaginatedRequestParams | None, + ) -> types.ListToolsResult: + return types.ListToolsResult( + tools=[ + types.Tool( + name="greet", + input_schema={ + "type": "object", + "properties": {"name": {"type": "string"}}, + "required": ["name"], + }, + ) + ] + ) + + async def call_tool( + ctx: ServerRequestContext[Any], + params: types.CallToolRequestParams, + ) -> types.CallToolResult | types.InputRequiredResult: + arguments = params.arguments + assert arguments is not None + calls.append((ctx.request_id, params.name, arguments, params.request_state)) + if params.request_state is None: + return types.InputRequiredResult(request_state="opaque-tool-state") + return types.CallToolResult(content=[types.TextContent(type="text", text=f"Hello, {arguments['name']}!")]) + + server = Server("mrtr-tool-server", on_list_tools=list_tools, on_call_tool=call_tool) + tool = _mcp_tool_for_in_process_server(server, load_tools=True, load_prompts=False) + + async with tool: + result = await tool.call_tool("greet", name="Ada") + + assert _mcp_result_to_text(result) == "Hello, Ada!" + assert [(name, arguments, state) for _, name, arguments, state in calls] == [ + ("greet", {"name": "Ada"}, None), + ("greet", {"name": "Ada"}, "opaque-tool-state"), + ] + assert calls[0][0] is not None + assert calls[1][0] is not None + assert calls[0][0] != calls[1][0] + + +async def test_tool_call_resolves_sampling_mrtr_input_request_through_existing_approval_surface() -> None: + from mcp.server import Server, ServerRequestContext + + calls: list[tuple[int | str | None, dict[str, Any] | None, str | None]] = [] + approvals: list[types.CreateMessageRequestParams] = [] + sampling_request = types.CreateMessageRequest( + params=types.CreateMessageRequestParams( + messages=[ + types.SamplingMessage( + role="user", + content=types.TextContent(type="text", text="What is the capital of France?"), + ) + ], + max_tokens=32, + ) + ) + + async def list_tools( + _ctx: ServerRequestContext[Any], + _params: types.PaginatedRequestParams | None, + ) -> types.ListToolsResult: + return types.ListToolsResult( + tools=[types.Tool(name="answer", input_schema={"type": "object", "properties": {}})] + ) + + async def call_tool( + ctx: ServerRequestContext[Any], + params: types.CallToolRequestParams, + ) -> types.CallToolResult | types.InputRequiredResult: + calls.append((ctx.request_id, params.input_responses, params.request_state)) + if params.input_responses is None: + return types.InputRequiredResult( + input_requests={"sample": sampling_request}, + request_state="opaque-sampling-state", + ) + sample = params.input_responses["sample"] + assert isinstance(sample, types.CreateMessageResult) + assert isinstance(sample.content, types.TextContent) + return types.CallToolResult(content=[types.TextContent(type="text", text=sample.content.text)]) + + def approve(params: types.CreateMessageRequestParams) -> bool: + approvals.append(params) + return True + + chat_client = AsyncMock() + chat_client.get_response = AsyncMock(return_value=_make_sampling_response("Paris")) + server = Server("mrtr-sampling-server", on_list_tools=list_tools, on_call_tool=call_tool) + with pytest.warns(DeprecationWarning, match="MCP sampling"): + tool = _mcp_tool_for_in_process_server( + server, + load_tools=True, + load_prompts=False, + client=chat_client, + sampling_approval_callback=approve, + ) + + async with tool: + result = await tool.call_tool("answer") + + assert _mcp_result_to_text(result) == "Paris" + assert approvals == [sampling_request.params] + assert [state for _, _, state in calls] == [None, "opaque-sampling-state"] + assert calls[0][0] is not None + assert calls[1][0] is not None + assert calls[0][0] != calls[1][0] + + +async def test_auto_mode_falls_back_to_legacy_sampling_backchannel() -> None: + from contextlib import asynccontextmanager + + import anyio + from mcp.shared.memory import create_client_server_memory_streams + from mcp.shared.message import SessionMessage + + sampling_params = types.CreateMessageRequestParams( + messages=[ + types.SamplingMessage( + role="user", + content=types.TextContent(type="text", text="What is the capital of France?"), + ) + ], + max_tokens=32, + ) + discover_calls = 0 + initialize_capabilities: dict[str, Any] | None = None + approvals: list[types.CreateMessageRequestParams] = [] + + @asynccontextmanager + async def legacy_transport() -> AsyncIterator[tuple[Any, Any]]: + async with create_client_server_memory_streams() as (client_streams, server_streams): + client_read, client_write = client_streams + server_read, server_write = server_streams + + async def run_server() -> None: + nonlocal discover_calls, initialize_capabilities + async for session_message in server_read: + if isinstance(session_message, Exception): + raise session_message + message = session_message.message + if isinstance(message, types.JSONRPCNotification): + continue + assert isinstance(message, types.JSONRPCRequest) + + if message.method == "server/discover": + discover_calls += 1 + response: types.JSONRPCResponse | types.JSONRPCError = types.JSONRPCError( + jsonrpc="2.0", + id=message.id, + error=types.ErrorData(code=types.METHOD_NOT_FOUND, message="Method not found"), + ) + elif message.method == "initialize": + assert message.params is not None + initialize_capabilities = message.params["capabilities"] + response = types.JSONRPCResponse( + jsonrpc="2.0", + id=message.id, + result={ + "protocolVersion": "2025-11-25", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "legacy-sampling-server", "version": "1.0"}, + }, + ) + elif message.method == "ping": + response = types.JSONRPCResponse(jsonrpc="2.0", id=message.id, result={}) + elif message.method == "tools/list": + response = types.JSONRPCResponse( + jsonrpc="2.0", + id=message.id, + result={"tools": [{"name": "answer", "inputSchema": {"type": "object", "properties": {}}}]}, + ) + elif message.method == "tools/call": + await server_write.send( + SessionMessage( + types.JSONRPCRequest( + jsonrpc="2.0", + id="sampling-1", + method="sampling/createMessage", + params=sampling_params.model_dump(by_alias=True, mode="json", exclude_none=True), + ) + ) + ) + sample_message = await server_read.receive() + assert not isinstance(sample_message, Exception) + sample_response = sample_message.message + assert isinstance(sample_response, types.JSONRPCResponse) + assert sample_response.id == "sampling-1" + sample_content = sample_response.result["content"] + assert isinstance(sample_content, dict) + response = types.JSONRPCResponse( + jsonrpc="2.0", + id=message.id, + result={ + "content": [{"type": "text", "text": sample_content["text"]}], + "isError": False, + }, + ) + else: + raise AssertionError(f"Unexpected legacy MCP method: {message.method}") + await server_write.send(SessionMessage(response)) + + async with anyio.create_task_group() as task_group: + task_group.start_soon(run_server) + try: + yield client_read, client_write + finally: + await client_write.aclose() + + def approve(params: types.CreateMessageRequestParams) -> bool: + approvals.append(params) + return True + + chat_client = AsyncMock() + chat_client.get_response = AsyncMock(return_value=_make_sampling_response("Paris")) + with pytest.warns(DeprecationWarning, match="MCP sampling"): + tool = _mcp_tool_for_in_process_server( + legacy_transport(), + load_tools=True, + load_prompts=False, + client=chat_client, + sampling_approval_callback=approve, + ) + + async with tool: + assert tool.session is not None + assert tool.session.protocol_version == "2025-11-25" + result = await tool.call_tool("answer") + + assert discover_calls == 1 + assert initialize_capabilities is not None + assert "sampling" in initialize_capabilities + assert _mcp_result_to_text(result) == "Paris" + assert approvals == [sampling_params] + chat_client.get_response.assert_awaited_once() + + +async def test_prompt_get_drives_state_only_mrtr_through_high_level_client() -> None: + from mcp.server import Server, ServerRequestContext + + calls: list[tuple[int | str | None, str, dict[str, str] | None, str | None]] = [] + + async def list_prompts( + _ctx: ServerRequestContext[Any], + _params: types.PaginatedRequestParams | None, + ) -> types.ListPromptsResult: + return types.ListPromptsResult( + prompts=[ + types.Prompt( + name="briefing", + arguments=[types.PromptArgument(name="topic", required=True)], + ) + ] + ) + + async def get_prompt( + ctx: ServerRequestContext[Any], + params: types.GetPromptRequestParams, + ) -> types.GetPromptResult | types.InputRequiredResult: + arguments = params.arguments + assert arguments is not None + calls.append((ctx.request_id, params.name, arguments, params.request_state)) + if params.request_state is None: + return types.InputRequiredResult(request_state="opaque-prompt-state") + return types.GetPromptResult( + messages=[ + types.PromptMessage( + role="user", + content=types.TextContent(type="text", text=f"Explain {arguments['topic']}"), + ) + ] + ) + + server = Server("mrtr-prompt-server", on_list_prompts=list_prompts, on_get_prompt=get_prompt) + tool = _mcp_tool_for_in_process_server(server, load_tools=False, load_prompts=True) + + async with tool: + result = await tool.functions[0].invoke(topic="Python") + + assert _mcp_result_to_text(result) == "Explain Python" + assert [(name, arguments, state) for _, name, arguments, state in calls] == [ + ("briefing", {"topic": "Python"}, None), + ("briefing", {"topic": "Python"}, "opaque-prompt-state"), + ] + assert calls[0][0] is not None + assert calls[1][0] is not None + assert calls[0][0] != calls[1][0] + + +async def test_supplied_session_does_not_drive_mrtr_automatically() -> None: + from mcp import Client + from mcp.server import Server, ServerRequestContext + + call_count = 0 + + async def list_tools( + _ctx: ServerRequestContext[Any], + _params: types.PaginatedRequestParams | None, + ) -> types.ListToolsResult: + return types.ListToolsResult( + tools=[types.Tool(name="greet", input_schema={"type": "object", "properties": {}})] + ) + + async def call_tool( + _ctx: ServerRequestContext[Any], + _params: types.CallToolRequestParams, + ) -> types.InputRequiredResult: + nonlocal call_count + call_count += 1 + return types.InputRequiredResult(request_state="caller-owned-state") + + server = Server("mrtr-session-server", on_list_tools=list_tools, on_call_tool=call_tool) + + async with Client(server) as client: + wrapper = MCPStdioTool( + name="mrtr-session", + command="unused", + session=client.session, + load_prompts=False, + ) + async with wrapper: + with pytest.raises(ToolExecutionException, match="input_required"): + await wrapper.call_tool("greet") + + assert call_count == 1 + + @pytest.mark.parametrize( ("mode", "expected_version"), [ diff --git a/python/packages/core/tests/core/test_mcp_skills.py b/python/packages/core/tests/core/test_mcp_skills.py index e2e35618392..f34eda45be5 100644 --- a/python/packages/core/tests/core/test_mcp_skills.py +++ b/python/packages/core/tests/core/test_mcp_skills.py @@ -604,6 +604,41 @@ async def read_resource( call("skill://unit-converter/SKILL.md", cache_mode="bypass"), ] + async def test_high_level_client_drives_state_only_mrtr_for_skill_resource(self) -> None: + from mcp import Client + from mcp.server import Server, ServerRequestContext + from mcp.types import InputRequiredResult, ReadResourceRequestParams + + calls: list[tuple[int | str | None, str, str | None]] = [] + + async def read_resource( + ctx: ServerRequestContext[Any], + params: ReadResourceRequestParams, + ) -> ReadResourceResult | InputRequiredResult: + uri = str(params.uri) + if uri == "skill://index.json": + return _make_text_result(SAMPLE_SKILL_INDEX, uri=uri) + calls.append((ctx.request_id, uri, params.request_state)) + if params.request_state is None: + return InputRequiredResult(request_state="opaque-resource-state") + return _make_text_result(SAMPLE_SKILL_MD, uri=uri) + + server = Server("mrtr-skills-server", on_read_resource=read_resource) + + async with Client(server) as client: + source = MCPSkillsSource(client=client) + skill = (await source.get_skills(_SOURCE_CTX))[0] + content = await skill.get_content() + + assert content == SAMPLE_SKILL_MD + assert [(uri, state) for _, uri, state in calls] == [ + ("skill://unit-converter/SKILL.md", None), + ("skill://unit-converter/SKILL.md", "opaque-resource-state"), + ] + assert calls[0][0] is not None + assert calls[1][0] is not None + assert calls[0][0] != calls[1][0] + @pytest.mark.parametrize( "uri", [ From 7a18ca9206ca25a6880fb92905800fcc67dfcf43 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 7 Oct 2026 15:33:42 +0200 Subject: [PATCH 37/42] logging migrated --- .../0045-python-mcp-v2-client-lifecycle.md | 11 ++- python/packages/core/agent_framework/_mcp.py | 31 ++++--- python/packages/core/tests/core/test_mcp.py | 85 ++++++++++++++----- 3 files changed, 88 insertions(+), 39 deletions(-) diff --git a/docs/decisions/0045-python-mcp-v2-client-lifecycle.md b/docs/decisions/0045-python-mcp-v2-client-lifecycle.md index a7f6199b03e..993c684bcc1 100644 --- a/docs/decisions/0045-python-mcp-v2-client-lifecycle.md +++ b/docs/decisions/0045-python-mcp-v2-client-lifecycle.md @@ -165,10 +165,10 @@ transport-specific wrappers. ### Legacy callbacks, logging, and liveness -The existing sampling-enabled path uses `mode="legacy"` because it relies on the legacy server-to-client -back-channel. Modern 2026-07-28 sampling, elicitation, and roots requests travel inside MRTR instead. Moving the -existing sampling option to auto mode requires a separate compatibility decision; it is not an incidental part of -routing standard calls through `Client`. +Framework-created connections use `mode="auto"` even when sampling is configured. The SDK routes the same +`sampling_callback` through the legacy server-to-client back-channel after a 2025 fallback and through embedded MRTR +requests after modern discovery. Agent Framework does not infer the protocol era from the presence of its sampling +ChatClient. Modern protocol logging is per-request metadata. Agent Framework should pass the selected log level to `Client` so the SDK stamps `io.modelcontextprotocol/logLevel`; `logging/setLevel` is retained only for a negotiated legacy peer. @@ -185,8 +185,7 @@ scope: - Hosting/server remains server-side and checked. - Tools, Tool refresh, Prompts, and Skills remain checked for their stated dual-era behavior. -- MRTR, Caching, Logging, Samples/docs, local validation, and live dual-era validation remain separate unchecked - work. +- MRTR, Caching, Logging, Samples/docs, local validation, and live dual-era validation remain separate checklist work. - The protocol-independent prompt snapshot bug remains in [microsoft/agent-framework#9115](https://github.com/microsoft/agent-framework/issues/9115). - Optional subscription stream recovery remains in diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 48df868caf3..d195b8efa4b 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -655,8 +655,9 @@ def _as_mcp_connection( # pyright: ignore[reportUnusedFunction] _DEFAULT_SAMPLING_MAX_TOKENS = 4096 _DEFAULT_SAMPLING_MAX_REQUESTS = 25 _MCP_SAMPLING_DEPRECATION_MESSAGE = ( - "MCP sampling is deprecated as of MCP specification version 2026-07-28 and will be removed no later than " - "2027-07-28. MCP servers should call LLM provider APIs directly." + "MCP sampling is deprecated as of MCP specification version 2026-07-28. Under the MCP feature lifecycle, " + "it remains supported for at least twelve months before becoming eligible for removal. " + "MCP servers should call LLM provider APIs directly." ) # A user-supplied gate invoked before each server-initiated sampling request is @@ -680,6 +681,15 @@ def _as_mcp_connection( # pyright: ignore[reportUnusedFunction] } +def _to_mcp_logging_level(level: int) -> types.LoggingLevel | None: + """Map a Python logging level to its MCP equivalent.""" + if level == logging.NOTSET: + return None + return cast( + "types.LoggingLevel | None", next((name for name, value in LOG_LEVEL_MAPPING.items() if value == level), None) + ) + + def _get_input_model_from_mcp_prompt(prompt: types.Prompt) -> dict[str, Any]: """Get the input model from an MCP prompt. @@ -2189,6 +2199,7 @@ async def _connect_on_owner( Raises: ToolException: If connection or session initialization fails. """ + log_level = _to_mcp_logging_level(logger.level) if reset: await self._cancel_capability_list_subscription() if reset_discovery: @@ -2226,6 +2237,7 @@ async def _connect_on_owner( ), message_handler=self.message_handler, logging_callback=self.logging_callback, + log_level=log_level, sampling_capabilities=sampling_capabilities, sampling_callback=self.sampling_callback, # pyright: ignore[reportDeprecated] ) @@ -2306,14 +2318,13 @@ async def _connect_on_owner( await self.load_prompts() self._prompts_loaded = True - if logger.level != logging.NOTSET and self._supports_logging is not False: - try: - level_name = cast( - Any, next(level for level, value in LOG_LEVEL_MAPPING.items() if value == logger.level) - ) - await self._require_connection().set_logging_level(level_name) - except Exception as exc: - logger.warning("Failed to set log level to %s", logger.level, exc_info=exc) + if log_level is not None: + connection = self._require_connection() + if connection.session.initialize_result is not None and self._supports_logging is not False: + try: + await connection.set_logging_level(log_level) + except Exception as exc: + logger.warning("Failed to set log level to %s", logger.level, exc_info=exc) except (Exception, asyncio.CancelledError): try: await self._close_on_owner() diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 875eccd44ef..a8c897eb1fc 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -4192,12 +4192,12 @@ async def test_mcp_tool_sampling_defaults_stay_silent_until_callback_is_used(): tool = MCPStdioTool(name="test_tool", command="python") callback = getattr(MCPTool, "sampling_callback") # noqa: B009 - assert "2027-07-28" in getattr(callback, "__deprecated__", "") + assert "eligible for removal" in getattr(callback, "__deprecated__", "") params = Mock() params.messages = [] - with pytest.warns(DeprecationWarning, match="MCP sampling.*2027-07-28"): + with pytest.warns(DeprecationWarning, match="MCP sampling.*eligible for removal"): result = await _invoke_sampling_callback(tool, params) assert isinstance(result, types.ErrorData) @@ -4205,7 +4205,7 @@ async def test_mcp_tool_sampling_defaults_stay_silent_until_callback_is_used(): async def test_mcp_tool_sampling_configuration_warns_once(): """Each sampling option warns at setup, without warning again on callback use.""" - with pytest.warns(DeprecationWarning, match="MCP sampling.*2027-07-28") as warning_info: + with pytest.warns(DeprecationWarning, match="MCP sampling.*eligible for removal") as warning_info: tool = MCPStdioTool( name="test_tool", command="python", @@ -4244,7 +4244,7 @@ async def test_mcp_tool_sampling_configuration_warns_once(): ) def test_mcp_tool_each_sampling_option_warns(sampling_option: dict[str, Any]): """Each non-default sampling option enables the setup warning.""" - with pytest.warns(DeprecationWarning, match="MCP sampling.*2027-07-28") as warning_info: + with pytest.warns(DeprecationWarning, match="MCP sampling.*eligible for removal") as warning_info: MCPStdioTool(name="test_tool", command="python", **sampling_option) assert len(warning_info) == 1 @@ -6843,7 +6843,7 @@ async def test_mcp_tool_safe_close_handles_cleanup_exception_group(): async def test_connect_sets_logging_level_when_logger_level_is_set(): - """Test that connect() sets the MCP server logging level when the logger level is not NOTSET.""" + """Test that connect() configures modern metadata and the legacy server logging level.""" tool = MCPStdioTool( name="test_server", @@ -6860,10 +6860,11 @@ async def test_connect_sets_logging_level_when_logger_level_is_set(): ) with ( - patch("mcp.Client", return_value=sdk_client), + patch("mcp.Client", return_value=sdk_client) as mock_client_class, patch.object(logger, "level", logging.DEBUG), # Set logger level to DEBUG ): async with tool: + assert mock_client_class.call_args.kwargs["log_level"] == "debug" mock_session.set_logging_level.assert_awaited_once_with("debug") @@ -6885,10 +6886,11 @@ async def test_connect_does_not_set_logging_level_when_logger_level_is_notset(): ) with ( - patch("mcp.Client", return_value=sdk_client), + patch("mcp.Client", return_value=sdk_client) as mock_client_class, patch.object(logger, "level", logging.NOTSET), # Set logger level to NOTSET ): async with tool: + assert mock_client_class.call_args.kwargs["log_level"] is None mock_session.set_logging_level.assert_not_called() @@ -6912,17 +6914,42 @@ async def test_connect_handles_set_logging_level_exception(): ) with ( - patch("mcp.Client", return_value=sdk_client), + patch("mcp.Client", return_value=sdk_client) as mock_client_class, patch.object(logger, "level", logging.INFO), # Set logger level to INFO patch.object(logger, "warning") as mock_warning, ): async with tool: + assert mock_client_class.call_args.kwargs["log_level"] == "info" mock_session.set_logging_level.assert_awaited_once_with("info") mock_warning.assert_called_once() call_args = mock_warning.call_args assert "Failed to set log level" in call_args[0][0] +async def test_connect_does_not_use_legacy_logging_method_for_modern_server() -> None: + tool = MCPStdioTool( + name="test_server", + command="test_command", + load_tools=False, + load_prompts=False, + ) + mock_session = Mock(spec=ClientSession) + mock_session.set_logging_level = AsyncMock() + sdk_client = _mock_sdk_client( + session=mock_session, + capabilities=types.ServerCapabilities(logging=types.LoggingCapability()), + protocol_version="2026-07-28", + ) + + with ( + patch("mcp.Client", return_value=sdk_client) as mock_client_class, + patch.object(logger, "level", logging.WARNING), + ): + async with tool: + assert mock_client_class.call_args.kwargs["log_level"] == "warning" + mock_session.set_logging_level.assert_not_awaited() + + def _mcp_tool_for_in_process_server( server: Any, *, @@ -9039,14 +9066,15 @@ async def test_mcp_streamble_http_tool_connects_to_v2_server() -> None: header_provider=lambda _kw: {"Authorization": "Bearer token-a"}, ) - async with tool_a: - assert tool_a.session is not None - assert tool_a.session.protocol_version == "2026-07-28" - assert [function.name for function in tool_a.functions] == ["greet"] + with patch.object(logger, "level", logging.INFO): + async with tool_a: + assert tool_a.session is not None + assert tool_a.session.protocol_version == "2026-07-28" + assert [function.name for function in tool_a.functions] == ["greet"] - result = await tool_a.call_tool("greet") - assert isinstance(result, list) - assert [item.text for item in result if item.type == "text"] == ["Hello!"] + result = await tool_a.call_tool("greet") + assert isinstance(result, list) + assert [item.text for item in result if item.type == "text"] == ["Hello!"] captured_methods = [body["method"] for body, _ in captured_requests] assert "server/discover" in captured_methods @@ -9060,15 +9088,18 @@ async def test_mcp_streamble_http_tool_connects_to_v2_server() -> None: params = body["params"] meta = params["_meta"] assert headers["mcp-protocol-version"] == meta["io.modelcontextprotocol/protocolVersion"] == "2026-07-28" + assert meta["io.modelcontextprotocol/logLevel"] == "info" assert headers["mcp-method"] == body["method"] == "tools/call" assert headers["mcp-name"] == params["name"] == "greet" assert isinstance(meta["io.modelcontextprotocol/clientCapabilities"], dict) async def test_mcp_streamable_http_tool_connects_to_legacy_server() -> None: + from mcp import MCPDeprecationWarning + transport, captured_requests = _make_mcp_protocol_server_mock( era="legacy", - capabilities={"tools": {}}, + capabilities={"tools": {}, "logging": {}}, endpoints={ "tools/list": { "cacheScope": "private", @@ -9080,6 +9111,7 @@ async def test_mcp_streamable_http_tool_connects_to_legacy_server() -> None: "content": [{"type": "text", "text": "Hello!"}], "isError": False, }, + "logging/setLevel": {}, }, ) user_client = AsyncClient(transport=transport) @@ -9091,19 +9123,26 @@ async def test_mcp_streamable_http_tool_connects_to_legacy_server() -> None: header_provider=lambda _kw: {"Authorization": "Bearer token-a"}, ) - async with tool_a: - assert tool_a.session is not None - assert tool_a.session.protocol_version == "2025-11-25" - assert [function.name for function in tool_a.functions] == ["greet"] + with ( + patch.object(logger, "level", logging.WARNING), + pytest.warns(MCPDeprecationWarning, match="logging capability"), + ): + async with tool_a: + assert tool_a.session is not None + assert tool_a.session.protocol_version == "2025-11-25" + assert [function.name for function in tool_a.functions] == ["greet"] - result = await tool_a.call_tool("greet") - assert isinstance(result, list) - assert [item.text for item in result if item.type == "text"] == ["Hello!"] + result = await tool_a.call_tool("greet") + assert isinstance(result, list) + assert [item.text for item in result if item.type == "text"] == ["Hello!"] captured_methods = [body["method"] for body, _ in captured_requests] assert "server/discover" in captured_methods assert "initialize" in captured_methods assert "tools/list" in captured_methods + logging_requests = [body for body, _ in captured_requests if body["method"] == "logging/setLevel"] + assert len(logging_requests) == 1 + assert logging_requests[0]["params"]["level"] == "warning" @pytest.mark.parametrize( From 0eca700054de3de9c5ec0ebea72908b5d6153f2d Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 7 Oct 2026 16:40:55 +0200 Subject: [PATCH 38/42] caching behaviour validated --- .../0045-python-mcp-v2-client-lifecycle.md | 26 +- python/packages/core/AGENTS.md | 1 + python/packages/core/agent_framework/_mcp.py | 15 +- .../packages/core/agent_framework/_skills.py | 5 + python/packages/core/tests/core/test_mcp.py | 285 +++++++++++++++++- .../core/tests/core/test_mcp_skills.py | 42 ++- python/packages/core/tests/test_security.py | 4 +- 7 files changed, 347 insertions(+), 31 deletions(-) diff --git a/docs/decisions/0045-python-mcp-v2-client-lifecycle.md b/docs/decisions/0045-python-mcp-v2-client-lifecycle.md index 993c684bcc1..f0f07c38533 100644 --- a/docs/decisions/0045-python-mcp-v2-client-lifecycle.md +++ b/docs/decisions/0045-python-mcp-v2-client-lifecycle.md @@ -113,20 +113,18 @@ Absent server hints use `CacheConfig.default_ttl_ms`, whose default is `0`. Unde request metadata forces a wire refresh. Although `server/discover` carries protocol cache hints, SDK 2.2.0 deliberately excludes it from the response cache; persisting or reusing `prior_discover` is caller-managed. -The Client-first cleanup and the separate Caching checklist item are staged deliberately: - -1. While preserving pre-caching Agent Framework behavior, explicit catalog and resource refresh paths use - `cache_mode="bypass"`. -2. The Caching migration later assigns an intentional policy per operation: - - `"refresh"` for an explicit authoritative refetch that must update or evict the SDK cache; - - `"use"` only where Agent Framework intentionally accepts server `ttlMs` / `cacheScope` freshness; - - `"bypass"` only where neither reading nor updating the SDK cache is desired. -3. Resource and catalog refresh tests must cover positive TTLs, pagination, changed metadata, empty snapshots, - reconnect, and authenticated identity changes. -4. A reconnect or effective header-identity change replaces the whole Client and its default per-client cache. A - shared cache store must be partitioned by a verified authorization identity. Because Agent Framework constructs - `Client` from a transport rather than a URL, a future shared store also requires an explicit stable - `CacheConfig.target_id`. +Framework-owned connections use the SDK's default `cache_mode="use"` for tool and prompt catalogs and resource +reads. Modern server-provided `ttlMs` / `cacheScope` hints therefore control freshness. Legacy peers provide no +hints and remain uncached under the SDK's default zero TTL; caller-supplied `ClientSession` connections bypass the +SDK response cache. + +Agent Framework does not expose a cache-mode or shared-cache configuration surface. A reconnect or effective +header-identity change replaces the whole Client and its default per-client cache. Exposing a shared store would +require a separate authorization-partition design and, because Agent Framework constructs `Client` from a transport +rather than a URL, an explicit stable `CacheConfig.target_id`. + +Resource and catalog cache tests cover positive TTLs, pagination, empty snapshots, reconnect, notifications, and +authenticated identity changes. MRTR-seeded and MRTR-resolved resource reads are not cached by the SDK. diff --git a/python/packages/core/AGENTS.md b/python/packages/core/AGENTS.md index d41b8b1f763..247fb6e8134 100644 --- a/python/packages/core/AGENTS.md +++ b/python/packages/core/AGENTS.md @@ -225,6 +225,7 @@ The vector store API is experimental under the shared `VECTOR_STORES` feature ID - **Streamable HTTP cookies** - Framework-created clients reject response-cookie persistence with or without a `header_provider`; explicit `Cookie` headers from `static_headers` or a provider remain supported. Caller-provided clients retain their cookie behavior and ownership. Applications requiring cookie persistence must scope clients and MCP sessions to one authenticated principal. Cookie rejection does not isolate MCP protocol sessions or other server-side state. - **`function_invocation_kwargs` and MCP servers** - That dict is shared across every tool in the run, including every attached `MCPTool`, and any name in it reaches a server that declares a matching `inputSchema` property. `header_provider` does not mitigate this — it reads the kwargs without consuming them. To keep a credential out of tool arguments, source it outside `function_invocation_kwargs`: read a `ContextVar` inside the provider (this still allows a different value per request), configure a custom `http_client`, or use `env` for `MCPStdioTool`. - **Sampling guardrails** (`sampling_callback`) - Passing `client=` advertises `SamplingCapability` so the server can send `sampling/createMessage`. Because remote servers are untrusted (confused-deputy risk), the default `sampling_callback` is **deny-by-default** and applies, in order: a per-session rate limit (`sampling_max_requests`, default `_DEFAULT_SAMPLING_MAX_REQUESTS`), an approval gate (`sampling_approval_callback`), and a `maxTokens` cap (`sampling_max_tokens`, default `_DEFAULT_SAMPLING_MAX_TOKENS`). The approval callback (constructor arg on all subclasses; exported type alias `SamplingApprovalCallback`) receives the raw `CreateMessageRequestParams`, may be sync or async, and must return truthy to approve. When it is `None` (the default) every sampling request is denied; pass `lambda params: True` to restore legacy auto-approve as an explicit opt-in. Requests and denials are logged at WARNING (content is not logged). The per-session counter resets in `_reset_session_state`. +- **MCP response caching** - Framework-owned high-level Clients use SDK `cache_mode="use"` for tool/prompt catalogs and resource reads, honoring modern server `ttlMs` / `cacheScope` hints. Legacy peers default to zero TTL, caller-supplied `ClientSession` connections stay uncached, and reconnect or effective header-identity changes replace the per-Client cache. Agent Framework does not expose shared-store or cache-mode configuration. - **`MCPTaskOptions`** (experimental, `MCP_LONG_RUNNING_TASKS` feature, **frozen**) - Per-tool-instance options controlling the SEP-2663 long-running task lifecycle. When the server advertises a tool with `execution.taskSupport == "required"`, `MCPTool.call_tool` transparently routes through `call_tool_as_task`, which sends an augmented `tools/call`, polls `tasks/get` until terminal, and reinterprets `tasks/result` as a normal `CallToolResult`. Instances are immutable; replace via `MCPTool.task_options = MCPTaskOptions(...)`. Fields: - `default_ttl: timedelta | None` — forwarded to the server as `params.task.ttl` (milliseconds). When `None`, the server's default applies. - `cancel_remote_task_on_local_cancellation: bool = True` — only gates the `CancelledError` path. Abandonment paths (see below) always cancel. diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index d195b8efa4b..8d3d4033a6f 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -566,17 +566,15 @@ async def get_prompt(self, name: str, arguments: dict[str, Any] | None) -> types return await self.client.get_prompt(name, arguments=cast("dict[str, str] | None", arguments)) async def list_tools_page(self, params: types.PaginatedRequestParams | None) -> types.ListToolsResult: - """List one tools page without changing existing cache behavior.""" + """List one tools page while honoring server cache hints.""" return await self.client.list_tools( cursor=params.cursor if params is not None else None, - cache_mode="bypass", ) async def list_prompts_page(self, params: types.PaginatedRequestParams | None) -> types.ListPromptsResult: - """List one prompts page without changing existing cache behavior.""" + """List one prompts page while honoring server cache hints.""" return await self.client.list_prompts( cursor=params.cursor if params is not None else None, - cache_mode="bypass", ) async def set_logging_level(self, level: Any) -> None: @@ -584,8 +582,8 @@ async def set_logging_level(self, level: Any) -> None: await self.session.set_logging_level(level) # pyright: ignore[reportDeprecated] async def read_resource(self, uri: str) -> types.ReadResourceResult: - """Read a resource through the high-level Client.""" - return await self.client.read_resource(uri, cache_mode="bypass") + """Read a resource through the high-level Client while honoring server cache hints.""" + return await self.client.read_resource(uri, cache_mode="use") @dataclass(frozen=True) @@ -1054,6 +1052,11 @@ class MCPTool: MCPTool cannot be instantiated directly. Use one of the subclasses: MCPStdioTool or MCPStreamableHTTPTool. + Caching: + Framework-owned modern connections honor server-provided ``ttlMs`` and ``cacheScope`` hints through the + MCP SDK's per-Client response cache. Legacy servers and caller-supplied ``ClientSession`` connections remain + uncached under the default zero-TTL policy. + Examples: See the subclass documentation for usage examples: diff --git a/python/packages/core/agent_framework/_skills.py b/python/packages/core/agent_framework/_skills.py index d8560b2735e..b7e3f7ec46d 100644 --- a/python/packages/core/agent_framework/_skills.py +++ b/python/packages/core/agent_framework/_skills.py @@ -5368,6 +5368,11 @@ class MCPSkillsSource(SkillsSource): already provides refresh/caching for any source, this source does not offer a separate refresh interval; wrap it in :class:`CachingSkillsSource` to cache. + When backed by a high-level MCP ``Client``, individual resource reads honor + modern server ``ttlMs`` / ``cacheScope`` hints through the SDK response cache. + A caller-supplied ``ClientSession`` remains uncached. This wire-response cache + is separate from :class:`CachingSkillsSource`, which caches the parsed skill list. + Archive digests: An archive entry's non-null ``digest`` must be ``sha256:`` followed by 64 lowercase hexadecimal characters. It is verified against the decoded diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index a8c897eb1fc..7e566e7004b 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -5033,8 +5033,8 @@ async def test_get_prompt_uses_sdk_client_for_framework_owned_connection() -> No session.get_prompt.assert_not_awaited() -async def test_catalog_loading_uses_sdk_client_without_cache() -> None: - """Test framework-owned catalog pagination uses the Client without caching.""" +async def test_catalog_loading_uses_sdk_client_cache() -> None: + """Test framework-owned catalog pagination uses the Client cache.""" capabilities = types.ServerCapabilities( tools=types.ToolsCapability(), prompts=types.PromptsCapability(), @@ -5079,17 +5079,290 @@ async def test_catalog_loading_uses_sdk_client_without_cache() -> None: ] assert [awaited.kwargs for awaited in sdk_client.list_tools.await_args_list] == [ - {"cursor": None, "cache_mode": "bypass"}, - {"cursor": "tools-next", "cache_mode": "bypass"}, + {"cursor": None, "cache_mode": "use"}, + {"cursor": "tools-next", "cache_mode": "use"}, ] assert [awaited.kwargs for awaited in sdk_client.list_prompts.await_args_list] == [ - {"cursor": None, "cache_mode": "bypass"}, - {"cursor": "prompts-next", "cache_mode": "bypass"}, + {"cursor": None, "cache_mode": "use"}, + {"cursor": "prompts-next", "cache_mode": "use"}, ] session.list_tools.assert_not_awaited() session.list_prompts.assert_not_awaited() +async def test_catalog_loading_honors_positive_server_ttl() -> None: + from mcp.server import Server, ServerRequestContext + + tool_list_count = 0 + prompt_list_count = 0 + + async def list_tools( + _ctx: ServerRequestContext[Any], + _params: types.PaginatedRequestParams | None, + ) -> types.ListToolsResult: + nonlocal tool_list_count + tool_list_count += 1 + return types.ListToolsResult( + tools=[types.Tool(name="cached_tool", input_schema={"type": "object", "properties": {}})], + ttl_ms=60_000, + ) + + async def list_prompts( + _ctx: ServerRequestContext[Any], + _params: types.PaginatedRequestParams | None, + ) -> types.ListPromptsResult: + nonlocal prompt_list_count + prompt_list_count += 1 + return types.ListPromptsResult( + prompts=[types.Prompt(name="cached_prompt", arguments=[])], + ttl_ms=60_000, + ) + + server = Server( + "cache-server", + on_list_tools=list_tools, + on_list_prompts=list_prompts, + ) + tool = _mcp_tool_for_in_process_server(server, load_tools=True, load_prompts=True) + + async with tool: + await tool.load_tools() + await tool.load_prompts() + + assert tool_list_count == 1 + assert prompt_list_count == 1 + assert [function.name for function in tool.functions] == ["cached_tool", "cached_prompt"] + + +async def test_tool_catalog_cache_reuses_only_first_page() -> None: + from mcp.server import Server, ServerRequestContext + + requested_cursors: list[str | None] = [] + + async def list_tools( + _ctx: ServerRequestContext[Any], + params: types.PaginatedRequestParams | None, + ) -> types.ListToolsResult: + cursor = params.cursor if params is not None else None + requested_cursors.append(cursor) + if cursor is None: + return types.ListToolsResult( + tools=[types.Tool(name="first", input_schema={"type": "object", "properties": {}})], + next_cursor="second", + ttl_ms=60_000, + ) + assert cursor == "second" + return types.ListToolsResult( + tools=[types.Tool(name="second", input_schema={"type": "object", "properties": {}})], + ttl_ms=60_000, + ) + + server = Server("paginated-cache-server", on_list_tools=list_tools) + tool = _mcp_tool_for_in_process_server(server, load_tools=True, load_prompts=False) + + async with tool: + await tool.load_tools() + + assert requested_cursors == [None, "second", "second"] + assert [function.name for function in tool.functions] == ["first", "second"] + + +async def test_tool_catalog_cache_preserves_empty_snapshot() -> None: + from mcp.server import Server, ServerRequestContext + + list_count = 0 + + async def list_tools( + _ctx: ServerRequestContext[Any], + _params: types.PaginatedRequestParams | None, + ) -> types.ListToolsResult: + nonlocal list_count + list_count += 1 + return types.ListToolsResult(tools=[], ttl_ms=60_000) + + server = Server("empty-cache-server", on_list_tools=list_tools) + tool = _mcp_tool_for_in_process_server(server, load_tools=True, load_prompts=False) + + async with tool: + await tool.load_tools() + + assert list_count == 1 + assert tool.functions == [] + + +async def test_reconnect_replaces_catalog_cache() -> None: + from mcp.server import Server, ServerRequestContext + + list_count = 0 + + async def list_tools( + _ctx: ServerRequestContext[Any], + _params: types.PaginatedRequestParams | None, + ) -> types.ListToolsResult: + nonlocal list_count + list_count += 1 + return types.ListToolsResult( + tools=[types.Tool(name=f"tool_{list_count}", input_schema={"type": "object", "properties": {}})], + ttl_ms=60_000, + ) + + server = Server("reconnect-cache-server", on_list_tools=list_tools) + tool = _mcp_tool_for_in_process_server(server, load_tools=True, load_prompts=False) + + async with tool: + await tool.load_tools() + assert list_count == 1 + assert [function.name for function in tool.functions] == ["tool_1"] + + await tool.connect(reset=True) + assert list_count == 2 + assert [function.name for function in tool.functions] == ["tool_2"] + + +async def test_modern_subscription_evicts_tool_catalog_cache() -> None: + from mcp.client._memory import InMemoryTransport + from mcp.server import Server, ServerRequestContext + from mcp.server.subscriptions import InMemorySubscriptionBus, ListenHandler, ToolsListChanged + + current_name = "first" + list_count = 0 + refreshed = asyncio.Event() + bus = InMemorySubscriptionBus() + listen_handler = ListenHandler(bus) + + async def list_tools( + _ctx: ServerRequestContext[Any], + _params: types.PaginatedRequestParams | None, + ) -> types.ListToolsResult: + nonlocal list_count + list_count += 1 + if list_count == 2: + refreshed.set() + return types.ListToolsResult( + tools=[types.Tool(name=current_name, input_schema={"type": "object", "properties": {}})], + ttl_ms=60_000, + ) + + server = Server( + "subscription-cache-server", + on_list_tools=list_tools, + on_subscriptions_listen=listen_handler, + ) + tool = _mcp_tool_for_in_process_server( + InMemoryTransport(server), + load_tools=True, + load_prompts=False, + ) + + async with tool: + assert tool.session is not None + assert tool.session.protocol_version == "2026-07-28" + await tool.load_tools() + assert list_count == 1 + + current_name = "second" + await bus.publish(ToolsListChanged()) + await asyncio.wait_for(refreshed.wait(), timeout=1) + await asyncio.gather(*tool._pending_reload_tasks) + + assert list_count == 2 + assert [function.name for function in tool.functions] == ["second"] + listen_handler.close() + + +async def test_caller_supplied_session_remains_uncached() -> None: + from mcp import Client + from mcp.server import Server, ServerRequestContext + + list_count = 0 + + async def list_tools( + _ctx: ServerRequestContext[Any], + _params: types.PaginatedRequestParams | None, + ) -> types.ListToolsResult: + nonlocal list_count + list_count += 1 + return types.ListToolsResult( + tools=[types.Tool(name="uncached", input_schema={"type": "object", "properties": {}})], + ttl_ms=60_000, + ) + + server = Server("session-cache-server", on_list_tools=list_tools) + + async with Client(server) as client: + tool = MCPStdioTool( + name="caller-owned", + command="unused", + session=client.session, + load_prompts=False, + ) + async with tool: + await tool.load_tools() + + assert list_count == 2 + + +async def test_authorization_identity_change_replaces_catalog_cache() -> None: + list_principals: list[str] = [] + + async def handle(request: Request) -> Response: + if request.method == "DELETE": + return Response(200) + if request.method == "GET": + return Response(405) + body = json.loads(request.content) + method = body["method"] + principal = request.headers.get("Authorization", "") + if method == "server/discover": + result = { + "supportedVersions": ["2026-07-28"], + "capabilities": {"tools": {}}, + } + elif method == "tools/list": + list_principals.append(principal) + result = { + "resultType": "complete", + "ttlMs": 60_000, + "cacheScope": "private", + "tools": [ + {"name": "noop", "inputSchema": {"type": "object", "properties": {}}}, + {"name": f"{principal}-only", "inputSchema": {"type": "object", "properties": {}}}, + ], + } + elif method == "tools/call": + result = { + "resultType": "complete", + "content": [{"type": "text", "text": principal}], + "isError": False, + } + else: + raise AssertionError(f"Unexpected MCP method: {method}") + return Response(200, json={"jsonrpc": "2.0", "id": body["id"], "result": result}) + + user_client = AsyncClient(transport=MockTransport(handle)) + tool = MCPStreamableHTTPTool( + name="identity-cache", + url="https://mcp.example/mcp", + http_client=user_client, + load_prompts=False, + header_provider=lambda kwargs: {"Authorization": kwargs.get("credential", "token-a")}, + ) + + try: + await tool.connect() + await tool.load_tools() + assert list_principals == ["token-a"] + + await tool.call_tool("noop", credential="token-b") + await tool.load_tools() + + assert list_principals == ["token-a", "token-b"] + assert {function.name for function in tool.functions} == {"noop", "token-b-only"} + finally: + await tool.close() + await user_client.aclose() + + # Test error handling in connect() method diff --git a/python/packages/core/tests/core/test_mcp_skills.py b/python/packages/core/tests/core/test_mcp_skills.py index f34eda45be5..a878cc9974e 100644 --- a/python/packages/core/tests/core/test_mcp_skills.py +++ b/python/packages/core/tests/core/test_mcp_skills.py @@ -573,7 +573,7 @@ def test_requires_exactly_one_of_client_or_session_provider(self) -> None: class TestMCPSkillsSource: """Tests for MCPSkillsSource.""" - async def test_high_level_client_reads_resources_with_cache_bypass(self) -> None: + async def test_high_level_client_reads_resources_with_cache(self) -> None: from mcp import Client from mcp.server import Server, ServerRequestContext from mcp.types import ReadResourceRequestParams @@ -600,10 +600,46 @@ async def read_resource( assert content == SAMPLE_SKILL_MD assert read_resource_mock.await_args_list == [ - call("skill://index.json", cache_mode="bypass"), - call("skill://unit-converter/SKILL.md", cache_mode="bypass"), + call("skill://index.json", cache_mode="use"), + call("skill://unit-converter/SKILL.md", cache_mode="use"), ] + async def test_high_level_client_honors_positive_resource_ttl(self) -> None: + from mcp import Client + from mcp.server import Server, ServerRequestContext + from mcp.types import ReadResourceRequestParams + + read_counts: dict[str, int] = {} + + async def read_resource( + _ctx: ServerRequestContext[Any], + params: ReadResourceRequestParams, + ) -> ReadResourceResult: + uri = str(params.uri) + read_counts[uri] = read_counts.get(uri, 0) + 1 + if uri == "skill://index.json": + return ReadResourceResult( + contents=[TextResourceContents(uri=uri, text=SAMPLE_SKILL_INDEX, mime_type="application/json")], + ttl_ms=60_000, + ) + return ReadResourceResult( + contents=[TextResourceContents(uri=uri, text=SAMPLE_SKILL_MD, mime_type="text/markdown")], + ttl_ms=60_000, + ) + + server = Server("cached-skills-server", on_read_resource=read_resource) + + async with Client(server) as client: + for _ in range(2): + source = MCPSkillsSource(client=client) + skill = (await source.get_skills(_SOURCE_CTX))[0] + assert await skill.get_content() == SAMPLE_SKILL_MD + + assert read_counts == { + "skill://index.json": 1, + "skill://unit-converter/SKILL.md": 1, + } + async def test_high_level_client_drives_state_only_mrtr_for_skill_resource(self) -> None: from mcp import Client from mcp.server import Server, ServerRequestContext diff --git a/python/packages/core/tests/test_security.py b/python/packages/core/tests/test_security.py index 8aaf253a36b..c7039c41f1f 100644 --- a/python/packages/core/tests/test_security.py +++ b/python/packages/core/tests/test_security.py @@ -6935,8 +6935,8 @@ async def test_apply_mcp_security_labels_uses_high_level_client_connection(self) await apply_mcp_security_labels(mcp_tool) assert sdk_client.list_tools.await_args_list == [ - call(cursor=None, cache_mode="bypass"), - call(cursor="next-page", cache_mode="bypass"), + call(cursor=None, cache_mode="use"), + call(cursor="next-page", cache_mode="use"), ] sdk_client.session.list_tools.assert_not_awaited() From e7bcc0add213667d619276365660b30759e57eb3 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Wed, 7 Oct 2026 17:26:58 +0200 Subject: [PATCH 39/42] Update MCP v2 samples and docs Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- python/samples/02-agents/mcp/README.md | 21 ++++++++------ .../02-agents/mcp/mcp_sampling_approval.py | 20 +++++++------ python/samples/04-hosting/mcp/README.md | 28 ++++++++++--------- .../mcp/{fastmcp_app.py => mcpserver_app.py} | 10 +++---- 4 files changed, 45 insertions(+), 34 deletions(-) rename python/samples/04-hosting/mcp/{fastmcp_app.py => mcpserver_app.py} (89%) diff --git a/python/samples/02-agents/mcp/README.md b/python/samples/02-agents/mcp/README.md index a409a8d6cf8..e7b665ae689 100644 --- a/python/samples/02-agents/mcp/README.md +++ b/python/samples/02-agents/mcp/README.md @@ -13,10 +13,20 @@ The Model Context Protocol (MCP) is an open standard for connecting AI agents to | **Agent as MCP Server** | [`agent_as_mcp_server.py`](agent_as_mcp_server.py) | Shows how to expose an Agent Framework agent as an MCP server that other AI applications can connect to | | **API Key Authentication** | [`mcp_api_key_auth.py`](mcp_api_key_auth.py) | Demonstrates API key authentication with MCP servers using `header_provider`, runtime invocation kwargs, and a command-line API key argument | | **GitHub Integration with PAT** | [`mcp_github_pat.py`](mcp_github_pat.py) | Demonstrates connecting to GitHub's MCP server using Personal Access Token (PAT) authentication | - - | **Progressive Disclosure** | [`mcp_progressive_disclosure.py`](mcp_progressive_disclosure.py) | Demonstrates `use_progressive_disclosure`, `always_load`, `allowed_tools`, and prefixed `list_mcp_tools` / `load_tool` / `unload_tool` names. `load_tool` and `unload_tool` can accept one tool name or multiple names. Self-spawns a stdio MCP child server | -| **Sampling Approval** | [`mcp_sampling_approval.py`](mcp_sampling_approval.py) | Demonstrates gating server-initiated `sampling/createMessage` requests with a `sampling_approval_callback`, plus the `sampling_max_tokens` and `sampling_max_requests` guardrails. MCP sampling is denied by default | +| **Sampling Approval** | [`mcp_sampling_approval.py`](mcp_sampling_approval.py) | Demonstrates gating legacy and modern MRTR `sampling/createMessage` requests with `sampling_approval_callback`. MCP sampling is deprecated and denied by default | + +The long-running Tasks sample remains deferred because MCP Python SDK 2.2 does +not yet expose the 2026 `io.modelcontextprotocol/tasks` runtime. + +## Protocol behavior + +- Framework-created connections use SDK auto negotiation: modern servers use + `server/discover`, while older servers fall back to `initialize`. +- The high-level SDK client owns MRTR retries, modern subscription refresh, and + modern server-directed response caching. +- Caller-supplied `ClientSession` objects remain caller-owned, low-level, and + uncached; they do not gain automatic MRTR. ## Prerequisites @@ -31,8 +41,3 @@ Run `mcp_api_key_auth.py` with the MCP API key as the first command-line argumen For `mcp_github_pat.py`: - `GITHUB_PAT` - Your GitHub Personal Access Token (create at https://github.com/settings/tokens) - -For `mcp_long_running_task.py` (uses Azure OpenAI via Entra-ID): -- Run `az login` once -- `AZURE_OPENAI_ENDPOINT` - your Azure OpenAI resource endpoint, e.g. `https://.openai.azure.com/` -- `AZURE_OPENAI_CHAT_MODEL` (or `AZURE_OPENAI_MODEL`) - the deployment name (e.g. `gpt-4o-mini`) diff --git a/python/samples/02-agents/mcp/mcp_sampling_approval.py b/python/samples/02-agents/mcp/mcp_sampling_approval.py index 75b32c13af8..0d86b0c75bc 100644 --- a/python/samples/02-agents/mcp/mcp_sampling_approval.py +++ b/python/samples/02-agents/mcp/mcp_sampling_approval.py @@ -13,11 +13,15 @@ """ MCP Sampling Approval Example -MCP servers can send the client a ``sampling/createMessage`` request, asking the -client to run an LLM completion on the server's behalf. Because remote MCP -servers are untrusted third parties, forwarding these server-controlled prompts -to your chat client without review is a confused-deputy risk: a malicious server -could exfiltrate context, force tool calls, or burn through your token budget. +MCP servers can ask the client to run an LLM completion on the server's behalf. +Legacy servers send ``sampling/createMessage`` over the back-channel; modern +servers carry the same request inside an MRTR ``InputRequiredResult``. The SDK +routes both forms through the same callback. + +Sampling is deprecated in MCP 2026-07-28, but remains supported for existing +servers. Because remote servers are untrusted third parties, forwarding their +prompts without review is a confused-deputy risk: a malicious server could +exfiltrate context, force tool calls, or burn through your token budget. For that reason Agent Framework **denies MCP sampling by default**. To allow it, pass a ``sampling_approval_callback`` to the MCP tool. The callback receives the @@ -30,13 +34,13 @@ - ``sampling_max_requests`` limits how many sampling requests a single session may make. -To restore the legacy "always approve" behavior (only do this for servers you -trust), pass ``sampling_approval_callback=lambda params: True``. +To always approve requests (only do this for servers you trust), pass +``sampling_approval_callback=lambda params: True``. """ async def approve_sampling(params: types.CreateMessageRequestParams) -> bool: - """Human-in-the-loop approval gate for server-initiated sampling. + """Human-in-the-loop approval gate for MCP sampling. Shows the server-supplied system prompt and messages, then asks the user to approve or deny. Returning ``False`` rejects the request. diff --git a/python/samples/04-hosting/mcp/README.md b/python/samples/04-hosting/mcp/README.md index 12e19922d68..aabd62ebd5b 100644 --- a/python/samples/04-hosting/mcp/README.md +++ b/python/samples/04-hosting/mcp/README.md @@ -7,8 +7,8 @@ There are two common ways to build that: native tool schema, validates or parses its arguments, routes the selected tool name, and returns its result. 2. Declare callable tools and let a higher-level server generate the list and - call handlers from those declarations. FastMCP follows this model by deriving - tool schemas and argument parsing from Python function signatures. + call handlers from those declarations. `MCPServer` follows this model by + deriving tool schemas and argument parsing from Python function signatures. Choose between them based on the other MCP features you want to expose and how much control you need over the server, schema, validation, lifecycle, and @@ -44,18 +44,18 @@ Use this when the application needs full control over a custom MCP contract. uv run manual_app.py ``` -### 2. FastMCP server +### 2. MCPServer -[`fastmcp_app.py`](fastmcp_app.py) keeps the same two conversion functions but -replaces the low-level server setup with FastMCP. FastMCP derives and validates +[`mcpserver_app.py`](mcpserver_app.py) keeps the same two conversion functions but +replaces the low-level server setup with `MCPServer`. `MCPServer` derives and validates the tool schema from the decorated `run_agent(...)` function and owns the streamable HTTP server. Use this when a normal Python function signature fully describes the MCP tool. -FastMCP keeps its generated schema and argument parsing aligned. +`MCPServer` keeps its generated schema and argument parsing aligned. ```bash -uv run fastmcp_app.py +uv run mcpserver_app.py ``` ### 3. Agent-derived tool @@ -65,7 +65,7 @@ derives the native tool name and description from the agent, owns the configured argument schema, runs the agent, and applies the same conversion boundary. Use this when one Agent Framework agent should be represented as one generated -MCP tool. Unlike the FastMCP sample, this adapter derives the contract from the +MCP tool. Unlike the `MCPServer` sample, this adapter derives the contract from the agent and its adapter configuration rather than a decorated function signature. ```bash @@ -106,8 +106,8 @@ uv run workflow_app.py | Concept | Responsibility | Used by | |---|---|---| -| `mcp_to_run(...)` | Converts validated MCP arguments into Agent Framework messages and selected chat options. | `manual_app.py`, `fastmcp_app.py` | -| `mcp_from_run(...)` | Converts a completed agent response into MCP result content blocks. | `manual_app.py`, `fastmcp_app.py` | +| `mcp_to_run(...)` | Converts validated MCP arguments into Agent Framework messages and selected chat options. | `manual_app.py`, `mcpserver_app.py` | +| `mcp_from_run(...)` | Converts a completed agent response into MCP result content blocks. | `manual_app.py`, `mcpserver_app.py` | | `AgentMCPTool` | Derives one native MCP tool from an agent and keeps schema, execution, and conversion aligned. | `agent_app.py`, `session_app.py` | | `WorkflowMCPTool` | Derives one native MCP tool from a workflow start executor and converts workflow outputs. | `workflow_app.py` | @@ -130,8 +130,10 @@ dependency set using PEP 723 inline script metadata. ## Common behavior -- **No framework choice:** the package does not select FastMCP, Starlette, +- **No framework choice:** the package does not select `MCPServer`, Starlette, Uvicorn, stdio, or streamable HTTP. +- **Protocol eras:** MCP SDK 2.2 serves modern 2026-07-28 requests and retains + legacy `initialize` compatibility. Modern server-to-client input uses MRTR. - **Chat options:** only explicitly selected MCP arguments are passed to the model client. The samples expose `reasoning_effort` as an example, but any option valid for the agent can be exposed. @@ -139,8 +141,8 @@ dependency set using PEP 723 inline script metadata. multimodal input content blocks. The samples do not present an application-specific image schema as protocol behavior. - **Streaming:** streamable HTTP can carry multiple MCP messages, but a tool - call still produces one final `CallToolResult`. Progress notifications and - experimental deferred tasks remain application-owned protocol features. + call still produces one final `CallToolResult`. The SDK does not yet expose + the 2026 Tasks extension runtime, so these samples do not implement deferred tasks. - **Authentication:** the local endpoints are intentionally unauthenticated. MCP transport session identifiers are not user authorization. Production servers must authenticate and authorize before loading tenant or user state. diff --git a/python/samples/04-hosting/mcp/fastmcp_app.py b/python/samples/04-hosting/mcp/mcpserver_app.py similarity index 89% rename from python/samples/04-hosting/mcp/fastmcp_app.py rename to python/samples/04-hosting/mcp/mcpserver_app.py index fc253e1b37a..d4b59d15557 100644 --- a/python/samples/04-hosting/mcp/fastmcp_app.py +++ b/python/samples/04-hosting/mcp/mcpserver_app.py @@ -7,13 +7,13 @@ # "mcp>=2.2.0,<3", # ] # /// -# Run with: uv run fastmcp_app.py +# Run with: uv run mcpserver_app.py # Copyright (c) Microsoft. All rights reserved. -"""Host an Agent Framework agent with FastMCP and the conversion helpers. +"""Host an Agent Framework agent with MCPServer and the conversion helpers. -FastMCP derives the native MCP tool schema from the decorated function +MCPServer derives the native MCP tool schema from the decorated function signature. The Agent Framework hosting package only converts the validated arguments and completed agent response at the protocol boundary. @@ -53,7 +53,7 @@ @asynccontextmanager async def lifespan(_server: MCPServer[None]) -> AsyncGenerator[None]: - """Close the model credential when the FastMCP server stops.""" + """Close the model credential when the MCP server stops.""" async with credential: yield @@ -74,7 +74,7 @@ async def run_agent( task: str, reasoning_effort: Literal["low", "medium", "high"] | None = None, ) -> types.CallToolResult: - """Run the agent with FastMCP-validated arguments.""" + """Run the agent with MCPServer-validated arguments.""" arguments: dict[str, object] = {"task": task} if reasoning_effort is not None: arguments["reasoning_effort"] = reasoning_effort From c3572caafd925c4da40374547689aa4b993a4099 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Thu, 8 Oct 2026 10:22:39 +0200 Subject: [PATCH 40/42] Test coverage for caching and documentation updates --- python/packages/core/AGENTS.md | 4 +- .../packages/core/agent_framework/_skills.py | 2 +- python/packages/core/tests/core/test_mcp.py | 431 ++++++++++-------- .../core/tests/core/test_mcp_http_auth.py | 18 +- .../core/tests/core/test_mcp_skills.py | 8 +- python/packages/core/tests/test_security.py | 27 +- .../_workflows/_executors_mcp.py | 4 +- .../_workflows/_mcp_handler.py | 4 +- .../tests/test_default_mcp_tool_handler.py | 5 +- .../_responses.py | 6 +- .../foundry_hosting/tests/test_responses.py | 19 +- .../claw_step04_production_ready/README.md | 2 +- 12 files changed, 302 insertions(+), 228 deletions(-) diff --git a/python/packages/core/AGENTS.md b/python/packages/core/AGENTS.md index 247fb6e8134..b8881d09714 100644 --- a/python/packages/core/AGENTS.md +++ b/python/packages/core/AGENTS.md @@ -232,8 +232,8 @@ The vector store API is experimental under the shared `VECTOR_STORES` feature ID - `max_task_wait: timedelta | None` — client-side deadline for the whole post-create lifecycle (poll + result fetch). When exceeded, raises `ToolExecutionException` and fires a best-effort `tasks/cancel`. `None` (default) means no client-side bound. Bounds sleeps, sends, AND reconnects via `asyncio.wait_for`. - **Permissive fallback**: servers that ignore the augmentation (return `CallToolResult` directly) or reject the unknown `task` field with `METHOD_NOT_FOUND` / `INVALID_PARAMS` fall back to the plain `session.call_tool(...)` path so legacy servers keep working. An unparseable success response (server accepted the augmented call but returned a payload that is neither `CreateTaskResult` nor `CallToolResult`) **does not** fall back — it raises `ToolExecutionException` to avoid double-executing a side-effecting tool. - **Submit-vs-track reconnect policy**: a dropped connection before a `task_id` is known raises `ToolExecutionException("connection lost; task state unknown")` without re-issuing the augmented `tools/call`, so a server that accepted the request but lost the response cannot be made to start the same operation twice; once a `task_id` exists, `tasks/get` / `tasks/result` reconnect once and retry against the same id (a shared `_send_with_one_reconnect` helper). -- **Cancel-on-abandonment vs terminal failure**: any path where the remote task may still be running (max-wait exceeded, hard `McpError` in poll, malformed `tasks/get`, second connection loss in poll/fetch, reconnect failure) fires best-effort `tasks/cancel` before raising. Terminal failures (`failed`/`cancelled`/`input_required` server-side, `completed+isError`, malformed `tasks/result` after server completed) do **not** cancel — the server is already done. `_MCPTaskAbandoned` is the private marker distinguishing the two. -- **Transient poll retry**: a slow `tasks/get` that surfaces as `McpError(code=408 REQUEST_TIMEOUT)` is retried (bounded by `max_task_wait`). All other non-connection `McpError`s during poll are treated as abandonment. `tasks/result` does not get transient retry — the server has already completed, so a slow payload fetch is anomalous. +- **Cancel-on-abandonment vs terminal failure**: any path where the remote task may still be running (max-wait exceeded, hard `MCPError` in poll, malformed `tasks/get`, second connection loss in poll/fetch, reconnect failure) fires best-effort `tasks/cancel` before raising. Terminal failures (`failed`/`cancelled`/`input_required` server-side, `completed+isError`, malformed `tasks/result` after server completed) do **not** cancel — the server is already done. `_MCPTaskAbandoned` is the private marker distinguishing the two. +- **Transient poll retry**: a slow `tasks/get` that surfaces as `MCPError(code=REQUEST_TIMEOUT)` is retried (bounded by `max_task_wait`). All other non-connection `MCPError`s during poll are treated as abandonment. `tasks/result` does not get transient retry — the server has already completed, so a slow payload fetch is anomalous. ### File Access Harness (`_harness/_file_access.py`) diff --git a/python/packages/core/agent_framework/_skills.py b/python/packages/core/agent_framework/_skills.py index b7e3f7ec46d..ac345893940 100644 --- a/python/packages/core/agent_framework/_skills.py +++ b/python/packages/core/agent_framework/_skills.py @@ -4441,7 +4441,7 @@ async def get_skills(self, context: SkillsSourceContext) -> list[Skill]: def _is_mcp_resource_not_found(ex: Exception) -> bool: - """Return ``True`` when *ex* is an :class:`McpError` indicating a missing resource. + """Return ``True`` when *ex* is an :class:`MCPError` indicating a missing resource. Two codes are treated as "not found": diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 7e566e7004b..1d1214940e3 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -33,7 +33,6 @@ FunctionTool, MCPStdioTool, MCPStreamableHTTPTool, - MCPWebsocketTool, Message, SupportsChatGetResponse, ) @@ -232,7 +231,7 @@ def _reset_progressive_mcp_warning_state() -> None: def _request_for_mcp_tool(tool: MCPStreamableHTTPTool, url: str = "http://example.com/mcp") -> Any: - import httpx + import httpx2 as httpx return httpx.Request( "POST", @@ -341,8 +340,9 @@ async def list_tools(params: types.PaginatedRequestParams | None = None) -> type return types.ListToolsResult(tools=advertised_tools[1:]) return types.ListToolsResult(tools=advertised_tools) - tool.session = AsyncMock() - tool.session.list_tools = AsyncMock(side_effect=list_tools) + mock_session = AsyncMock() + tool.session = mock_session + mock_session.list_tools = AsyncMock(side_effect=list_tools) with pytest.raises(ToolExecutionException, match="configuration name 'docs_search' is ambiguous"): await tool.load_tools() @@ -373,13 +373,14 @@ async def test_ambiguous_policy_reload_preserves_previous_discovery( tool = MCPTool( # type: ignore[abstract] # ty: ignore[call-non-callable] name="docs", tool_name_prefix="docs", allowed_tools=allowed_tools, approval_mode=approval_mode ) - tool.session = AsyncMock() + mock_session = AsyncMock() + tool.session = mock_session original = types.Tool( name=remote_names[0], input_schema={"type": "object", "properties": {"query": {"type": "string"}}}, _meta={"original": True}, ) - tool.session.list_tools = AsyncMock(return_value=types.ListToolsResult(tools=[original])) + mock_session.list_tools = AsyncMock(return_value=types.ListToolsResult(tools=[original])) await tool.load_tools() original_functions = list(tool._functions) original_meta = dict(tool._tool_call_meta_by_name) @@ -395,7 +396,7 @@ async def list_tools(params: types.PaginatedRequestParams | None = None) -> type ) return types.ListToolsResult(tools=[types.Tool(name=remote_names[1], input_schema={"type": "object"})]) - tool.session.list_tools = AsyncMock(side_effect=list_tools) + mock_session.list_tools = AsyncMock(side_effect=list_tools) with caplog.at_level(logging.WARNING, logger=logger.name): await tool.message_handler(types.ToolListChangedNotification()) await asyncio.gather(*tool._pending_reload_tasks) @@ -427,8 +428,9 @@ async def test_tool_refresh_accepts_unambiguous_rename( tool = MCPTool( # type: ignore[abstract] # ty: ignore[call-non-callable] name="docs", tool_name_prefix="docs", allowed_tools=allowed_tools, approval_mode=approval_mode ) - tool.session = AsyncMock() - tool.session.list_tools = AsyncMock( + mock_session = AsyncMock() + tool.session = mock_session + mock_session.list_tools = AsyncMock( side_effect=[ types.ListToolsResult(tools=[types.Tool(name=name, input_schema={"type": "object"})]) for name in remote_names @@ -440,17 +442,18 @@ async def test_tool_refresh_accepts_unambiguous_rename( assert [function.name for function in tool.functions] == [f"docs_{remote_names[1]}"] assert tool.functions[0].approval_mode == expected_approval - tool.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="ok")]) ) await tool.functions[0].invoke(arguments={}) - assert tool.session.call_tool.call_args.args[0] == remote_names[1] + assert mock_session.call_tool.call_args.args[0] == remote_names[1] @pytest.mark.parametrize("empty_snapshot", [False, True]) async def test_tool_refresh_replaces_snapshot_and_preserves_other_functions(empty_snapshot: bool) -> None: tool = MCPTool(name="docs", tool_name_prefix="docs") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = AsyncMock() + mock_session = AsyncMock() + tool.session = mock_session keep = types.Tool( name="keep", input_schema={"type": "object", "properties": {"query": {"type": "string"}}}, @@ -462,12 +465,12 @@ async def test_tool_refresh_replaces_snapshot_and_preserves_other_functions(empt _meta={"version": 1}, execution=types.ToolExecution(task_support="required"), ) - tool.session.list_tools = AsyncMock(return_value=types.ListToolsResult(tools=[keep, removed])) + mock_session.list_tools = AsyncMock(return_value=types.ListToolsResult(tools=[keep, removed])) await tool.load_tools() kept_function = tool._functions[0] parser = Mock(return_value="custom result") kept_function.result_parser = parser - tool.session.list_prompts = AsyncMock(return_value=types.ListPromptsResult(prompts=[types.Prompt(name="summary")])) + mock_session.list_prompts = AsyncMock(return_value=types.ListPromptsResult(prompts=[types.Prompt(name="summary")])) await tool.load_prompts() prompt_function = tool._functions[-1] custom_function = FunctionTool(name="custom", func=lambda: "custom") @@ -485,7 +488,7 @@ async def list_tools(params: types.PaginatedRequestParams | None = None) -> type assert params.cursor == "second" return types.ListToolsResult(tools=[] if empty_snapshot else [keep]) - tool.session.list_tools = AsyncMock(side_effect=list_tools) + mock_session.list_tools = AsyncMock(side_effect=list_tools) await tool.load_tools() assert tool._functions is original_list @@ -504,11 +507,12 @@ async def list_tools(params: types.PaginatedRequestParams | None = None) -> type async def test_tool_refresh_preserves_prompt_when_server_advertises_same_raw_name() -> None: tool = MCPTool(name="docs") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = AsyncMock() - tool.session.list_prompts = AsyncMock(return_value=types.ListPromptsResult(prompts=[types.Prompt(name="summary")])) + mock_session = AsyncMock() + tool.session = mock_session + mock_session.list_prompts = AsyncMock(return_value=types.ListPromptsResult(prompts=[types.Prompt(name="summary")])) await tool.load_prompts() prompt_function = tool._functions[0] - tool.session.list_tools = AsyncMock( + mock_session.list_tools = AsyncMock( side_effect=[ types.ListToolsResult(tools=[types.Tool(name="summary", input_schema={"type": "object"})]), types.ListToolsResult(tools=[]), @@ -578,8 +582,9 @@ async def test_overlapping_names_accept_unambiguous_policy( "never_require_approval": ["search", "docs_docs_search"], }, ) - tool.session = AsyncMock() - tool.session.list_tools = AsyncMock( + mock_session = AsyncMock() + tool.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool(name=name, input_schema={"type": "object", "properties": {}}) @@ -604,8 +609,9 @@ async def test_prompts_reject_ambiguous_policy_names(load_order: str) -> None: tool_name_prefix="docs", allowed_tools=["docs_search"], ) - tool.session = AsyncMock() - tool.session.list_prompts = AsyncMock( + mock_session = AsyncMock() + tool.session = mock_session + mock_session.list_prompts = AsyncMock( side_effect=[ types.ListPromptsResult(prompts=[types.Prompt(name="search")], next_cursor="second"), types.ListPromptsResult(prompts=[types.Prompt(name="docs_search")]), @@ -617,12 +623,12 @@ async def test_prompts_reject_ambiguous_policy_names(load_order: str) -> None: assert tool._functions == [] return - tool.session.list_tools = AsyncMock( + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[types.Tool(name="search", input_schema={"type": "object", "properties": {}})] ) ) - tool.session.list_prompts = AsyncMock( + mock_session.list_prompts = AsyncMock( return_value=types.ListPromptsResult(prompts=[types.Prompt(name="docs_search")]) ) first_load, second_load = ( @@ -1057,8 +1063,9 @@ async def test_generated_mcp_tool_preserves_complete_host_payload_once() -> None _meta={"widget": "image"}, ) tool = MCPTool(name="helper") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function_result = await _call_generated_mcp_tool(tool, "widget") assert function_result.items is not None @@ -1092,8 +1099,9 @@ async def test_generated_mcp_host_payload_replaces_duplicate_private_markers() - Content.from_text("two", additional_properties=stale), ], ) - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function_result = await _call_generated_mcp_tool(tool, "widget") @@ -1112,8 +1120,9 @@ async def test_custom_mcp_result_parser_preserves_direct_shape_and_generated_hos _meta={"source": "server"}, ) tool = MCPTool(name="helper", parse_tool_results=lambda _: "Custom model summary") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) direct_result = await tool.call_tool("widget") function_result = await _call_generated_mcp_tool(tool, "widget") @@ -1142,8 +1151,9 @@ async def test_oversized_mcp_host_payload_is_omitted_without_changing_model_resu parse_tool_results=lambda _: "Bounded model summary", max_host_payload_size_bytes=128, ) - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) with caplog.at_level(logging.WARNING): function_result = await _call_generated_mcp_tool(tool, "widget") @@ -1208,8 +1218,9 @@ async def test_generated_mcp_error_preserves_complete_host_payload_on_function_r _meta={"source": "server"}, ) tool = MCPTool(name="helper") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function_result = await _call_generated_mcp_tool(tool, "widget") host_payload = function_result.additional_properties[_MCP_TOOL_RESULT_HOST_PAYLOAD_KEY] @@ -1229,8 +1240,9 @@ async def test_generated_mcp_parser_failure_preserves_complete_host_payload_on_f _meta={"source": "server"}, ) tool = MCPTool(name="helper", parse_tool_results=_raise_result_parser) # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function_result = await _call_generated_mcp_tool(tool, "widget") @@ -1249,8 +1261,9 @@ async def test_direct_mcp_calls_do_not_materialize_host_payload(monkeypatch: pyt name="helper", parse_tool_results=lambda result: cast(types.TextContent, result.content[0]).text, ) - tool.session = Mock() - tool.session.call_tool = AsyncMock(side_effect=[success, error]) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(side_effect=[success, error]) def fail_if_captured(*_args: Any, **_kwargs: Any) -> Any: raise AssertionError("direct calls must not materialize Host metadata") @@ -1265,8 +1278,9 @@ def fail_if_captured(*_args: Any, **_kwargs: Any) -> Any: async def test_direct_mcp_parser_failure_does_not_materialize_host_payload(monkeypatch: pytest.MonkeyPatch) -> None: tool = MCPTool(name="helper", parse_tool_results=_raise_result_parser) # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() - tool.session.call_tool = AsyncMock( + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult( content=[types.TextContent(type="text", text="ok")], _meta={"source": "server"}, @@ -1290,8 +1304,9 @@ async def test_function_tool_result_parser_cannot_discard_mcp_host_payload() -> _meta={"source": "server"}, ) tool = MCPTool(name="helper", parse_tool_results=lambda _: "MCP parser projection") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function_result = await _call_generated_mcp_tool( tool, @@ -1318,8 +1333,9 @@ async def test_empty_custom_parser_projection_remains_empty(parser_layer: str) - name="helper", parse_tool_results=(lambda _: []) if parser_layer == "mcp" else None, ) - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function_result = await _call_generated_mcp_tool( tool, @@ -1342,8 +1358,9 @@ async def test_oversized_mcp_error_preserves_independently_bounded_meta() -> Non _meta={"source": "small"}, ) tool = MCPTool(name="helper", max_host_payload_size_bytes=128) # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function_result = await _call_generated_mcp_tool(tool, "widget") @@ -1387,8 +1404,9 @@ async def test_exhausted_aggregate_budget_rejects_meta_before_copy(monkeypatch: parse_tool_results=lambda _: "model projection", max_host_payload_size_bytes=128, ) - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) budget = _FunctionResultPayloadBudget() assert budget.reserve(128, 128) @@ -1410,8 +1428,9 @@ async def test_oversized_mcp_meta_is_omitted_from_host_and_model_items() -> None _meta={"large": "x" * 1024}, ) tool = MCPTool(name="helper", max_host_payload_size_bytes=128) # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function_result = await _call_generated_mcp_tool(tool, "widget") @@ -1433,8 +1452,9 @@ async def test_mcp_host_payload_survives_real_function_loop( _meta={"source": "server"}, ) tool = MCPTool(name="helper") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function = FunctionTool( name="widget", description="", @@ -1528,7 +1548,8 @@ async def test_mcp_host_payload_has_aggregate_request_budget( expected_markers: int, ) -> None: tool = MCPTool(name="helper", max_host_payload_size_bytes=size_limit) # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() + mock_session = Mock() + tool.session = mock_session async def call_tool(tool_name: str, **_kwargs: Any) -> types.CallToolResult: return types.CallToolResult( @@ -1537,7 +1558,7 @@ async def call_tool(tool_name: str, **_kwargs: Any) -> types.CallToolResult: _meta={"source": tool_name * 12}, ) - tool.session.call_tool = AsyncMock(side_effect=call_tool) + mock_session.call_tool = AsyncMock(side_effect=call_tool) functions = [ FunctionTool( name=name, @@ -1602,8 +1623,9 @@ async def test_secure_mcp_auto_hide_preserves_outer_host_payload() -> None: Content.from_text("untrusted payload", additional_properties={"_meta": result.meta}) ], ) - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function = FunctionTool( name="widget", description="", @@ -1659,8 +1681,9 @@ async def test_secure_mcp_builtin_parser_restricts_all_result_shapes(result_shap _meta={"ifc": {"integrity": "trusted", "confidentiality": "public"}}, ) tool = MCPTool(name="helper") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function = FunctionTool( name="widget", description="", @@ -1701,8 +1724,9 @@ async def test_secure_mcp_builtin_parser_honors_locally_trusted_server_ifc() -> _meta={"ifc": {"integrity": "trusted", "confidentiality": "public"}}, ) tool = MCPTool(name="helper") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function_result = await _call_generated_mcp_tool( tool, @@ -1730,8 +1754,9 @@ async def test_custom_mcp_parser_cannot_make_meta_authoritative() -> None: name="helper", parse_tool_results=lambda _: [Content.from_text("projection", additional_properties={"_meta": forged_meta})], ) - tool.session = Mock() - tool.session.call_tool = AsyncMock(return_value=mcp_result) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(return_value=mcp_result) function_result = await _call_generated_mcp_tool( tool, @@ -2362,9 +2387,10 @@ async def test_local_mcp_server_load_functions(): class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) + mock_session = Mock(spec=ClientSession) + self.session = mock_session # Mock tools list response - self.session.list_tools = AsyncMock( + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -2398,9 +2424,10 @@ async def test_local_mcp_server_load_prompts(): class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) + mock_session = Mock(spec=ClientSession) + self.session = mock_session # Mock prompts list response - self.session.list_prompts = AsyncMock( + mock_session.list_prompts = AsyncMock( return_value=types.ListPromptsResult( prompts=[ types.Prompt( @@ -2427,8 +2454,9 @@ async def test_mcp_tool_call_tool_with_meta_integration(): class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -2450,7 +2478,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri _meta={"executionTime": 1.5, "cost": {"usd": 0.002}, "isError": False, "toolVersion": "1.2.3"}, ) - self.session.call_tool = AsyncMock(return_value=tool_result) + mock_session.call_tool = AsyncMock(return_value=tool_result) def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: return None # type: ignore[return-value] # pyrefly: ignore[bad-return] # ty: ignore[invalid-return-type] @@ -2472,8 +2500,9 @@ async def test_local_mcp_server_function_execution(): class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -2488,7 +2517,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult( content=[types.TextContent(type="text", text="Tool executed successfully")] ) @@ -2512,8 +2541,9 @@ async def test_local_mcp_server_function_execution_with_nested_object(): class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -2534,7 +2564,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult( content=[types.TextContent(type="text", text='{"name": "John Doe", "id": 251}')] ) @@ -2565,8 +2595,9 @@ async def test_local_mcp_server_function_execution_error(): class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -2582,7 +2613,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ) ) # Mock a tool call that raises an MCP error - self.session.call_tool = AsyncMock(side_effect=MCPError(-1, "Tool execution failed")) + mock_session.call_tool = AsyncMock(side_effect=MCPError(-1, "Tool execution failed")) def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: return None # type: ignore[return-value] # pyrefly: ignore[bad-return] # ty: ignore[invalid-return-type] @@ -2607,12 +2638,13 @@ def __init__(self, **kwargs: Any) -> None: async def connect(self, *, reset: bool = False) -> None: self.connect_count += 1 - self.session = Mock(spec=ClientSession) + mock_session = Mock(spec=ClientSession) + self.session = mock_session self.sessions.append(self.session) if self.connect_count == 1: - self.session.call_tool = AsyncMock(side_effect=MCPError(-32000, "Session terminated")) + mock_session.call_tool = AsyncMock(side_effect=MCPError(-32000, "Session terminated")) else: - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="recovered")]) ) self.is_connected = True @@ -2636,8 +2668,9 @@ async def test_mcp_tool_call_tool_raises_on_is_error(): class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -2652,7 +2685,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult( content=[types.TextContent(type="text", text="Something went wrong")], is_error=True, @@ -2676,8 +2709,9 @@ async def test_mcp_tool_call_tool_succeeds_when_is_error_false(): class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -2692,7 +2726,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult( content=[types.TextContent(type="text", text="Success")], is_error=False, @@ -2726,8 +2760,9 @@ async def process(self, context: FunctionInvocationContext, call_next): class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -2742,7 +2777,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult( content=[types.TextContent(type="text", text="MCP error occurred")], is_error=True, @@ -2778,8 +2813,9 @@ async def test_local_mcp_server_prompt_execution(): class TestMCPTool(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_prompts = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_prompts = AsyncMock( return_value=types.ListPromptsResult( prompts=[ types.Prompt( @@ -2790,7 +2826,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.get_prompt = AsyncMock( + mock_session.get_prompt = AsyncMock( return_value=types.GetPromptResult( description="Generated prompt", messages=[ @@ -2841,8 +2877,9 @@ async def test_mcp_tool_approval_mode(approval_mode, expected_approvals): class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -2918,8 +2955,9 @@ async def test_mcp_tool_allowed_tools(allowed_tools, expected_count, expected_na class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -3075,9 +3113,10 @@ async def _load_progressive_test_server( approval_mode=approval_mode, use_progressive_disclosure=True, ) - server.session = AsyncMock() - server.session.list_tools = AsyncMock(return_value=_progressive_tool_list_page(tools=tools)) - server.session.call_tool = AsyncMock( + mock_session = AsyncMock() + server.session = mock_session + mock_session.list_tools = AsyncMock(return_value=_progressive_tool_list_page(tools=tools)) + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="ok")]) ) await server.load_tools() @@ -3625,8 +3664,10 @@ def test_local_mcp_streamable_http_tool_init(): def test_mcp_websocket_tool_is_deprecated() -> None: + from agent_framework import MCPWebsocketTool # ty: ignore[deprecated] + with pytest.warns(DeprecationWarning, match="MCP WebSocket transport was removed in MCP v2"): - tool = MCPWebsocketTool(name="test", url="ws://localhost:8080") # pyright: ignore[reportDeprecated] + tool = MCPWebsocketTool(name="test", url="ws://localhost:8080") # pyright: ignore[reportDeprecated] # ty: ignore[deprecated] with pytest.raises(RuntimeError, match="Use MCPStreamableHTTPTool instead"): tool.get_mcp_client() @@ -3807,8 +3848,8 @@ async def listen_context() -> AsyncIterator[AsyncIterator[ToolsListChanged | Pro sdk_client = _mock_sdk_client(capabilities=capabilities, protocol_version="2026-07-28") sdk_client.listen = Mock(return_value=listen_context()) tool = MCPStdioTool(name="test_tool", command="unused") - tool.load_tools = load_tools # type: ignore[method-assign] - tool.load_prompts = load_prompts # type: ignore[method-assign] + tool.load_tools = load_tools # type: ignore[method-assign] # ty: ignore[invalid-assignment] + tool.load_prompts = load_prompts # type: ignore[method-assign] # ty: ignore[invalid-assignment] with patch("mcp.Client", return_value=sdk_client): async with tool: @@ -3974,8 +4015,8 @@ async def unsupported_listen() -> AsyncIterator[Any]: sdk_client = _mock_sdk_client(capabilities=capabilities) sdk_client.listen = Mock(return_value=unsupported_listen()) tool = MCPStdioTool(name="test_tool", command="unused") - tool.load_tools = load_tools # type: ignore[method-assign] - tool.load_prompts = load_prompts # type: ignore[method-assign] + tool.load_tools = load_tools # type: ignore[method-assign] # ty: ignore[invalid-assignment] + tool.load_prompts = load_prompts # type: ignore[method-assign] # ty: ignore[invalid-assignment] with patch("mcp.Client", return_value=sdk_client): async with tool: @@ -4963,7 +5004,9 @@ async def test_connect_retains_and_close_clears_sdk_client() -> None: with pytest.raises(RuntimeError, match="framework-owned MCP Client"): tool.session = None await tool.connect(reset=True) - assert tool._connection.client is sdk_client + reset_connection = tool._connection + assert reset_connection is not None + assert reset_connection.client is sdk_client assert tool.session is sdk_client.session finally: await tool.close() @@ -5079,12 +5122,12 @@ async def test_catalog_loading_uses_sdk_client_cache() -> None: ] assert [awaited.kwargs for awaited in sdk_client.list_tools.await_args_list] == [ - {"cursor": None, "cache_mode": "use"}, - {"cursor": "tools-next", "cache_mode": "use"}, + {"cursor": None}, + {"cursor": "tools-next"}, ] assert [awaited.kwargs for awaited in sdk_client.list_prompts.await_args_list] == [ - {"cursor": None, "cache_mode": "use"}, - {"cursor": "prompts-next", "cache_mode": "use"}, + {"cursor": None}, + {"cursor": "prompts-next"}, ] session.list_tools.assert_not_awaited() session.list_prompts.assert_not_awaited() @@ -5313,6 +5356,7 @@ async def handle(request: Request) -> Response: body = json.loads(request.content) method = body["method"] principal = request.headers.get("Authorization", "") + result: dict[str, Any] if method == "server/discover": result = { "supportedVersions": ["2026-07-28"], @@ -6730,8 +6774,9 @@ async def test_mcp_tool_call_tool_requires_loaded_tools() -> None: async def test_generated_mcp_function_ignores_model_supplied_remote_tool_name() -> None: """A model-supplied argument must not be able to redirect the call to another remote tool.""" tool = MCPTool(name="test_tool") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock(spec=ClientSession) - tool.session.list_tools = AsyncMock( # ty: ignore[unresolved-attribute] + mock_session = Mock(spec=ClientSession) + tool.session = mock_session + mock_session.list_tools = AsyncMock( # ty: ignore[unresolved-attribute] return_value=types.ListToolsResult( tools=[ types.Tool( @@ -6755,7 +6800,7 @@ async def test_generated_mcp_function_ignores_model_supplied_remote_tool_name() ] ) ) - tool.session.call_tool = AsyncMock( # ty: ignore[unresolved-attribute] + mock_session.call_tool = AsyncMock( # ty: ignore[unresolved-attribute] return_value=types.CallToolResult(content=[types.TextContent(type="text", text="ok")]) ) @@ -6766,8 +6811,8 @@ async def test_generated_mcp_function_ignores_model_supplied_remote_tool_name() arguments={"query": "quarterly report", "_remote_tool_name": "delete_repo", "repo": "corp/prod"} ) - tool.session.call_tool.assert_awaited_once() # ty: ignore[unresolved-attribute] - await_args = tool.session.call_tool.await_args # ty: ignore[unresolved-attribute] + mock_session.call_tool.assert_awaited_once() # ty: ignore[unresolved-attribute] + await_args = mock_session.call_tool.await_args # ty: ignore[unresolved-attribute] assert await_args is not None assert await_args.args[0] == "search_docs" assert await_args.kwargs["arguments"] == {"query": "quarterly report"} @@ -6816,9 +6861,10 @@ async def test_mcp_tool_get_prompt_raises_after_reconnection_still_fails() -> No async def test_mcp_tool_wraps_unexpected_call_tool_and_get_prompt_errors() -> None: tool = MCPTool(name="test_tool", load_tools=True, load_prompts=True) # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock() - tool.session.call_tool = AsyncMock(side_effect=RuntimeError("tool boom")) - tool.session.get_prompt = AsyncMock(side_effect=RuntimeError("prompt boom")) + mock_session = Mock() + tool.session = mock_session + mock_session.call_tool = AsyncMock(side_effect=RuntimeError("tool boom")) + mock_session.get_prompt = AsyncMock(side_effect=RuntimeError("prompt boom")) with pytest.raises(ToolExecutionException, match="Failed to call tool 'remote_tool'"): await tool.call_tool("remote_tool") @@ -7788,13 +7834,14 @@ async def test_connect_sets_logging_level_when_server_advertises_logging() -> No async def test_ensure_connected_skips_future_pings_when_ping_is_not_available() -> None: tool = MCPTool(name="test_tool") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = Mock(send_ping=AsyncMock(side_effect=MCPError(-32601, "Method 'ping' is not available."))) + mock_session = Mock(send_ping=AsyncMock(side_effect=MCPError(-32601, "Method 'ping' is not available."))) + tool.session = mock_session with patch.object(tool, "_reconnect_without_loading", AsyncMock()) as mock_reconnect: await tool._ensure_connected() await tool._ensure_connected() - tool.session.send_ping.assert_awaited_once() + mock_session.send_ping.assert_awaited_once() mock_reconnect.assert_not_awaited() assert tool._ping_available is False @@ -7888,8 +7935,9 @@ async def test_mcp_tool_filters_framework_kwargs(): class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -7905,7 +7953,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ) ) # Mock call_tool to capture the arguments it receives - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="Success")]) ) @@ -7972,8 +8020,9 @@ async def test_mcp_tool_call_tool_otel_meta(use_span, expect_traceparent, span_e class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -7988,7 +8037,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="result")]) ) @@ -8033,8 +8082,9 @@ async def test_mcp_tool_call_tool_forwards_tool_list_meta(): class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -8050,10 +8100,10 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="result")]) ) - self.session.list_prompts = AsyncMock(return_value=types.ListPromptsResult(prompts=[])) + mock_session.list_prompts = AsyncMock(return_value=types.ListPromptsResult(prompts=[])) def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: return None # type: ignore[return-value] # pyrefly: ignore[bad-return] # ty: ignore[invalid-return-type] @@ -8078,8 +8128,9 @@ async def test_mcp_tool_call_tool_user_meta_merges_with_tool_list_meta(): class TestServer(MCPTool): async def connect(self) -> None: # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -8091,7 +8142,7 @@ async def connect(self) -> None: # type: ignore[override] # pyrefly: ignore[ba ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="result")]) ) @@ -8121,8 +8172,9 @@ async def test_mcp_tool_function_invocation_strips_model_supplied_meta() -> None class TestServer(MCPTool): async def connect(self) -> None: # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -8133,7 +8185,7 @@ async def connect(self) -> None: # type: ignore[override] # pyrefly: ignore[ba ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="result")]) ) @@ -8165,8 +8217,9 @@ async def test_mcp_tool_function_invocation_preserves_trusted_meta_over_model_me class TestServer(MCPTool): async def connect(self) -> None: # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -8177,7 +8230,7 @@ async def connect(self) -> None: # type: ignore[override] # pyrefly: ignore[ba ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="result")]) ) @@ -8216,8 +8269,9 @@ async def test_mcp_tool_call_tool_otel_meta_overrides_user_meta_but_not_tool_lis class TestServer(MCPTool): async def connect(self) -> None: # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -8229,7 +8283,7 @@ async def connect(self) -> None: # type: ignore[override] # pyrefly: ignore[ba ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="result")]) ) @@ -8335,7 +8389,9 @@ def provider(kwargs): await server.call_tool("greet", name="Alice", some_token="my-secret") # Verify the high-level Client.call_tool was called. - sdk_client = server._connection.client + connection = server._connection + assert connection is not None + sdk_client = connection.client assert sdk_client is not None cast(AsyncMock, sdk_client.call_tool).assert_awaited_once() @@ -8470,7 +8526,9 @@ def get_mcp_client(self): # pyrefly: ignore[bad-override] async with server: await server.load_tools() await server.call_tool("greet", name="Alice") - sdk_client = server._connection.client + connection = server._connection + assert connection is not None + sdk_client = connection.client assert sdk_client is not None cast(AsyncMock, sdk_client.call_tool).assert_awaited_once() @@ -8689,7 +8747,7 @@ def failing_provider(kw: dict[str, Any]) -> dict[str, str]: async def test_mcp_streamable_http_tool_header_provider_skips_cross_origin_redirect(): """The request hook must not re-add caller headers after a cross-origin redirect.""" - import httpx + import httpx2 as httpx from agent_framework._mcp import _mcp_call_headers @@ -8735,7 +8793,7 @@ async def test_mcp_streamable_http_tool_header_provider_skips_cross_origin_redir async def test_mcp_streamable_http_tool_keeps_bound_headers_on_same_origin_redirect(): """A redirected request must retain the header set bound to its session.""" - import httpx + import httpx2 as httpx provider_headers = {"X-Previous": "old"} tool = MCPStreamableHTTPTool( @@ -8784,7 +8842,7 @@ async def test_mcp_streamable_http_tool_keeps_bound_headers_on_same_origin_redir @pytest.mark.parametrize("use_header_provider", [False, True]) async def test_mcp_streamable_http_tool_header_provider_with_user_httpx_client(use_header_provider: bool): """Supplied clients preserve configuration, cookie persistence, and ownership.""" - import httpx + import httpx2 as httpx from agent_framework._mcp import _mcp_call_headers @@ -8835,7 +8893,7 @@ def handle(request: httpx.Request) -> httpx.Response: async def test_mcp_streamable_http_tool_header_provider_isolated_on_shared_httpx_client(): """Each MCP transport must use its own headers when sharing an httpx client.""" - import httpx + import httpx2 as httpx captured_headers: list[dict[str, str]] = [] @@ -8893,7 +8951,7 @@ async def handler(request: httpx.Request) -> httpx.Response: async def test_mcp_streamable_http_tool_removes_header_hook_on_close(): """Closing one tool must remove only its hook, and reconnecting must restore it.""" - import httpx + import httpx2 as httpx user_client = httpx.AsyncClient() tool_a = MCPStreamableHTTPTool( @@ -8930,7 +8988,7 @@ async def test_mcp_streamable_http_tool_removes_header_hook_on_close(): async def test_mcp_streamable_http_tool_removes_hook_without_mutating_active_hook_list(): """Closing one tool must not disrupt an in-progress iteration over shared hooks.""" - import httpx + import httpx2 as httpx from agent_framework._mcp import _MCP_INJECTED_HEADER_KEYS_EXTENSION @@ -8994,7 +9052,7 @@ async def delayed_hook(_request: httpx.Request) -> None: async def test_mcp_streamable_http_tool_keeps_header_hook_until_cancelled_close_finishes(): """Caller cancellation must not remove the hook while lifecycle teardown continues.""" - import httpx + import httpx2 as httpx user_client = httpx.AsyncClient() tool = MCPStreamableHTTPTool( @@ -9045,7 +9103,7 @@ async def delayed_close() -> None: async def test_mcp_header_scoped_client_tags_send_requests(): """The transport wrapper must identify requests sent through AsyncClient.send.""" - import httpx + import httpx2 as httpx from agent_framework._mcp import _MCP_HEADER_OWNER_EXTENSION, _MCPHeaderScopedClient @@ -9068,7 +9126,7 @@ async def handle(request: httpx.Request) -> httpx.Response: async def test_mcp_header_scoped_client_delegates_unwrapped_attributes(): """The transport wrapper must stay a drop-in for the caller's httpx client.""" - import httpx + import httpx2 as httpx from agent_framework._mcp import _MCPHeaderScopedClient @@ -9169,7 +9227,9 @@ def provider(kwargs): assert provider_received[0]["some_token"] == "my-secret" # Verify Client.call_tool was called with the tool arguments (not the runtime kwargs). - sdk_client = server._connection.client + connection = server._connection + assert connection is not None + sdk_client = connection.client assert sdk_client is not None call_tool = cast(AsyncMock, sdk_client.call_tool) call_tool.assert_awaited_once() @@ -9187,7 +9247,7 @@ async def test_agent_run_supplies_mcp_connect_headers( in-process mock HTTP endpoint and asserts that header_provider can use those credentials on the initialize request, before any tool invocation occurs. """ - import httpx + import httpx2 as httpx captured_requests: list[tuple[str, str, dict[str, str]]] = [] @@ -9601,7 +9661,7 @@ async def test_agent_context_manager_authenticates_connect_with_closure_provider Pins that ``header_provider`` already covers construction-time credentials: entering the agent context connects before any run exists, and the server rejects unauthenticated calls. """ - import httpx + import httpx2 as httpx captured_requests: list[tuple[str, dict[str, str]]] = [] @@ -9669,7 +9729,7 @@ async def test_constructor_supplied_mcp_tool_uses_run_credentials_on_lazy_connec Without the agent context manager the handshake is deferred to ``run()``, so the run's credentials are available and must reach ``header_provider``. """ - import httpx + import httpx2 as httpx captured_requests: list[tuple[str, dict[str, str]]] = [] @@ -9739,7 +9799,7 @@ async def test_mcp_streamable_http_tool_header_provider_applies_across_transport drives the real transport against an in-process mock server and asserts the per-call Authorization header arrives on the tools/call HTTP request. """ - import httpx + import httpx2 as httpx captured_requests: list[tuple[str, str, dict[str, str]]] = [] @@ -9866,8 +9926,10 @@ async def blocking_call_tool(tool_name, *, arguments=None, meta=None): class _TestServer(MCPStreamableHTTPTool): async def connect(self, *, reset: bool = False) -> None: - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + sdk_client = AsyncMock() + sdk_client.session = mock_session + sdk_client.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -9878,8 +9940,9 @@ async def connect(self, *, reset: bool = False) -> None: ] ) ) - self.session.call_tool = AsyncMock(side_effect=blocking_call_tool) - self.session.send_ping = AsyncMock() + sdk_client.call_tool = AsyncMock(side_effect=blocking_call_tool) + mock_session.send_ping = AsyncMock() + self._connection = _ClientMCPConnection(sdk_client) self.is_connected = True async def _reconnect_for_identity_change(self) -> None: @@ -10042,7 +10105,8 @@ async def test_task_options_rejects_non_positive_default_ttl() -> None: async def test_load_tools_captures_task_support() -> None: tool = MCPTool(name="lro") # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = AsyncMock() + mock_session = AsyncMock() + tool.session = mock_session tool.load_tools_flag = True page = Mock() @@ -10060,7 +10124,7 @@ async def test_load_tools_captures_task_support() -> None: ), ] page.next_cursor = None - tool.session.list_tools = AsyncMock(return_value=page) + mock_session.list_tools = AsyncMock(return_value=page) await tool.load_tools() @@ -10117,10 +10181,11 @@ async def call_tool_as_task(self, tool_name: str, **kwargs: Any) -> str | list[C return await super().call_tool_as_task(tool_name, **kwargs) tool = OverriddenTaskTool() # type: ignore[abstract] # ty: ignore[call-non-callable] - tool.session = AsyncMock(spec=ClientSession) + mock_session = AsyncMock(spec=ClientSession) + tool.session = mock_session tool._tool_task_support_by_name["slow_op"] = "required" fallback_result = types.CallToolResult(content=[types.TextContent(type="text", text="fallback")]) - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] + mock_session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] return_value=types.Result.model_validate(fallback_result.model_dump(by_alias=True, exclude_none=True)) ) @@ -10849,7 +10914,7 @@ async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> An async def test_call_tool_as_task_poll_transient_request_timeout_keeps_polling( monkeypatch: pytest.MonkeyPatch, ) -> None: - import httpx + import httpx2 as httpx from agent_framework import _mcp as _mcp_module @@ -11368,8 +11433,9 @@ async def test_call_tool_forwards_only_declared_arguments() -> None: class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -11384,7 +11450,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="ok")]) ) @@ -11422,8 +11488,9 @@ async def test_call_tool_forwards_runtime_kwargs_the_server_declares() -> None: class TestServer(MCPTool): async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-override] # ty: ignore[invalid-method-override] - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + self.session = mock_session + mock_session.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -11442,7 +11509,7 @@ async def connect(self): # type: ignore[override] # pyrefly: ignore[bad-overri ] ) ) - self.session.call_tool = AsyncMock( + mock_session.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="ok")]) ) @@ -11482,8 +11549,10 @@ async def test_header_provider_reading_contextvar_keeps_credential_out_of_argume class TestServer(MCPStreamableHTTPTool): async def connect(self, *, reset: bool = False) -> None: - self.session = Mock(spec=ClientSession) - self.session.list_tools = AsyncMock( + mock_session = Mock(spec=ClientSession) + sdk_client = AsyncMock() + sdk_client.session = mock_session + sdk_client.list_tools = AsyncMock( return_value=types.ListToolsResult( tools=[ types.Tool( @@ -11498,10 +11567,11 @@ async def connect(self, *, reset: bool = False) -> None: ] ) ) - self.session.call_tool = AsyncMock( + sdk_client.call_tool = AsyncMock( return_value=types.CallToolResult(content=[types.TextContent(type="text", text="sunny")]) ) - self.session.send_ping = AsyncMock() + mock_session.send_ping = AsyncMock() + self._connection = _ClientMCPConnection(sdk_client) self.is_connected = True async def _reconnect_for_identity_change(self) -> None: @@ -11524,7 +11594,10 @@ def provider(_kwargs: dict[str, Any]) -> dict[str, str]: context = FunctionInvocationContext(function=tool, arguments={"city": "Seattle"}, kwargs={}) await tool.invoke(arguments={"city": "Seattle"}, context=context) - _, call_kwargs = server.session.call_tool.call_args # type: ignore[union-attr] # ty: ignore[unresolved-attribute] + connection = server._connection + assert connection is not None + assert connection.client is not None + _, call_kwargs = cast(AsyncMock, connection.client.call_tool).call_args assert seen_headers[-1] == {"Authorization": "Bearer secret-1"} assert call_kwargs["arguments"] == {"city": "Seattle"} @@ -11543,7 +11616,7 @@ def provider(_kwargs: dict[str, Any]) -> dict[str, str]: def _unauthorized_http_client() -> Any: """An HTTP client whose every response is 401, so `initialize` always fails.""" - import httpx + import httpx2 as httpx return httpx.AsyncClient(transport=httpx.MockTransport(lambda request: httpx.Response(401, request=request))) diff --git a/python/packages/core/tests/core/test_mcp_http_auth.py b/python/packages/core/tests/core/test_mcp_http_auth.py index f0d9059e6ed..31eb3a80b80 100644 --- a/python/packages/core/tests/core/test_mcp_http_auth.py +++ b/python/packages/core/tests/core/test_mcp_http_auth.py @@ -11,7 +11,7 @@ from typing import Any, Literal, TypeAlias from unittest.mock import AsyncMock, Mock, patch -import httpx +import httpx2 as httpx import pytest from mcp.client.session import ClientSession @@ -182,7 +182,7 @@ def create_owned_client(**kwargs: Any) -> httpx.AsyncClient: load_prompts=False, header_provider=lambda _: {"Authorization": principal.get()}, ) - with patch("httpx.AsyncClient", side_effect=create_owned_client): + with patch("httpx2.AsyncClient", side_effect=create_owned_client): async with tool: await tool.call_tool("record") response_cookies.clear() @@ -303,7 +303,7 @@ async def transport(**kwargs: Any) -> AsyncGenerator[tuple[()]]: else patch("agent_framework._mcp.streamable_http_client", side_effect=transport) ) try: - with transport_patch, patch("httpx.AsyncClient", return_value=client), pytest.raises(error): + with transport_patch, patch("httpx2.AsyncClient", return_value=client), pytest.raises(error): await tool.connect() assert client.event_hooks["request"] == original_hooks assert client.is_closed is owned_client @@ -336,7 +336,7 @@ def create_owned_client(**kwargs: Any) -> httpx.AsyncClient: if header_source == "provider" else None, ) - with patch("httpx.AsyncClient", side_effect=create_owned_client): + with patch("httpx2.AsyncClient", side_effect=create_owned_client): async with tool: await tool.call_tool("record") await tool.call_tool("record") @@ -844,7 +844,7 @@ async def test_discovery_failure_cleans_up_resources( failure = failure_type("discovery failed") try: with ( - patch("httpx.AsyncClient", return_value=client), + patch("httpx2.AsyncClient", return_value=client), patch.object(tool, discovery_method, new=AsyncMock(side_effect=failure)), pytest.raises(failure_type, match="discovery failed") as error, ): @@ -928,12 +928,12 @@ def create_client(**kwargs: Any) -> httpx.AsyncClient: http_client=None if owned_client else create_client(), header_provider=lambda _: {"Authorization": "token-a"}, ) - from mcp.shared.exceptions import McpError + from mcp.shared.exceptions import MCPError try: - with patch("httpx.AsyncClient", side_effect=create_client): + with patch("httpx2.AsyncClient", side_effect=create_client): for _ in range(2): - with pytest.raises(McpError, match="discovery failed"): + with pytest.raises(MCPError, match="discovery failed"): await tool.connect() assert tool.session is None assert not tool.is_connected @@ -996,7 +996,7 @@ async def enter() -> None: async with tool: pytest.fail("Cancelled setup must not enter the context manager") - with patch("httpx.AsyncClient", return_value=client), patch.object(tool, "_close_on_owner", record_cleanup): + with patch("httpx2.AsyncClient", return_value=client), patch.object(tool, "_close_on_owner", record_cleanup): caller = asyncio.create_task(enter()) try: await asyncio.wait_for(setup_started.wait(), timeout=5) diff --git a/python/packages/core/tests/core/test_mcp_skills.py b/python/packages/core/tests/core/test_mcp_skills.py index a878cc9974e..1f8fa73d450 100644 --- a/python/packages/core/tests/core/test_mcp_skills.py +++ b/python/packages/core/tests/core/test_mcp_skills.py @@ -94,8 +94,8 @@ def _modern_resource_not_found(uri: str) -> MCPError: def _make_client( - *, resource_not_found_error: Callable[[str], MCPError] = _legacy_resource_not_found, + /, **read_resource_responses: ReadResourceResult, ) -> AsyncMock: """Create a mock ClientSession whose read_resource returns different results per URI. @@ -786,7 +786,7 @@ async def test_archive_missing_resource_is_skipped( ], }) client = _make_client( - resource_not_found_error=resource_not_found_error, + resource_not_found_error, **{"skill://index.json": _make_text_result(index_json, uri="skill://index.json")}, ) source = MCPSkillsSource(client=client) @@ -910,7 +910,7 @@ async def test_index_resource_not_found_returns_empty( resource_not_found_error: Callable[[str], MCPError], ) -> None: """Either valid missing-resource error shape means the server has no skill index.""" - client = _make_client(resource_not_found_error=resource_not_found_error) + client = _make_client(resource_not_found_error) source = MCPSkillsSource(client=client) skills = await source.get_skills(_SOURCE_CTX) assert skills == [] @@ -989,7 +989,7 @@ async def test_get_resource_not_found_returns_none( """Either valid missing-resource error shape on get_resource returns None.""" from agent_framework import SkillFrontmatter - client = _make_client(resource_not_found_error=resource_not_found_error) + client = _make_client(resource_not_found_error) fm = SkillFrontmatter(name="test-skill", description="Test.") skill = MCPSkill(frontmatter=fm, skill_md_uri="skill://test/SKILL.md", client=client) result = await skill.get_resource("references/file.md") diff --git a/python/packages/core/tests/test_security.py b/python/packages/core/tests/test_security.py index c7039c41f1f..683846d094d 100644 --- a/python/packages/core/tests/test_security.py +++ b/python/packages/core/tests/test_security.py @@ -10,7 +10,7 @@ from datetime import timedelta from types import MappingProxyType, SimpleNamespace from typing import Annotated, Any, cast -from unittest.mock import AsyncMock, Mock, call +from unittest.mock import AsyncMock, call import pytest from pydantic import AfterValidator, BaseModel, field_validator @@ -6141,7 +6141,7 @@ async def overlapping_call(_tool: MCPTool, tool_name: str, **_kwargs: Any) -> st async def test_headers_are_sent_only_to_the_configured_origin(self) -> None: from unittest.mock import patch - import httpx + import httpx2 as httpx from agent_framework._mcp import _MCPHeaderScopedClient from agent_framework.security import SecureMCPToolProxy @@ -6215,7 +6215,7 @@ def create_client(*_args: Any, **kwargs: Any) -> httpx.AsyncClient: client.follow_redirects = kwargs.get("follow_redirects", client.follow_redirects) return client - with patch("httpx.AsyncClient", side_effect=create_client): + with patch("httpx2.AsyncClient", side_effect=create_client): proxy = SecureMCPToolProxy( url="https://mcp.example/mcp", headers=configured_headers, @@ -6258,7 +6258,7 @@ def test_headers_require_an_absolute_http_origin(self, url: str) -> None: from agent_framework.security import SecureMCPToolProxy - with patch("httpx.AsyncClient") as create_client: + with patch("httpx2.AsyncClient") as create_client: proxy = SecureMCPToolProxy(url=url, headers={"X-Custom-Credential": "custom-value"}) with pytest.raises(ValueError, match="absolute HTTP.*URL with a host"): @@ -6345,8 +6345,9 @@ async def fake_call(**kwargs: Any) -> list[Content]: ) mcp_tool = MCPTool(name="helper") # type: ignore[abstract] # ty: ignore[call-non-callable] mcp_tool.is_connected = True - mcp_tool.session = AsyncMock() - mcp_tool.session.list_tools = AsyncMock( # type: ignore[method-assign] + mock_session = AsyncMock() + mcp_tool.session = mock_session + mock_session.list_tools = AsyncMock( # type: ignore[method-assign] return_value=SimpleNamespace( tools=[SimpleNamespace(name="remote_tool", annotations=annotations)], next_cursor=None, @@ -6390,8 +6391,9 @@ def get_mcp_client(self): always_load=always_load, ) mcp_tool.is_connected = True - mcp_tool.session = AsyncMock() - mcp_tool.session.call_tool = AsyncMock( + mock_session = AsyncMock() + mcp_tool.session = mock_session + mock_session.call_tool = AsyncMock( return_value=mcp_types.CallToolResult( content=[mcp_types.TextContent(type="text", text="payload")], _meta=result_meta or {"ifc": {"integrity": "trusted", "confidentiality": "private"}}, @@ -6922,7 +6924,8 @@ async def test_apply_mcp_security_labels_uses_high_level_client_connection(self) annotations = SimpleNamespace(read_only_hint=True, open_world_hint=False) mcp_tool, _ = _make_connected_mcp_tool_for_ifc(annotations=annotations, server_meta={}) sdk_client = AsyncMock() - sdk_client.session = AsyncMock() + mock_session = AsyncMock() + sdk_client.session = mock_session sdk_client.list_tools.side_effect = [ SimpleNamespace( tools=[SimpleNamespace(name="remote_tool", annotations=annotations)], @@ -6935,10 +6938,10 @@ async def test_apply_mcp_security_labels_uses_high_level_client_connection(self) await apply_mcp_security_labels(mcp_tool) assert sdk_client.list_tools.await_args_list == [ - call(cursor=None, cache_mode="use"), - call(cursor="next-page", cache_mode="use"), + call(cursor=None), + call(cursor="next-page"), ] - sdk_client.session.list_tools.assert_not_awaited() + mock_session.list_tools.assert_not_awaited() async def test_framework_stamped_mcp_label_remains_authoritative_through_tracking(self) -> None: from agent_framework.security import apply_mcp_security_labels diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py index 7d3647d00bd..99bbb7b8769 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py @@ -503,10 +503,10 @@ async def _invoke_with_narrow_catch(self, invocation: MCPToolInvocation) -> MCPT ) except Exception as exc: try: - from mcp.shared.exceptions import McpError + from mcp.shared.exceptions import MCPError except ImportError: # pragma: no cover - mcp is a hard dep raise - if isinstance(exc, McpError): + if isinstance(exc, MCPError): message = str(exc) or type(exc).__name__ return MCPToolResult( outputs=[Content.from_text(f"Error: {message}")], diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py b/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py index 64eb59750fb..88fd2eded93 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py @@ -345,10 +345,10 @@ async def _invoke_entry(self, entry: _CacheEntry, invocation: MCPToolInvocation) # Be defensive about MCP errors that may bubble up without being # wrapped in ToolExecutionException by custom parsers. try: - from mcp.shared.exceptions import McpError + from mcp.shared.exceptions import MCPError except ImportError: # pragma: no cover - mcp is a hard dep but stay defensive raise - if isinstance(exc, McpError): + if isinstance(exc, MCPError): message = str(exc) or type(exc).__name__ return MCPToolResult( outputs=[Content.from_text(f"Error: {message}")], diff --git a/python/packages/declarative/tests/test_default_mcp_tool_handler.py b/python/packages/declarative/tests/test_default_mcp_tool_handler.py index 0826302cf69..b3d8d7d51ae 100644 --- a/python/packages/declarative/tests/test_default_mcp_tool_handler.py +++ b/python/packages/declarative/tests/test_default_mcp_tool_handler.py @@ -581,10 +581,9 @@ async def connect(tool: FakeTool) -> None: @pytest.mark.parametrize("tool_name", ["search", "tools/list"]) async def test_mcp_error_mapping_cleans_up(self, tool_name: str) -> None: - from mcp.shared.exceptions import McpError - from mcp.types import ErrorData + from mcp.shared.exceptions import MCPError - error = McpError(ErrorData(code=-32603, message="operation stopped")) + error = MCPError(-32603, "operation stopped") with ( _patch_tool(), patch.object(FakeTool, "call_tool", new_callable=AsyncMock, side_effect=error), diff --git a/python/packages/foundry_hosting/agent_framework_foundry_hosting/_responses.py b/python/packages/foundry_hosting/agent_framework_foundry_hosting/_responses.py index 1d4d684e8ac..0fb6407ced7 100644 --- a/python/packages/foundry_hosting/agent_framework_foundry_hosting/_responses.py +++ b/python/packages/foundry_hosting/agent_framework_foundry_hosting/_responses.py @@ -97,7 +97,7 @@ TextContentBuilder, ) from azure.ai.agentserver.responses.streaming._checkpoint import ResponseCheckpointEvent -from mcp import McpError +from mcp import MCPError from typing_extensions import Any from ._agent_source import is_agent, resolve_agent, validate_agent_source @@ -574,8 +574,8 @@ def consent_url_from_error(exc: BaseException) -> list[ConsentError] | None: Returns: The consent URL(s) extracted from the error, or ``None`` if no consent error was found. """ - inner_exception = next((arg for arg in exc.args if isinstance(arg, McpError)), None) - if inner_exception is not None and inner_exception.error.code == CONSENT_ERROR_CODE: + inner_exception = next((arg for arg in exc.args if isinstance(arg, MCPError)), None) + if inner_exception is not None and inner_exception.code == CONSENT_ERROR_CODE: # Parse the error message # The error message is structured with the following format: # "tools/list failed for 1 tool source(s), succeeded for 0 tool source(s) {"errors":[{"name": ..." diff --git a/python/packages/foundry_hosting/tests/test_responses.py b/python/packages/foundry_hosting/tests/test_responses.py index d3c518dfe98..99407948174 100644 --- a/python/packages/foundry_hosting/tests/test_responses.py +++ b/python/packages/foundry_hosting/tests/test_responses.py @@ -79,8 +79,7 @@ ResponseObject, ) from azure.ai.agentserver.responses.streaming._checkpoint import ResponseCheckpointEvent -from mcp import McpError -from mcp.types import ErrorData +from mcp import MCPError from openai import AsyncOpenAI, DefaultAsyncHttpxClient from openai.types.responses.response_input_item_param import ResponseInputItemParam from openai.types.responses.response_usage import ResponseUsage as OpenAIResponseUsage @@ -7497,12 +7496,12 @@ def _make_consent_error( """Build an exception wrapping a Foundry MCP gateway consent error. Mirrors the real-world wrapping produced by ``MCPStreamableHTTPTool.__aenter__``, - which catches connection-time ``McpError``s and re-raises them as a + which catches connection-time ``MCPError``s and re-raises them as a ``ToolExecutionException`` (an ``AgentFrameworkException`` subclass) with the original error attached via ``inner_exception``. ``consent_url_from_error`` - then finds the wrapped ``McpError`` in ``exc.args``. + then finds the wrapped ``MCPError`` in ``exc.args``. - The McpError message uses the structured Foundry MCP gateway format: + The MCPError message uses the structured Foundry MCP gateway format: a human-readable prefix followed by a JSON document describing each failed tool source and its consent URL. """ @@ -7521,7 +7520,7 @@ def _make_consent_error( ] }) message = f"tools/list failed for 1 tool source(s), succeeded for 0 tool source(s) {payload}" - inner = McpError(ErrorData(code=CONSENT_ERROR_CODE, message=message)) + inner = MCPError(CONSENT_ERROR_CODE, message) return ToolExecutionException("MCP consent required", inner_exception=inner) @@ -7544,21 +7543,21 @@ def test_returns_none_when_no_mcp_error_in_args(self) -> None: assert consent_url_from_error(Exception("boom")) is None def test_returns_none_when_mcp_error_has_different_code(self) -> None: - inner = McpError(ErrorData(code=-32000, message="some other error")) + inner = MCPError(-32000, "some other error") exc = Exception("wrapped", inner) assert consent_url_from_error(exc) is None def test_returns_none_for_bare_mcp_error_without_wrapping(self) -> None: - # `args` of a bare McpError holds the message string, not an McpError + # `args` of a bare MCPError holds the message string, not an MCPError # instance, so it does not match the wrapping pattern produced by the # MCP client when it bubbles consent errors up. - bare = McpError(ErrorData(code=CONSENT_ERROR_CODE, message="https://x")) + bare = MCPError(CONSENT_ERROR_CODE, "https://x") assert consent_url_from_error(bare) is None def test_returns_none_when_message_has_no_json(self) -> None: from agent_framework.exceptions import ToolExecutionException - inner = McpError(ErrorData(code=CONSENT_ERROR_CODE, message="no json here")) + inner = MCPError(CONSENT_ERROR_CODE, "no json here") exc = ToolExecutionException("MCP consent required", inner_exception=inner) assert consent_url_from_error(exc) is None diff --git a/python/samples/02-agents/harness/build_your_own_claw/claw_step04_production_ready/README.md b/python/samples/02-agents/harness/build_your_own_claw/claw_step04_production_ready/README.md index 48b462ac3ed..a3f18afae64 100644 --- a/python/samples/02-agents/harness/build_your_own_claw/claw_step04_production_ready/README.md +++ b/python/samples/02-agents/harness/build_your_own_claw/claw_step04_production_ready/README.md @@ -51,7 +51,7 @@ export OTEL_EXPORTER_OTLP_ENDPOINT="http://localhost:4317" > likely reason a Toolbox skill fails to load, and the failure actively misleads you: connecting to > the toolbox and **discovering** skills both succeed (`skill://index.json` is toolbox metadata, which > needs no role), so the skill is advertised to the model exactly as expected. Only the first -> `load_skill` fails — with `McpError('Failed to read resource.')` — because reading a skill's *body* +> `load_skill` fails — with `MCPError('Failed to read resource.')` — because reading a skill's *body* > dereferences the project-level skill resource, which does require the role. The toolbox answers > with a bare JSON-RPC `-32603` and no `data`, so nothing in the error names the cause. > From 14fca468c1c5667d86174f724ff8299f64aba058 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Thu, 8 Oct 2026 11:08:45 +0200 Subject: [PATCH 41/42] fixed issue with headers leaking through shared http_client --- python/packages/core/agent_framework/_mcp.py | 23 +++++++++++++------- python/packages/core/tests/core/test_mcp.py | 19 +++++++++++----- 2 files changed, 28 insertions(+), 14 deletions(-) diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 8d3d4033a6f..8132b7e9954 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -166,7 +166,8 @@ class MCPSpecificApproval(TypedDict, total=False): "response_format", "_meta", }) -_mcp_call_headers: contextvars.ContextVar[dict[str, str]] = contextvars.ContextVar("_mcp_call_headers") +# object is a reference to the owner of the headers, so we keep track which headers belong to which tool +_mcp_call_headers: contextvars.ContextVar[tuple[object, dict[str, str]]] = contextvars.ContextVar("_mcp_call_headers") _mcp_tool_runtime_context: contextvars.ContextVar[tuple[object, Mapping[str, Any]] | None] = contextvars.ContextVar( "_mcp_tool_runtime_context", default=None ) @@ -4315,11 +4316,17 @@ async def _inject_headers(request: Request) -> None: # ruff:ignore[unused-async if self._header_provider is not None: # The transport may send this request from a task whose context was # captured before call_tool set the ContextVar; fall back to the - # instance-level snapshot of the active call's headers. Both are None - # only when this is an ambient request outside call_tool; an active - # call that legitimately produced no headers yields an empty dict and - # must not trigger the ambient fallback below. - dynamic_headers = _mcp_call_headers.get(None) + # instance-level snapshot of the active call's headers. Context values + # are owner-tagged so a nested tool sharing the client cannot inherit + # another tool's credentials. Both are None only when this is an ambient + # request outside call_tool; an active call that legitimately produced + # no headers yields an empty dict and must not trigger the ambient fallback. + call_header_context = _mcp_call_headers.get(None) + dynamic_headers = ( + call_header_context[1] + if call_header_context is not None and call_header_context[0] is self._header_request_owner + else None + ) if dynamic_headers is None: dynamic_headers = self._active_call_headers else: @@ -4538,7 +4545,7 @@ async def _call_prompt_with_runtime_kwargs( headers = self._effective_headers(runtime_kwargs) async with self._call_headers_lock: await self._ensure_session_identity(headers, runtime_kwargs) - token = _mcp_call_headers.set(headers) + token = _mcp_call_headers.set((self._header_request_owner, headers)) self._active_call_headers = headers try: return await super()._call_prompt_with_runtime_kwargs( @@ -4591,7 +4598,7 @@ async def call_tool(self, tool_name: str, **kwargs: Any) -> str | list[Content]: headers = self._effective_headers(header_kwargs) async with self._call_headers_lock: await self._ensure_session_identity(headers, header_kwargs) - token = _mcp_call_headers.set(headers) + token = _mcp_call_headers.set((self._header_request_owner, headers)) self._active_call_headers = headers try: return await super().call_tool(tool_name, **kwargs) diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 1d1214940e3..a56322cf4e3 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -8406,7 +8406,9 @@ async def test_mcp_streamable_http_tool_header_provider_sets_contextvar(): async def spy_call_tool(self, tool_name, **kwargs): # Capture the contextvar value during the super call try: - observed_headers.append(_mcp_call_headers.get()) + call_header_context = _mcp_call_headers.get() + assert isinstance(call_header_context, tuple) + observed_headers.append(call_header_context[1]) except LookupError: observed_headers.append({}) return await original_call_tool(self, tool_name, **kwargs) @@ -8560,7 +8562,7 @@ async def test_mcp_streamable_http_tool_header_provider_with_httpx_event_hook(): assert len(hooks) == 1, "Expected one request event hook" # Simulate what happens during a call_tool: contextvar is set - token = _mcp_call_headers.set({"X-Custom": "test-value"}) + token = _mcp_call_headers.set((tool._header_request_owner, {"X-Custom": "test-value"})) try: request = _request_for_mcp_tool(tool) await hooks[0](request) @@ -8662,7 +8664,7 @@ def provider(kw: dict[str, Any]) -> dict[str, str]: assert len(hooks) == 1 # Simulate an active call whose provider returned {} (both ContextVar and snapshot set). - token = _mcp_call_headers.set({}) + token = _mcp_call_headers.set((tool._header_request_owner, {})) tool._active_call_headers = {} try: call_count = 0 @@ -8768,7 +8770,10 @@ async def test_mcp_streamable_http_tool_header_provider_skips_cross_origin_redir hooks = tool._httpx_client.event_hooks.get("request", []) assert len(hooks) == 1 - token = _mcp_call_headers.set({"Authorization": "Bearer secret", "X-API-Key": "api-secret"}) + token = _mcp_call_headers.set(( + tool._header_request_owner, + {"Authorization": "Bearer secret", "X-API-Key": "api-secret"}, + )) try: same_origin = _request_for_mcp_tool(tool, "http://example.com/redirected") await hooks[0](same_origin) @@ -8873,7 +8878,7 @@ def handle(request: httpx.Request) -> httpx.Response: hooks = user_client.event_hooks["request"] assert len(hooks) == int(use_header_provider) if use_header_provider: - token = _mcp_call_headers.set({"X-Dynamic": "per-request"}) + token = _mcp_call_headers.set((tool._header_request_owner, {"X-Dynamic": "per-request"})) try: request = _request_for_mcp_tool(tool) await hooks[0](request) @@ -9155,7 +9160,9 @@ async def spy_call_tool(self, tool_name, **kwargs): # Capture the contextvar value set by call_tool before delegating result = await original_call_tool(self, tool_name, **kwargs) try: - observed_headers.append(_mcp_call_headers.get()) + call_header_context = _mcp_call_headers.get() + assert isinstance(call_header_context, tuple) + observed_headers.append(call_header_context[1]) except LookupError: observed_headers.append({}) return result From 73f7f98e327c48c60e3ff742324018d042da9480 Mon Sep 17 00:00:00 2001 From: Jose Alvarez Date: Thu, 8 Oct 2026 15:20:00 +0200 Subject: [PATCH 42/42] Removed MCP Tasks to start work on new mcp v2 API --- python/PACKAGE_STATUS.md | 4 - python/packages/core/AGENTS.md | 11 +- .../packages/core/agent_framework/__init__.py | 2 - .../core/agent_framework/__init__.pyi | 2 - .../core/agent_framework/_feature_stage.py | 1 - python/packages/core/agent_framework/_mcp.py | 568 +------ python/packages/core/tests/core/test_mcp.py | 1324 +---------------- .../core/tests/core/test_mcp_http_auth.py | 3 - 8 files changed, 63 insertions(+), 1852 deletions(-) diff --git a/python/PACKAGE_STATUS.md b/python/PACKAGE_STATUS.md index b8701ea2d8c..6aa7c77fb43 100644 --- a/python/PACKAGE_STATUS.md +++ b/python/PACKAGE_STATUS.md @@ -129,10 +129,6 @@ listed below. - `agent-framework-core`: experimental harness APIs for background agents, file access, looping, memory, and file-backed todo storage under `agent_framework/_harness/` -#### `MCP_LONG_RUNNING_TASKS` - -- `agent-framework-core`: `MCPTaskOptions` from `agent_framework/_mcp.py` - #### `MCP_SKILLS` - `agent-framework-core`: `MCPSkillResource`, `MCPSkill`, and `MCPSkillsSource` from diff --git a/python/packages/core/AGENTS.md b/python/packages/core/AGENTS.md index b8881d09714..5abda50d2aa 100644 --- a/python/packages/core/AGENTS.md +++ b/python/packages/core/AGENTS.md @@ -217,7 +217,7 @@ The vector store API is experimental under the shared `VECTOR_STORES` feature ID transport; it cannot create a connection. - **Argument allowlist (`_prepare_call_kwargs`)** - Before each `tools/call`, kwargs are filtered to an **allowlist** built from the tool's declared parameters (`inputSchema.properties`) plus any user-configured extras. **The declared half comes from the server's advertised schema**, and runtime kwargs (`FunctionInvocationContext.kwargs`, seeded from `function_invocation_kwargs`) are merged with the model-supplied arguments upstream in `_call_tool_with_runtime_kwargs`, so provenance is gone by the time the filter runs. A runtime kwarg is therefore forwarded whenever the server declares a property of that name, without the model mentioning it — the server, not the caller, decides which runtime kwarg names it receives. A tool that declares no usable `properties` (including schemas with `additionalProperties: true`) forwards only the configured extras. `_MCP_FRAMEWORK_DENYLIST` is a narrow safety net covering only non-serializable framework objects a server *declares* in its schema (those are dropped); it does not generalize to arbitrary caller-chosen names, and explicit extras always win. The reserved `_meta` key is never forwarded as an argument; trusted caller/runtime `_meta` is validated as MCP request metadata, model-supplied `_meta` is discarded in generated MCP functions, and metadata precedence is caller/runtime < OpenTelemetry < tools/list metadata. - **`allowed_tools`** (constructor arg on all `MCPTool` subclasses) - Restricts exposed MCP tools by raw remote MCP tool identity. Prefixed local names remain accepted only when the raw remote name already matches its normalized form; normalized/local aliases do not authorize a different raw remote name. Configured allow/approval names must identify at most one raw remote name across loaded tools and prompts: a raw name that overlaps another tool's prefixed alias raises `ToolExecutionException`. Discovery validates all pages before publishing new functions; a failed reload retains the previous functions and metadata. The allowlist is also revalidated when exposing functions so runtime changes cannot select an ambiguous name. If multiple raw remote tool names map to the same local function name, tool loading raises `ToolExecutionException` instead of first-one-wins shadowing. -- **Progressive MCP disclosure** (`use_progressive_disclosure`, `always_load`) - When enabled on any `MCPTool` subclass, the initial model-facing surface is loader tools (`list_mcp_tools` / `load_tool` / `unload_tool`, prefixed by `tool_name_prefix` when configured) plus allowed tools selected by `always_load` and tools loaded earlier on the same `MCPTool` instance. `list_mcp_tools` only reports tools that pass `allowed_tools`; filtered tools are not listed or loadable. Loader tool names are reserved in progressive mode: remote MCP tools whose local generated name collides with a loader name are omitted from the initial/listed surface, and explicit `load_tool` calls return a model-visible message pointing callers to `tool_name_prefix` or excluding the colliding tool. `load_tool` accepts one tool name or a list of tool names and uses `FunctionInvocationContext.add_tools(...)` so the selected generated MCP `FunctionTool`s become available on the next function-calling iteration while keeping existing approval mode, argument filtering, header-provider runtime kwargs, result parsing, OTel, and task behavior. `unload_tool` accepts one dynamically loaded tool name or a list of names and removes them from the live tool list and persisted progressive surface, but it does not remove tools configured in `always_load`. Invalid `always_load` entries are ignored like unmatched `allowed_tools` entries. +- **Progressive MCP disclosure** (`use_progressive_disclosure`, `always_load`) - When enabled on any `MCPTool` subclass, the initial model-facing surface is loader tools (`list_mcp_tools` / `load_tool` / `unload_tool`, prefixed by `tool_name_prefix` when configured) plus allowed tools selected by `always_load` and tools loaded earlier on the same `MCPTool` instance. `list_mcp_tools` only reports tools that pass `allowed_tools`; filtered tools are not listed or loadable. Loader tool names are reserved in progressive mode: remote MCP tools whose local generated name collides with a loader name are omitted from the initial/listed surface, and explicit `load_tool` calls return a model-visible message pointing callers to `tool_name_prefix` or excluding the colliding tool. `load_tool` accepts one tool name or a list of tool names and uses `FunctionInvocationContext.add_tools(...)` so the selected generated MCP `FunctionTool`s become available on the next function-calling iteration while keeping existing approval mode, argument filtering, header-provider runtime kwargs, result parsing, and OTel. `unload_tool` accepts one dynamically loaded tool name or a list of names and removes them from the live tool list and persisted progressive surface, but it does not remove tools configured in `always_load`. Invalid `always_load` entries are ignored like unmatched `allowed_tools` entries. - **Tool discovery refresh** - Validate the complete current tool snapshot alongside retained non-tool functions, not stale tool wrappers. Successful refreshes atomically replace the previous tool set and metadata, including empty snapshots, while preserving prompts, custom functions, and unchanged tool wrappers (including caller customizations). Removed or replaced tool identities lose their persisted progressive-loaded state; failed refreshes preserve the previous state. - **`additional_tool_argument_names`** (constructor arg on all `MCPTool` subclasses) - Opt extra argument names back into the allowlist. Accepts a `Sequence[str]` (applied to every tool) or a `Mapping[str, Sequence[str]]` keyed by **remote tool name**, where the reserved key `"*"` denotes global extras. It is configured only in user code at construction; there is **no per-call/runtime override**, so a model-issued tool call cannot change which names pass through — but note this constrains the *model*, not the *server*, which still widens the effective allowlist through its schema. To use a server that accepts `additionalProperties: true`, list the extra names here and then either (1) manually extend that tool's `inputSchema` (via the `.functions` list after connecting) so the model is prompted to supply them, or (2) supply the values yourself via `function_invocation_kwargs`. If a normal forwarded argument name is supplied by both the model and `function_invocation_kwargs`, the model-supplied value wins; `_meta` is the exception and only trusted runtime/caller metadata is used. - **MCP HTTP header request scoping** - `static_headers` supplies fixed, origin-scoped headers without serializing concurrent calls. Generated tool and prompt calls give `header_provider` only host runtime kwargs, separately from model arguments; direct `call_tool` calls give it caller kwargs. Model-over-runtime precedence is confined to outbound tool arguments. Fixed and dynamic headers form the session's complete effective identity (case-insensitive names, case-sensitive values), with dynamic values overriding fixed values of the same name. Agent runs reconcile a connected tool before copying its functions; generated tool and prompt calls perform the same check at invocation. Connected run preparation defers when a provider needs runtime values supplied by invocation middleware; invocation-time resolution remains strict. A different identity waits for in-flight discovery, reconnects framework-created sessions without loading, then refreshes session-derived discovery under the same lock. Because a caller-supplied session's established identity is unknown and the wrapper cannot reconnect it, dynamic header resolution on that wrapper is rejected even while disconnected. Ambient requests retain the bound set. With a shared `http_client`, processing stays scoped to the originating `MCPStreamableHTTPTool`, and every injected header is stripped from cross-origin redirects. The session exit stack removes its request hook after transport shutdown, including failed initialization or discovery, and closes framework-created HTTP clients. Failed discovery resets the connection and discovery flags and rolls back partial function/metadata additions. Framework-created sessions are discarded; constructor-supplied sessions remain caller-owned and reusable across cleanup, close, and reset. @@ -226,14 +226,7 @@ The vector store API is experimental under the shared `VECTOR_STORES` feature ID - **`function_invocation_kwargs` and MCP servers** - That dict is shared across every tool in the run, including every attached `MCPTool`, and any name in it reaches a server that declares a matching `inputSchema` property. `header_provider` does not mitigate this — it reads the kwargs without consuming them. To keep a credential out of tool arguments, source it outside `function_invocation_kwargs`: read a `ContextVar` inside the provider (this still allows a different value per request), configure a custom `http_client`, or use `env` for `MCPStdioTool`. - **Sampling guardrails** (`sampling_callback`) - Passing `client=` advertises `SamplingCapability` so the server can send `sampling/createMessage`. Because remote servers are untrusted (confused-deputy risk), the default `sampling_callback` is **deny-by-default** and applies, in order: a per-session rate limit (`sampling_max_requests`, default `_DEFAULT_SAMPLING_MAX_REQUESTS`), an approval gate (`sampling_approval_callback`), and a `maxTokens` cap (`sampling_max_tokens`, default `_DEFAULT_SAMPLING_MAX_TOKENS`). The approval callback (constructor arg on all subclasses; exported type alias `SamplingApprovalCallback`) receives the raw `CreateMessageRequestParams`, may be sync or async, and must return truthy to approve. When it is `None` (the default) every sampling request is denied; pass `lambda params: True` to restore legacy auto-approve as an explicit opt-in. Requests and denials are logged at WARNING (content is not logged). The per-session counter resets in `_reset_session_state`. - **MCP response caching** - Framework-owned high-level Clients use SDK `cache_mode="use"` for tool/prompt catalogs and resource reads, honoring modern server `ttlMs` / `cacheScope` hints. Legacy peers default to zero TTL, caller-supplied `ClientSession` connections stay uncached, and reconnect or effective header-identity changes replace the per-Client cache. Agent Framework does not expose shared-store or cache-mode configuration. -- **`MCPTaskOptions`** (experimental, `MCP_LONG_RUNNING_TASKS` feature, **frozen**) - Per-tool-instance options controlling the SEP-2663 long-running task lifecycle. When the server advertises a tool with `execution.taskSupport == "required"`, `MCPTool.call_tool` transparently routes through `call_tool_as_task`, which sends an augmented `tools/call`, polls `tasks/get` until terminal, and reinterprets `tasks/result` as a normal `CallToolResult`. Instances are immutable; replace via `MCPTool.task_options = MCPTaskOptions(...)`. Fields: - - `default_ttl: timedelta | None` — forwarded to the server as `params.task.ttl` (milliseconds). When `None`, the server's default applies. - - `cancel_remote_task_on_local_cancellation: bool = True` — only gates the `CancelledError` path. Abandonment paths (see below) always cancel. - - `max_task_wait: timedelta | None` — client-side deadline for the whole post-create lifecycle (poll + result fetch). When exceeded, raises `ToolExecutionException` and fires a best-effort `tasks/cancel`. `None` (default) means no client-side bound. Bounds sleeps, sends, AND reconnects via `asyncio.wait_for`. -- **Permissive fallback**: servers that ignore the augmentation (return `CallToolResult` directly) or reject the unknown `task` field with `METHOD_NOT_FOUND` / `INVALID_PARAMS` fall back to the plain `session.call_tool(...)` path so legacy servers keep working. An unparseable success response (server accepted the augmented call but returned a payload that is neither `CreateTaskResult` nor `CallToolResult`) **does not** fall back — it raises `ToolExecutionException` to avoid double-executing a side-effecting tool. -- **Submit-vs-track reconnect policy**: a dropped connection before a `task_id` is known raises `ToolExecutionException("connection lost; task state unknown")` without re-issuing the augmented `tools/call`, so a server that accepted the request but lost the response cannot be made to start the same operation twice; once a `task_id` exists, `tasks/get` / `tasks/result` reconnect once and retry against the same id (a shared `_send_with_one_reconnect` helper). -- **Cancel-on-abandonment vs terminal failure**: any path where the remote task may still be running (max-wait exceeded, hard `MCPError` in poll, malformed `tasks/get`, second connection loss in poll/fetch, reconnect failure) fires best-effort `tasks/cancel` before raising. Terminal failures (`failed`/`cancelled`/`input_required` server-side, `completed+isError`, malformed `tasks/result` after server completed) do **not** cancel — the server is already done. `_MCPTaskAbandoned` is the private marker distinguishing the two. -- **Transient poll retry**: a slow `tasks/get` that surfaces as `MCPError(code=REQUEST_TIMEOUT)` is retried (bounded by `max_task_wait`). All other non-connection `MCPError`s during poll are treated as abandonment. `tasks/result` does not get transient retry — the server has already completed, so a slow payload fetch is anomalous. +- **MCP Tasks extension** - MCP Python SDK 2.2 does not implement the 2026 `io.modelcontextprotocol/tasks` runtime. Tools that advertise `execution.taskSupport == "required"` are skipped with a warning; optional and task-free tools continue through ordinary `tools/call`. Restoration of the experimental Tasks API and removed sample is tracked in microsoft/agent-framework#8245. ### File Access Harness (`_harness/_file_access.py`) diff --git a/python/packages/core/agent_framework/__init__.py b/python/packages/core/agent_framework/__init__.py index 3953187578e..2bd169460e6 100644 --- a/python/packages/core/agent_framework/__init__.py +++ b/python/packages/core/agent_framework/__init__.py @@ -168,7 +168,6 @@ "._mcp": ( "MCPStdioTool", "MCPStreamableHTTPTool", - "MCPTaskOptions", "MCPWebsocketTool", "SamplingApprovalCallback", ), @@ -559,7 +558,6 @@ "MCPSkillsSource", "MCPStdioTool", "MCPStreamableHTTPTool", - "MCPTaskOptions", "MCPWebsocketTool", "MemoryContextProvider", "MemoryFileStore", diff --git a/python/packages/core/agent_framework/__init__.pyi b/python/packages/core/agent_framework/__init__.pyi index e8d1bc6e411..86f1fbd9692 100644 --- a/python/packages/core/agent_framework/__init__.pyi +++ b/python/packages/core/agent_framework/__init__.pyi @@ -127,7 +127,6 @@ from ._in_memory import InMemoryCollection, InMemoryStore from ._mcp import ( MCPStdioTool, MCPStreamableHTTPTool, - MCPTaskOptions, MCPWebsocketTool, SamplingApprovalCallback, ) @@ -514,7 +513,6 @@ __all__ = [ "MCPSkillsSource", "MCPStdioTool", "MCPStreamableHTTPTool", - "MCPTaskOptions", "MCPWebsocketTool", "MemoryContextProvider", "MemoryFileStore", diff --git a/python/packages/core/agent_framework/_feature_stage.py b/python/packages/core/agent_framework/_feature_stage.py index 24309e0b144..9421c12cc27 100644 --- a/python/packages/core/agent_framework/_feature_stage.py +++ b/python/packages/core/agent_framework/_feature_stage.py @@ -60,7 +60,6 @@ class ExperimentalFeature(str, Enum): FOUNDRY_PREVIEW_TOOLS = "FOUNDRY_PREVIEW_TOOLS" FUNCTIONAL_WORKFLOWS = "FUNCTIONAL_WORKFLOWS" HARNESS = "HARNESS" - MCP_LONG_RUNNING_TASKS = "MCP_LONG_RUNNING_TASKS" MCP_SKILLS = "MCP_SKILLS" PROGRESSIVE_TOOLS = "PROGRESSIVE_TOOLS" SESSION_STORE = "SESSION_STORE" diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 8132b7e9954..c1dbfe528d7 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -4,7 +4,6 @@ import asyncio import base64 -import contextlib import contextvars import json import logging @@ -35,7 +34,6 @@ ExperimentalFeature, ExperimentalWarning, _warn_on_feature_use, # pyright: ignore[reportPrivateUsage] - experimental, ) from ._serialization import make_json_safe from ._telemetry import FeatureIndex, mark_feature_used @@ -895,73 +893,6 @@ def _mcp_header_identity(headers: Mapping[str, str]) -> _MCPHeaderIdentity: return tuple(sorted(normalized.items())) -# Internal polling bounds for MCP long-running tasks. Not user-tunable today; -# promote to MCPTaskOptions if a concrete need arises. -_MCP_TASK_MIN_POLL_INTERVAL = timedelta(milliseconds=500) -_MCP_TASK_MAX_POLL_INTERVAL = timedelta(seconds=5) -_MCP_TASK_CANCEL_TIMEOUT = timedelta(seconds=5) -_MCP_TASK_TERMINAL_STATUSES: frozenset[str] = frozenset({"completed", "failed", "cancelled", "input_required"}) - -# Total send attempts for a Phase 2 request (initial try + one reconnect-and-retry). -# A single transient disconnect should not abort a long-running task; sustained outages -# surface as ``_MCPTaskAbandoned`` after the second failure. -_MCP_RECONNECT_ATTEMPTS = 2 - - -class _MCPTaskAbandoned(ToolExecutionException): - """Raised when the remote MCP task may still be running and must be cancelled. - - Subclass of ToolExecutionException so callers see a normal tool failure. - """ - - -class _MCPDeadlineExpired(Exception): - """Internal marker for ``max_task_wait`` expiry; distinct from inner TimeoutError.""" - - -@experimental(feature_id=ExperimentalFeature.MCP_LONG_RUNNING_TASKS) -@dataclass(frozen=True) -class MCPTaskOptions: - """Options controlling how MCPTool drives the MCP long-running task lifecycle. - - When an MCP server advertises a tool with ``execution.taskSupport == "required"``, - the framework transparently drives the SEP-2663 ``tools/call`` → ``tasks/get`` - (polled) → ``tasks/result`` lifecycle so the agent sees a normal tool result. - - Instances are immutable; replace the whole object via - ``MCPTool.task_options = MCPTaskOptions(...)`` to change behavior. - - Attributes: - default_ttl: Optional task-record retention time forwarded to the server as - ``params.task.ttl`` (milliseconds, integer). The server keeps the task - record around this long after the task reaches a terminal status so the - client can still call ``tasks/get`` / ``tasks/result``; it does not - cancel a running task. When ``None``, the server applies its own default. - Must be positive if set (zero would expire the record before any client - could read it). - cancel_remote_task_on_local_cancellation: If True (default), a local - cancellation of the awaiting coroutine triggers a best-effort - ``tasks/cancel`` on the server before re-raising ``CancelledError``. - Only gates ``CancelledError``; abandonment paths (max-wait, - unrecoverable poll errors, lost connection after task_id is known) - always cancel regardless of this flag. - max_task_wait: Optional client-side deadline for the whole post-create - lifecycle (poll + result fetch). When exceeded, raises - ``ToolExecutionException`` and fires a best-effort ``tasks/cancel``. - ``None`` (default) means no client-side bound. Must be positive if set. - """ - - default_ttl: timedelta | None = None - cancel_remote_task_on_local_cancellation: bool = True - max_task_wait: timedelta | None = None - - def __post_init__(self) -> None: - if self.default_ttl is not None and self.default_ttl.total_seconds() <= 0: - raise ValueError("MCPTaskOptions.default_ttl must be positive.") - if self.max_task_wait is not None and self.max_task_wait.total_seconds() <= 0: - raise ValueError("MCPTaskOptions.max_task_wait must be positive.") - - def streamable_http_client(*args: Any, **kwargs: Any) -> _AsyncGeneratorContextManager[Any, None]: """Lazily import the MCP streamable HTTP transport.""" try: @@ -1083,7 +1014,6 @@ def __init__( sampling_max_tokens: int | None = _DEFAULT_SAMPLING_MAX_TOKENS, sampling_max_requests: int | None = _DEFAULT_SAMPLING_MAX_REQUESTS, additional_properties: dict[str, Any] | None = None, - task_options: MCPTaskOptions | None = None, additional_tool_argument_names: Sequence[str] | Mapping[str, Sequence[str]] | None = None, use_progressive_disclosure: bool = False, always_load: Collection[str] | None = None, @@ -1144,9 +1074,6 @@ def __init__( connection; further requests are rejected. The counter resets on reconnect. Set to ``None`` to disable the limit. Defaults to ``_DEFAULT_SAMPLING_MAX_REQUESTS``. additional_properties: Additional properties for the tool. - task_options: Options controlling how long-running MCP tasks are driven for - tools that advertise ``execution.taskSupport == "required"``. When ``None``, - the defaults from :class:`MCPTaskOptions` are used. additional_tool_argument_names: Extra argument names to forward to the MCP server in addition to each tool's declared parameters. A ``Sequence[str]`` applies to every tool; a ``Mapping[str, Sequence[str]]`` is keyed by remote tool name with @@ -1197,10 +1124,6 @@ def __init__( self.load_prompts_flag = load_prompts self.parse_prompt_results = parse_prompt_results self.max_host_payload_size_bytes = max_host_payload_size_bytes - # Defer constructing the default MCPTaskOptions so the experimental warning - # only fires when LRO is actually engaged (lazy-resolved by _effective_task_options). - self._task_options_explicit: MCPTaskOptions | None = task_options - self._task_options_default: MCPTaskOptions | None = None self._exit_stack = AsyncExitStack() self._lifecycle_lock = asyncio.Lock() self._lifecycle_request_lock = asyncio.Lock() @@ -1232,7 +1155,6 @@ def __init__( self._progressive_loader_functions: list[FunctionTool] | None = None self._progressive_loaded_tool_names: set[str] = set() self._tool_call_meta_by_name: dict[str, dict[str, Any]] = {} - self._tool_task_support_by_name: dict[str, str] = {} self._tool_param_names_by_name: dict[str, set[str]] = {} self._global_extra_arg_names, self._tool_extra_arg_names = _normalize_additional_tool_argument_names( additional_tool_argument_names @@ -2141,7 +2063,6 @@ def _reset_session_discovery_state(self) -> None: if not isinstance((function.additional_properties or {}).get(_MCP_REMOTE_NAME_KEY), str) ] self._tool_call_meta_by_name.clear() - self._tool_task_support_by_name.clear() self._tool_param_names_by_name.clear() self._progressive_loaded_tool_names.clear() @@ -2308,7 +2229,6 @@ async def _connect_on_owner( raise functions_before_discovery = self._functions.copy() call_meta_before_discovery = self._tool_call_meta_by_name - task_support_before_discovery = self._tool_task_support_by_name param_names_before_discovery = self._tool_param_names_by_name try: logger.debug("Connected to MCP server: %s", self.session) @@ -2335,7 +2255,6 @@ async def _connect_on_owner( finally: self._functions[:] = functions_before_discovery self._tool_call_meta_by_name = call_meta_before_discovery - self._tool_task_support_by_name = task_support_before_discovery self._tool_param_names_by_name = param_names_before_discovery raise @@ -2812,7 +2731,6 @@ async def _load_tools_locked(self) -> None: if isinstance(remote_name, str): existing_remote_by_local[func.name] = remote_name tool_call_meta_by_name: dict[str, dict[str, Any]] = {} - tool_task_support_by_name: dict[str, str] = {} tool_param_names_by_name: dict[str, set[str]] = {} tool_annotations_by_name: dict[str, Any] = {} @@ -2850,14 +2768,19 @@ async def _load_tools_locked(self) -> None: raise ToolExecutionException("Failed to load tools.") for tool in tool_list.tools: + task_support = getattr(getattr(tool, "execution", None), "task_support", None) + if task_support == "required": + logger.warning( + "Skipping MCP tool %r because it requires the Tasks extension, " + "which MCP Python SDK 2.2 does not implement.", + tool.name, + ) + continue + tool_annotations_by_name[tool.name] = tool.annotations if tool.meta is not None: tool_call_meta_by_name[tool.name] = _validate_mcp_meta(tool.meta) or {} - task_support = getattr(getattr(tool, "execution", None), "task_support", None) - if task_support is not None: - tool_task_support_by_name[tool.name] = task_support - # Normalize inputSchema: ensure "properties" exists for object schemas. # Some MCP servers (e.g. zero-argument tools) omit "properties", # which causes OpenAI API to reject the schema with a 400 error. @@ -2869,8 +2792,7 @@ async def _load_tools_locked(self) -> None: # Register declared param names before the existing-tool skip below so that # reloads (e.g. notifications/tools/list_changed) preserve the allowlist for - # tools that are already loaded, consistent with tool_call_meta_by_name and - # tool_task_support_by_name above. + # tools that are already loaded, consistent with tool_call_meta_by_name above. schema_properties = input_schema.get("properties") tool_param_names_by_name[tool.name] = ( set(cast(dict[str, Any], schema_properties)) if isinstance(schema_properties, dict) else set() @@ -2947,7 +2869,6 @@ async def _load_tools_locked(self) -> None: self._function_load_callback(function, tool_annotations_by_name[remote_name]) self._functions[:] = current_functions self._tool_call_meta_by_name = tool_call_meta_by_name - self._tool_task_support_by_name = tool_task_support_by_name self._tool_param_names_by_name = tool_param_names_by_name self._progressive_loaded_tool_names.difference_update(existing_tools.keys() - reused_tool_names) @@ -3037,29 +2958,6 @@ async def _ensure_connected(self) -> None: inner_exception=ex, ) from ex - def _effective_task_options(self) -> MCPTaskOptions: - """Return the effective MCPTaskOptions, lazily constructing defaults on first use. - - Defers the implicit ``MCPTaskOptions()`` so the experimental warning only - fires when LRO is actually engaged (server advertises ``taskSupport=required``). - """ - explicit = self._task_options_explicit - if explicit is not None: - return explicit - if self._task_options_default is None: - self._task_options_default = MCPTaskOptions() - return self._task_options_default - - @property - def task_options(self) -> MCPTaskOptions: - """The effective MCPTaskOptions for this tool (lazy defaults).""" - return self._effective_task_options() - - @task_options.setter - def task_options(self, value: MCPTaskOptions | None) -> None: - self._task_options_explicit = value - self._task_options_default = None - async def call_tool(self, tool_name: str, **kwargs: Any) -> str | list[Content]: """Call a tool with the given arguments. @@ -3098,12 +2996,6 @@ async def _call_tool( raise ToolExecutionException( "Tools are not loaded for this server, please set load_tools=True in the constructor." ) - - # Tools advertising taskSupport == "required" cannot complete via plain tools/call; - # route through the long-running task lifecycle transparently. - if self._tool_task_support_by_name.get(tool_name) == "required": - return await self.call_tool_as_task(tool_name, **kwargs) - filtered_kwargs, meta = self._prepare_call_kwargs(tool_name, kwargs) parser = self.parse_tool_results or self._parse_tool_result_from_mcp @@ -3241,430 +3133,6 @@ def _prepare_call_kwargs( request_meta = {**(request_meta or {}), **tool_meta} return filtered_kwargs, request_meta - async def call_tool_as_task(self, tool_name: str, **kwargs: Any) -> str | list[Content]: - """Call an MCP tool via the long-running task lifecycle (SEP-2663). - - Issues an augmented ``tools/call`` with ``params.task`` set from - ``self.task_options``, then polls ``tasks/get`` until the server reports a - terminal status. On ``completed`` the payload is fetched via ``tasks/result``, - validated as a ``CallToolResult`` and parsed identically to :meth:`call_tool`. - - Local cancellation triggers a best-effort ``tasks/cancel`` (controlled by - :attr:`MCPTaskOptions.cancel_remote_task_on_local_cancellation`) before - ``asyncio.CancelledError`` is re-raised. - - Args: - tool_name: The remote MCP tool name. - - Keyword Args: - kwargs: Arguments forwarded to the tool. See :meth:`call_tool` for the - framework kwargs that are filtered out. - - Returns: - A list of Content items (or a string when a custom ``parse_tool_results`` - callback is configured). - """ - return await self._call_tool_as_task(tool_name, kwargs) - - async def _call_tool_as_task( - self, - tool_name: str, - kwargs: dict[str, Any], - ) -> str | list[Content]: - from anyio import ClosedResourceError - from mcp import MCPError - - if not self.load_tools_flag: - raise ToolExecutionException( - "Tools are not loaded for this server, please set load_tools=True in the constructor." - ) - - filtered_kwargs, meta = self._prepare_call_kwargs(tool_name, kwargs) - parser = self.parse_tool_results or self._parse_tool_result_from_mcp - - # Submit the task: issue augmented tools/call. Do NOT retry on connection loss here: - # the server may have accepted the request and created a task before the - # response was lost, so retrying could start the long-running operation twice. - # Reconnect-and-retry is only safe after the task_id is known. - try: - task_id, fallback_result = await self._call_tool_as_task_create(tool_name, filtered_kwargs, meta) - except (ClosedResourceError, MCPError) as ex: - if not self._is_connection_lost(ex): - error_message = ex.error.message if isinstance(ex, MCPError) else str(ex) - raise ToolExecutionException(error_message, inner_exception=ex) from ex - raise ToolExecutionException( - f"Failed to call tool '{tool_name}' - connection lost; task state unknown.", - inner_exception=ex, - ) from ex - except ToolExecutionException: - raise - except Exception as ex: - raise ToolExecutionException(f"Failed to call tool '{tool_name}'.", inner_exception=ex) from ex - - # Server returned a CallToolResult (no task created) or fell back to plain tools/call. - if fallback_result is not None: - _capture_mcp_tool_result(fallback_result) - if fallback_result.is_error: - parsed = parser(fallback_result) - text = ( - "\n".join(c.text for c in parsed if c.type == "text" and c.text) - if isinstance(parsed, list) - else str(parsed) - ) - raise ToolExecutionException(text or str(parsed)) - return parser(fallback_result) - - if task_id is None: - raise ToolExecutionException(f"MCP server did not return a task_id or fallback result for '{tool_name}'.") - - # Track to completion: poll until terminal, then fetch payload. Never re-issue - # tools/call past this point; reconnect-and-retry only against the same task_id. - opts = self._effective_task_options() - max_wait_s = opts.max_task_wait.total_seconds() if opts.max_task_wait is not None else None - - async def _await_task_completion() -> str | list[Content]: - terminal = await self._poll_task_until_terminal(task_id) - return await self._handle_terminal_task( - tool_name, - task_id, - terminal, - parser, - ) - - try: - if max_wait_s is not None: - try: - result = await self._await_with_deadline(_await_task_completion(), max_wait_s) - return cast("str | list[Content]", result) - except _MCPDeadlineExpired as ex: - self._spawn_best_effort_cancel(task_id) - raise ToolExecutionException( - f"MCP task '{task_id}' exceeded max_task_wait of {max_wait_s}s.", - inner_exception=ex, - ) from ex - else: - return await _await_task_completion() - except asyncio.CancelledError: - if opts.cancel_remote_task_on_local_cancellation: - self._spawn_best_effort_cancel(task_id) - raise - except _MCPTaskAbandoned: - # Pre-terminal abandonment (hard poll error, malformed get, second - # disconnect, reconnect failure): cancel + re-raise as plain - # ToolExecutionException to the function-calling loop. - self._spawn_best_effort_cancel(task_id) - raise - # Plain ToolExecutionException from terminal failures (failed/cancelled/ - # input_required, completed+isError, malformed result post-completion) - # propagates without cancel — server is already done. - - async def _call_tool_as_task_create( - self, tool_name: str, arguments: dict[str, Any], meta: dict[str, Any] | None - ) -> tuple[str | None, types.CallToolResult | None]: - """Send the augmented tools/call. - - Returns ``(task_id, None)`` when the server created a task, - ``(None, CallToolResult)`` when it returned a non-task result, falling back - to plain ``tools/call`` if the server rejects the ``task`` field outright. - """ - from mcp import MCPError, types - from pydantic import ValidationError - - opts = self._effective_task_options() - ttl_ms: int | None = None - if opts.default_ttl is not None: - ttl_ms = int(opts.default_ttl.total_seconds() * 1000) - # Always send TaskMetadata to mark the call as task-augmented; ttl may be omitted. - task_metadata = types.TaskMetadata(ttl=ttl_ms) - - request_meta = types.RequestParams.Meta(**meta) if meta else None - params = types.CallToolRequestParams( - name=tool_name, - arguments=arguments, - task=task_metadata, - _meta=request_meta, - ) - request = types.ClientRequest(types.CallToolRequest(params=params)) - - # Use the lenient Result type so we can extract the task_id even when - # the strict CreateTaskResult schema rejects the payload (the MCP Python - # SDK requires Task.ttl, but servers may legitimately omit it). - try: - lenient = await self.session.send_request( # type: ignore[union-attr] - request, - types.Result, - ) - except MCPError as ex: - if ex.error.code not in (types.METHOD_NOT_FOUND, types.INVALID_PARAMS): - raise - logger.debug( - "Server rejected augmented tools/call for '%s' (code=%s); falling back.", - tool_name, - ex.error.code, - ) - fallback = await self.session.call_tool(tool_name, arguments=arguments, meta=meta) # type: ignore[union-attr] - return None, fallback - - # Inspect the raw payload: a CreateTaskResult carries `task.taskId`; - # a legacy CallToolResult carries `content` and/or `isError`. - raw: dict[str, Any] = lenient.model_dump(by_alias=True, exclude_none=True) - - task_field = raw.get("task") - if isinstance(task_field, dict): - task_id_val = cast(dict[str, Any], task_field).get("taskId") - if isinstance(task_id_val, str): - return task_id_val, None - - try: - legacy = types.CallToolResult.model_validate(raw) - except ValidationError as ex: - # Augmented call succeeded server-side; re-issuing a plain tools/call - # could double-execute a side-effecting tool. - raise ToolExecutionException( - f"MCP server returned an unparseable response to augmented tools/call " - f"for '{tool_name}'; cannot safely retry (server may have started the operation).", - inner_exception=ex, - ) from ex - - return None, legacy - - async def _poll_task_until_terminal(self, task_id: str) -> types.GetTaskResult: - """Poll ``tasks/get`` until the task reaches a terminal status.""" - import httpx - from mcp import MCPError, types - - # SDK raises MCPError(code=httpx.REQUEST_TIMEOUT=408) on session read timeout. - transient_codes: frozenset[int] = frozenset({int(httpx.codes.REQUEST_TIMEOUT)}) - - while True: - request = types.ClientRequest(types.GetTaskRequest(params=types.GetTaskRequestParams(task_id=task_id))) - try: - # GetTaskResult.ttl is required-but-Optional in the SDK; coerce below. - lenient = await self._send_with_one_reconnect( - request, types.Result, operation="tasks/get", task_id=task_id - ) - except MCPError as ex: - if ex.error.code in transient_codes: - logger.debug("Transient %s on tasks/get for '%s'; will retry.", ex.error.code, task_id) - await asyncio.sleep(_MCP_TASK_MIN_POLL_INTERVAL.total_seconds()) - continue - # Hard server error mid-poll: task may still be running. - raise _MCPTaskAbandoned(ex.error.message, inner_exception=ex) from ex - - try: - snapshot = self._coerce_get_task_result(lenient, task_id) - except ToolExecutionException as ex: - # Malformed tasks/get response; task may still be running. - raise _MCPTaskAbandoned(str(ex), inner_exception=ex) from ex - - if snapshot.status in _MCP_TASK_TERMINAL_STATUSES: - return snapshot - - await asyncio.sleep(self._compute_poll_delay(snapshot.poll_interval).total_seconds()) - - @staticmethod - def _coerce_get_task_result(lenient: types.Result, task_id: str) -> types.GetTaskResult: - """Coerce a lenient Result into GetTaskResult, defaulting ``ttl`` when absent.""" - from mcp import types - - raw = lenient.model_dump(by_alias=True, exclude_none=True) - raw.pop("_meta", None) - raw.setdefault("ttl", None) - try: - return types.GetTaskResult.model_validate(raw) - except Exception as ex: - raise ToolExecutionException( - f"MCP server returned a malformed tasks/get response for task '{task_id}'.", - inner_exception=ex, - ) from ex - - @staticmethod - def _compute_poll_delay(server_interval_ms: int | None) -> timedelta: - """Clamp the server-suggested poll interval to ``[min, max]``.""" - if server_interval_ms is None or server_interval_ms <= 0: - return _MCP_TASK_MIN_POLL_INTERVAL - suggested = timedelta(milliseconds=server_interval_ms) - if suggested < _MCP_TASK_MIN_POLL_INTERVAL: - return _MCP_TASK_MIN_POLL_INTERVAL - if suggested > _MCP_TASK_MAX_POLL_INTERVAL: - return _MCP_TASK_MAX_POLL_INTERVAL - return suggested - - async def _handle_terminal_task( - self, - tool_name: str, - task_id: str, - snapshot: types.GetTaskResult, - parser: Callable[[types.CallToolResult], str | list[Content]], - ) -> str | list[Content]: - """Map a terminal task snapshot to either a parsed result or an exception.""" - status = snapshot.status - if status == "completed": - payload = await self._fetch_task_result(task_id) - _capture_mcp_tool_result(payload) - if payload.is_error: - parsed = parser(payload) - text = ( - "\n".join(c.text for c in parsed if c.type == "text" and c.text) - if isinstance(parsed, list) - else str(parsed) - ) - raise ToolExecutionException(text or str(parsed)) - return parser(payload) - - # Non-completed terminal statuses surface as ToolExecutionException so the - # function-calling loop sees a normal failure for tool_name. - message = snapshot.status_message or f"MCP task ended with status '{status}'." - if status == "input_required": - # Spec-non-terminal; treated as terminal here because the framework does - # not implement the interactive input flow. - message = snapshot.status_message or "MCP task requires additional input and cannot continue." - raise ToolExecutionException(f"Tool '{tool_name}' task {status}: {message}") - - async def _fetch_task_result(self, task_id: str) -> types.CallToolResult: - """Send ``tasks/result`` and reinterpret the open-typed payload as a CallToolResult.""" - from mcp import MCPError, types - from pydantic import ValidationError - - request = types.ClientRequest( - types.GetTaskPayloadRequest(params=types.GetTaskPayloadRequestParams(task_id=task_id)) - ) - # Connection-loss retry only via the helper; no transient-code retry — server - # has already completed the task, so a slow payload fetch is anomalous. - try: - payload = await self._send_with_one_reconnect( - request, types.GetTaskPayloadResult, operation="tasks/result", task_id=task_id - ) - except MCPError as ex: - # Server reported completed; a hard fetch error is a plain failure (no cancel). - raise ToolExecutionException(ex.error.message, inner_exception=ex) from ex - - # GetTaskPayloadResult carries the tool result via extra fields; reinterpret as CallToolResult. - payload_dict = payload.model_dump(by_alias=True, exclude_none=True) - try: - return types.CallToolResult.model_validate(payload_dict) - except ValidationError as ex: - # Server reported completed; malformed payload is a plain failure (no cancel needed). - raise ToolExecutionException( - f"MCP task '{task_id}' result payload could not be parsed as a CallToolResult.", - inner_exception=ex, - ) from ex - - async def _send_with_one_reconnect( - self, - request: types.ClientRequest, - result_type: type[Any], - *, - operation: str, - task_id: str, - ) -> Any: - """Send ``request`` with one reconnect-and-retry on connection loss. - - After a second loss (or reconnect failure), raise ``_MCPTaskAbandoned``. - Non-connection errors propagate unchanged. - """ - from anyio import ClosedResourceError - from mcp import MCPError - - for attempt in range(_MCP_RECONNECT_ATTEMPTS): - try: - return await self.session.send_request(request, result_type) # type: ignore[union-attr] - except (ClosedResourceError, MCPError) as ex: - if not self._is_connection_lost(ex): - raise - if attempt < _MCP_RECONNECT_ATTEMPTS - 1: - logger.info("MCP connection lost during %s; reconnecting (task_id=%s).", operation, task_id) - try: - await self.connect(reset=True) - except Exception as reconn_ex: - # Reconnect failure: task may still be running. - raise _MCPTaskAbandoned( - "Failed to reconnect to MCP server.", inner_exception=reconn_ex - ) from reconn_ex - continue - # Final attempt also lost the connection: task may still be running. - raise _MCPTaskAbandoned( - f"MCP connection lost; task state unknown (task_id={task_id}).", - inner_exception=ex, - ) from ex - raise AssertionError(f"unreachable: {operation} for {task_id}") # pragma: no cover - - @staticmethod - async def _await_with_deadline(coro: Coroutine[Any, Any, Any], timeout_s: float) -> Any: - """Await ``coro`` with a deadline; raise ``_MCPDeadlineExpired`` only on deadline. - - Unlike ``asyncio.wait_for``, an ``asyncio.TimeoutError`` raised by ``coro`` - itself propagates unchanged so callers can distinguish their own deadline - from a stray inner timeout. - """ - inner = asyncio.ensure_future(coro) - try: - done, _pending = await asyncio.wait({inner}, timeout=timeout_s) - except BaseException: - # Outer caller cancelled (or another exception): cancel inner + drain. - inner.cancel() - with contextlib.suppress(BaseException): - await inner - raise - if inner in done: - return inner.result() - # Deadline fired before inner finished. - inner.cancel() - with contextlib.suppress(BaseException): - await inner - raise _MCPDeadlineExpired - - def _spawn_best_effort_cancel(self, task_id: str) -> None: - """Fire-and-forget ``tasks/cancel`` so local cancellation propagates server-side.""" - try: - loop = asyncio.get_running_loop() - except RuntimeError: - return - cancel_task = loop.create_task(self._try_cancel_task(task_id)) - # Reuse pending-reload bookkeeping so close-on-owner waits/cancels these too. - self._pending_reload_tasks.add(cancel_task) - cancel_task.add_done_callback(self._pending_reload_tasks.discard) - - async def _try_cancel_task(self, task_id: str) -> None: - """Send ``tasks/cancel``; bounded by ``_MCP_TASK_CANCEL_TIMEOUT``. - - Failures log at warning so unattributed orphan tasks are debuggable. - """ - from mcp import types - - request = types.ClientRequest(types.CancelTaskRequest(params=types.CancelTaskRequestParams(task_id=task_id))) - try: - await asyncio.wait_for( - self.session.send_request(request, types.CancelTaskResult), # type: ignore[union-attr] - timeout=_MCP_TASK_CANCEL_TIMEOUT.total_seconds(), - ) - except asyncio.CancelledError: - raise - except asyncio.TimeoutError: - logger.warning( - "Best-effort tasks/cancel for '%s' timed out after %.1fs; remote task may still be running.", - task_id, - _MCP_TASK_CANCEL_TIMEOUT.total_seconds(), - ) - except Exception: - logger.warning( - "Best-effort tasks/cancel for '%s' failed; remote task may still be running.", - task_id, - exc_info=True, - ) - - @staticmethod - def _is_connection_lost(ex: BaseException) -> bool: - """Return True if *ex* indicates the MCP transport was torn down.""" - from anyio import ClosedResourceError - from mcp import MCPError - - if isinstance(ex, ClosedResourceError): - return True - if isinstance(ex, MCPError): - return "session terminated" in ex.error.message.lower() - return False - async def _call_prompt_with_runtime_kwargs( self, prompt_name: str, @@ -3828,7 +3296,6 @@ def __init__( sampling_max_tokens: int | None = _DEFAULT_SAMPLING_MAX_TOKENS, sampling_max_requests: int | None = _DEFAULT_SAMPLING_MAX_REQUESTS, additional_properties: dict[str, Any] | None = None, - task_options: MCPTaskOptions | None = None, additional_tool_argument_names: Sequence[str] | Mapping[str, Sequence[str]] | None = None, max_host_payload_size_bytes: int | None = _DEFAULT_MCP_HOST_PAYLOAD_SIZE_BYTES, tool_result_content: MCPToolResultContentMode = "structured_first", @@ -3900,8 +3367,6 @@ def __init__( (``min(requested, cap)``); ``None`` disables it. sampling_max_requests: Per-session cap on the number of sampling requests; further requests are rejected. Resets on reconnect. ``None`` disables it. - task_options: Options for tools that advertise - ``execution.taskSupport == "required"``. See :class:`MCPTaskOptions`. additional_tool_argument_names: Extra argument names to forward to the MCP server in addition to each tool's declared parameters (from its ``inputSchema.properties``). By default only declared parameters and these extras are sent. Accepts either a @@ -3946,7 +3411,6 @@ def __init__( load_prompts=load_prompts, parse_prompt_results=parse_prompt_results, request_timeout=request_timeout, - task_options=task_options, additional_tool_argument_names=additional_tool_argument_names, sampling_approval_callback=sampling_approval_callback, sampling_max_tokens=sampling_max_tokens, @@ -4037,7 +3501,6 @@ def __init__( http_client: httpx2.AsyncClient | None = None, static_headers: Mapping[str, str] | None = None, header_provider: Callable[[dict[str, Any]], dict[str, str]] | None = None, - task_options: MCPTaskOptions | None = None, additional_tool_argument_names: Sequence[str] | Mapping[str, Sequence[str]] | None = None, max_host_payload_size_bytes: int | None = _DEFAULT_MCP_HOST_PAYLOAD_SIZE_BYTES, tool_result_content: MCPToolResultContentMode = "structured_first", @@ -4147,8 +3610,8 @@ def __init__( values case-sensitively. Before an Agent exposes a connected tool's functions for a run, it reconciles the run's effective headers and reconnects when they differ; generated tool and prompt calls perform the same check before sending. Connection-lifetime - requests - including discovery, background pings, resource and prompt reloads, - and long-running task polling - always use the session-bound header set. A + requests - including discovery, background pings, and resource and prompt reloads - + always use the session-bound header set. A caller-supplied session's established identity is unknown and cannot be reconnected by this wrapper, so dynamic header resolution raises ``ToolExecutionException`` and requires a separate framework-managed tool instance. @@ -4176,8 +3639,6 @@ def __init__( values continue on to the outbound argument filter, so reading a credential here does not withhold it from the server. See ``additional_tool_argument_names`` below. - task_options: Options for tools that advertise - ``execution.taskSupport == "required"``. See :class:`MCPTaskOptions`. additional_tool_argument_names: Extra argument names to forward to the MCP server in addition to each tool's declared parameters (from its ``inputSchema.properties``). By default only declared parameters and these extras are sent. Accepts either a @@ -4225,7 +3686,6 @@ def __init__( load_prompts=load_prompts, parse_prompt_results=parse_prompt_results, request_timeout=request_timeout, - task_options=task_options, additional_tool_argument_names=additional_tool_argument_names, sampling_approval_callback=sampling_approval_callback, sampling_max_tokens=sampling_max_tokens, @@ -4639,7 +4099,6 @@ def __init__( sampling_max_tokens: int | None = _DEFAULT_SAMPLING_MAX_TOKENS, sampling_max_requests: int | None = _DEFAULT_SAMPLING_MAX_REQUESTS, additional_properties: dict[str, Any] | None = None, - task_options: MCPTaskOptions | None = None, additional_tool_argument_names: Sequence[str] | Mapping[str, Sequence[str]] | None = None, max_host_payload_size_bytes: int | None = _DEFAULT_MCP_HOST_PAYLOAD_SIZE_BYTES, tool_result_content: MCPToolResultContentMode = "structured_first", @@ -4708,8 +4167,6 @@ def __init__( (``min(requested, cap)``); ``None`` disables it. sampling_max_requests: Per-session cap on the number of sampling requests; further requests are rejected. Resets on reconnect. ``None`` disables it. - task_options: Options for tools that advertise - ``execution.taskSupport == "required"``. See :class:`MCPTaskOptions`. additional_tool_argument_names: Extra argument names to forward to the MCP server in addition to each tool's declared parameters (from its ``inputSchema.properties``). By default only declared parameters and these extras are sent. Accepts either a @@ -4754,7 +4211,6 @@ def __init__( load_prompts=load_prompts, parse_prompt_results=parse_prompt_results, request_timeout=request_timeout, - task_options=task_options, additional_tool_argument_names=additional_tool_argument_names, sampling_approval_callback=sampling_approval_callback, sampling_max_tokens=sampling_max_tokens, diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index a56322cf4e3..49e39ce0925 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -11,7 +11,6 @@ from collections.abc import AsyncIterator, Callable, Mapping from contextlib import AbstractAsyncContextManager, _AsyncGeneratorContextManager # type: ignore from contextvars import ContextVar -from datetime import timedelta from textwrap import dedent from types import TracebackType from typing import Any, cast @@ -385,7 +384,6 @@ async def test_ambiguous_policy_reload_preserves_previous_discovery( original_functions = list(tool._functions) original_meta = dict(tool._tool_call_meta_by_name) original_params = dict(tool._tool_param_names_by_name) - original_tasks = dict(tool._tool_task_support_by_name) async def list_tools(params: types.PaginatedRequestParams | None = None) -> types.ListToolsResult: assert tool._functions == original_functions @@ -406,7 +404,6 @@ async def list_tools(params: types.PaginatedRequestParams | None = None) -> type assert tool.functions == original_functions assert tool._tool_call_meta_by_name == original_meta assert tool._tool_param_names_by_name == original_params - assert tool._tool_task_support_by_name == original_tasks @pytest.mark.parametrize( @@ -502,7 +499,55 @@ async def list_tools(params: types.PaginatedRequestParams | None = None) -> type assert kept_function.result_parser is parser assert tool._tool_call_meta_by_name == ({} if empty_snapshot else {"keep": {"version": 2}}) assert tool._tool_param_names_by_name == ({} if empty_snapshot else {"keep": {"query"}}) - assert tool._tool_task_support_by_name == {} + + +async def test_load_tools_skips_only_tools_that_require_tasks(caplog: pytest.LogCaptureFixture) -> None: + tool = MCPTool(name="docs") # type: ignore[abstract] # ty: ignore[call-non-callable] + mock_session = AsyncMock() + tool.session = mock_session + mock_session.list_tools = AsyncMock( + return_value=types.ListToolsResult( + tools=[ + types.Tool( + name="required", + input_schema={"type": "object", "properties": {"secret": {"type": "string"}}}, + execution=types.ToolExecution(task_support="required"), + _meta={"must-not-leak": True}, + ), + types.Tool( + name="optional", + input_schema={"type": "object", "properties": {}}, + execution=types.ToolExecution(task_support="optional"), + _meta={"mode": "optional"}, + ), + types.Tool( + name="forbidden", + input_schema={"type": "object", "properties": {}}, + execution=types.ToolExecution(task_support="forbidden"), + ), + types.Tool(name="ordinary", input_schema={"type": "object", "properties": {}}), + ] + ) + ) + mock_session.call_tool = AsyncMock( + return_value=types.CallToolResult(content=[types.TextContent(type="text", text="ok")]) + ) + + with caplog.at_level(logging.WARNING, logger=logger.name): + await tool.load_tools() + + assert [function.name for function in tool.functions] == ["optional", "forbidden", "ordinary"] + assert tool._tool_call_meta_by_name == {"optional": {"mode": "optional"}} + assert set(tool._tool_param_names_by_name) == {"optional", "forbidden", "ordinary"} + assert ( + "Skipping MCP tool 'required' because it requires the Tasks extension, " + "which MCP Python SDK 2.2 does not implement." + ) in caplog.messages + + for function in tool.functions: + await function.invoke(arguments={}) + + assert [call.args[0] for call in mock_session.call_tool.await_args_list] == ["optional", "forbidden", "ordinary"] async def test_tool_refresh_preserves_prompt_when_server_advertises_same_raw_name() -> None: @@ -9986,1277 +10031,6 @@ def get_mcp_client(self): # pyrefly: ignore[bad-override] # endregion -# region: MCP long-running task (SEP-2663) tests - - -def _utc_now() -> str: - from datetime import datetime, timezone - - return datetime.now(timezone.utc).isoformat() - - -def _make_task_snapshot( - *, - task_id: str = "task-1", - status: str = "working", - status_message: str | None = None, - poll_interval_ms: int | None = None, -) -> types.GetTaskResult: - now = _utc_now() - return types.GetTaskResult( - task_id=task_id, - status=status, # type: ignore[arg-type] # ty: ignore[invalid-argument-type] - status_message=status_message, - created_at=now, - last_updated_at=now, - ttl=None, - poll_interval=poll_interval_ms, - ) - - -def _make_create_task_result(task_id: str = "task-1") -> types.CreateTaskResult: - now = _utc_now() - return types.CreateTaskResult( - task=types.Task( - task_id=task_id, - status="working", - status_message=None, - created_at=now, - last_updated_at=now, - ttl=None, - ) - ) - - -def _make_payload( - text: str = "done!", - is_error: bool = False, - structured_content: dict[str, Any] | None = None, - meta: dict[str, Any] | None = None, -) -> types.GetTaskPayloadResult: - payload: dict[str, Any] = { - "content": [{"type": "text", "text": text}], - "isError": is_error, - } - if structured_content is not None: - payload["structuredContent"] = structured_content - if meta is not None: - payload["_meta"] = meta - return types.GetTaskPayloadResult.model_validate(payload) - - -def _make_task_tool( - tool_name: str = "slow_op", - *, - task_support: str | None = "required", - task_options: Any = None, -) -> MCPTool: - from agent_framework import MCPTaskOptions - - tool = MCPTool( # type: ignore[abstract] # ty: ignore[call-non-callable] - name="lro", - task_options=task_options if task_options is not None else MCPTaskOptions(), - ) - tool.session = AsyncMock(spec=ClientSession) - if task_support is not None: - tool._tool_task_support_by_name[tool_name] = task_support - return tool - - -def _send_request_dispatcher(*responses_by_method: tuple[str, Any]) -> Any: - """Build a send_request side_effect that returns responses keyed by request method. - - Each tuple is ``(method_name, response_or_exception_or_callable)``. The dispatcher - advances a per-method queue on every call. A callable response is invoked with no - args so tests can raise exceptions deterministically. - """ - from collections import defaultdict - - queues: dict[str, list[Any]] = defaultdict(list) - for method, response in responses_by_method: - queues[method].append(response) - - async def _dispatch(request: Any, _result_type: Any, *_args: Any, **_kw: Any) -> Any: - method = getattr(request, "method", None) or getattr(request, "method", None) - queue = queues.get(method) # type: ignore[arg-type, call-overload] # pyrefly: ignore[bad-argument-type] - if not queue: - raise AssertionError(f"No mocked send_request response for method '{method}'.") - item = queue.pop(0) - if callable(item): - return item() - if isinstance(item, BaseException): - raise item - return item - - return _dispatch - - -async def test_task_options_defaults_are_sane() -> None: - from agent_framework import MCPTaskOptions - - opts = MCPTaskOptions() - assert opts.default_ttl is None - assert opts.cancel_remote_task_on_local_cancellation is True - - -async def test_task_options_rejects_non_positive_default_ttl() -> None: - from datetime import timedelta - - from agent_framework import MCPTaskOptions - - with pytest.raises(ValueError, match="positive"): - MCPTaskOptions(default_ttl=timedelta(seconds=-1)) - with pytest.raises(ValueError, match="positive"): - MCPTaskOptions(default_ttl=timedelta(0)) - - -async def test_load_tools_captures_task_support() -> None: - tool = MCPTool(name="lro") # type: ignore[abstract] # ty: ignore[call-non-callable] - mock_session = AsyncMock() - tool.session = mock_session - tool.load_tools_flag = True - - page = Mock() - page.tools = [ - types.Tool( - name="slow_op", - description="slow", - input_schema={"type": "object", "properties": {}}, - execution=types.ToolExecution(task_support="required"), - ), - types.Tool( - name="fast_op", - description="fast", - input_schema={"type": "object", "properties": {}}, - ), - ] - page.next_cursor = None - mock_session.list_tools = AsyncMock(return_value=page) - - await tool.load_tools() - - assert tool._tool_task_support_by_name == {"slow_op": "required"} - - -async def test_call_tool_routes_required_through_task_lifecycle(monkeypatch: pytest.MonkeyPatch) -> None: - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - tool = _make_task_tool() - tool.parse_tool_results = lambda _: "custom task summary" - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=_send_request_dispatcher( - ("tools/call", _make_create_task_result()), - ("tasks/get", _make_task_snapshot(status="working")), - ("tasks/get", _make_task_snapshot(status="completed")), - ( - "tasks/result", - _make_payload( - "hello task", - structured_content={"widget": "task"}, - meta={"source": "completed-task"}, - ), - ), - ) - ) - - function_result = await _call_generated_mcp_tool(tool, "slow_op", x=1) - - assert function_result.result == "custom task summary" - assert function_result.items is not None - assert function_result.additional_properties[_MCP_TOOL_RESULT_HOST_PAYLOAD_KEY]["structuredContent"] == { - "widget": "task" - } - assert "_meta" not in function_result.items[0].additional_properties - assert function_result.additional_properties["_meta"] == {"source": "completed-task"} - assert function_result.additional_properties[_MCP_TOOL_RESULT_HOST_PAYLOAD_KEY]["_meta"] == { - "source": "completed-task" - } - # Plain session.call_tool must NOT be used for required tools. - tool.session.call_tool.assert_not_called() # type: ignore[union-attr] # ty: ignore[unresolved-attribute] - - -async def test_call_tool_routes_required_through_public_task_override() -> None: - class OverriddenTaskTool(MCPTool): - def __init__(self) -> None: - super().__init__(name="override") - self.override_called = False - - async def call_tool_as_task(self, tool_name: str, **kwargs: Any) -> str | list[Content]: - self.override_called = True - return await super().call_tool_as_task(tool_name, **kwargs) - - tool = OverriddenTaskTool() # type: ignore[abstract] # ty: ignore[call-non-callable] - mock_session = AsyncMock(spec=ClientSession) - tool.session = mock_session - tool._tool_task_support_by_name["slow_op"] = "required" - fallback_result = types.CallToolResult(content=[types.TextContent(type="text", text="fallback")]) - mock_session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - return_value=types.Result.model_validate(fallback_result.model_dump(by_alias=True, exclude_none=True)) - ) - - function_result = await _call_generated_mcp_tool(tool, "slow_op") - - assert function_result.result == "fallback" - assert _MCP_TOOL_RESULT_HOST_PAYLOAD_KEY in function_result.additional_properties - assert tool.override_called is True - - -async def test_call_tool_as_task_fallback_preserves_custom_parser_host_payload() -> None: - """A legacy non-task response retains the Host payload after custom parsing.""" - tool = _make_task_tool() - tool.parse_tool_results = lambda _: "custom fallback summary" - fallback_result = types.CallToolResult( - content=[types.TextContent(type="text", text="fallback")], - structured_content={"widget": "fallback"}, - _meta={"source": "fallback"}, - ) - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - return_value=types.Result.model_validate(fallback_result.model_dump(by_alias=True, exclude_none=True)) - ) - - function_result = await _call_generated_mcp_tool(tool, "slow_op") - - assert function_result.result == "custom fallback summary" - assert function_result.items is not None - assert function_result.additional_properties[_MCP_TOOL_RESULT_HOST_PAYLOAD_KEY]["structuredContent"] == { - "widget": "fallback" - } - assert "_meta" not in function_result.items[0].additional_properties - assert function_result.additional_properties["_meta"] == {"source": "fallback"} - assert function_result.additional_properties[_MCP_TOOL_RESULT_HOST_PAYLOAD_KEY]["_meta"] == {"source": "fallback"} - - -@pytest.mark.parametrize("result_path", ["fallback", "completed"], ids=["task-fallback", "completed-task"]) -async def test_secure_mcp_task_results_cannot_relax_local_label(result_path: str) -> None: - from agent_framework.security import LabelTrackingFunctionMiddleware - - tool = _make_task_tool() - result_meta = {"ifc": {"integrity": "trusted", "confidentiality": "public"}} - structured_content = {"widget": result_path} - if result_path == "fallback": - raw_result = types.CallToolResult( - content=[types.TextContent(type="text", text="fallback")], - structured_content=structured_content, - _meta=result_meta, - ) - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - return_value=types.Result.model_validate(raw_result.model_dump(by_alias=True, exclude_none=True)) - ) - else: - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=_send_request_dispatcher( - ("tools/call", _make_create_task_result()), - ("tasks/get", _make_task_snapshot(status="completed")), - ( - "tasks/result", - _make_payload( - "completed", - structured_content=structured_content, - meta=result_meta, - ), - ), - ) - ) - - function_result = await _call_generated_mcp_tool( - tool, - "slow_op", - middleware_pipeline=FunctionMiddlewarePipeline(LabelTrackingFunctionMiddleware(auto_hide_untrusted=True)), - host_payload_budget=_FunctionResultPayloadBudget(), - mcp_local_label=("untrusted", "private"), - ) - - assert function_result.items is not None - assert len(function_result.items) == 1 - for item in function_result.items: - assert item.additional_properties["_variable_reference"] is True - assert item.additional_properties["security_label"]["integrity"] == "untrusted" - assert item.additional_properties["security_label"]["confidentiality"] == "private" - assert item.additional_properties["_meta"] == result_meta - assert function_result.additional_properties["_meta"] == result_meta - assert function_result.additional_properties[_MCP_TOOL_RESULT_HOST_PAYLOAD_KEY]["structuredContent"] == ( - structured_content - ) - - -@pytest.mark.parametrize("result_path", ["fallback", "completed"], ids=["task-fallback", "completed-task"]) -async def test_task_parser_failure_preserves_complete_host_payload(result_path: str) -> None: - tool = _make_task_tool() - tool.parse_tool_results = _raise_result_parser - result_meta = {"source": result_path} - if result_path == "fallback": - raw_result = types.CallToolResult( - content=[types.TextContent(type="text", text="fallback")], - structured_content={"widget": result_path}, - _meta=result_meta, - ) - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - return_value=types.Result.model_validate(raw_result.model_dump(by_alias=True, exclude_none=True)) - ) - else: - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=_send_request_dispatcher( - ("tools/call", _make_create_task_result()), - ("tasks/get", _make_task_snapshot(status="completed")), - ( - "tasks/result", - _make_payload( - "completed", - structured_content={"widget": result_path}, - meta=result_meta, - ), - ), - ) - ) - - function_result = await _call_generated_mcp_tool(tool, "slow_op") - - assert function_result.result == "Error: Function failed." - assert function_result.additional_properties["_meta"] == result_meta - assert function_result.additional_properties[_MCP_TOOL_RESULT_HOST_PAYLOAD_KEY]["structuredContent"] == { - "widget": result_path - } - - -async def test_call_tool_as_task_default_ttl_propagates() -> None: - from datetime import timedelta - - from agent_framework import MCPTaskOptions - - tool = _make_task_tool(task_options=MCPTaskOptions(default_ttl=timedelta(minutes=7))) - - captured: list[Any] = [] - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - captured.append(request) - method = request.method - if method == "tools/call": - return _make_create_task_result() - if method == "tasks/get": - return _make_task_snapshot(status="completed") - if method == "tasks/result": - return _make_payload("ok") - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - await tool.call_tool("slow_op") - - create_req = captured[0] - assert create_req.method == "tools/call" - assert create_req.params.task is not None - assert create_req.params.task.ttl == 7 * 60 * 1000 - - -async def test_call_tool_as_task_sends_empty_task_metadata_when_ttl_none() -> None: - # Without a TTL we still mark the call as task-augmented (servers require - # the `task` field to route through the lifecycle). - tool = _make_task_tool() - - captured: list[Any] = [] - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - captured.append(request) - method = request.method - if method == "tools/call": - return _make_create_task_result() - if method == "tasks/get": - return _make_task_snapshot(status="completed") - if method == "tasks/result": - return _make_payload("ok") - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - await tool.call_tool("slow_op") - - create_req = captured[0] - assert create_req.method == "tools/call" - assert create_req.params.task is not None - assert create_req.params.task.ttl is None - - -async def test_call_tool_skips_task_path_for_optional_and_forbidden() -> None: - for support in ("optional", "forbidden", None): - tool = _make_task_tool(task_support=support) - tool.session.call_tool = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - return_value=types.CallToolResult(content=[types.TextContent(type="text", text="plain")]) - ) - tool.session.send_request = AsyncMock(side_effect=AssertionError("task path should not be used")) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - result = await tool.call_tool("slow_op") - assert _mcp_result_to_text(result) == "plain" - - -async def test_call_tool_as_task_cancelled_status_raises() -> None: - tool = _make_task_tool() - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=_send_request_dispatcher( - ("tools/call", _make_create_task_result()), - ("tasks/get", _make_task_snapshot(status="cancelled", status_message="server stop")), - ) - ) - - with pytest.raises(ToolExecutionException, match="cancelled.*server stop"): - await tool.call_tool("slow_op") - - -async def test_call_tool_as_task_failed_status_raises() -> None: - tool = _make_task_tool() - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=_send_request_dispatcher( - ("tools/call", _make_create_task_result()), - ("tasks/get", _make_task_snapshot(status="failed", status_message="boom")), - ) - ) - - with pytest.raises(ToolExecutionException, match="failed.*boom"): - await tool.call_tool("slow_op") - - -async def test_call_tool_as_task_input_required_raises() -> None: - tool = _make_task_tool() - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=_send_request_dispatcher( - ("tools/call", _make_create_task_result()), - ("tasks/get", _make_task_snapshot(status="input_required", status_message="need more")), - ) - ) - - with pytest.raises(ToolExecutionException, match="input_required.*need more"): - await tool.call_tool("slow_op") - - -async def test_call_tool_as_task_payload_iserror_raises() -> None: - tool = _make_task_tool() - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=_send_request_dispatcher( - ("tools/call", _make_create_task_result()), - ("tasks/get", _make_task_snapshot(status="completed")), - ( - "tasks/result", - _make_payload( - "payload exploded", - is_error=True, - structured_content={"reason": "task failed"}, - meta={"source": "failed-task"}, - ), - ), - ) - ) - - function_result = await _call_generated_mcp_tool(tool, "slow_op") - - assert function_result.exception is not None - assert function_result.additional_properties["_meta"] == {"source": "failed-task"} - assert function_result.additional_properties[_MCP_TOOL_RESULT_HOST_PAYLOAD_KEY]["structuredContent"] == { - "reason": "task failed" - } - - -async def test_call_tool_as_task_malformed_payload_raises() -> None: - tool = _make_task_tool() - bad_payload = types.GetTaskPayloadResult.model_validate({"random": "stuff"}) - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=_send_request_dispatcher( - ("tools/call", _make_create_task_result(task_id="abc")), - ("tasks/get", _make_task_snapshot(task_id="abc", status="completed")), - ("tasks/result", bad_payload), - ) - ) - - with pytest.raises(ToolExecutionException, match="task 'abc' result payload"): - await tool.call_tool("slow_op") - - -async def test_call_tool_as_task_method_not_found_falls_back() -> None: - tool = _make_task_tool() - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=MCPError(types.METHOD_NOT_FOUND, "no tasks here") - ) - tool.session.call_tool = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - return_value=types.CallToolResult(content=[types.TextContent(type="text", text="fell back")]) - ) - - result = await tool.call_tool("slow_op") - - assert _mcp_result_to_text(result) == "fell back" - tool.session.call_tool.assert_awaited_once() # type: ignore[union-attr] # ty: ignore[unresolved-attribute] - - -async def test_call_tool_as_task_invalid_params_falls_back() -> None: - tool = _make_task_tool() - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=MCPError(types.INVALID_PARAMS, "unknown field") - ) - tool.session.call_tool = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - return_value=types.CallToolResult(content=[types.TextContent(type="text", text="plain ok")]) - ) - - result = await tool.call_tool("slow_op") - - assert _mcp_result_to_text(result) == "plain ok" - - -async def test_call_tool_as_task_legacy_calltoolresult_response_used_directly() -> None: - """Server may ignore augmentation and return CallToolResult; treat it as the result.""" - # Build a lenient Result whose extras match a CallToolResult shape. - legacy_payload = types.Result.model_validate({ - "content": [{"type": "text", "text": "legacy ok"}], - "isError": False, - }) - - tool = _make_task_tool() - tool.session.send_request = AsyncMock(return_value=legacy_payload) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - result = await tool.call_tool("slow_op") - - assert _mcp_result_to_text(result) == "legacy ok" - # Polling must not occur: a single tools/call was enough. - assert tool.session.send_request.call_count == 1 # type: ignore[union-attr] # ty: ignore[unresolved-attribute] - - -async def test_call_tool_as_task_poll_interval_is_clamped(monkeypatch: pytest.MonkeyPatch) -> None: - from datetime import timedelta as _td - - from agent_framework import _mcp as _mcp_module - - # Stub asyncio.sleep so we can capture delays without actually sleeping. - delays: list[float] = [] - - async def fake_sleep(delay: float) -> None: - delays.append(delay) - - monkeypatch.setattr(_mcp_module.asyncio, "sleep", fake_sleep) - - tool = _make_task_tool() - tool.session.send_request = AsyncMock( # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - side_effect=_send_request_dispatcher( - ("tools/call", _make_create_task_result()), - ("tasks/get", _make_task_snapshot(status="working", poll_interval_ms=50)), # below 500ms min - ("tasks/get", _make_task_snapshot(status="working", poll_interval_ms=10_000)), # above 5s max - ("tasks/get", _make_task_snapshot(status="working", poll_interval_ms=None)), # default to min - ("tasks/get", _make_task_snapshot(status="working", poll_interval_ms=0)), # invalid -> min - ("tasks/get", _make_task_snapshot(status="working", poll_interval_ms=2_000)), # in-band - ("tasks/get", _make_task_snapshot(status="completed")), - ("tasks/result", _make_payload("ok")), - ) - ) - - await tool.call_tool("slow_op") - - expected = [ - _td(milliseconds=500).total_seconds(), # clamp up - _td(seconds=5).total_seconds(), # clamp down - _td(milliseconds=500).total_seconds(), # missing -> min - _td(milliseconds=500).total_seconds(), # zero -> min - _td(milliseconds=2_000).total_seconds(), - ] - assert delays == expected - - -async def test_call_tool_as_task_local_cancellation_fires_remote_cancel( - monkeypatch: pytest.MonkeyPatch, -) -> None: - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - tool = _make_task_tool() - - cancel_seen = asyncio.Event() - create_seen = asyncio.Event() - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - method = request.method - if method == "tools/call": - create_seen.set() - return _make_create_task_result() - if method == "tasks/get": - await asyncio.sleep(0) - return _make_task_snapshot(status="working") - if method == "tasks/cancel": - cancel_seen.set() - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - task = asyncio.create_task(tool.call_tool("slow_op")) - await asyncio.wait_for(create_seen.wait(), timeout=1.0) - # Let polling iterate a few times. - await asyncio.sleep(0.02) - task.cancel() - with pytest.raises(asyncio.CancelledError): - await task - - # Wait for the fire-and-forget cancel to complete. - await asyncio.wait_for(cancel_seen.wait(), timeout=1.0) - # Drain any tracked background tasks. - pending = list(tool._pending_reload_tasks) - if pending: - await asyncio.gather(*pending, return_exceptions=True) - - -async def test_call_tool_as_task_cancellation_suppressed_when_disabled( - monkeypatch: pytest.MonkeyPatch, -) -> None: - from agent_framework import MCPTaskOptions - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - tool = _make_task_tool( - task_options=MCPTaskOptions(cancel_remote_task_on_local_cancellation=False), - ) - - cancel_called = False - create_seen = asyncio.Event() - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal cancel_called - method = request.method - if method == "tools/call": - create_seen.set() - return _make_create_task_result() - if method == "tasks/get": - await asyncio.sleep(0) - return _make_task_snapshot(status="working") - if method == "tasks/cancel": - cancel_called = True - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - task = asyncio.create_task(tool.call_tool("slow_op")) - await asyncio.wait_for(create_seen.wait(), timeout=1.0) - await asyncio.sleep(0.02) - task.cancel() - with pytest.raises(asyncio.CancelledError): - await task - - # Let any (incorrect) background work settle, then verify cancel was NOT sent. - await asyncio.sleep(0.02) - assert cancel_called is False - - -async def test_call_tool_as_task_reconnects_during_poll(monkeypatch: pytest.MonkeyPatch) -> None: - from anyio import ClosedResourceError - - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - tool = _make_task_tool() - - poll_calls = 0 - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal poll_calls - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="abc") - if method == "tasks/get": - poll_calls += 1 - assert request.params.task_id == "abc" - if poll_calls == 1: - raise ClosedResourceError - return _make_task_snapshot(task_id="abc", status="completed") - if method == "tasks/result": - return _make_payload("recovered") - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - reconnect_calls = 0 - - async def fake_connect(reset: bool = False) -> None: - nonlocal reconnect_calls - reconnect_calls += 1 - assert reset is True - - with patch.object(MCPTool, "connect", side_effect=fake_connect): - result = await tool.call_tool("slow_op") - - assert _mcp_result_to_text(result) == "recovered" - assert reconnect_calls == 1 - # Critically, tools/call must NOT be re-issued after task_id is known. - assert ( - sum( - 1 # type: ignore[misc] - for c in tool.session.send_request.await_args_list # type: ignore[union-attr] # ty: ignore[unresolved-attribute] - if c.args[0].method == "tools/call" - ) - == 1 - ) - - -async def test_call_tool_as_task_second_disconnect_raises_connection_lost( - monkeypatch: pytest.MonkeyPatch, -) -> None: - from anyio import ClosedResourceError - - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - tool = _make_task_tool() - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="abc") - if method == "tasks/get": - raise ClosedResourceError - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - with ( - patch.object(MCPTool, "connect", new=AsyncMock(return_value=None)), - pytest.raises(ToolExecutionException, match="task state unknown"), - ): - await tool.call_tool("slow_op") - - -async def test_call_tool_as_task_create_disconnect_does_not_retry() -> None: - """A connection loss during the augmented tools/call must NOT retry. - - Retrying could spawn a duplicate long-running task on the server, because the - first request may have been accepted before the response was lost. - """ - from anyio import ClosedResourceError - - tool = _make_task_tool() - - send_calls = 0 - - async def fake_send(_request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal send_calls - send_calls += 1 - raise ClosedResourceError - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - reconnect_mock = AsyncMock(return_value=None) - with ( - patch.object(MCPTool, "connect", new=reconnect_mock), - pytest.raises(ToolExecutionException, match="task state unknown"), - ): - await tool.call_tool("slow_op") - - # Exactly one tools/call was issued — the server-side task state is unknown, - # so retry is unsafe and must be skipped. - assert send_calls == 1 - reconnect_mock.assert_not_awaited() - - -async def test_fetch_task_result_reconnects_during_fetch() -> None: - from anyio import ClosedResourceError - - tool = _make_task_tool() - - fetch_calls = 0 - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal fetch_calls - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="r1") - if method == "tasks/get": - return _make_task_snapshot(task_id="r1", status="completed") - if method == "tasks/result": - fetch_calls += 1 - if fetch_calls == 1: - raise ClosedResourceError - return _make_payload("fetched after reconnect") - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - reconnect_calls = 0 - - async def fake_connect(reset: bool = False) -> None: - nonlocal reconnect_calls - reconnect_calls += 1 - assert reset is True - - with patch.object(MCPTool, "connect", side_effect=fake_connect): - result = await tool.call_tool("slow_op") - - assert _mcp_result_to_text(result) == "fetched after reconnect" - assert reconnect_calls == 1 - assert fetch_calls == 2 - - -async def test_fetch_task_result_second_disconnect_raises_task_state_unknown_and_cancels() -> None: - from anyio import ClosedResourceError - - tool = _make_task_tool() - - cancel_called = False - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal cancel_called - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="r2") - if method == "tasks/get": - return _make_task_snapshot(task_id="r2", status="completed") - if method == "tasks/result": - raise ClosedResourceError - if method == "tasks/cancel": - cancel_called = True - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - with ( - patch.object(MCPTool, "connect", new=AsyncMock(return_value=None)), - pytest.raises(ToolExecutionException, match="task state unknown"), - ): - await tool.call_tool("slow_op") - - # Drain the fire-and-forget cancel so the assertion is deterministic. - pending = list(tool._pending_reload_tasks) - if pending: - await asyncio.gather(*pending, return_exceptions=True) - assert cancel_called is True - - -async def test_call_tool_as_task_create_unparseable_success_raises() -> None: - """An unparseable success-shaped response must NOT silently retry tools/call.""" - # Result with neither task.taskId nor a valid CallToolResult shape. - unparseable = types.Result.model_validate({"foo": "bar"}) - - tool = _make_task_tool() - tool.session.send_request = AsyncMock(return_value=unparseable) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - tool.session.call_tool = AsyncMock(return_value=types.CallToolResult(content=[])) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - with pytest.raises(ToolExecutionException, match="unparseable response"): - await tool.call_tool("slow_op") - - # Critically: no plain tools/call fallback (would risk double execution). - tool.session.call_tool.assert_not_called() # type: ignore[union-attr] # ty: ignore[unresolved-attribute] - - -async def test_call_tool_as_task_max_wait_exceeded_raises_and_cancels(monkeypatch: pytest.MonkeyPatch) -> None: - from agent_framework import MCPTaskOptions - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - tool = _make_task_tool(task_options=MCPTaskOptions(max_task_wait=timedelta(milliseconds=50))) - - cancel_called = False - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal cancel_called - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="mw") - if method == "tasks/get": - return _make_task_snapshot(task_id="mw", status="working") - if method == "tasks/cancel": - cancel_called = True - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - with pytest.raises(ToolExecutionException, match="exceeded max_task_wait"): - await tool.call_tool("slow_op") - - pending = list(tool._pending_reload_tasks) - if pending: - await asyncio.gather(*pending, return_exceptions=True) - assert cancel_called is True - - -async def test_call_tool_as_task_max_wait_cancels_even_when_local_cancel_option_disabled( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Locks contract: max_task_wait abandonment ignores the local-cancel option.""" - from agent_framework import MCPTaskOptions - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - tool = _make_task_tool( - task_options=MCPTaskOptions( - cancel_remote_task_on_local_cancellation=False, - max_task_wait=timedelta(milliseconds=50), - ), - ) - - cancel_called = False - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal cancel_called - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="mw2") - if method == "tasks/get": - return _make_task_snapshot(task_id="mw2", status="working") - if method == "tasks/cancel": - cancel_called = True - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - with pytest.raises(ToolExecutionException, match="exceeded max_task_wait"): - await tool.call_tool("slow_op") - - pending = list(tool._pending_reload_tasks) - if pending: - await asyncio.gather(*pending, return_exceptions=True) - assert cancel_called is True - - -async def test_call_tool_as_task_poll_transient_request_timeout_keeps_polling( - monkeypatch: pytest.MonkeyPatch, -) -> None: - import httpx2 as httpx - - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - tool = _make_task_tool() - - poll_calls = 0 - cancel_called = False - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal poll_calls, cancel_called - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="t1") - if method == "tasks/get": - poll_calls += 1 - if poll_calls == 1: - raise MCPError(int(httpx.codes.REQUEST_TIMEOUT), "slow poll") - return _make_task_snapshot(task_id="t1", status="completed") - if method == "tasks/result": - return _make_payload("recovered after transient") - if method == "tasks/cancel": - cancel_called = True - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - result = await tool.call_tool("slow_op") - assert _mcp_result_to_text(result) == "recovered after transient" - assert poll_calls == 2 - # Transient retry must not fire cancel. - pending = list(tool._pending_reload_tasks) - if pending: - await asyncio.gather(*pending, return_exceptions=True) - assert cancel_called is False - - -async def test_call_tool_as_task_poll_hard_mcperror_cancels_and_raises( - monkeypatch: pytest.MonkeyPatch, -) -> None: - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - tool = _make_task_tool() - - cancel_called = False - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal cancel_called - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="h1") - if method == "tasks/get": - raise MCPError(types.INVALID_PARAMS, "bad task id") - if method == "tasks/cancel": - cancel_called = True - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - with pytest.raises(ToolExecutionException, match="bad task id"): - await tool.call_tool("slow_op") - - pending = list(tool._pending_reload_tasks) - if pending: - await asyncio.gather(*pending, return_exceptions=True) - assert cancel_called is True - - -async def test_call_tool_as_task_malformed_tasks_get_response_cancels_and_raises( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Malformed tasks/get response counts as abandonment (task may still be running).""" - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - tool = _make_task_tool() - - # Result without a valid GetTaskResult shape (no taskId/status/etc.). - malformed = types.Result.model_validate({"some": "junk"}) - - cancel_called = False - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal cancel_called - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="m1") - if method == "tasks/get": - return malformed - if method == "tasks/cancel": - cancel_called = True - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - with pytest.raises(ToolExecutionException, match="malformed tasks/get"): - await tool.call_tool("slow_op") - - pending = list(tool._pending_reload_tasks) - if pending: - await asyncio.gather(*pending, return_exceptions=True) - assert cancel_called is True - - -async def test_call_tool_as_task_failed_terminal_does_not_cancel(monkeypatch: pytest.MonkeyPatch) -> None: - """Terminal failures (server already done) must NOT fire tasks/cancel.""" - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - tool = _make_task_tool() - - cancel_called = False - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal cancel_called - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="f1") - if method == "tasks/get": - return _make_task_snapshot(task_id="f1", status="failed", status_message="boom") - if method == "tasks/cancel": - cancel_called = True - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - with pytest.raises(ToolExecutionException, match="task failed: boom"): - await tool.call_tool("slow_op") - - # Let any (incorrect) background work settle, then verify no cancel. - await asyncio.sleep(0.02) - assert cancel_called is False - - -async def test_try_cancel_task_logs_warning_on_timeout( - caplog: pytest.LogCaptureFixture, monkeypatch: pytest.MonkeyPatch -) -> None: - from agent_framework import _mcp as _mcp_module - - # Shorten cancel timeout so the test is fast. - monkeypatch.setattr(_mcp_module, "_MCP_TASK_CANCEL_TIMEOUT", _mcp_module.timedelta(milliseconds=20)) - - tool = _make_task_tool() - - async def hang(*_a: Any, **_kw: Any) -> Any: - await asyncio.sleep(10.0) - - tool.session.send_request = AsyncMock(side_effect=hang) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - with caplog.at_level(logging.WARNING, logger=_mcp_module.logger.name): - await tool._try_cancel_task("hang-1") - - assert any("timed out" in r.getMessage() and "hang-1" in r.getMessage() for r in caplog.records) - - -async def test_mcp_task_options_is_frozen() -> None: - from dataclasses import FrozenInstanceError - - from agent_framework import MCPTaskOptions - - opts = MCPTaskOptions() - with pytest.raises(FrozenInstanceError): - opts.default_ttl = timedelta(seconds=5) # type: ignore[misc] # ty: ignore[invalid-assignment] - - -async def test_mcp_task_options_max_task_wait_rejects_non_positive() -> None: - from agent_framework import MCPTaskOptions - - with pytest.raises(ValueError, match="positive"): - MCPTaskOptions(max_task_wait=timedelta(0)) - with pytest.raises(ValueError, match="positive"): - MCPTaskOptions(max_task_wait=timedelta(seconds=-1)) - - -async def test_fetch_task_result_hard_mcperror_raises_without_cancel() -> None: - """tasks/result hard MCPError must wrap as ToolExecutionException without cancel (server done).""" - tool = _make_task_tool() - - cancel_called = False - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal cancel_called - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="hf") - if method == "tasks/get": - return _make_task_snapshot(task_id="hf", status="completed") - if method == "tasks/result": - raise MCPError(types.INTERNAL_ERROR, "payload vanished") - if method == "tasks/cancel": - cancel_called = True - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - with pytest.raises(ToolExecutionException, match="payload vanished"): - await tool.call_tool("slow_op") - - # No raw MCPError leak and no cancel — server already reported the task as done. - await asyncio.sleep(0.02) - assert cancel_called is False - - -async def test_completion_wait_timeout_without_max_wait_is_not_translated(monkeypatch: pytest.MonkeyPatch) -> None: - """Stray asyncio.TimeoutError during the completion wait must not pretend the deadline - expired when max_task_wait is None (and must not fire a spurious tasks/cancel). - """ - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - tool = _make_task_tool() - - def boom_parser(_: Any) -> list[Content]: - raise asyncio.TimeoutError - - tool.parse_tool_results = boom_parser - - cancel_called = False - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal cancel_called - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="t2") - if method == "tasks/get": - return _make_task_snapshot(task_id="t2", status="completed") - if method == "tasks/result": - return _make_payload("ok") - if method == "tasks/cancel": - cancel_called = True - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - with pytest.raises(asyncio.TimeoutError): - await tool.call_tool("slow_op") - - # Must NOT translate to max_task_wait expiry and must NOT cancel. - await asyncio.sleep(0.02) - assert cancel_called is False - - -async def test_completion_wait_inner_timeout_with_max_wait_set_propagates( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """An asyncio.TimeoutError raised by the completion wait itself must propagate - unchanged even when max_task_wait IS set, and must NOT fire a spurious cancel. - """ - from agent_framework import MCPTaskOptions - from agent_framework import _mcp as _mcp_module - - monkeypatch.setattr(_mcp_module, "_MCP_TASK_MIN_POLL_INTERVAL", _mcp_module.timedelta(milliseconds=1)) - - # Deadline set comfortably above the actual test run time. - tool = _make_task_tool(task_options=MCPTaskOptions(max_task_wait=timedelta(seconds=5))) - - def boom_parser(_: Any) -> list[Content]: - raise asyncio.TimeoutError("inner parser timeout") - - tool.parse_tool_results = boom_parser - - cancel_called = False - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - nonlocal cancel_called - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="t3") - if method == "tasks/get": - return _make_task_snapshot(task_id="t3", status="completed") - if method == "tasks/result": - return _make_payload("ok") - if method == "tasks/cancel": - cancel_called = True - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - with pytest.raises(asyncio.TimeoutError, match="inner parser timeout"): - await tool.call_tool("slow_op") - - # Inner TimeoutError must NOT be translated into "exceeded max_task_wait" and must NOT cancel. - await asyncio.sleep(0.02) - assert cancel_called is False - - -async def test_max_wait_interrupts_long_poll_sleep(monkeypatch: pytest.MonkeyPatch) -> None: - """Deadline must cancel through a long ``asyncio.sleep`` (clamped to MAX), not wait it out.""" - from agent_framework import MCPTaskOptions - - tool = _make_task_tool(task_options=MCPTaskOptions(max_task_wait=timedelta(milliseconds=100))) - - async def fake_send(request: Any, _result_type: Any, *_a: Any, **_kw: Any) -> Any: - method = request.method - if method == "tools/call": - return _make_create_task_result(task_id="ds") - if method == "tasks/get": - # Suggest a 5s poll interval (gets clamped to MAX=5s); wait_for must cut through it. - return _make_task_snapshot(task_id="ds", status="working", poll_interval_ms=5000) - if method == "tasks/cancel": - return types.CancelTaskResult() # type: ignore[call-arg] # pyrefly: ignore[missing-argument] # ty: ignore[missing-argument] - raise AssertionError(method) - - tool.session.send_request = AsyncMock(side_effect=fake_send) # type: ignore[method-assign, union-attr] # ty: ignore[invalid-assignment] - - loop = asyncio.get_running_loop() - started = loop.time() - with pytest.raises(ToolExecutionException, match="exceeded max_task_wait"): - await tool.call_tool("slow_op") - elapsed = loop.time() - started - - # Should fire near the 100ms deadline, well below the 5s clamped sleep. - assert elapsed < 1.0, f"deadline did not interrupt long sleep (elapsed={elapsed:.3f}s)" - - pending = list(tool._pending_reload_tasks) - if pending: - await asyncio.gather(*pending, return_exceptions=True) - - -# endregion - - # region additional_tool_argument_names / allowlist filtering diff --git a/python/packages/core/tests/core/test_mcp_http_auth.py b/python/packages/core/tests/core/test_mcp_http_auth.py index 31eb3a80b80..4fa62f7ff3e 100644 --- a/python/packages/core/tests/core/test_mcp_http_auth.py +++ b/python/packages/core/tests/core/test_mcp_http_auth.py @@ -737,7 +737,6 @@ async def test_cancelled_identity_switch_keeps_the_existing_session_bound(mcp_ht existing_session = tool.session existing_functions = list(tool.functions) existing_call_meta = dict(tool._tool_call_meta_by_name) - existing_task_support = dict(tool._tool_task_support_by_name) existing_param_names = {name: set(params) for name, params in tool._tool_param_names_by_name.items()} async def cancel_reconnect() -> None: @@ -753,7 +752,6 @@ async def cancel_reconnect() -> None: assert tool.session is existing_session assert tool.functions == existing_functions assert tool._tool_call_meta_by_name == existing_call_meta - assert tool._tool_task_support_by_name == existing_task_support assert tool._tool_param_names_by_name == existing_param_names await tool.call_tool("record", credential="token-b") @@ -941,7 +939,6 @@ def create_client(**kwargs: Any) -> httpx.AsyncClient: assert not tool._prompts_loaded assert tool.functions == [] assert tool._tool_call_meta_by_name == {} - assert tool._tool_task_support_by_name == {} assert tool._tool_param_names_by_name == {} assert all(not client.event_hooks["request"] for client in clients) assert all(client.is_closed is owned_client for client in clients)