"""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)