/
man4j
/
agent-server
Обзор
Документация
Войти
/
man4j
/
agent-server
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
src/agent_server/mcp_client/client.py
207 строк
6 KB
Vladimir
fix mcp client
18 апр 2026, 23:09
18 апр 2026, 23:09
882ad59
Код
Авторство
О чём код?
import logging from typing import Any, Awaitable, Callable, TypeVar from agent_server.chat.types import McpServerConnection, McpState, OpenAIToolSpec from mcp import ClientSession from mcp.client.streamable_http import streamable_http_client from .payloads import extract_prompt_text, normalize_mcp_call_result logger = logging.getLogger(__name__) T = TypeVar("T") def mcp_tools_to_openai(mcp_tools: list[dict[str, Any]]) -> list[OpenAIToolSpec]: return [ { "type": "function", "function": { "name": tool["name"], "description": tool.get("description", ""), "parameters": tool.get("input_schema") or {"type": "object", "properties": {}}, }, } for tool in mcp_tools ] async def _with_initialized_session( url: str, callback: Callable[[ClientSession], Awaitable[T]], ) -> T: # Keep MCP sessions short-lived per operation. This avoids task-affinity # issues inside anyio/mcp and keeps connect/close behavior predictable. async with streamable_http_client(url) as (read, write, _): async with ClientSession(read, write) as session: await session.initialize() return await callback(session) async def load_all_server_prompts(session: ClientSession) -> list[tuple[str, str]]: try: listed = await session.list_prompts() except Exception: logger.exception("Failed to list MCP prompts") return [] prompts: list[tuple[str, str]] = [] for prompt in listed.prompts: name = getattr(prompt, "name", None) if not name: continue try: prompt_result = await session.get_prompt(name, arguments={}) text = extract_prompt_text(prompt_result.messages) if text: prompts.append((name, text)) except Exception: logger.exception("Failed to load MCP prompt %s", name) continue return prompts def build_prompt_from_servers(prompt_to_server: dict[str, list[tuple[str, str]]]) -> str: parts: list[str] = [] for server_name in sorted(prompt_to_server): for prompt_name, text in prompt_to_server[server_name]: if text.strip(): parts.append(f"\n\n### {server_name} MCP prompt: \n\n{text}") return "\n\n".join(parts) def build_mcp_state( *, servers: dict[str, McpServerConnection], tool_to_server: dict[str, str], prompt_to_server: dict[str, list[tuple[str, str]]], errors: dict[str, str], ) -> McpState: all_tools = [tool for conn in servers.values() for tool in conn["tools"]] prompt = build_prompt_from_servers(prompt_to_server) return { "servers": servers, "tool_to_server": tool_to_server, "tools": all_tools, "prompt_to_server": prompt_to_server, "prompt": prompt, "errors": errors, } async def bootstrap_one_mcp(name: str, url: str) -> McpServerConnection: async def _bootstrap(session: ClientSession) -> McpServerConnection: # Bootstrap is the only place where we fetch tool/prompt metadata and # cache it into McpState for later routing and system-prompt building. listed = await session.list_tools() tools = [ { "name": t.name, "description": t.description, "input_schema": t.inputSchema, } for t in listed.tools ] prompts = await load_all_server_prompts(session) return { "name": name, "url": url, "tools": tools, "prompts": prompts, } try: return await _with_initialized_session(url, _bootstrap) except Exception: logger.exception("Failed to bootstrap MCP server %s (%s)", name, url) raise async def connect_mcp_servers(targets: list[tuple[str, str]]) -> McpState: servers: dict[str, McpServerConnection] = {} tool_to_server: dict[str, str] = {} prompt_to_server: dict[str, list[tuple[str, str]]] = {} errors: dict[str, str] = {} for name, url in targets: try: conn = await bootstrap_one_mcp(name, url) servers[name] = conn prompt_to_server[name] = conn["prompts"] for tool in conn["tools"]: tool_to_server[tool["name"]] = name except Exception as e: errors[name] = f"{type(e).__name__}: {e}" return build_mcp_state( servers=servers, tool_to_server=tool_to_server, prompt_to_server=prompt_to_server, errors=errors, ) async def disconnect_mcp_servers(mcp_state: McpState | None) -> None: return def all_openai_tools_from_state(mcp_state: McpState | None) -> list[OpenAIToolSpec]: if not mcp_state: return [] return mcp_tools_to_openai(mcp_state.get("tools", [])) def all_prompts_from_state(mcp_state: McpState | None) -> str: if not mcp_state: return "" return mcp_state.get("prompt", "") or "" def find_mcp_for_tool( mcp_state: McpState | None, tool_name: str, ) -> McpServerConnection | None: if not mcp_state: return None server_name = (mcp_state.get("tool_to_server") or {}).get(tool_name) if not server_name: return None return (mcp_state.get("servers") or {}).get(server_name) async def call_mcp_tool( mcp_state: McpState | None, tool_name: str, tool_input: dict, ) -> dict[str, Any]: conn = find_mcp_for_tool(mcp_state, tool_name) if conn is None: return {"error": f"Tool not found in MCP connections: {tool_name}"} server_name = conn.get("name") or "unknown" url = conn.get("url") if not url: return {"error": f"MCP server URL is missing for tool: {tool_name}"} async def _call(session: ClientSession) -> dict[str, Any]: # Tool calls should execute only the requested tool. Do not re-fetch # tool lists or prompts here, or each invocation becomes unnecessarily heavy. result = await session.call_tool(tool_name, tool_input) return normalize_mcp_call_result(result) try: return await _with_initialized_session(url, _call) except Exception as e: logger.exception("Failed to call MCP tool %s on server %s", tool_name, server_name) return {"error": f"{type(e).__name__}: {e}"}