Skip to content

Commit 413d8f1

Browse files
tylerpayneTyler Payneekzhuvictordibia
authored
Expand MCP Workbench to support more MCP Client features (#6785)
Co-authored-by: Tyler Payne <tylerpayne@microsoft.com> Co-authored-by: Eric Zhu <ekzhu@users.noreply.github.com> Co-authored-by: Victor Dibia <victordibia@microsoft.com>
1 parent aa131bb commit 413d8f1

8 files changed

Lines changed: 3583 additions & 18 deletions

File tree

python/packages/autogen-ext/src/autogen_ext/tools/mcp/_actor.py

Lines changed: 173 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1,27 +1,87 @@
11
import asyncio
22
import atexit
3+
import base64
4+
import io
5+
import logging
36
from typing import Any, Coroutine, Dict, Mapping, TypedDict
47

5-
from autogen_core import Component, ComponentBase
8+
from autogen_core import Component, ComponentBase, ComponentModel, Image
9+
from autogen_core.models import (
10+
AssistantMessage,
11+
ChatCompletionClient,
12+
LLMMessage,
13+
ModelInfo,
14+
SystemMessage,
15+
UserMessage,
16+
)
17+
from PIL import Image as PILImage
618
from pydantic import BaseModel
719
from typing_extensions import Self
820

9-
from mcp.types import CallToolResult, ListToolsResult
21+
from mcp import types as mcp_types
22+
from mcp.client.session import ClientSession
23+
from mcp.shared.context import RequestContext
1024

1125
from ._config import McpServerParams
1226
from ._session import create_mcp_server_session
1327

14-
McpResult = Coroutine[Any, Any, ListToolsResult] | Coroutine[Any, Any, CallToolResult]
28+
logger = logging.getLogger(__name__)
29+
30+
McpResult = (
31+
Coroutine[Any, Any, mcp_types.ListToolsResult]
32+
| Coroutine[Any, Any, mcp_types.CallToolResult]
33+
| Coroutine[Any, Any, mcp_types.ListPromptsResult]
34+
| Coroutine[Any, Any, mcp_types.ListResourcesResult]
35+
| Coroutine[Any, Any, mcp_types.ListResourceTemplatesResult]
36+
| Coroutine[Any, Any, mcp_types.ReadResourceResult]
37+
| Coroutine[Any, Any, mcp_types.GetPromptResult]
38+
)
1539
McpFuture = asyncio.Future[McpResult]
1640

1741

42+
def _parse_sampling_content(
43+
content: mcp_types.TextContent | mcp_types.ImageContent | mcp_types.AudioContent, model_info: ModelInfo
44+
) -> str | Image:
45+
"""Convert MCP content types to Autogen content types."""
46+
if content.type == "text":
47+
return content.text
48+
elif content.type == "image":
49+
if not model_info["vision"]:
50+
raise ValueError("Sampling model does not support image content.")
51+
# Decode base64 image data and create PIL Image
52+
image_data = base64.b64decode(content.data)
53+
pil_image = PILImage.open(io.BytesIO(image_data))
54+
return Image.from_pil(pil_image)
55+
else:
56+
raise ValueError(f"Unsupported content type: {content.type}")
57+
58+
59+
def _parse_sampling_message(message: mcp_types.SamplingMessage, model_info: ModelInfo) -> LLMMessage:
60+
"""Convert MCP sampling messages to Autogen messages."""
61+
content = _parse_sampling_content(message.content, model_info=model_info)
62+
if message.role == "user":
63+
return UserMessage(
64+
source="user",
65+
content=[content],
66+
)
67+
elif message.role == "assistant":
68+
assert isinstance(content, str), "Assistant messages only support string content."
69+
return AssistantMessage(
70+
source="assistant",
71+
content=content,
72+
)
73+
else:
74+
raise ValueError(f"Unrecognized message role: {message.role}")
75+
76+
1877
class McpActorArgs(TypedDict):
1978
name: str | None
2079
kargs: Mapping[str, Any]
2180

2281

2382
class McpSessionActorConfig(BaseModel):
2483
server_params: McpServerParams
84+
model_client: ComponentModel | Dict[str, Any] | None = None
2585

2686

2787
class McpSessionActor(ComponentBase[BaseModel], Component[McpSessionActorConfig]):
@@ -33,16 +93,22 @@ class McpSessionActor(ComponentBase[BaseModel], Component[McpSessionActorConfig]
3393

3494
# model_config = ConfigDict(arbitrary_types_allowed=True)
3595

36-
def __init__(self, server_params: McpServerParams) -> None:
96+
def __init__(self, server_params: McpServerParams, model_client: ChatCompletionClient | None = None) -> None:
3797
self.server_params: McpServerParams = server_params
98+
self._model_client = model_client
3899
self.name = "mcp_session_actor"
39100
self.description = "MCP session actor"
40101
self._command_queue: asyncio.Queue[Dict[str, Any]] = asyncio.Queue()
41102
self._actor_task: asyncio.Task[Any] | None = None
42103
self._shutdown_future: asyncio.Future[Any] | None = None
43104
self._active = False
105+
self._initialize_result: mcp_types.InitializeResult | None = None
44106
atexit.register(self._sync_shutdown)
45107

108+
@property
109+
def initialize_result(self) -> mcp_types.InitializeResult | None:
110+
return self._initialize_result
111+
46112
async def initialize(self) -> None:
47113
if not self._active:
48114
self._active = True
@@ -54,17 +120,28 @@ async def call(self, type: str, args: McpActorArgs | None = None) -> McpFuture:
54120
if self._actor_task and self._actor_task.done():
55121
raise RuntimeError("MCP actor task crashed", self._actor_task.exception())
56122
fut: asyncio.Future[McpFuture] = asyncio.Future()
57-
if type in {"list_tools", "shutdown"}:
123+
if type in {"list_tools", "list_prompts", "list_resources", "list_resource_templates", "shutdown"}:
58124
await self._command_queue.put({"type": type, "future": fut})
59125
res = await fut
60-
elif type == "call_tool":
126+
elif type in {"call_tool", "read_resource", "get_prompt"}:
61127
if args is None:
62-
raise ValueError("args is required for call_tool")
128+
raise ValueError(f"args is required for {type}")
63129
name = args.get("name", None)
64130
kwargs = args.get("kargs", {})
65-
if name is None:
131+
if type == "call_tool" and name is None:
66132
raise ValueError("name is required for call_tool")
67-
await self._command_queue.put({"type": type, "name": name, "args": kwargs, "future": fut})
133+
elif type == "read_resource":
134+
uri = kwargs.get("uri", None)
135+
if uri is None:
136+
raise ValueError("uri is required for read_resource")
137+
await self._command_queue.put({"type": type, "uri": uri, "future": fut})
138+
elif type == "get_prompt":
139+
if name is None:
140+
raise ValueError("name is required for get_prompt")
141+
prompt_args = kwargs.get("arguments", None)
142+
await self._command_queue.put({"type": type, "name": name, "args": prompt_args, "future": fut})
143+
else: # call_tool
144+
await self._command_queue.put({"type": type, "name": name, "args": kwargs, "future": fut})
68145
res = await fut
69146
else:
70147
raise ValueError(f"Unknown command type: {type}")
@@ -79,11 +156,64 @@ async def close(self) -> None:
79156
await self._actor_task
80157
self._active = False
81158

159+
async def _sampling_callback(
160+
self,
161+
context: RequestContext[ClientSession, Any],
162+
params: mcp_types.CreateMessageRequestParams,
163+
) -> mcp_types.CreateMessageResult | mcp_types.ErrorData:
164+
"""Handle sampling requests using the provided model client."""
165+
if self._model_client is None:
166+
# Return an error when no model client is available
167+
return mcp_types.ErrorData(
168+
code=mcp_types.INVALID_REQUEST,
169+
message="No model client available for sampling.",
170+
data=None,
171+
)
172+
173+
llm_messages: list[LLMMessage] = []
174+
175+
try:
176+
if params.systemPrompt:
177+
llm_messages.append(SystemMessage(content=params.systemPrompt))
178+
179+
for mcp_message in params.messages:
180+
llm_messages.append(_parse_sampling_message(mcp_message, model_info=self._model_client.model_info))
181+
182+
except Exception as e:
183+
return mcp_types.ErrorData(
184+
code=mcp_types.INVALID_PARAMS,
185+
message="Error processing sampling messages.",
186+
data=f"{type(e).__name__}: {e}",
187+
)
188+
189+
try:
190+
result = await self._model_client.create(messages=llm_messages)
191+
192+
content = result.content
193+
if not isinstance(content, str):
194+
content = str(content)
195+
196+
return mcp_types.CreateMessageResult(
197+
role="assistant",
198+
content=mcp_types.TextContent(type="text", text=content),
199+
model=self._model_client.model_info["family"],
200+
stopReason=result.finish_reason,
201+
)
202+
except Exception as e:
203+
return mcp_types.ErrorData(
204+
code=mcp_types.INTERNAL_ERROR,
205+
message="Error sampling from model client.",
206+
data=f"{type(e).__name__}: {e}",
207+
)
208+
82209
async def _run_actor(self) -> None:
83210
result: McpResult
84211
try:
85-
async with create_mcp_server_session(self.server_params) as session:
86-
await session.initialize()
212+
async with create_mcp_server_session(
213+
self.server_params, sampling_callback=self._sampling_callback
214+
) as session:
215+
# Save the initialize result
216+
self._initialize_result = await session.initialize()
87217
while True:
88218
cmd = await self._command_queue.get()
89219
if cmd["type"] == "shutdown":
@@ -95,15 +225,47 @@ async def _run_actor(self) -> None:
95225
cmd["future"].set_result(result)
96226
except Exception as e:
97227
cmd["future"].set_exception(e)
228+
elif cmd["type"] == "read_resource":
229+
try:
230+
result = session.read_resource(uri=cmd["uri"])
231+
cmd["future"].set_result(result)
232+
except Exception as e:
233+
cmd["future"].set_exception(e)
234+
elif cmd["type"] == "get_prompt":
235+
try:
236+
result = session.get_prompt(name=cmd["name"], arguments=cmd["args"])
237+
cmd["future"].set_result(result)
238+
except Exception as e:
239+
cmd["future"].set_exception(e)
98240
elif cmd["type"] == "list_tools":
99241
try:
100242
result = session.list_tools()
101243
cmd["future"].set_result(result)
102244
except Exception as e:
103245
cmd["future"].set_exception(e)
246+
elif cmd["type"] == "list_prompts":
247+
try:
248+
result = session.list_prompts()
249+
cmd["future"].set_result(result)
250+
except Exception as e:
251+
cmd["future"].set_exception(e)
252+
elif cmd["type"] == "list_resources":
253+
try:
254+
result = session.list_resources()
255+
cmd["future"].set_result(result)
256+
except Exception as e:
257+
cmd["future"].set_exception(e)
258+
elif cmd["type"] == "list_resource_templates":
259+
try:
260+
result = session.list_resource_templates()
261+
cmd["future"].set_result(result)
262+
except Exception as e:
263+
cmd["future"].set_exception(e)
104264
except Exception as e:
105265
if self._shutdown_future and not self._shutdown_future.done():
106266
self._shutdown_future.set_exception(e)
267+
else:
268+
logger.exception("Exception in MCP actor task")
107269
finally:
108270
self._active = False
109271
self._actor_task = None

python/packages/autogen-ext/src/autogen_ext/tools/mcp/_session.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
from typing import AsyncGenerator
44

55
from mcp import ClientSession
6+
from mcp.client.session import SamplingFnT
67
from mcp.client.sse import sse_client
78
from mcp.client.stdio import stdio_client
89
from mcp.client.streamable_http import streamablehttp_client
@@ -12,7 +13,7 @@
1213

1314
@asynccontextmanager
1415
async def create_mcp_server_session(
15-
server_params: McpServerParams,
16+
server_params: McpServerParams, sampling_callback: SamplingFnT | None = None
1617
) -> AsyncGenerator[ClientSession, None]:
1718
"""Create an MCP client session for the given server parameters."""
1819
if isinstance(server_params, StdioServerParams):
@@ -21,6 +22,7 @@ async def create_mcp_server_session(
2122
read_stream=read,
2223
write_stream=write,
2324
read_timeout_seconds=timedelta(seconds=server_params.read_timeout_seconds),
25+
sampling_callback=sampling_callback,
2426
) as session:
2527
yield session
2628
elif isinstance(server_params, SseServerParams):
@@ -29,6 +31,7 @@ async def create_mcp_server_session(
2931
read_stream=read,
3032
write_stream=write,
3133
read_timeout_seconds=timedelta(seconds=server_params.sse_read_timeout),
34+
sampling_callback=sampling_callback,
3235
) as session:
3336
yield session
3437
elif isinstance(server_params, StreamableHttpServerParams):
@@ -47,5 +50,6 @@ async def create_mcp_server_session(
4750
read_stream=read,
4851
write_stream=write,
4952
read_timeout_seconds=timedelta(seconds=server_params.sse_read_timeout),
53+
sampling_callback=sampling_callback,
5054
) as session:
5155
yield session

0 commit comments

Comments
 (0)