Coverage for agentos/mcp/__init__.py: 26%

314 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-08 10:59 +0800

1"""MCP (Model Context Protocol) client implementation for AgentOS. 

2 

3Full MCP client with JSON-RPC 2.0, initialize handshake, tool/resource/prompt 

4discovery, dual transport (stdio + SSE), Sampling, Logging, Roots. 

5Designed to be used as async context manager. 

6 

7v1.14.0: Added Sampling, Resource Templates, Logging, Roots support. 

8""" 

9 

10from __future__ import annotations 

11 

12import asyncio 

13import json 

14import logging 

15import subprocess 

16from abc import ABC, abstractmethod 

17from dataclasses import dataclass, field 

18from typing import Any 

19 

20import httpx 

21 

22logger = logging.getLogger(__name__) 

23 

24# ── Data Models ──────────────────────────── 

25 

26 

27@dataclass 

28class MCPServerConfig: 

29 """Configuration for connecting to an MCP server.""" 

30 

31 name: str 

32 transport: str = "stdio" # stdio | sse 

33 command: str | None = None 

34 args: list[str] = field(default_factory=list) 

35 url: str | None = None 

36 env: dict[str, str] = field(default_factory=dict) 

37 timeout: int = 30 

38 capabilities: dict[str, Any] = field(default_factory=dict) 

39 

40 

41@dataclass 

42class MCPToolInfo: 

43 """Metadata for a discovered MCP tool.""" 

44 

45 name: str 

46 description: str = "" 

47 input_schema: dict[str, Any] = field(default_factory=dict) 

48 server_name: str = "" 

49 

50 

51@dataclass 

52class MCPResourceInfo: 

53 """Metadata for a discovered MCP resource.""" 

54 

55 uri: str 

56 name: str = "" 

57 description: str = "" 

58 mime_type: str = "" 

59 server_name: str = "" 

60 

61 

62@dataclass 

63class MCPPromptInfo: 

64 """Metadata for a discovered MCP prompt.""" 

65 

66 name: str 

67 description: str = "" 

68 arguments: list[dict[str, Any]] = field(default_factory=list) 

69 server_name: str = "" 

70 

71 

72# ── JSON-RPC 2.0 Transport ────────────────── 

73 

74 

75class MCPTransport(ABC): 

76 """Abstract transport layer for MCP JSON-RPC 2.0 communication.""" 

77 

78 @abstractmethod 

79 async def connect(self, config: MCPServerConfig) -> None: ... 

80 

81 @abstractmethod 

82 async def send_request(self, method: str, params: dict | None = None) -> dict[str, Any]: ... 

83 

84 @abstractmethod 

85 async def send_notification(self, method: str, params: dict | None = None) -> None: ... 

86 

87 @abstractmethod 

88 async def close(self) -> None: ... 

89 

90 

91class StdioTransport(MCPTransport): 

92 """MCP transport over subprocess stdio. 

93 

94 Communicates with an MCP server launched as a child process 

95 using newline-delimited JSON-RPC 2.0 messages. 

96 """ 

97 

98 def __init__(self): 

99 self._proc: subprocess.Popen | None = None 

100 self._lock = asyncio.Lock() 

101 self._request_id = 0 

102 self._pending: dict[int, asyncio.Future] = {} 

103 self._reader_task: asyncio.Task | None = None 

104 

105 async def connect(self, config: MCPServerConfig) -> None: 

106 cmd = config.command or "npx" 

107 full_args = [cmd] + list(config.args) 

108 env = {**__import__("os").environ, **config.env} 

109 

110 self._proc = subprocess.Popen( 

111 full_args, 

112 stdin=subprocess.PIPE, 

113 stdout=subprocess.PIPE, 

114 stderr=subprocess.PIPE, 

115 env=env, 

116 text=False, 

117 ) 

118 self._reader_task = asyncio.create_task(self._read_loop()) 

119 

120 async def _read_loop(self) -> None: 

121 """Continuously read JSON-RPC responses from stdout.""" 

122 loop = asyncio.get_event_loop() 

123 while self._proc and self._proc.poll() is None: 

124 try: 

125 line = await loop.run_in_executor(None, self._proc.stdout.readline) 

126 if not line: 

127 break 

128 data = json.loads(line.decode("utf-8")) 

129 req_id = data.get("id") 

130 if req_id is not None and req_id in self._pending: 

131 future = self._pending.pop(req_id) 

132 if "error" in data: 

133 future.set_exception( 

134 MCPError( 

135 data["error"].get("code", -1), 

136 data["error"].get("message", "Unknown error"), 

137 ) 

138 ) 

139 else: 

140 future.set_result(data.get("result", {})) 

141 except Exception: 

142 continue 

143 

144 async def send_request(self, method: str, params: dict | None = None) -> dict[str, Any]: 

145 """Send a JSON-RPC 2.0 request and await the response.""" 

146 if not self._proc or self._proc.poll() is not None: 

147 raise MCPError(-32000, "MCP server process is not running") 

148 

149 async with self._lock: 

150 self._request_id += 1 

151 req_id = self._request_id 

152 request = { 

153 "jsonrpc": "2.0", 

154 "id": req_id, 

155 "method": method, 

156 "params": params or {}, 

157 } 

158 

159 future: asyncio.Future = asyncio.get_event_loop().create_future() 

160 self._pending[req_id] = future 

161 

162 payload = json.dumps(request).encode("utf-8") + b"\n" 

163 loop = asyncio.get_event_loop() 

164 await loop.run_in_executor(None, self._proc.stdin.write, payload) 

165 await loop.run_in_executor(None, self._proc.stdin.flush) 

166 

167 try: 

168 return await asyncio.wait_for(future, timeout=30) 

169 except TimeoutError: 

170 self._pending.pop(req_id, None) 

171 raise MCPError(-32001, "Request timed out") 

172 

173 async def send_notification(self, method: str, params: dict | None = None) -> None: 

174 """Send a JSON-RPC 2.0 notification (no response expected).""" 

175 if not self._proc or self._proc.poll() is not None: 

176 return 

177 

178 async with self._lock: 

179 notification = { 

180 "jsonrpc": "2.0", 

181 "method": method, 

182 "params": params or {}, 

183 } 

184 payload = json.dumps(notification).encode("utf-8") + b"\n" 

185 loop = asyncio.get_event_loop() 

186 await loop.run_in_executor(None, self._proc.stdin.write, payload) 

187 await loop.run_in_executor(None, self._proc.stdin.flush) 

188 

189 async def close(self) -> None: 

190 if self._reader_task: 

191 self._reader_task.cancel() 

192 try: 

193 await self._reader_task 

194 except asyncio.CancelledError: 

195 pass 

196 if self._proc: 

197 self._proc.terminate() 

198 try: 

199 self._proc.wait(timeout=5) 

200 except subprocess.TimeoutExpired: 

201 self._proc.kill() 

202 self._proc = None 

203 

204 

205class SSETransport(MCPTransport): 

206 """MCP transport over HTTP SSE (Server-Sent Events). 

207 

208 Connects to a remote MCP server via HTTP POST for requests 

209 and SSE stream for responses. 

210 """ 

211 

212 def __init__(self): 

213 self._client: httpx.AsyncClient | None = None 

214 self._request_id = 0 

215 self._pending: dict[int, asyncio.Future] = {} 

216 self._sse_task: asyncio.Task | None = None 

217 self._response_queue: asyncio.Queue = asyncio.Queue() 

218 self._message_endpoint: str = "" 

219 self._sse_endpoint: str = "" 

220 

221 async def connect(self, config: MCPServerConfig) -> None: 

222 if not config.url: 

223 raise MCPError(-32602, "URL required for SSE transport") 

224 

225 self._message_endpoint = config.url.rstrip("/") + "/message" 

226 self._sse_endpoint = config.url.rstrip("/") + "/sse" 

227 self._client = httpx.AsyncClient(timeout=config.timeout) 

228 self._sse_task = asyncio.create_task(self._sse_loop()) 

229 

230 async def _sse_loop(self) -> None: 

231 """Read SSE events and route to pending futures.""" 

232 while self._client: 

233 try: 

234 async with self._client.stream("GET", self._sse_endpoint) as response: 

235 async for line in response.aiter_lines(): 

236 if line.startswith("data: "): 

237 data_str = line[6:].strip() 

238 try: 

239 data = json.loads(data_str) 

240 req_id = data.get("id") 

241 if req_id is not None and req_id in self._pending: 

242 future = self._pending.pop(req_id) 

243 if "error" in data: 

244 future.set_exception( 

245 MCPError( 

246 data["error"].get("code", -1), 

247 data["error"].get("message", ""), 

248 ) 

249 ) 

250 else: 

251 future.set_result(data.get("result", {})) 

252 except json.JSONDecodeError: 

253 continue 

254 except Exception: 

255 await asyncio.sleep(1) 

256 

257 async def send_request(self, method: str, params: dict | None = None) -> dict[str, Any]: 

258 if not self._client: 

259 raise MCPError(-32000, "SSE transport not connected") 

260 

261 self._request_id += 1 

262 req_id = self._request_id 

263 request = { 

264 "jsonrpc": "2.0", 

265 "id": req_id, 

266 "method": method, 

267 "params": params or {}, 

268 } 

269 

270 future: asyncio.Future = asyncio.get_event_loop().create_future() 

271 self._pending[req_id] = future 

272 

273 resp = await self._client.post(self._message_endpoint, json=request) 

274 resp.raise_for_status() 

275 

276 try: 

277 return await asyncio.wait_for(future, timeout=30) 

278 except TimeoutError: 

279 self._pending.pop(req_id, None) 

280 raise MCPError(-32001, "SSE request timed out") 

281 

282 async def send_notification(self, method: str, params: dict | None = None) -> None: 

283 if not self._client: 

284 return 

285 notification = { 

286 "jsonrpc": "2.0", 

287 "method": method, 

288 "params": params or {}, 

289 } 

290 await self._client.post(self._message_endpoint, json=notification) 

291 

292 async def close(self) -> None: 

293 if self._sse_task: 

294 self._sse_task.cancel() 

295 try: 

296 await self._sse_task 

297 except asyncio.CancelledError: 

298 pass 

299 if self._client: 

300 await self._client.aclose() 

301 self._client = None 

302 

303 

304# ── Error ──────────────────────────────────── 

305 

306 

307class MCPError(Exception): 

308 """MCP protocol error.""" 

309 

310 def __init__(self, code: int, message: str, data: Any = None): 

311 self.code = code 

312 self.message = message 

313 self.data = data 

314 super().__init__(f"MCP Error [{code}]: {message}") 

315 

316 

317# ── Full Client ───────────────────────────── 

318 

319 

320class MCPClient: 

321 """Full MCP client for connecting to and using MCP servers. 

322 

323 Supports stdio (local process) and SSE (remote HTTP) transports. 

324 

325 Usage: 

326 async with MCPClient() as client: 

327 await client.connect_server(MCPServerConfig( 

328 name="filesystem", 

329 command="npx", 

330 args=["-y", "@modelcontextprotocol/server-filesystem", "/tmp"], 

331 )) 

332 tools = client.list_tools() 

333 result = await client.call_tool("filesystem", "read_file", {"path": "/tmp/test.txt"}) 

334 """ 

335 

336 TRANSPORTS = { 

337 "stdio": StdioTransport, 

338 "sse": SSETransport, 

339 } 

340 

341 def __init__(self): 

342 self._servers: dict[str, MCPTransport] = {} 

343 self._server_configs: dict[str, MCPServerConfig] = {} 

344 self._tools: dict[str, MCPToolInfo] = {} 

345 self._resources: dict[str, MCPResourceInfo] = {} 

346 self._prompts: dict[str, MCPPromptInfo] = {} 

347 self._server_capabilities: dict[str, dict[str, Any]] = {} 

348 

349 async def __aenter__(self) -> MCPClient: 

350 return self 

351 

352 async def __aexit__(self, *args) -> None: 

353 await self.close_all() 

354 

355 async def connect_server(self, config: MCPServerConfig) -> dict[str, Any]: 

356 """Connect to an MCP server and perform initialization handshake. 

357 

358 Returns the server's capabilities dict. 

359 """ 

360 transport_cls = self.TRANSPORTS.get(config.transport) 

361 if not transport_cls: 

362 raise MCPError(-32601, f"Unknown transport: {config.transport}") 

363 

364 transport = transport_cls() 

365 await transport.connect(config) 

366 

367 # MCP Initialize handshake 

368 init_result = await transport.send_request( 

369 "initialize", 

370 { 

371 "protocolVersion": "2024-11-05", 

372 "capabilities": config.capabilities or {}, 

373 "clientInfo": { 

374 "name": "agentos-mcp-client", 

375 "version": "1.0.0", 

376 }, 

377 }, 

378 ) 

379 

380 # Send initialized notification 

381 await transport.send_notification("notifications/initialized") 

382 

383 self._servers[config.name] = transport 

384 self._server_configs[config.name] = config 

385 self._server_capabilities[config.name] = init_result.get("capabilities", {}) 

386 

387 # Discover tools, resources, prompts 

388 await self._discover_server(config.name) 

389 

390 return init_result 

391 

392 async def _discover_server(self, server_name: str) -> None: 

393 """Discover all capabilities of a connected server.""" 

394 transport = self._servers[server_name] 

395 caps = self._server_capabilities.get(server_name, {}) 

396 

397 # Discover tools 

398 if caps.get("tools"): 

399 try: 

400 result = await transport.send_request("tools/list") 

401 for tool in result.get("tools", []): 

402 full_name = f"mcp__{server_name}__{tool['name']}" 

403 self._tools[full_name] = MCPToolInfo( 

404 name=tool["name"], 

405 description=tool.get("description", ""), 

406 input_schema=tool.get("inputSchema", {}), 

407 server_name=server_name, 

408 ) 

409 except MCPError: 

410 logger.debug(f"Server '{server_name}' tools/list not supported") 

411 

412 # Discover resources 

413 if caps.get("resources"): 

414 try: 

415 result = await transport.send_request("resources/list") 

416 for res in result.get("resources", []): 

417 self._resources[res["uri"]] = MCPResourceInfo( 

418 uri=res["uri"], 

419 name=res.get("name", ""), 

420 description=res.get("description", ""), 

421 mime_type=res.get("mimeType", ""), 

422 server_name=server_name, 

423 ) 

424 except MCPError: 

425 logger.debug(f"Server '{server_name}' resources/list not supported") 

426 

427 # Discover prompts 

428 if caps.get("prompts"): 

429 try: 

430 result = await transport.send_request("prompts/list") 

431 for prompt in result.get("prompts", []): 

432 key = f"{server_name}__{prompt['name']}" 

433 self._prompts[key] = MCPPromptInfo( 

434 name=prompt["name"], 

435 description=prompt.get("description", ""), 

436 arguments=prompt.get("arguments", []), 

437 server_name=server_name, 

438 ) 

439 except MCPError: 

440 logger.debug(f"Server '{server_name}' prompts/list not supported") 

441 

442 async def call_tool( 

443 self, 

444 server_name: str, 

445 tool_name: str, 

446 arguments: dict[str, Any] | None = None, 

447 ) -> Any: 

448 """Call a tool on a connected MCP server. 

449 

450 Args: 

451 server_name: Name of the MCP server. 

452 tool_name: Name of the tool to call. 

453 arguments: Tool arguments dict. 

454 

455 Returns: 

456 Tool result content (text or structured data). 

457 """ 

458 if server_name not in self._servers: 

459 raise MCPError(-32602, f"Server '{server_name}' not connected") 

460 

461 transport = self._servers[server_name] 

462 result = await transport.send_request( 

463 "tools/call", 

464 { 

465 "name": tool_name, 

466 "arguments": arguments or {}, 

467 }, 

468 ) 

469 

470 content = result.get("content", []) 

471 if not content: 

472 return "" 

473 

474 # Extract text from content blocks 

475 texts = [] 

476 for block in content: 

477 if block.get("type") == "text": 

478 texts.append(block.get("text", "")) 

479 elif block.get("type") == "resource": 

480 texts.append(f"[Resource: {block.get('resource', {}).get('uri', '')}]") 

481 elif block.get("type") == "image": 

482 texts.append(f"[Image: {block.get('data', '')[:50]}...]") 

483 

484 return "\n".join(texts) if texts else content 

485 

486 async def read_resource(self, server_name: str, uri: str) -> dict[str, Any]: 

487 """Read a resource from a connected MCP server.""" 

488 if server_name not in self._servers: 

489 raise MCPError(-32602, f"Server '{server_name}' not connected") 

490 

491 transport = self._servers[server_name] 

492 result = await transport.send_request("resources/read", {"uri": uri}) 

493 return result.get("contents", [{}])[0] if result.get("contents") else {} 

494 

495 async def get_prompt( 

496 self, 

497 server_name: str, 

498 prompt_name: str, 

499 arguments: dict[str, str] | None = None, 

500 ) -> dict[str, Any]: 

501 """Get a prompt template from a connected MCP server.""" 

502 if server_name not in self._servers: 

503 raise MCPError(-32602, f"Server '{server_name}' not connected") 

504 

505 transport = self._servers[server_name] 

506 result = await transport.send_request( 

507 "prompts/get", 

508 { 

509 "name": prompt_name, 

510 "arguments": arguments or {}, 

511 }, 

512 ) 

513 return result 

514 

515 def list_tools(self, server_name: str | None = None) -> list[MCPToolInfo]: 

516 """List discovered tools, optionally filtered by server.""" 

517 tools = list(self._tools.values()) 

518 if server_name: 

519 tools = [t for t in tools if t.server_name == server_name] 

520 return tools 

521 

522 def list_resources(self, server_name: str | None = None) -> list[MCPResourceInfo]: 

523 """List discovered resources, optionally filtered by server.""" 

524 resources = list(self._resources.values()) 

525 if server_name: 

526 resources = [r for r in resources if r.server_name == server_name] 

527 return resources 

528 

529 def list_prompts(self, server_name: str | None = None) -> list[MCPPromptInfo]: 

530 """List discovered prompts, optionally filtered by server.""" 

531 prompts = list(self._prompts.values()) 

532 if server_name: 

533 prompts = [p for p in prompts if p.server_name == server_name] 

534 return prompts 

535 

536 def get_server_capabilities(self, server_name: str) -> dict[str, Any]: 

537 """Get the capabilities reported by a server.""" 

538 return self._server_capabilities.get(server_name, {}) 

539 

540 def get_tool_schemas( 

541 self, 

542 server_name: str | None = None, 

543 format: str = "openai", 

544 ) -> list[dict[str, Any]]: 

545 """Export tool schemas in OpenAI or Anthropic function format. 

546 

547 Args: 

548 server_name: Optional filter by server. 

549 format: 'openai' or 'anthropic'. 

550 

551 Returns: 

552 List of function/tool schema dicts. 

553 """ 

554 tools = self.list_tools(server_name) 

555 schemas = [] 

556 

557 for tool in tools: 

558 params = tool.input_schema 

559 if format == "openai": 

560 schemas.append( 

561 { 

562 "type": "function", 

563 "function": { 

564 "name": f"mcp__{tool.server_name}__{tool.name}", 

565 "description": tool.description, 

566 "parameters": params, 

567 }, 

568 } 

569 ) 

570 elif format == "anthropic": 

571 schemas.append( 

572 { 

573 "name": f"mcp__{tool.server_name}__{tool.name}", 

574 "description": tool.description, 

575 "input_schema": params, 

576 } 

577 ) 

578 

579 return schemas 

580 

581 @property 

582 def connected_servers(self) -> list[str]: 

583 """List names of connected servers.""" 

584 return list(self._servers.keys()) 

585 

586 async def disconnect_server(self, server_name: str) -> None: 

587 """Disconnect from a specific MCP server.""" 

588 if server_name in self._servers: 

589 await self._servers[server_name].close() 

590 del self._servers[server_name] 

591 self._server_configs.pop(server_name, None) 

592 self._server_capabilities.pop(server_name, None) 

593 # Remove associated tools/resources/prompts 

594 self._tools = {k: v for k, v in self._tools.items() if v.server_name != server_name} 

595 self._resources = { 

596 k: v for k, v in self._resources.items() if v.server_name != server_name 

597 } 

598 self._prompts = {k: v for k, v in self._prompts.items() if v.server_name != server_name} 

599 

600 async def close_all(self) -> None: 

601 """Disconnect from all MCP servers.""" 

602 for name in list(self._servers.keys()): 

603 await self.disconnect_server(name) 

604 

605 

606# ── Convenience Functions ──────────────────── 

607 

608 

609async def connect_mcp_servers( 

610 configs: list[MCPServerConfig], 

611) -> MCPClient: 

612 """Connect to multiple MCP servers at once. 

613 

614 Usage: 

615 client = await connect_mcp_servers([ 

616 MCPServerConfig(name="filesystem", command="npx", 

617 args=["-y", "@modelcontextprotocol/server-filesystem", "/tmp"]), 

618 MCPServerConfig(name="github", command="npx", 

619 args=["-y", "@modelcontextprotocol/server-github"], 

620 env={"GITHUB_PERSONAL_ACCESS_TOKEN": os.environ["GITHUB_TOKEN"]}), 

621 ]) 

622 """ 

623 client = MCPClient() 

624 for config in configs: 

625 await client.connect_server(config) 

626 return client 

627 

628 

629# ── MCP Server (v1.5.2) ───────────────────── 

630 

631from agentos.mcp.server import ( # noqa: E402 

632 MCPPromptDef, 

633 MCPResource, 

634 MCPServer, 

635 MCPToolDef, 

636 create_default_server, 

637 start_mcp_server, 

638) 

639 

640__all__ = [ 

641 "MCPServerConfig", 

642 "MCPToolInfo", 

643 "MCPResourceInfo", 

644 "MCPPromptInfo", 

645 "MCPError", 

646 "MCPTransport", 

647 "StdioTransport", 

648 "SSETransport", 

649 "MCPClient", 

650 "connect_mcp_servers", 

651 # MCP Server (v1.5.2) 

652 "MCPServer", 

653 "MCPToolDef", 

654 "MCPResource", 

655 "MCPPromptDef", 

656 "create_default_server", 

657 "start_mcp_server", 

658 # MCP Sampling, Resource Templates, Logging, Roots (v1.14.0) 

659 "MCPClientSampling", 

660 "SamplingRequest", 

661 "SamplingResponse", 

662 "SamplingMessage", 

663 "SamplingContentBlock", 

664 "SamplingRole", 

665 "SamplingError", 

666 "mock_llm_call", 

667 "MCPResourceTemplate", 

668 "MCPLogLevel", 

669 "MCPLoggingHandler", 

670 "MCPRoot", 

671 # MCP Tool Adapter (v1.16.10) 

672 "MCPToolAdapter", 

673 "MCPAdapter", 

674 # Built-in MCP Servers (v1.16.11) 

675 "FilesystemServer", 

676 "WebFetchServer", 

677 "MemoryServer", 

678 "SearchServer", 

679 "GitServer", 

680 "ShellServer", 

681 "CodeServer", 

682 "TextServer", 

683 "BuiltinMCPRegistry", 

684] 

685 

686# ── Convenience alias ── 

687from agentos.mcp.adapter import MCPToolAdapter # noqa: E402, F811 

688 

689MCPAdapter = MCPToolAdapter 

690 

691# ── Built-in MCP Servers ── 

692from agentos.mcp.builtin_servers import ( # noqa: E402 

693 BuiltinMCPRegistry, 

694 CodeServer, 

695 FilesystemServer, 

696 GitServer, 

697 MemoryServer, 

698 SearchServer, 

699 ShellServer, 

700 TextServer, 

701 WebFetchServer, 

702)