173 lines
7.2 KiB
Python
173 lines
7.2 KiB
Python
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 [] |