220 lines
10 KiB
Python
220 lines
10 KiB
Python
from pydantic import BaseModel, Field, create_model
|
||
from langchain_core.tools import StructuredTool, ToolException
|
||
|
||
# ----------------- 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 = []
|
||
|
||
# Проверяем, не вернул ли сам MCP-пакет флаг ошибки (isError)
|
||
# Если да, мы тоже должны это воспринимать как падение
|
||
is_mcp_error = getattr(result, 'isError', False)
|
||
|
||
for c in result.content:
|
||
if hasattr(c, 'text'):
|
||
text_outputs.append(c.text)
|
||
else:
|
||
text_outputs.append(str(c))
|
||
|
||
output_text = "\n".join(text_outputs)
|
||
|
||
# Если сервер MCP явно сказал об ошибке, кидаем Exception
|
||
if is_mcp_error:
|
||
raise Exception(output_text)
|
||
|
||
return output_text
|
||
|
||
# Оборачиваем попытку запуска в try..except
|
||
try:
|
||
raw_output = asyncio.run(_run())
|
||
except Exception as e:
|
||
def extract_root_errors(exc):
|
||
# Если ошибка содержит вложенные ошибки (ExceptionGroup)
|
||
if hasattr(exc, 'exceptions'):
|
||
msgs = []
|
||
for child_exc in exc.exceptions:
|
||
msgs.append(extract_root_errors(child_exc))
|
||
# Объединяем сообщения, убирая пустые
|
||
return " | ".join(filter(bool, msgs))
|
||
# Если это базовая ошибка (например, McpError), возвращаем ее текст
|
||
return str(exc)
|
||
|
||
clean_error_msg = extract_root_errors(e)
|
||
|
||
# 1. Поймали обрыв связи (сервер не запущен, отвалился stdio и т.д.)
|
||
error_msg = f"Отсутствует связь с сервером или ошибка выполнения: {clean_error_msg}"
|
||
|
||
# 2. Фиксируем kwargs (запрос), чтобы фронтенд в блоке Response/Request показал JSON!
|
||
if debug_callback:
|
||
debug_callback(kwargs, error_msg)
|
||
|
||
# 3. ВАЖНО: выбрасываем ошибку дальше. Бэкенд LangChain ее перехватит
|
||
# и отправит на фронтенд SSE-событие с `is_error: true` -> появится КРАСНЫЙ КРЕСТ.
|
||
raise ToolException(error_msg)
|
||
|
||
# Выполняется только при успехе (зеленый чекмарк)
|
||
if debug_callback:
|
||
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,
|
||
handle_tool_error=True # Разрешает агенту "выжить" после ошибки тула и сгенерировать ответ
|
||
)
|
||
|
||
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:
|
||
import re
|
||
# Регулярка ищет: КЛЮЧ = "ЗНАЧЕНИЕ" | 'ЗНАЧЕНИЕ' | ЗНАЧЕНИЕ_ДО_ЗАПЯТОЙ
|
||
pattern = r'([^,= \t]+)\s*=\s*(?:"([^"]*)"|\'([^\']*)\'|([^,]*))'
|
||
|
||
for m in re.finditer(pattern, env_str):
|
||
k = m.group(1)
|
||
# Берем то совпадение, которое сработало (2 - двойные кавычки, 3 - одинарные, 4 - без кавычек)
|
||
v = m.group(2) if m.group(2) is not None else \
|
||
m.group(3) if m.group(3) is not None else \
|
||
m.group(4) if m.group(4) is not None else ""
|
||
|
||
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 [] |