llm-agent-backend/app/mcp_tools.py
2026-06-14 01:33:20 +03:00

173 lines
7.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

from pydantic import BaseModel, Field, create_model
from langchain_core.tools import StructuredTool
# ----------------- MCP MANAGER -----------------
try:
from mcp.client.stdio import stdio_client, StdioServerParameters
from mcp.client.session import ClientSession
MCP_AVAILABLE = True
except ImportError:
MCP_AVAILABLE = False
_MCP_CACHED_TOOLS = {}
def json_schema_to_pydantic(schema: dict, model_name: str) -> type[BaseModel]:
"""Преобразует JSON Schema от MCP сервера в Pydantic модель для LangChain"""
fields = {}
properties = schema.get("properties", {})
required = schema.get("required", [])
for key, val in properties.items():
t = val.get("type", "string")
py_type = str
if t == "integer": py_type = int
elif t == "number": py_type = float
elif t == "boolean": py_type = bool
elif t == "array": py_type = list
elif t == "object": py_type = dict
desc = val.get("description", "")
if key in required:
fields[key] = (py_type, Field(..., description=desc))
else:
fields[key] = (py_type, Field(default=None, description=desc))
if not fields:
fields["kwargs"] = (dict, Field(default_factory=dict, description="Аргументы"))
return create_model(model_name, **fields)
def create_mcp_tool(server_config, tool_name, tool_desc, json_schema, full_env, args, debug_callback=None):
"""Обертка, которая поднимает контейнер/процесс ровно на 1 вызов тула и убивает его"""
args_schema = json_schema_to_pydantic(json_schema, f"MCP_{tool_name.replace('-','_')}_Schema")
command = server_config.get("command")
def mcp_tool_runner(**kwargs):
import asyncio
import json
input_json = json.dumps(kwargs, indent=2, ensure_ascii=False)
async def _run():
server_params = StdioServerParameters(command=command, args=args, env=full_env)
async with stdio_client(server_params) as (read, write):
async with ClientSession(read, write) as session:
await session.initialize()
result = await session.call_tool(tool_name, arguments=kwargs)
text_outputs = []
for c in result.content:
if hasattr(c, 'text'):
text_outputs.append(c.text)
else:
text_outputs.append(str(c))
return "\n".join(text_outputs)
raw_output = asyncio.run(_run())
if debug_callback:
# debug_callback(f"MCP: {tool_name}\n" + input_json, raw_output)
debug_callback(kwargs, raw_output)
return f"Инструмент '{tool_name}' вернул следующий результат:\n<tool_output>\n{raw_output}\n</tool_output>"
return StructuredTool.from_function(
func=mcp_tool_runner,
name=tool_name.replace('-','_'), # Langchain не любит тире в именах
description=tool_desc or f"MCP Tool {tool_name}",
args_schema=args_schema
)
def fetch_mcp_tools(server_config, debug_callback=None):
"""Один раз читает список тулов от сервера и кеширует их схемы"""
if not MCP_AVAILABLE:
print("⚠️ MCP серверы настроены, но библиотека 'mcp' не установлена. Выполните: pip install mcp")
return []
import asyncio
import shlex
import os
config_hash = str(server_config) # primitive hash
if config_hash in _MCP_CACHED_TOOLS:
return _MCP_CACHED_TOOLS[config_hash]
async def _fetch():
command = server_config.get("command", "")
args_str = server_config.get("args", "")
args = shlex.split(args_str) if args_str else []
env_dict = {}
env_str = server_config.get("envString", "")
if env_str:
for pair in env_str.split(','):
if '=' in pair:
k, v = pair.split('=', 1)
env_dict[k.strip()] = v.strip()
full_env = {**os.environ.copy(), **env_dict}
server_params = StdioServerParameters(command=command, args=args, env=full_env)
async with stdio_client(server_params) as (read, write):
async with ClientSession(read, write) as session:
await session.initialize()
tools_resp = await session.list_tools()
return tools_resp.tools, full_env, args
try:
raw_tools, full_env, args = asyncio.run(_fetch())
wrapped_tools = []
# Получаем кастомные описания из конфига
fetched_tools_ui = server_config.get("fetchedTools", [])
custom_descriptions = {t.get("name"): t.get("customDescription") for t in fetched_tools_ui}
for t in raw_tools:
# Если есть кастомное описание и оно не пустое - используем его
c_desc = custom_descriptions.get(t.name)
final_description = c_desc if c_desc and c_desc.strip() else t.description
wrapped_tools.append(create_mcp_tool(server_config, t.name, final_description, t.inputSchema, full_env, args, debug_callback))
_MCP_CACHED_TOOLS[config_hash] = wrapped_tools
print(f"🔌 Успешно загружено {len(wrapped_tools)} инструментов от MCP-сервера '{server_config.get('name')}'")
return wrapped_tools
except Exception as e:
print(f"❌ Ошибка инициализации MCP сервера {server_config.get('name')}: {e}")
return []
def get_raw_mcp_tools_list(server_config):
"""Метод для UI: просто возвращает список доступных инструментов в JSON-friendly формате"""
if not MCP_AVAILABLE:
return []
import asyncio
import shlex
import os
async def _fetch_raw():
command = server_config.get("command", "")
args_str = server_config.get("args", "")
args = shlex.split(args_str) if args_str else []
env_dict = {}
env_str = server_config.get("envString", "")
if env_str:
for pair in env_str.split(','):
if '=' in pair:
k, v = pair.split('=', 1)
env_dict[k.strip()] = v.strip()
full_env = {**os.environ.copy(), **env_dict}
server_params = StdioServerParameters(command=command, args=args, env=full_env)
async with stdio_client(server_params) as (read, write):
async with ClientSession(read, write) as session:
await session.initialize()
tools_resp = await session.list_tools()
return tools_resp.tools
try:
raw_tools = asyncio.run(_fetch_raw())
return [{"name": t.name, "description": t.description} for t in raw_tools]
except Exception as e:
print(f"Ошибка получения списка тулов: {e}")
return []