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
« 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"""
6from __future__ import annotations
8import asyncio
9import json
10import subprocess
11from abc import ABC, abstractmethod
12from dataclasses import dataclass, field
13from typing import Any
16@dataclass
17class MCPServerConfig:
18 """MCP 服务端配置。"""
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)
28@dataclass
29class MCPToolSchema:
30 """MCP 工具 Schema。"""
32 name: str
33 description: str
34 input_schema: dict
37# ── 传输层 ──────────────────────────────────────
40class MCPTransport(ABC):
41 """MCP 传输协议。"""
43 @abstractmethod
44 async def connect(self, config: MCPServerConfig): ...
46 @abstractmethod
47 async def send(self, method: str, params: dict | None = None) -> dict: ...
49 @abstractmethod
50 async def close(self): ...
53class StdioTransport(MCPTransport):
54 """通过 subprocess 与 MCP Server 通信。"""
56 def __init__(self):
57 self._proc: subprocess.Popen | None = None
58 self._lock = asyncio.Lock()
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 )
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)
77 async def close(self):
78 if self._proc:
79 self._proc.terminate()
82class SSETransport(MCPTransport):
83 """通过 HTTP SSE 与远程 MCP Server 通信。"""
85 async def connect(self, config: MCPServerConfig):
86 import httpx
88 self._client = httpx.AsyncClient(base_url=config.url, timeout=30)
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()
96 async def close(self):
97 if hasattr(self, "_client"):
98 await self._client.aclose()
101# ── MCP 客户端 ──────────────────────────────────
104class MCPClient:
105 """MCP 协议客户端,管理多个 MCP Server 连接。"""
107 TRANSPORTS = {"stdio": StdioTransport, "sse": SSETransport}
109 def __init__(self):
110 self._servers: dict[str, MCPTransport] = {}
111 self._tools: dict[str, MCPToolSchema] = {}
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
128 async def _list_tools(self, server_name: str) -> dict:
129 transport = self._servers[server_name]
130 return await transport.send("tools/list")
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}")
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
161 async def close_all(self):
162 for transport in self._servers.values():
163 await transport.close()