11import asyncio
22import atexit
3+ import base64
4+ import io
5+ import logging
36from 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
618from pydantic import BaseModel
719from 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
1125from ._config import McpServerParams
1226from ._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+ )
1539McpFuture = 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+
1877class McpActorArgs (TypedDict ):
1978 name : str | None
2079 kargs : Mapping [str , Any ]
2180
2281
2382class McpSessionActorConfig (BaseModel ):
2483 server_params : McpServerParams
84+ model_client : ComponentModel | Dict [str , Any ] | None = None
2585
2686
2787class 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
0 commit comments