Coverage for agentos/protocols/mcp.py: 47%

81 statements  

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

1""" 

2AgentOS v0.20 MCP (Model Context Protocol) 客户端。 

3支持 stdio / SSE / WebSocket 三种传输方式。 

4""" 

5 

6from __future__ import annotations 

7 

8import asyncio 

9import json 

10import subprocess 

11from abc import ABC, abstractmethod 

12from dataclasses import dataclass, field 

13from typing import Any 

14 

15 

16@dataclass 

17class MCPServerConfig: 

18 """MCP 服务端配置。""" 

19 

20 name: str 

21 transport: str = "stdio" # stdio | sse | ws 

22 command: str | None = None 

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

24 url: str | None = None 

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

26 

27 

28@dataclass 

29class MCPToolSchema: 

30 """MCP 工具 Schema。""" 

31 

32 name: str 

33 description: str 

34 input_schema: dict 

35 

36 

37# ── 传输层 ────────────────────────────────────── 

38 

39 

40class MCPTransport(ABC): 

41 """MCP 传输协议。""" 

42 

43 @abstractmethod 

44 async def connect(self, config: MCPServerConfig): ... 

45 

46 @abstractmethod 

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

48 

49 @abstractmethod 

50 async def close(self): ... 

51 

52 

53class StdioTransport(MCPTransport): 

54 """通过 subprocess 与 MCP Server 通信。""" 

55 

56 def __init__(self): 

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

58 self._lock = asyncio.Lock() 

59 

60 async def connect(self, config: MCPServerConfig): 

61 self._proc = subprocess.Popen( 

62 [config.command or "npx"] + config.args, 

63 stdin=subprocess.PIPE, 

64 stdout=subprocess.PIPE, 

65 stderr=subprocess.PIPE, 

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

67 ) 

68 

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

70 async with self._lock: 

71 msg = json.dumps({"jsonrpc": "2.0", "method": method, "params": params or {}, "id": 1}) 

72 self._proc.stdin.write((msg + "\n").encode()) 

73 self._proc.stdin.flush() 

74 line = self._proc.stdout.readline() 

75 return json.loads(line) 

76 

77 async def close(self): 

78 if self._proc: 

79 self._proc.terminate() 

80 

81 

82class SSETransport(MCPTransport): 

83 """通过 HTTP SSE 与远程 MCP Server 通信。""" 

84 

85 async def connect(self, config: MCPServerConfig): 

86 import httpx 

87 

88 self._client = httpx.AsyncClient(base_url=config.url, timeout=30) 

89 

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

91 resp = await self._client.post( 

92 "/message", json={"jsonrpc": "2.0", "method": method, "params": params or {}, "id": 1} 

93 ) 

94 return resp.json() 

95 

96 async def close(self): 

97 if hasattr(self, "_client"): 

98 await self._client.aclose() 

99 

100 

101# ── MCP 客户端 ────────────────────────────────── 

102 

103 

104class MCPClient: 

105 """MCP 协议客户端,管理多个 MCP Server 连接。""" 

106 

107 TRANSPORTS = {"stdio": StdioTransport, "sse": SSETransport} 

108 

109 def __init__(self): 

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

111 self._tools: dict[str, MCPToolSchema] = {} 

112 

113 async def connect_server(self, config: MCPServerConfig): 

114 transport_cls = self.TRANSPORTS.get(config.transport, StdioTransport) 

115 transport = transport_cls() 

116 await transport.connect(config) 

117 self._servers[config.name] = transport 

118 # 拉取工具列表 

119 result = await self._list_tools(config.name) 

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

121 schema = MCPToolSchema( 

122 name=f"mcp_{config.name}_{tool['name']}", 

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

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

125 ) 

126 self._tools[schema.name] = schema 

127 

128 async def _list_tools(self, server_name: str) -> dict: 

129 transport = self._servers[server_name] 

130 return await transport.send("tools/list") 

131 

132 async def call_tool(self, full_name: str, arguments: dict) -> Any: 

133 full_name.replace("mcp_", "", 1) 

134 # 找到所属server 

135 for name in self._servers: 

136 if full_name.startswith(f"mcp_{name}_"): 

137 tool_name = full_name[len(f"mcp_{name}_") + 1 :] 

138 transport = self._servers[name] 

139 result = await transport.send( 

140 "tools/call", {"name": tool_name, "arguments": arguments} 

141 ) 

142 return result.get("content", [{}])[0].get("text", "") 

143 raise ValueError(f"Unknown MCP tool: {full_name}") 

144 

145 def get_mcp_tool_schemas(self) -> list[dict]: 

146 """转为 OpenAI function 格式。""" 

147 schemas = [] 

148 for tool in self._tools.values(): 

149 schemas.append( 

150 { 

151 "type": "function", 

152 "function": { 

153 "name": tool.name, 

154 "description": tool.description, 

155 "parameters": tool.input_schema, 

156 }, 

157 } 

158 ) 

159 return schemas 

160 

161 async def close_all(self): 

162 for transport in self._servers.values(): 

163 await transport.close()