mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-07-29 09:26:00 +00:00
* fix(mcp): migrate local MCP-produced files into sandbox outputs (#3597) Stdio MCP servers (e.g. Playwright) write files to host paths that the sandbox/artifact API cannot resolve, since it only serves paths under /mnt/user-data. Copy local files referenced by ResourceLink results into the thread's sandbox outputs dir and rewrite their URIs to /mnt/user-data/outputs/... so they become readable. Also scope pooled MCP sessions by user_id:thread_id instead of thread_id alone, matching the per-(user_id, thread_id) filesystem isolation. * fix(mcp): restrict file migration to trusted source roots (#3597) Add a source-root allowlist to the MCP file-migration path so a malicious or buggy MCP server cannot have us copy arbitrary host files (e.g. /etc/passwd) into a thread's outputs directory, from where the artifact API would serve them. Files are migrated only when located under the OS temp dir (Playwright's default), the thread's own user-data tree, or an operator-configured root via DEERFLOW_MCP_MIGRATION_SOURCE_ROOTS. Expand test coverage with allowlist/security cases (path escape refusal, trusted-root acceptance), URL-encoded file:// paths, converter content branches (image/embedded/error/structured), and copy/resolve failure fallbacks. * fix(mcp): harden local file migration into sandbox outputs Address robustness and security gaps in the MCP ResourceLink file migration: - Set migrated files to 0o644 so a differently-UID sandbox container can read them, instead of inheriting the source's (possibly 0o600) mode. - Enforce the 100MB size cap during the copy (chunked, byte-counted) rather than from a prior stat(), closing the grow-after-stat TOCTOU. - Create the destination atomically with O_CREAT|O_EXCL to remove the check-then-create name-collision race. - Document the shared-$TMPDIR multi-tenant read surface and mitigation. Add regression tests: symlink escape refusal, explicit $TMPDIR source migration, 0o644 mode, and the outputs/user-data resolve() OSError fallback branches. * fix(mcp): migrate playwright text file outputs * fix(mcp): translate MCP file outputs to virtual paths instead of copying (#3597) Pin stdio MCP subprocess cwd and TMPDIR/TMP/TEMP under the thread workspace so produced files always land in the mounted user-data tree, then rewrite returned references via deterministic host->virtual path translation. Free text is best-effort only: a reference is rewritten only when it resolves to an existing file inside the thread's tree, and bare filenames are matched against files created/modified by the same tool call. Replaces the previous copy-into-outputs + regex approach (which missed cases like temp/page-*.yml). * style(mcp): apply ruff format to mcp path translation tests * perf(mcp): offload stdio FS work off event loop and gate on transport Address review on #3600: - Wrap the workspace dir prep, snapshot diff, and per-token path resolution in asyncio.to_thread so they no longer block the event loop (matches the repo's blocking-IO gate convention). - Gate the cwd/temp pinning and snapshots on stdio transport only; SSE/HTTP servers skip the filesystem work entirely. - Skip the post-call snapshot diff when the result has no text content. * test(mcp): cover stdio transport gating and text-content after-walk skip Add unit/integration coverage for the new review-driven behavior: - _prepare_stdio_workspace dir/temp/snapshot bundle - _result_has_text_content detection (text, embedded text, image, empty) - non-stdio transport skips cwd/temp pinning and touches no workspace dirs - post-call snapshot diff is skipped without text content and runs with it * fix(mcp): address stdio path rewrite review feedback - Restrict the stdio MCP temp directory to 0700 instead of 0777. - Preserve operator-provided stdio cwd values while keeping injected cwd values as strings. - Add debug logging for deterministic path rewrites and bare-filename rewrite decisions. - Document the stdio cwd/temp pinning, virtual-path translation, and user/thread session scope. - Cover explicit cwd preservation and temp-dir permissions in session-pool tests. --------- Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
654 lines
28 KiB
Python
654 lines
28 KiB
Python
"""Load MCP tools using langchain-mcp-adapters with stdio session pooling."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import asyncio
|
||
import logging
|
||
import re
|
||
from collections.abc import Iterable, Mapping
|
||
from pathlib import Path
|
||
from typing import Any
|
||
from urllib.parse import unquote, urlparse
|
||
|
||
from langchain_core.tools import BaseTool, StructuredTool
|
||
from langgraph.config import get_config
|
||
|
||
from deerflow.config.extensions_config import ExtensionsConfig
|
||
from deerflow.config.paths import VIRTUAL_PATH_PREFIX, Paths, get_paths
|
||
from deerflow.mcp.client import build_servers_config
|
||
from deerflow.mcp.oauth import build_oauth_tool_interceptor, get_initial_oauth_headers
|
||
from deerflow.mcp.session_pool import get_session_pool
|
||
from deerflow.reflection import resolve_variable
|
||
from deerflow.runtime.user_context import resolve_runtime_user_id
|
||
from deerflow.tools.sync import make_sync_tool_wrapper
|
||
from deerflow.tools.types import Runtime
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
# Subdirectory under the thread's workspace used as the temp dir for stdio MCP
|
||
# subprocesses. Pinning the process temp dir here (alongside its cwd) makes
|
||
# tools that write to ``os.tmpdir()`` / ``tempfile.gettempdir()`` land inside
|
||
# the mounted user-data tree, where their output is resolvable by the
|
||
# sandbox/artifact API — instead of on an unreachable host temp path.
|
||
_MCP_TMP_SUBDIR = ".mcp/tmp"
|
||
|
||
# Matches local-file references embedded in free text returned by an MCP server.
|
||
# Some servers (notably Playwright's ``browser_take_screenshot``) report saved
|
||
# files only as text/markdown links rather than ``ResourceLink`` blocks. Those
|
||
# references may be absolute paths, ``file://`` URIs, or paths relative to the
|
||
# server process cwd (e.g. ``temp/page.yml``, ``./shot.png``). Each match is
|
||
# only rewritten when it resolves to an existing file inside the thread's
|
||
# user-data tree, so an over-eager match is harmless (left untouched).
|
||
_LOCAL_PATH_IN_TEXT_RE = re.compile(r"(?:file://)?/[^\s'\"<>|*?]+|(?:\.{0,2}/|[\w.-]+/)[^\s'\"<>|*?]+")
|
||
|
||
# Trailing characters that are punctuation/markup rather than part of a path.
|
||
_TEXT_PATH_TRAILING_CHARS = ".,;:!?)]}>\"'`"
|
||
|
||
_FILE_SNAPSHOT = dict[Path, tuple[int, int]]
|
||
|
||
|
||
def _local_path_from_uri(uri: str, *, base_dir: Path | None = None) -> Path | None:
|
||
"""Return an absolute local filesystem ``Path`` if *uri* points to a local
|
||
file, otherwise ``None``.
|
||
|
||
Accepts bare paths and ``file://`` URIs. Remote URIs
|
||
(``http``/``https``/``data``/...) return ``None`` so the caller leaves them
|
||
untouched. Relative paths are resolved only when *base_dir* is supplied.
|
||
"""
|
||
if not uri:
|
||
return None
|
||
parsed = urlparse(uri)
|
||
if parsed.scheme == "file":
|
||
raw = unquote(parsed.path)
|
||
elif parsed.scheme == "":
|
||
raw = uri
|
||
else:
|
||
return None
|
||
if not raw:
|
||
return None
|
||
path = Path(raw)
|
||
if not path.is_absolute():
|
||
if base_dir is None:
|
||
return None
|
||
path = base_dir / path
|
||
return path
|
||
|
||
|
||
def _local_uri_to_virtual_path(
|
||
uri: str,
|
||
*,
|
||
thread_id: str,
|
||
user_id: str,
|
||
source_base_dir: Path | None = None,
|
||
) -> str | None:
|
||
"""Translate a local file reference into its ``/mnt/user-data/...`` virtual path.
|
||
|
||
Stdio MCP servers run with their cwd and temp dir pinned inside the thread's
|
||
mounted user-data tree (see :func:`_make_session_pool_tool`), so the files
|
||
they produce already live somewhere the sandbox/artifact API can serve — the
|
||
only thing missing is the virtual prefix the rest of DeerFlow addresses them
|
||
by. This performs that purely deterministic host→virtual mapping: no copy, no
|
||
trusted-root list, and no exposure of files outside the thread's own tree.
|
||
|
||
Returns ``None`` (so the caller leaves the reference untouched) when the URI
|
||
is remote, cannot be resolved, points outside this thread's user-data tree,
|
||
or does not name an existing file. Relative references are resolved against
|
||
*source_base_dir* (the server's cwd).
|
||
"""
|
||
src = _local_path_from_uri(uri, base_dir=source_base_dir)
|
||
if src is None:
|
||
return None
|
||
|
||
try:
|
||
real = src.resolve()
|
||
except OSError:
|
||
return None
|
||
if not real.is_file():
|
||
return None
|
||
|
||
try:
|
||
user_data_root = get_paths().sandbox_user_data_dir(thread_id, user_id=user_id).resolve()
|
||
except OSError:
|
||
return None
|
||
|
||
try:
|
||
relative = real.relative_to(user_data_root)
|
||
except ValueError:
|
||
# The file lives outside this thread's user-data mount; we cannot
|
||
# express it as a virtual path, so leave the original reference as-is.
|
||
logger.debug("MCP path rewrite skipped outside user-data tree: %s", real)
|
||
return None
|
||
|
||
virtual_path = f"{VIRTUAL_PATH_PREFIX}/{relative.as_posix()}"
|
||
logger.debug("MCP path rewrite: %s -> %s", real, virtual_path)
|
||
return virtual_path
|
||
|
||
|
||
def _snapshot_workspace_files(root: Path) -> _FILE_SNAPSHOT:
|
||
"""Return a lightweight snapshot of regular files under *root*."""
|
||
snapshot: _FILE_SNAPSHOT = {}
|
||
if not root.exists():
|
||
return snapshot
|
||
|
||
try:
|
||
candidates = root.rglob("*")
|
||
for path in candidates:
|
||
try:
|
||
stat = path.stat()
|
||
except OSError:
|
||
continue
|
||
if path.is_file():
|
||
snapshot[path] = (stat.st_mtime_ns, stat.st_size)
|
||
except OSError:
|
||
return snapshot
|
||
return snapshot
|
||
|
||
|
||
def _changed_workspace_files(root: Path, before: _FILE_SNAPSHOT) -> list[Path]:
|
||
"""Return files under *root* that were created or modified since *before*."""
|
||
after = _snapshot_workspace_files(root)
|
||
return [path for path, signature in after.items() if before.get(path) != signature]
|
||
|
||
|
||
def _prepare_stdio_workspace(paths: Paths, *, thread_id: str, user_id: str) -> tuple[Path, Path, _FILE_SNAPSHOT]:
|
||
"""Prepare the thread workspace for a pinned stdio MCP subprocess.
|
||
|
||
Bundles all the synchronous filesystem work (dir creation, temp-dir prep,
|
||
and the pre-call snapshot) into one helper so the caller can run it off the
|
||
event loop via :func:`asyncio.to_thread`. Returns the workspace cwd, the
|
||
pinned temp dir, and the pre-call file snapshot.
|
||
"""
|
||
paths.ensure_thread_dirs(thread_id, user_id=user_id)
|
||
source_base_dir = paths.sandbox_work_dir(thread_id, user_id=user_id)
|
||
tmp_dir = source_base_dir / _MCP_TMP_SUBDIR
|
||
try:
|
||
tmp_dir.mkdir(parents=True, exist_ok=True)
|
||
tmp_dir.chmod(0o700)
|
||
except OSError:
|
||
logger.warning("Failed to prepare MCP temp dir: %s", tmp_dir, exc_info=True)
|
||
before_files = _snapshot_workspace_files(source_base_dir)
|
||
return source_base_dir, tmp_dir, before_files
|
||
|
||
|
||
def _result_has_text_content(call_tool_result: Any) -> bool:
|
||
"""Return ``True`` when the MCP result carries any text content.
|
||
|
||
The after-call snapshot diff only feeds bare-filename correlation in free
|
||
text. When the result has no text blocks there is nothing to rewrite, so the
|
||
caller can skip the second recursive walk entirely.
|
||
"""
|
||
from mcp.types import EmbeddedResource, TextContent, TextResourceContents
|
||
|
||
content = getattr(call_tool_result, "content", None)
|
||
if not content:
|
||
return False
|
||
for item in content:
|
||
if isinstance(item, TextContent):
|
||
return True
|
||
if isinstance(item, EmbeddedResource) and isinstance(item.resource, TextResourceContents):
|
||
return True
|
||
return False
|
||
|
||
|
||
def _rewrite_unique_bare_filenames(
|
||
text: str,
|
||
*,
|
||
changed_files: Iterable[Path],
|
||
thread_id: str,
|
||
user_id: str,
|
||
source_base_dir: Path | None = None,
|
||
) -> str:
|
||
"""Rewrite bare filenames only when this call produced a unique match.
|
||
|
||
A response like ``Saved as page-2026.yml`` is not structurally a path. The
|
||
only safe way to interpret it is to correlate the filename with files
|
||
created/modified by this exact tool call, and rewrite only when the basename
|
||
maps to exactly one file inside this thread's mounted user-data tree.
|
||
"""
|
||
candidates: dict[str, list[str]] = {}
|
||
for path in changed_files:
|
||
virtual_path = _local_uri_to_virtual_path(
|
||
str(path),
|
||
thread_id=thread_id,
|
||
user_id=user_id,
|
||
source_base_dir=source_base_dir,
|
||
)
|
||
if virtual_path is None:
|
||
continue
|
||
candidates.setdefault(path.name, []).append(virtual_path)
|
||
|
||
unique = {name: paths[0] for name, paths in candidates.items() if len(set(paths)) == 1}
|
||
if not unique:
|
||
if candidates:
|
||
logger.debug("MCP bare filename rewrite skipped: no unique candidate in %s", sorted(candidates))
|
||
else:
|
||
logger.debug("MCP bare filename rewrite skipped: no snapshot candidates")
|
||
return text
|
||
|
||
rewritten = text
|
||
for name in sorted(unique, key=len, reverse=True):
|
||
# Do not rewrite inside longer paths/words. A final sentence period is
|
||
# allowed, but ".bak" or another path segment is not.
|
||
pattern = re.compile(rf"(?<![\w./-]){re.escape(name)}(?!(?:[\w/-]|\.[\w]))")
|
||
rewritten_text, count = pattern.subn(unique[name], rewritten)
|
||
if count:
|
||
logger.debug("MCP bare filename rewrite: %s -> %s", name, unique[name])
|
||
rewritten = rewritten_text
|
||
return rewritten
|
||
|
||
|
||
def _rewrite_local_paths_in_text(
|
||
text: str,
|
||
*,
|
||
thread_id: str,
|
||
user_id: str,
|
||
source_base_dir: Path | None = None,
|
||
changed_files: Iterable[Path] | None = None,
|
||
) -> str:
|
||
"""Best-effort rewrite of local file references found in free text.
|
||
|
||
Some MCP servers (notably Playwright's ``browser_take_screenshot``) report
|
||
the saved file only as free text — e.g. ``Took the screenshot and saved it
|
||
as temp/page-2026.png`` — instead of a ``ResourceLink``. Free text is not a
|
||
reliable protocol, so this is deliberately conservative: every candidate
|
||
token is handed to :func:`_local_uri_to_virtual_path`, which only rewrites
|
||
it when it resolves to an existing file inside this thread's user-data tree.
|
||
Tokens that are not real paths (or point elsewhere) are left exactly as they
|
||
were, so an over-eager regex match has no harmful effect.
|
||
"""
|
||
translated_by_source: dict[str, str | None] = {}
|
||
|
||
def _replace(match: re.Match[str]) -> str:
|
||
token = match.group(0)
|
||
# A path can end a sentence ("saved as temp/a.png."); strip trailing
|
||
# punctuation and restore it after the (possibly rewritten) path.
|
||
stripped = token.rstrip(_TEXT_PATH_TRAILING_CHARS)
|
||
trailing = token[len(stripped) :]
|
||
if stripped not in translated_by_source:
|
||
translated_by_source[stripped] = _local_uri_to_virtual_path(
|
||
stripped,
|
||
thread_id=thread_id,
|
||
user_id=user_id,
|
||
source_base_dir=source_base_dir,
|
||
)
|
||
rewritten = translated_by_source[stripped]
|
||
if rewritten is None:
|
||
return token
|
||
return f"{rewritten}{trailing}"
|
||
|
||
rewritten = _LOCAL_PATH_IN_TEXT_RE.sub(_replace, text)
|
||
if changed_files is None:
|
||
return rewritten
|
||
return _rewrite_unique_bare_filenames(
|
||
rewritten,
|
||
changed_files=changed_files,
|
||
thread_id=thread_id,
|
||
user_id=user_id,
|
||
source_base_dir=source_base_dir,
|
||
)
|
||
|
||
|
||
def _extract_thread_id(runtime: Runtime | None) -> str:
|
||
"""Extract thread_id from the injected tool runtime or LangGraph config."""
|
||
if runtime is not None:
|
||
tid = runtime.context.get("thread_id") if runtime.context else None
|
||
if tid is not None:
|
||
return str(tid)
|
||
config = runtime.config or {}
|
||
tid = config.get("configurable", {}).get("thread_id")
|
||
if tid is not None:
|
||
return str(tid)
|
||
|
||
try:
|
||
tid = get_config().get("configurable", {}).get("thread_id")
|
||
return str(tid) if tid is not None else "default"
|
||
except RuntimeError:
|
||
return "default"
|
||
|
||
|
||
def _convert_call_tool_result(
|
||
call_tool_result: Any,
|
||
*,
|
||
thread_id: str | None = None,
|
||
user_id: str | None = None,
|
||
source_base_dir: Path | None = None,
|
||
changed_files: Iterable[Path] | None = None,
|
||
) -> Any:
|
||
"""Convert an MCP CallToolResult to the LangChain ``content_and_artifact`` format.
|
||
|
||
Implements the same conversion logic as the adapter without relying on
|
||
the private ``langchain_mcp_adapters.tools._convert_call_tool_result`` symbol.
|
||
|
||
When ``thread_id`` and ``user_id`` are provided, local files referenced by
|
||
``ResourceLink`` blocks or plain text (e.g. screenshots saved by Playwright
|
||
MCP) have their references translated from the host path to the
|
||
``/mnt/user-data/...`` virtual path so they can be resolved by the sandbox
|
||
and artifact API. The files themselves are not copied — stdio servers run
|
||
with their cwd/temp pinned inside the mounted tree, so they already live in
|
||
a servable location. Remote URIs and files outside the thread's user-data
|
||
tree are left untouched.
|
||
"""
|
||
from langchain_core.messages import ToolMessage
|
||
from langchain_core.messages.content import create_file_block, create_image_block, create_text_block
|
||
from langchain_core.tools import ToolException
|
||
from mcp.types import EmbeddedResource, ImageContent, ResourceLink, TextContent, TextResourceContents
|
||
|
||
# Pass ToolMessage through directly (interceptor short-circuit).
|
||
if isinstance(call_tool_result, ToolMessage):
|
||
return call_tool_result, None
|
||
|
||
# Pass LangGraph Command through directly when langgraph is installed.
|
||
try:
|
||
from langgraph.types import Command
|
||
|
||
if isinstance(call_tool_result, Command):
|
||
return call_tool_result, None
|
||
except ImportError:
|
||
# langgraph is optional; if unavailable, continue with standard MCP content conversion.
|
||
pass
|
||
|
||
def _resolve_link_url(uri: str) -> str:
|
||
if thread_id is None or user_id is None:
|
||
return uri
|
||
rewritten = _local_uri_to_virtual_path(uri, thread_id=thread_id, user_id=user_id, source_base_dir=source_base_dir)
|
||
return rewritten if rewritten is not None else uri
|
||
|
||
def _resolve_text(text: str) -> str:
|
||
# Servers like Playwright report saved files only as plain text, with no
|
||
# ResourceLink to hook into. Scan the text for local paths and translate
|
||
# them so the produced files are readable through the sandbox/artifact API.
|
||
if thread_id is None or user_id is None:
|
||
return text
|
||
return _rewrite_local_paths_in_text(
|
||
text,
|
||
thread_id=thread_id,
|
||
user_id=user_id,
|
||
source_base_dir=source_base_dir,
|
||
changed_files=changed_files,
|
||
)
|
||
|
||
# Convert MCP content blocks to LangChain content blocks.
|
||
lc_content = []
|
||
for item in call_tool_result.content:
|
||
if isinstance(item, TextContent):
|
||
lc_content.append(create_text_block(text=_resolve_text(item.text)))
|
||
elif isinstance(item, ImageContent):
|
||
lc_content.append(create_image_block(base64=item.data, mime_type=item.mimeType))
|
||
elif isinstance(item, ResourceLink):
|
||
mime = item.mimeType or None
|
||
url = _resolve_link_url(str(item.uri))
|
||
if mime and mime.startswith("image/"):
|
||
lc_content.append(create_image_block(url=url, mime_type=mime))
|
||
else:
|
||
lc_content.append(create_file_block(url=url, mime_type=mime))
|
||
elif isinstance(item, EmbeddedResource):
|
||
from mcp.types import BlobResourceContents
|
||
|
||
res = item.resource
|
||
if isinstance(res, TextResourceContents):
|
||
lc_content.append(create_text_block(text=_resolve_text(res.text)))
|
||
elif isinstance(res, BlobResourceContents):
|
||
mime = res.mimeType or None
|
||
if mime and mime.startswith("image/"):
|
||
lc_content.append(create_image_block(base64=res.blob, mime_type=mime))
|
||
else:
|
||
lc_content.append(create_file_block(base64=res.blob, mime_type=mime))
|
||
else:
|
||
lc_content.append(create_text_block(text=str(res)))
|
||
else:
|
||
lc_content.append(create_text_block(text=str(item)))
|
||
|
||
if call_tool_result.isError:
|
||
error_parts = [item["text"] for item in lc_content if isinstance(item, dict) and item.get("type") == "text"]
|
||
raise ToolException("\n".join(error_parts) if error_parts else str(lc_content))
|
||
|
||
artifact = None
|
||
if call_tool_result.structuredContent is not None:
|
||
artifact = {"structured_content": call_tool_result.structuredContent}
|
||
|
||
return lc_content, artifact
|
||
|
||
|
||
def _make_session_pool_tool(
|
||
tool: BaseTool,
|
||
server_name: str,
|
||
connection: dict[str, Any],
|
||
tool_interceptors: list[Any] | None = None,
|
||
) -> BaseTool:
|
||
"""Wrap an MCP tool so it reuses a persistent session from the pool.
|
||
|
||
Replaces the per-call session creation with pool-managed sessions scoped
|
||
by ``(server_name, user_id:thread_id)``. This ensures stateful MCP servers
|
||
(e.g. Playwright) keep their state across tool calls within the same thread
|
||
while staying isolated per user.
|
||
|
||
The configured ``tool_interceptors`` (OAuth, custom) are preserved and
|
||
applied on every call before invoking the pooled session.
|
||
"""
|
||
# Strip the server-name prefix to recover the original MCP tool name.
|
||
original_name = tool.name
|
||
prefix = f"{server_name}_"
|
||
if original_name.startswith(prefix):
|
||
original_name = original_name[len(prefix) :]
|
||
|
||
pool = get_session_pool()
|
||
|
||
async def call_with_persistent_session(
|
||
runtime: Runtime | None = None,
|
||
**arguments: Any,
|
||
) -> Any:
|
||
thread_id = _extract_thread_id(runtime)
|
||
user_id = resolve_runtime_user_id(runtime)
|
||
# Scope the pooled session by user *and* thread. Filesystem isolation is
|
||
# per-(user_id, thread_id), so a thread_id alone could otherwise let two
|
||
# users with a colliding thread_id share one stateful MCP session.
|
||
scope_key = f"{user_id}:{thread_id}"
|
||
session_connection = dict(connection)
|
||
# cwd/temp pinning and the workspace snapshot only matter for stdio
|
||
# servers, which run as local subprocesses writing to a real filesystem.
|
||
# SSE/HTTP servers have no local cwd to pin, so skip the filesystem work
|
||
# entirely for them (avoids needless dir creation and recursive walks).
|
||
is_stdio = session_connection.get("transport", "stdio") == "stdio"
|
||
source_base_dir: Path | None = None
|
||
process_cwd: Path | None = None
|
||
before_files: _FILE_SNAPSHOT | None = None
|
||
if is_stdio:
|
||
paths = get_paths()
|
||
# Bundle the synchronous filesystem prep (dir creation, temp-dir
|
||
# setup, pre-call snapshot) and run it off the event loop — the
|
||
# snapshot walks the whole workspace and would otherwise block.
|
||
source_base_dir, tmp_dir, before_files = await asyncio.to_thread(_prepare_stdio_workspace, paths, thread_id=thread_id, user_id=user_id)
|
||
# Stdio MCP servers resolve relative output links against their
|
||
# process cwd. Keep that cwd inside the thread's mounted user-data
|
||
# tree so files produced by tools like Playwright land where the
|
||
# sandbox/artifact API can serve them and their references can be
|
||
# translated to virtual paths.
|
||
configured_cwd = session_connection.get("cwd", str(source_base_dir))
|
||
session_connection["cwd"] = str(configured_cwd)
|
||
process_cwd = Path(configured_cwd)
|
||
# Pin the subprocess temp dir under the same mounted tree. Tools that
|
||
# default to the OS temp dir (Node's os.tmpdir(), Python's tempfile,
|
||
# many CLIs) then write inside user-data instead of an unreachable
|
||
# host path — the tool-agnostic counterpart to fixing the cwd. Merge
|
||
# rather than replace any operator-provided env.
|
||
session_env = dict(session_connection.get("env") or {})
|
||
session_env.setdefault("TMPDIR", str(tmp_dir))
|
||
session_env.setdefault("TMP", str(tmp_dir))
|
||
session_env.setdefault("TEMP", str(tmp_dir))
|
||
session_connection["env"] = session_env
|
||
session = await pool.get_session(server_name, scope_key, session_connection)
|
||
|
||
if tool_interceptors:
|
||
from langchain_mcp_adapters.interceptors import MCPToolCallRequest
|
||
|
||
async def base_handler(request: MCPToolCallRequest) -> Any:
|
||
# Preserve interceptor-injected headers for stdio MCP calls by
|
||
# forwarding them through MCP call meta.
|
||
call_kwargs: dict[str, Any] = {}
|
||
if request.headers:
|
||
if isinstance(request.headers, Mapping):
|
||
call_kwargs["meta"] = {"headers": dict(request.headers)}
|
||
else:
|
||
logger.warning("Ignoring MCP interceptor headers with unsupported type: %s", type(request.headers).__name__)
|
||
return await session.call_tool(request.name, request.args, **call_kwargs)
|
||
|
||
handler = base_handler
|
||
for interceptor in reversed(tool_interceptors):
|
||
outer = handler
|
||
|
||
async def wrapped(req: Any, _i: Any = interceptor, _h: Any = outer) -> Any:
|
||
return await _i(req, _h)
|
||
|
||
handler = wrapped
|
||
|
||
request = MCPToolCallRequest(
|
||
name=original_name,
|
||
args=arguments,
|
||
server_name=server_name,
|
||
runtime=runtime,
|
||
)
|
||
call_tool_result = await handler(request)
|
||
else:
|
||
call_tool_result = await session.call_tool(original_name, arguments)
|
||
|
||
# The after-call snapshot diff only feeds bare-filename correlation in
|
||
# free text, so skip the second recursive walk when there is no text
|
||
# content to rewrite. Both the diff and the per-token path resolution
|
||
# inside _convert_call_tool_result touch the filesystem, so run them off
|
||
# the event loop.
|
||
changed_files: list[Path] | None = None
|
||
if is_stdio and before_files is not None and _result_has_text_content(call_tool_result):
|
||
changed_files = await asyncio.to_thread(_changed_workspace_files, source_base_dir, before_files)
|
||
return await asyncio.to_thread(
|
||
_convert_call_tool_result,
|
||
call_tool_result,
|
||
thread_id=thread_id,
|
||
user_id=user_id,
|
||
source_base_dir=process_cwd,
|
||
changed_files=changed_files,
|
||
)
|
||
|
||
return StructuredTool(
|
||
name=tool.name,
|
||
description=tool.description,
|
||
args_schema=tool.args_schema,
|
||
coroutine=call_with_persistent_session,
|
||
response_format="content_and_artifact",
|
||
metadata=tool.metadata,
|
||
)
|
||
|
||
|
||
async def get_mcp_tools() -> list[BaseTool]:
|
||
"""Get all tools from enabled MCP servers.
|
||
|
||
Tools using stdio transport are wrapped with persistent-session logic so
|
||
consecutive calls within the same thread reuse the same MCP session.
|
||
HTTP/SSE tools are returned unwrapped to avoid cross-task TaskGroup
|
||
cleanup errors.
|
||
|
||
Returns:
|
||
List of LangChain tools from all enabled MCP servers.
|
||
"""
|
||
try:
|
||
from langchain_mcp_adapters.client import MultiServerMCPClient
|
||
except ImportError:
|
||
logger.warning("langchain-mcp-adapters not installed. Install it to enable MCP tools: pip install langchain-mcp-adapters")
|
||
return []
|
||
|
||
# NOTE: We use ExtensionsConfig.from_file() instead of get_extensions_config()
|
||
# to always read the latest configuration from disk. This ensures that changes
|
||
# made through the Gateway API (which runs in a separate process) are immediately
|
||
# reflected when initializing MCP tools.
|
||
extensions_config = ExtensionsConfig.from_file()
|
||
servers_config = build_servers_config(extensions_config)
|
||
|
||
if not servers_config:
|
||
logger.info("No enabled MCP servers configured")
|
||
return []
|
||
|
||
try:
|
||
# Create the multi-server MCP client
|
||
logger.info(f"Initializing MCP client with {len(servers_config)} server(s)")
|
||
|
||
# Inject initial OAuth headers for server connections (tool discovery/session init)
|
||
initial_oauth_headers = await get_initial_oauth_headers(extensions_config)
|
||
for server_name, auth_header in initial_oauth_headers.items():
|
||
if server_name not in servers_config:
|
||
continue
|
||
if servers_config[server_name].get("transport") in ("sse", "http"):
|
||
existing_headers = dict(servers_config[server_name].get("headers", {}))
|
||
existing_headers["Authorization"] = auth_header
|
||
servers_config[server_name]["headers"] = existing_headers
|
||
|
||
tool_interceptors: list[Any] = []
|
||
oauth_interceptor = build_oauth_tool_interceptor(extensions_config)
|
||
if oauth_interceptor is not None:
|
||
tool_interceptors.append(oauth_interceptor)
|
||
|
||
# Load custom interceptors declared in extensions_config.json
|
||
# Format: "mcpInterceptors": ["pkg.module:builder_func", ...]
|
||
raw_interceptor_paths = (extensions_config.model_extra or {}).get("mcpInterceptors")
|
||
if isinstance(raw_interceptor_paths, str):
|
||
raw_interceptor_paths = [raw_interceptor_paths]
|
||
elif not isinstance(raw_interceptor_paths, list):
|
||
if raw_interceptor_paths is not None:
|
||
logger.warning(f"mcpInterceptors must be a list of strings, got {type(raw_interceptor_paths).__name__}; skipping")
|
||
raw_interceptor_paths = []
|
||
for interceptor_path in raw_interceptor_paths:
|
||
try:
|
||
builder = resolve_variable(interceptor_path)
|
||
interceptor = builder()
|
||
if callable(interceptor):
|
||
tool_interceptors.append(interceptor)
|
||
logger.info(f"Loaded MCP interceptor: {interceptor_path}")
|
||
elif interceptor is not None:
|
||
logger.warning(f"Builder {interceptor_path} returned non-callable {type(interceptor).__name__}; skipping")
|
||
except Exception as e:
|
||
logger.warning(
|
||
f"Failed to load MCP interceptor {interceptor_path}: {e}",
|
||
exc_info=True,
|
||
)
|
||
|
||
client = MultiServerMCPClient(
|
||
servers_config,
|
||
tool_interceptors=tool_interceptors,
|
||
tool_name_prefix=True,
|
||
)
|
||
|
||
# Get all tools from all servers (discovers tool definitions via
|
||
# temporary sessions – the persistent-session wrapping is applied below).
|
||
tools = await client.get_tools()
|
||
logger.info(f"Successfully loaded {len(tools)} tool(s) from MCP servers")
|
||
|
||
# Wrap each tool with persistent-session logic.
|
||
# Only pool stdio sessions. HTTP/SSE transports use anyio TaskGroups
|
||
# internally which cannot be closed from a different async task, so
|
||
# pooling them causes RuntimeError on cleanup (see #3203).
|
||
wrapped_tools: list[BaseTool] = []
|
||
for tool in tools:
|
||
tool_server: str | None = None
|
||
for name in servers_config:
|
||
if tool.name.startswith(f"{name}_"):
|
||
tool_server = name
|
||
break
|
||
|
||
if tool_server is not None:
|
||
transport = servers_config[tool_server].get("transport", "stdio")
|
||
if transport == "stdio":
|
||
wrapped_tools.append(_make_session_pool_tool(tool, tool_server, servers_config[tool_server], tool_interceptors))
|
||
else:
|
||
wrapped_tools.append(tool)
|
||
else:
|
||
wrapped_tools.append(tool)
|
||
|
||
# Patch tools to support sync invocation, as deerflow client streams synchronously
|
||
for tool in wrapped_tools:
|
||
if getattr(tool, "func", None) is None and getattr(tool, "coroutine", None) is not None:
|
||
tool.func = make_sync_tool_wrapper(tool.coroutine, tool.name)
|
||
|
||
return wrapped_tools
|
||
|
||
except Exception as e:
|
||
logger.error(f"Failed to load MCP tools: {e}", exc_info=True)
|
||
return []
|