mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-28 23:46:21 +00:00
McpRoutingMiddleware picked the latest user message with is_real_user_message, which rejects every hide_from_ui message. A Human Input Card reply is hidden but is still the user's current request, so when the routing keyword lived only in the clarification answer the deferred MCP tool was never auto-promoted and the model had to call tool_search by hand. Switch _latest_user_message to is_genuine_user_message, the same predicate summarization_middleware already uses for this reason (#5416). Fixes #5425
156 lines
6.6 KiB
Python
156 lines
6.6 KiB
Python
"""Auto-promote deferred MCP tools from routing metadata before model calls."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
from collections.abc import Mapping, Sequence
|
|
from typing import Any, TypedDict, override
|
|
|
|
from langchain.agents import AgentState
|
|
from langchain.agents.middleware import AgentMiddleware
|
|
from langchain_core.messages import HumanMessage
|
|
from langgraph.runtime import Runtime
|
|
|
|
from deerflow.agents.middlewares.message_utils import is_genuine_user_message
|
|
from deerflow.agents.middlewares.tool_promotion_audit_middleware import record_tool_promotion
|
|
from deerflow.config.tool_search_config import clamp_auto_promote_top_k
|
|
from deerflow.utils.messages import get_original_user_content_text
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
class McpRoutingIndexEntry(TypedDict):
|
|
priority: int
|
|
keywords: list[str]
|
|
|
|
|
|
McpRoutingIndex = Mapping[str, McpRoutingIndexEntry]
|
|
|
|
|
|
class McpRoutingMiddleware(AgentMiddleware[AgentState]):
|
|
"""Write minimal deferred-tool promotion state from latest user text.
|
|
|
|
The middleware intentionally receives only serialized routing data. It does
|
|
not hold ``BaseTool`` objects, does not execute tools, and does not filter
|
|
tool calls. ``DeferredToolFilterMiddleware`` remains responsible for hiding
|
|
unpromoted schemas and blocking unpromoted deferred tool calls.
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
routing_index: McpRoutingIndex,
|
|
catalog_hash: str | None,
|
|
top_k: int,
|
|
) -> None:
|
|
super().__init__()
|
|
self._catalog_hash = catalog_hash
|
|
self._top_k = clamp_auto_promote_top_k(top_k)
|
|
self._routing_index = self._normalize_index(routing_index)
|
|
|
|
@staticmethod
|
|
def _normalize_index(routing_index: McpRoutingIndex) -> dict[str, tuple[int, tuple[str, ...]]]:
|
|
# Defensive re-normalization: this middleware is built to accept arbitrary
|
|
# serialized routing data, not only the output of
|
|
# tool_search._routing_priority / _routing_keywords. In practice it is a
|
|
# no-op over the builder's output; keep the coercion rules aligned with
|
|
# those two helpers if either side changes.
|
|
normalized: dict[str, tuple[int, tuple[str, ...]]] = {}
|
|
for raw_name, raw_entry in routing_index.items():
|
|
name = str(raw_name)
|
|
if not name:
|
|
continue
|
|
try:
|
|
priority = int(raw_entry.get("priority", 0))
|
|
except (TypeError, ValueError):
|
|
priority = 0
|
|
raw_keywords = raw_entry.get("keywords") or []
|
|
if not isinstance(raw_keywords, Sequence) or isinstance(raw_keywords, (str, bytes)):
|
|
raw_keywords = []
|
|
keywords = tuple(keyword for keyword in (str(item).strip() for item in raw_keywords) if keyword)
|
|
if not keywords:
|
|
continue
|
|
normalized[name] = (priority, keywords)
|
|
return normalized
|
|
|
|
@staticmethod
|
|
def _latest_user_message(messages: list[Any]) -> HumanMessage | None:
|
|
"""Latest user-authored message, including a hidden Human Input Card reply.
|
|
|
|
The card reply is hidden from the UI but is still the user's current
|
|
request, so ``is_genuine_user_message`` — not ``is_real_user_message``,
|
|
which drops every ``hide_from_ui`` message — decides this.
|
|
"""
|
|
for message in reversed(messages):
|
|
if is_genuine_user_message(message):
|
|
return message
|
|
return None
|
|
|
|
def _matched_names(self, state: Mapping[str, Any] | None) -> list[str]:
|
|
if not self._catalog_hash or not self._routing_index:
|
|
return []
|
|
messages = list((state or {}).get("messages") or [])
|
|
target = self._latest_user_message(messages)
|
|
if target is None:
|
|
return []
|
|
|
|
text = get_original_user_content_text(target.content, target.additional_kwargs)
|
|
if not text:
|
|
return []
|
|
|
|
haystack = text.casefold()
|
|
matched: list[tuple[int, str]] = []
|
|
for name, (priority, keywords) in self._routing_index.items():
|
|
if any(keyword.casefold() in haystack for keyword in keywords):
|
|
matched.append((priority, name))
|
|
|
|
if not matched:
|
|
return []
|
|
|
|
matched.sort(key=lambda item: (-item[0], item[1]))
|
|
return [name for _, name in matched[: self._top_k]]
|
|
|
|
def _state_update(self, state: Mapping[str, Any] | None, runtime: Runtime | None) -> dict[str, Any] | None:
|
|
names = self._matched_names(state)
|
|
if not names:
|
|
return None
|
|
promoted = (state or {}).get("promoted")
|
|
raw_promoted_names = promoted.get("names") if isinstance(promoted, Mapping) and promoted.get("catalog_hash") == self._catalog_hash else None
|
|
already_promoted = {name for name in raw_promoted_names if isinstance(name, str)} if isinstance(raw_promoted_names, Sequence) and not isinstance(raw_promoted_names, (str, bytes)) else set()
|
|
record_tool_promotion(
|
|
runtime,
|
|
producer=type(self).__name__,
|
|
hook="before_model",
|
|
source="routing_hint",
|
|
tool_names=set(names) - already_promoted,
|
|
)
|
|
logger.debug(
|
|
"McpRoutingMiddleware auto-promoted %d deferred tool schema(s) catalog=%s names=%s",
|
|
len(names),
|
|
(self._catalog_hash or "")[:8],
|
|
names,
|
|
)
|
|
return {
|
|
"promoted": {
|
|
"catalog_hash": self._catalog_hash,
|
|
"names": names,
|
|
}
|
|
}
|
|
|
|
@override
|
|
def before_model(self, state: AgentState, runtime: Runtime) -> dict[str, Any] | None:
|
|
return self._state_update(state, runtime)
|
|
|
|
@override
|
|
async def abefore_model(self, state: AgentState, runtime: Runtime) -> dict[str, Any] | None:
|
|
return self._state_update(state, runtime)
|
|
|
|
|
|
def assert_mcp_routing_before_deferred_filter(middlewares: Sequence[AgentMiddleware]) -> None:
|
|
"""Fail fast if auto-promote would run after deferred schema filtering."""
|
|
from deerflow.agents.middlewares.deferred_tool_filter_middleware import DeferredToolFilterMiddleware
|
|
|
|
routing_idx = next((idx for idx, middleware in enumerate(middlewares) if isinstance(middleware, McpRoutingMiddleware)), None)
|
|
filter_idx = next((idx for idx, middleware in enumerate(middlewares) if isinstance(middleware, DeferredToolFilterMiddleware)), None)
|
|
if routing_idx is not None and filter_idx is not None and routing_idx > filter_idx:
|
|
raise RuntimeError(f"McpRoutingMiddleware must be installed before DeferredToolFilterMiddleware (routing index {routing_idx}, deferred filter index {filter_idx})")
|