"""ReAct Engine — the core Think → Act → Observe loop.""" import asyncio import json import time import uuid from typing import AsyncIterator, Optional from agentic_rag.data.models import ( AgentEvent, AgentEventType, AgentInput, AgentOutput, LLMChunk, LLMResponse, Message, ToolCall, ToolCallResult, ToolDefinition, ) from agentic_rag.agent.react_parser import ( ReActStep, extract_final_answer, format_observation, parse_react_output, ) from agentic_rag.agent.react_prompt import build_react_prompt, build_tools_description from agentic_rag.services.llm.base import ( BaseLLMProvider, ReasoningStreamFilter, strip_reasoning, ) class ReActEngine: """ReAct (Reasoning + Acting) reasoning engine. Executes the Think → Act → Observe loop that powers the agent. """ def __init__( self, llm: BaseLLMProvider, tools: list, system_prompt_template: str, max_iterations: int = 10, stop_on_error: bool = False, enable_native_tool_calls: bool = True, require_tool_call: bool = False, ): """ Args: llm: LLM provider for generation. tools: List of BaseTool instances available to the agent. system_prompt_template: Template string for the system prompt. max_iterations: Maximum ReAct loop iterations. stop_on_error: If True, stop on first tool error. enable_native_tool_calls: If False, tools are NOT sent to the LLM (pure ReAct text mode for models without function calling). require_tool_call: If True, reject a final answer until at least one available tool has completed successfully. Used for freshness- sensitive queries where model memory is not acceptable evidence. """ self.llm = llm self.tools = tools self.system_prompt_template = system_prompt_template self.max_iterations = max_iterations self.stop_on_error = stop_on_error self.enable_native_tool_calls = enable_native_tool_calls self.require_tool_call = require_tool_call self._tool_map = {t.name: t for t in tools} async def run(self, input: AgentInput, turn_id: str = "") -> AgentOutput: """Execute the ReAct loop (non-streaming).""" if not turn_id: turn_id = uuid.uuid4().hex messages = self._build_initial_messages(input) tool_calls_made: list[ToolCallResult] = [] total_usage: dict[str, int] = {"prompt_tokens": 0, "completion_tokens": 0} invalid_output_count = 0 # Track consecutive invalid outputs max_invalid_attempts = 2 # Force final answer after 2 consecutive invalid outputs executed_sigs: set[str] = set() # (tool, args) dedupe — mirrors stream() for iteration in range(self.max_iterations): # Contract for the NEXT turn: a native function-call round ends with # "continue", so if no new tool result was appended after this point # the model must answer in text ReAct format (never emit format-only # text alongside a function call, or the final tool result gets # ignored and the last observation is lost). n_msgs_before = len(messages) messages.append(Message.user( "【重要】如果你决定调用工具,请只发起工具调用,不要同时输出任何 Thought/Action 格式行;" "如果你不调用工具,请按格式输出:Thought: ... Final Answer: ...(或 Thought/Action/Action Input 发起文本式工具调用)。" )) response = await self.llm.agenerate(messages, self._get_llm_tool_definitions()) if len(messages) > n_msgs_before: messages.pop() # remove the contract — don't pollute history total_usage["prompt_tokens"] += response.usage.get("prompt_tokens", 0) total_usage["completion_tokens"] += response.usage.get("completion_tokens", 0) # Handle native tool calls (from providers that support function calling) if response.tool_calls: invalid_output_count = 0 # Reset on valid tool call for tc in response.tool_calls: args = tc.arguments or {} if not args: messages.append(Message.tool( content=f"Error: '{tc.name}' 缺少参数", tool_call_id=tc.id, )) continue sig = (tc.name, json.dumps(args, sort_keys=True, ensure_ascii=False)) if sig in executed_sigs: messages.append(Message.tool( content=f"Error: 禁止重复调用 '{tc.name}'(参数相同)。请基于已有 Observation 直接输出 Final Answer。", tool_call_id=tc.id, )) continue executed_sigs.add(sig) result = await self._execute_tool(tc.name, args) tool_calls_made.append(result) messages.append(Message.tool( content=str(result.result) if not result.error else f"Error: {result.error}", tool_call_id=tc.id, )) continue # Parse ReAct format output step = parse_react_output(response.content, list(self._tool_map.keys())) if step.is_final: if self.require_tool_call and not self._has_successful_tool_result(tool_calls_made): invalid_output_count += 1 messages.append(Message.assistant(strip_reasoning(response.content))) messages.append(Message.user( "该问题必须先调用可用的网络搜索工具并获得成功的 Observation。" "不得依据模型记忆直接回答,也不得自行生成引用。请立即调用工具。" )) continue return AgentOutput( messages=messages, final_answer=step.final_answer, tool_calls_made=tool_calls_made, usage=total_usage, iterations=iteration + 1, ) if step.action: invalid_output_count = 0 # Reset on valid action # Guard: skip if action_input is empty (broken parse) if not step.action_input: messages.append(Message.user( f"'{step.action}' 需要参数,请在 Action Input 中提供 JSON。" )) continue sig = (step.action, json.dumps(step.action_input, sort_keys=True, ensure_ascii=False)) if sig in executed_sigs: messages.append(Message.user( f"你已经用相同参数调用过 '{step.action}' 了。请基于已有 Observation 直接输出 Final Answer。" )) continue executed_sigs.add(sig) # Execute the tool result = await self._execute_tool(step.action, step.action_input) tool_calls_made.append(result) # Format observation and append to messages observation = format_observation( step.action, str(result.result) if result.result else "", result.error, ) # Append the assistant's ReAct output + observation as a single message messages.append(Message.assistant(response.content)) messages.append(Message.user(observation)) else: # No valid action and no final answer — LLM output is unparseable invalid_output_count += 1 if invalid_output_count >= max_invalid_attempts: if self.require_tool_call and not self._has_successful_tool_result(tool_calls_made): return AgentOutput( messages=messages, final_answer=self._live_search_failure_answer(), tool_calls_made=tool_calls_made, usage=total_usage, iterations=iteration + 1, ) # Force exit after repeated invalid outputs to prevent infinite loop final = extract_final_answer(response.content) if not final: # Use accumulated content as fallback answer final = response.content.strip() or "抱歉,我暂时无法回答这个问题。" return AgentOutput( messages=messages, final_answer=final, tool_calls_made=tool_calls_made, usage=total_usage, iterations=iteration + 1, ) # Prompt LLM to continue with correct format messages.append(Message.user( "你的输出格式不正确。请严格按照以下格式之一输出:\n" "1. 调用工具:Thought: ...\nAction: tool_name\nAction Input: {\"param\": \"value\"}\n" "2. 最终答案:Thought: ...\nFinal Answer: ..." )) # Max iterations reached if self.require_tool_call and not self._has_successful_tool_result(tool_calls_made): final = self._live_search_failure_answer() else: final = await self._force_final_answer(messages) return AgentOutput( messages=messages, final_answer=final, tool_calls_made=tool_calls_made, usage=total_usage, iterations=self.max_iterations, ) async def stream(self, input: AgentInput, turn_id: str = "") -> AsyncIterator[AgentEvent]: """Execute the ReAct loop with streaming events.""" if not turn_id: turn_id = uuid.uuid4().hex messages = self._build_initial_messages(input) tool_calls_made: list[ToolCallResult] = [] invalid_output_count = 0 # Track consecutive invalid outputs max_invalid_attempts = 2 # Force final answer after 2 consecutive invalid outputs executed_sigs: set[tuple[str, str]] = set() for iteration in range(self.max_iterations): # Emit thought event yield AgentEvent( event_type=AgentEventType.THOUGHT, data={"iteration": iteration}, turn_id=turn_id, ) # Stream LLM generation — collect both text deltas AND native tool calls. full_content = "" # Native function-calling state (accumulated across streaming chunks) native_tool_name = "" native_tool_args = "" has_native_tool_call = False stream_error = None # Incremental reasoning stripper — thinking models ( blocks) # must not leak their reasoning to the UI or into the ReAct parser. think_filter = ReasoningStreamFilter() # Answer gating: only stream the actual answer text to the UI. # "Thought:" lines and any pre-answer monologue are buffered and # dropped — the final answer is extracted from full_content below. pending_delta = "" answer_streaming = False try: async for chunk in self.llm.agenerate_stream(messages, self._get_llm_tool_definitions()): if chunk.content_delta: delta = think_filter.feed(chunk.content_delta) full_content += delta # After tool call detection, further text is likely # post-tool narration — suppress to reduce noise. if not delta or has_native_tool_call: continue if answer_streaming: emit = delta else: pending_delta += delta if "Final Answer:" in pending_delta: emit = pending_delta.split("Final Answer:", 1)[1].lstrip() answer_streaming = True else: continue if emit: yield AgentEvent( event_type=AgentEventType.TEXT_DELTA, data={"content": emit}, turn_id=turn_id, ) # Collect native function-call deltas (Qwen / OpenAI function calling) if chunk.tool_call_delta: has_native_tool_call = True if chunk.tool_call_delta.get("name"): native_tool_name = chunk.tool_call_delta["name"] if chunk.tool_call_delta.get("arguments"): native_tool_args += chunk.tool_call_delta["arguments"] # Flush text held back for partial-tag detection tail = think_filter.flush() if tail: full_content += tail if not has_native_tool_call: if answer_streaming: yield AgentEvent( event_type=AgentEventType.TEXT_DELTA, data={"content": tail}, turn_id=turn_id, ) else: pending_delta += tail if "Final Answer:" in pending_delta: emit = pending_delta.split("Final Answer:", 1)[1].lstrip() if emit: answer_streaming = True yield AgentEvent( event_type=AgentEventType.TEXT_DELTA, data={"content": emit}, turn_id=turn_id, ) except Exception as e: stream_error = str(e) import sys print(f" [ReAct] ⚠ LLM stream error (iteration {iteration}): {e}", flush=True) sys.stdout.flush() # ── Handle stream error ── if stream_error and not full_content.strip(): yield AgentEvent( event_type=AgentEventType.ERROR, data={"error": f"LLM stream failed: {stream_error}"}, turn_id=turn_id, ) yield AgentEvent( event_type=AgentEventType.DONE, data={"final_answer": f"抱歉,模型服务连接中断:{stream_error},请稍后重试。"}, turn_id=turn_id, ) return # ── Resolve action: native function calling takes priority ── if has_native_tool_call and native_tool_name: invalid_output_count = 0 # Reset on valid tool call try: action_input = json.loads(native_tool_args) if native_tool_args else {} except json.JSONDecodeError: action_input = {"query": native_tool_args} if native_tool_args else {} sig = (native_tool_name, json.dumps(action_input, sort_keys=True, ensure_ascii=False)) if sig in executed_sigs: messages.append(Message.user( f"你已经用相同参数调用过 '{native_tool_name}'。请根据已有 Observation 输出 Final Answer。" )) continue executed_sigs.add(sig) yield AgentEvent( event_type=AgentEventType.TOOL_CALL_START, data={"tool": native_tool_name, "input": action_input}, turn_id=turn_id, ) result = await self._execute_tool(native_tool_name, action_input) tool_calls_made.append(result) yield AgentEvent( event_type=AgentEventType.TOOL_CALL_RESULT, data={ "tool": native_tool_name, "success": not result.error, "result": str(result.result)[:500] if result.result else "", "error": result.error, }, turn_id=turn_id, ) observation = format_observation( native_tool_name, str(result.result) if result.result else "", result.error, ) messages.append(Message.assistant( strip_reasoning(full_content).strip() or f"Thought: 调用 {native_tool_name}\nAction: {native_tool_name}\nAction Input: {json.dumps(action_input, ensure_ascii=False)}" )) messages.append(Message.user(observation)) continue # ── Fallback: parse ReAct text format ── step = parse_react_output(full_content, list(self._tool_map.keys())) if step.is_final: if self.require_tool_call and not self._has_successful_tool_result(tool_calls_made): invalid_output_count += 1 messages.append(Message.assistant(strip_reasoning(full_content))) messages.append(Message.user( "该问题必须先调用可用的网络搜索工具并获得成功的 Observation。" "不得依据模型记忆直接回答,也不得自行生成引用。请立即调用工具。" )) continue yield AgentEvent( event_type=AgentEventType.DONE, data={"final_answer": step.final_answer, "iterations": iteration + 1}, turn_id=turn_id, ) return if step.action: invalid_output_count = 0 # Reset on valid action # Guard: if action_input is empty and the tool requires args, skip if not step.action_input: messages.append(Message.user( f"'{step.action}' 需要参数,请在 Action Input 中提供 JSON。" )) continue sig = (step.action, json.dumps(step.action_input, sort_keys=True, ensure_ascii=False)) if sig in executed_sigs: messages.append(Message.user( f"你已经用相同参数调用过 '{step.action}'。请根据已有 Observation 输出 Final Answer。" )) continue executed_sigs.add(sig) yield AgentEvent( event_type=AgentEventType.TOOL_CALL_START, data={"tool": step.action, "input": step.action_input}, turn_id=turn_id, ) result = await self._execute_tool(step.action, step.action_input) tool_calls_made.append(result) yield AgentEvent( event_type=AgentEventType.TOOL_CALL_RESULT, data={ "tool": step.action, "success": not result.error, "result": str(result.result)[:500] if result.result else "", "error": result.error, }, turn_id=turn_id, ) observation = format_observation( step.action, str(result.result) if result.result else "", result.error, ) messages.append(Message.assistant(strip_reasoning(full_content))) messages.append(Message.user(observation)) else: # No valid action and no final answer — LLM output is unparseable invalid_output_count += 1 if invalid_output_count >= max_invalid_attempts: if self.require_tool_call and not self._has_successful_tool_result(tool_calls_made): yield AgentEvent( event_type=AgentEventType.DONE, data={"final_answer": self._live_search_failure_answer(), "iterations": iteration + 1}, turn_id=turn_id, ) return # Force exit after repeated invalid outputs to prevent infinite loop final = extract_final_answer(full_content) if not final: final = await self._force_final_answer(messages) if not final or ("Thought:" in final and "Action" in final): final = full_content.strip() yield AgentEvent( event_type=AgentEventType.DONE, data={"final_answer": final or "抱歉,我暂时无法回答这个问题。", "iterations": iteration + 1}, turn_id=turn_id, ) return messages.append(Message.user( "你的输出格式不正确。请严格按照以下格式之一输出:\n" "1. 调用工具:Thought: ...\nAction: tool_name\nAction Input: {\"param\": \"value\"}\n" "2. 最终答案:Thought: ...\nFinal Answer: ..." )) # Max iterations — force LLM to give a final answer with what it has if self.require_tool_call and not self._has_successful_tool_result(tool_calls_made): final = self._live_search_failure_answer() else: final = await self._force_final_answer(messages) yield AgentEvent( event_type=AgentEventType.DONE, data={"final_answer": final or "Max iterations reached. Could not complete the task.", "iterations": self.max_iterations}, turn_id=turn_id, ) def _build_initial_messages(self, input: AgentInput) -> list[Message]: """Build the initial message list for the ReAct loop.""" tools_desc = build_tools_description(self._get_tool_definitions()) memory_context = self._format_messages(input.messages) system_prompt = build_react_prompt( tools_description=tools_desc, memory_context=memory_context, ) messages = [Message.system(system_prompt)] # Check if input already has a multimodal user message has_multimodal_query = any( isinstance(m.content, list) and m.role.value == "user" for m in input.messages ) # Add conversation history (excluding system messages) for msg in input.messages: if msg.role.value != "system": messages.append(msg) # Add current query — skip if already included as multimodal message if not has_multimodal_query: query = input.query if input.multimodal and input.multimodal.text: query = input.multimodal.text messages.append(Message.user(query)) return messages async def _execute_tool(self, name: str, arguments: dict) -> ToolCallResult: """Execute a tool by name with arguments. Falls back to the global tool registry when the tool isn't in this engine's filtered set: the model sometimes emits a tool it knows from context (e.g. an MCP tool) that the router didn't hand it, and answering "tool not found" wastes a turn when the tool actually exists in the process. """ call_id = f"call_{uuid.uuid4().hex[:12]}" tool = self._tool_map.get(name) if tool is None: try: from agentic_rag.orchestration.l1_tools.registry import get_tool_registry tool = get_tool_registry().get(name) except Exception: tool = None if tool is None: return ToolCallResult( call_id=call_id, name=name, result=None, error=f"Tool '{name}' not found. Available: {list(self._tool_map.keys())}", ) try: # Add timeout to prevent hanging on slow tools result = await asyncio.wait_for(tool.execute(**arguments), timeout=60.0) return ToolCallResult( call_id=call_id, name=name, result=result, ) except asyncio.TimeoutError: return ToolCallResult( call_id=call_id, name=name, result=None, error=f"Tool '{name}' execution timed out after 60 seconds", ) except Exception as e: return ToolCallResult( call_id=call_id, name=name, result=None, error=str(e), ) @staticmethod def _has_successful_tool_result(results: list[ToolCallResult]) -> bool: """Return whether at least one tool produced usable evidence.""" return any(not result.error and result.result for result in results) @staticmethod def _live_search_failure_answer() -> str: """Fail closed when fresh information could not be retrieved.""" return ( "当前问题需要实时网络信息,但本次未能获得有效的搜索结果," "因此无法可靠确认。请稍后重试;为避免误导,我不会使用模型记忆猜测答案或编造来源。" ) def _get_tool_definitions(self) -> list[ToolDefinition]: """Get all tool definitions for prompt construction and parsing.""" return [t.to_definition() for t in self.tools] def _get_llm_tool_definitions(self) -> list[ToolDefinition]: """Get definitions passed through the provider's native tools API.""" if not self.enable_native_tool_calls: return [] return self._get_tool_definitions() @staticmethod def _format_messages(messages: list[Message]) -> str: """Format conversation history for the prompt.""" if not messages: return "" lines = [] for msg in messages[-10:]: # Last 10 messages content = msg.content if isinstance(content, list): # Extract text parts for history summary texts = [p.get("text", "") for p in content if isinstance(p, dict) and p.get("text")] img_count = sum(1 for p in content if isinstance(p, dict) and p.get("type") == "image_url") parts = texts if img_count: parts.append(f"[{img_count} image(s)]") content = " ".join(parts) if parts else "[multimodal content]" lines.append(f"{msg.role.value}: {str(content)[:200]}") return "\n".join(lines) async def _force_final_answer(self, messages: list[Message]) -> str: """Force the LLM to produce a final answer when max iterations are reached.""" messages.append(Message.user( "已达到最大步数。请基于已有信息给出 Final Answer。" )) response = await self.llm.agenerate(messages) final = extract_final_answer(response.content) return final or response.content