334 lines
13 KiB
Python
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)
|