react_agent.py 6.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174
  1. """
  2. 数据库Agent - 基于ReAct框架的智能数据库查询助手
  3. """
  4. import re
  5. from typing import Optional, List
  6. from hello_agents import ReActAgent, HelloAgentsLLM, Config, Message, ToolRegistry
  7. from tools import OracleQueryTool, SQLGeneratorTool, format_query_result
  8. from config import DatabaseConfig
  9. DATABASE_AGENT_PROMPT = """你是一个专业的数据库查询助手。你可以理解用户的自然语言查询,将其转换为SQL语句,从Oracle数据库中获取数据并格式化输出。
  10. ## 可用工具
  11. {tools}
  12. ## 工作流程
  13. 请严格按照以下格式进行回应:
  14. Thought: 你的思考过程,分析用户需求并规划下一步行动。
  15. Action: 你决定采取的行动,必须是以下格式之一:
  16. - `{{tool_name}}[{{tool_input}}]` - 调用指定工具
  17. - `Finish[最终答案]` - 当你有足够信息给出最终答案时
  18. ## 使用指南
  19. 1. 当用户提出查询需求时,首先使用 GetSchema 工具获取数据库表结构
  20. 2. 使用 GenerateSQL 工具将自然语言转换为SQL语句
  21. 3. 使用 ExecuteQuery 工具执行SQL并获取结果
  22. ## 当前任务
  23. **Question:** {question}
  24. ## 执行历史
  25. {history}
  26. 现在开始你的推理和行动:
  27. """
  28. class DatabaseAgent(ReActAgent):
  29. """数据库查询Agent"""
  30. def __init__(
  31. self,
  32. name: str,
  33. llm: HelloAgentsLLM,
  34. db_config: DatabaseConfig,
  35. system_prompt: Optional[str] = None,
  36. config: Optional[Config] = None,
  37. max_steps: int = 5
  38. ):
  39. super().__init__(name, llm, system_prompt, config)
  40. self.db_config = db_config
  41. self.max_steps = max_steps
  42. self.current_history: List[str] = []
  43. self.prompt_template = DATABASE_AGENT_PROMPT
  44. self.oracle_tool = OracleQueryTool(db_config)
  45. self.sql_generator = SQLGeneratorTool(llm)
  46. self.tool_registry = ToolRegistry()
  47. self.tool_registry.register_function(
  48. "GetSchema",
  49. "获取数据库表结构信息,包括所有表名和字段定义。",
  50. self._get_schema
  51. )
  52. self.tool_registry.register_function(
  53. "GenerateSQL",
  54. "将自然语言查询转换为Oracle SQL语句。",
  55. self._generate_sql
  56. )
  57. self.tool_registry.register_function(
  58. "ExecuteQuery",
  59. "执行SQL查询并返回结果。",
  60. self._execute_query
  61. )
  62. self.schema_cache = None
  63. print(f"✅ {name} 初始化完成,最大步数: {max_steps}")
  64. def _get_schema(self, input_text: str) -> str:
  65. """获取数据库表结构信息,包括所有表名和字段定义"""
  66. schema_info = self.oracle_tool.get_schema_info()
  67. self.schema_cache = schema_info
  68. return schema_info
  69. def _generate_sql(self, input_text: str) -> str:
  70. """将自然语言查询转换为Oracle SQL语句"""
  71. if not self.schema_cache:
  72. self.schema_cache = self.oracle_tool.get_schema_info()
  73. sql = self.sql_generator.generate_sql(input_text, self.schema_cache)
  74. is_valid, msg = self.sql_generator.validate_sql(sql)
  75. if not is_valid:
  76. return f"SQL生成失败: {msg}"
  77. return f"生成的SQL: {sql}"
  78. def _execute_query(self, input_text: str) -> str:
  79. """执行SQL查询并返回结果"""
  80. sql = input_text.strip()
  81. if sql.startswith("生成的SQL: "):
  82. sql = sql.replace("生成的SQL: ", "")
  83. result = self.oracle_tool.execute_query(sql)
  84. if not result["success"]:
  85. return f"查询执行失败: {result['error']}"
  86. formatted_result = format_query_result(result)
  87. return formatted_result
  88. def run(self, input_text: str, **kwargs) -> str:
  89. """运行数据库Agent"""
  90. self.current_history = []
  91. current_step = 0
  92. print(f"\n🤖 {self.name} 开始处理问题: {input_text}")
  93. while current_step < self.max_steps:
  94. current_step += 1
  95. print(f"\n--- 第 {current_step} 步 ---")
  96. # 1. 构建提示词
  97. tools_desc = self.tool_registry.get_tools_description()
  98. history_str = "\n".join(self.current_history)
  99. prompt = self.prompt_template.format(
  100. tools=tools_desc,
  101. question=input_text,
  102. history=history_str
  103. )
  104. # 2. 调用LLM
  105. messages = [{"role": "user", "content": prompt}]
  106. response_text = self.llm.invoke(messages, **kwargs)
  107. # 3. 解析输出
  108. thought, action = self._parse_output(response_text)
  109. if thought:
  110. print(f"🤔 思考: {thought}")
  111. if action and action.startswith("Finish"):
  112. final_answer = self._parse_action_input(action)
  113. self.add_message(Message(input_text, "user"))
  114. self.add_message(Message(final_answer, "assistant"))
  115. return final_answer
  116. if action:
  117. tool_name, tool_input = self._parse_action(action)
  118. observation = self.tool_registry.execute_tool(tool_name, tool_input)
  119. print(f"🎬 行动: {tool_name}[{tool_input}]")
  120. print(f"👀 观察: {observation}")
  121. self.current_history.append(f"Action: {action}")
  122. self.current_history.append(f"Observation: {observation}")
  123. final_answer = "抱歉,我无法在限定步数内完成这个任务。"
  124. self.add_message(Message(input_text, "user"))
  125. self.add_message(Message(final_answer, "assistant"))
  126. return final_answer
  127. def _parse_output(self, text: str):
  128. thought_match = re.search(r"Thought:\s*(.*?)(?=\nAction:|$)", text, re.DOTALL)
  129. action_match = re.search(r"Action:\s*(.*?)$", text, re.DOTALL)
  130. thought = thought_match.group(1).strip() if thought_match else None
  131. action = action_match.group(1).strip() if action_match else None
  132. return thought, action
  133. def _parse_action(self, action_text: str):
  134. match = re.match(r"(\w+)\[(.*)\]", action_text, re.DOTALL)
  135. return (match.group(1), match.group(2)) if match else (None, None)
  136. def _parse_action_input(self, action_text: str):
  137. match = re.match(r"\w+\[(.*)\]", action_text, re.DOTALL)
  138. return match.group(1) if match else ""