deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/agents/middlewares/clarification_middleware.py
2026-09-07 18:24:55 +08:00

334 lines
13 KiB
Python

"""Middleware for intercepting clarification requests and presenting them to the user."""
import json
import logging
from collections.abc import Callable
from hashlib import sha256
from typing import Any, override
from deerflow.tools.builtins.clarification_utils import resolve_allow_multiple
from langchain.agents import AgentState
from langchain.agents.middleware import AgentMiddleware, hook_config
from langchain_core.messages import AIMessage, ToolMessage
from langgraph.prebuilt.tool_node import ToolCallRequest
from langgraph.types import Command
logger = logging.getLogger(__name__)
class ClarificationMiddlewareState(AgentState):
"""Compatible with the `ThreadState` schema."""
pass
class ClarificationMiddleware(AgentMiddleware[ClarificationMiddlewareState]):
"""Intercepts clarification tool calls and interrupts execution to present questions to the user.
When the model calls the `ask_clarification` tool, this middleware:
1. Intercepts the tool call before execution
2. Extracts the clarification question and metadata
3. Formats a user-friendly message
4. Returns a Command that interrupts execution and presents the question
5. Waits for user response before continuing
This replaces the tool-based approach where clarification continued the conversation flow.
"""
state_schema = ClarificationMiddlewareState
def _stable_message_id(self, tool_call_id: str, formatted_message: str) -> str:
"""Build a deterministic message ID so retried clarification calls replace, not append."""
if tool_call_id:
return f"clarification:{tool_call_id}"
digest = sha256(formatted_message.encode("utf-8")).hexdigest()[:16]
return f"clarification:{digest}"
def _is_chinese(self, text: str) -> bool:
"""Check if text contains Chinese characters.
Args:
text: Text to check
Returns:
True if text contains Chinese characters
"""
return any("\u4e00" <= char <= "\u9fff" for char in text)
def _format_clarification_message(self, args: dict) -> str:
"""Format the clarification arguments into a user-friendly message.
Args:
args: The tool call arguments containing clarification details
Returns:
Formatted message string
"""
question = args.get("question", "")
clarification_type = args.get("clarification_type", "missing_info")
context = args.get("context")
options = self._normalize_options(args.get("options", []))
# Type-specific icons
type_icons = {
"missing_info": "❓",
"ambiguous_requirement": "🤔",
"approach_choice": "🔀",
"risk_confirmation": "⚠️",
"suggestion": "💡",
}
icon = type_icons.get(clarification_type, "❓")
# Build the message naturally
message_parts = []
# Add icon and question together for a more natural flow
if context:
# If there's context, present it first as background
message_parts.append(f"{icon} {context}")
message_parts.append(f"\n{question}")
else:
# Just the question with icon
message_parts.append(f"{icon} {question}")
# Add options in a cleaner format
if options and len(options) > 0:
message_parts.append("") # blank line for spacing
for i, option in enumerate(options, 1):
label = option.get("label") or option.get("id") or ""
message_parts.append(f" {i}. {label}")
return "\n".join(message_parts)
def _normalize_options(self, raw_options: object) -> list[dict[str, str]]:
"""Normalize model-produced options into a stable UI-friendly shape."""
options = raw_options
# Some models serialize array parameters as JSON strings instead of native arrays.
if isinstance(options, str):
try:
options = json.loads(options)
except (json.JSONDecodeError, TypeError):
options = [options]
if options is None:
values: list[object] = []
elif isinstance(options, list):
values = options
else:
values = [options]
normalized: list[dict[str, str]] = []
for index, value in enumerate(values, 1):
if isinstance(value, dict):
raw_label = value.get("label") or value.get("title") or value.get("name") or value.get("value") or value.get("id")
label = str(raw_label).strip() if raw_label is not None else ""
if not label:
continue
raw_id = value.get("id") or value.get("value") or f"option-{index}"
option = {
"id": str(raw_id).strip() or f"option-{index}",
"label": label,
}
description = value.get("description")
if description is not None and str(description).strip():
option["description"] = str(description).strip()
normalized.append(option)
else:
label = str(value).strip()
if label:
normalized.append({"id": f"option-{index}", "label": label})
return normalized
def _build_payload(self, args: dict) -> dict:
"""Build structured clarification data for rich clients."""
return {
"question": str(args.get("question", "")).strip(),
"clarification_type": args.get("clarification_type", "missing_info"),
"context": args.get("context"),
"options": self._normalize_options(args.get("options", [])),
"allow_custom": True,
"allow_multiple": resolve_allow_multiple(args),
}
def _parse_tool_args(self, tool_call: Any) -> dict[str, Any]:
args = self._tool_field(tool_call, "args") or {}
if isinstance(args, dict):
return args
if isinstance(args, str):
try:
parsed = json.loads(args)
return parsed if isinstance(parsed, dict) else {}
except (json.JSONDecodeError, TypeError):
return {}
return {}
def _handle_clarification(self, request: ToolCallRequest) -> ToolMessage:
"""Return the clarification ToolMessage without Command(goto=END).
Stopping the turn is ``after_model`` + ``jump_to=end``. Returning a
Command here used to mix with parallel tool results and loop the graph.
If the tools node still runs (jump_to not honored), a ToolMessage with
the same stable id replaces rather than appends.
"""
args = self._parse_tool_args(request.tool_call)
tool_call_id = str(self._tool_field(request.tool_call, "id") or "")
logger.info("Intercepted clarification request in wrap_tool_call")
return self._clarification_tool_message(tool_call_id, args)
def _tool_field(self, tool_call: Any, key: str, default: Any = None) -> Any:
if isinstance(tool_call, dict):
return tool_call.get(key, default)
getter = getattr(tool_call, "get", None)
if callable(getter):
return getter(key, default)
return getattr(tool_call, key, default)
def _tool_name(self, tool_call: Any) -> str:
"""Normalize tool names from dict / object / OpenAI-style nested function."""
raw = self._tool_field(tool_call, "name")
if raw is None and isinstance(tool_call, dict):
fn = tool_call.get("function")
if isinstance(fn, dict):
raw = fn.get("name")
name = str(raw or "").strip()
if "." in name:
name = name.rsplit(".", 1)[-1]
return name
def _is_clarification_call(self, tool_call: Any) -> bool:
return self._tool_name(tool_call).casefold() == "ask_clarification"
def _clarification_tool_message(self, tool_call_id: str, args: dict[str, Any]) -> ToolMessage:
formatted_message = self._format_clarification_message(args)
return ToolMessage(
id=self._stable_message_id(tool_call_id, formatted_message),
content=formatted_message,
tool_call_id=tool_call_id,
name="ask_clarification",
additional_kwargs={"clarification": self._build_payload(args)},
)
def _after_model_intercept(self, last: AIMessage) -> dict[str, Any] | None:
"""Stop the run as soon as the model emits ask_clarification.
Applies to **all** lead-agent chats (通用问答 / 智能体 / 写作配置 / 报告结构…)。
``wrap_tool_call`` + ``Command(goto=END)`` is unreliable with parallel
tool calls (weak models often emit two clarifications, or clarification
plus another tool): the Command mixes with other tool results and the
graph loops back to the model, which then treats the question ToolMessage
as "the user still hasn't answered" and keeps generating.
"""
tool_calls = list(last.tool_calls or [])
clarification_calls = [tc for tc in tool_calls if self._is_clarification_call(tc)]
if not clarification_calls:
return None
messages: list[ToolMessage] = []
for tc in tool_calls:
name = self._tool_name(tc) or "unknown"
tool_call_id = str(self._tool_field(tc, "id") or "")
args = self._parse_tool_args(tc)
if self._is_clarification_call(tc):
messages.append(self._clarification_tool_message(tool_call_id, args))
else:
messages.append(
ToolMessage(
id=f"clarification-skip:{tool_call_id or name}",
content="已暂停:需要先等待用户回答澄清问题,本工具未执行。",
tool_call_id=tool_call_id,
name=name,
)
)
logger.info(
"Intercepted %s clarification tool call(s) in after_model; ending turn",
len(clarification_calls),
)
return {"messages": messages, "jump_to": "end"}
@hook_config(can_jump_to=["end"])
@override
def after_model(
self,
state: ClarificationMiddlewareState,
runtime: object,
) -> dict[str, Any] | None:
messages = state.get("messages") or []
if not messages:
return None
last = messages[-1]
if not isinstance(last, AIMessage):
return None
return self._after_model_intercept(last)
@hook_config(can_jump_to=["end"])
@override
async def aafter_model(
self,
state: ClarificationMiddlewareState,
runtime: object,
) -> dict[str, Any] | None:
return self.after_model(state, runtime)
@override
def _stub_sibling_tool(self, request: ToolCallRequest) -> ToolMessage:
name = self._tool_name(request.tool_call) or "unknown"
tool_call_id = str(self._tool_field(request.tool_call, "id") or "")
return ToolMessage(
id=f"clarification-skip:{tool_call_id or name}",
content="已暂停:需要先等待用户回答澄清问题,本工具未执行。",
tool_call_id=tool_call_id,
name=name,
)
def _request_has_sibling_clarification(self, request: ToolCallRequest) -> bool:
state = getattr(request, "state", None)
messages: Any = None
if isinstance(state, dict):
messages = state.get("messages")
elif state is not None:
getter = getattr(state, "get", None)
if callable(getter):
messages = getter("messages")
else:
messages = getattr(state, "messages", None)
if not messages:
return False
last = messages[-1]
tool_calls = getattr(last, "tool_calls", None) or []
return any(self._is_clarification_call(tc) for tc in tool_calls)
@override
def wrap_tool_call(
self,
request: ToolCallRequest,
handler: Callable[[ToolCallRequest], ToolMessage | Command],
) -> ToolMessage | Command:
"""Complete clarification tool calls; do not Command(goto=END).
Turn-stop is ``after_model`` + ``jump_to=end``. If the tools node still
runs, return a ToolMessage so parallel tools do not mix Commands.
"""
if self._is_clarification_call(request.tool_call):
return self._handle_clarification(request)
if self._request_has_sibling_clarification(request):
return self._stub_sibling_tool(request)
return handler(request)
@override
async def awrap_tool_call(
self,
request: ToolCallRequest,
handler: Callable[[ToolCallRequest], ToolMessage | Command],
) -> ToolMessage | Command:
if self._is_clarification_call(request.tool_call):
return self._handle_clarification(request)
if self._request_has_sibling_clarification(request):
return self._stub_sibling_tool(request)
return await handler(request)