sql_agent.py 1.6 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950
  1. """
  2. NL2SQL Agent:将自然语言转换为 SQL 查询,自动执行并返回分析结果。
  3. """
  4. import os
  5. import sys
  6. # 确保父目录可导入(支持直接运行 python agents/sql_agent.py)
  7. sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
  8. from langchain_community.agent_toolkits import SQLDatabaseToolkit
  9. from langchain.agents import create_agent
  10. from config import llm, db
  11. # ---- 创建 SQL 工具包 ----
  12. toolkit = SQLDatabaseToolkit(db=db, llm=llm)
  13. tools = toolkit.get_tools()
  14. # ---- System Prompt ----
  15. SQL_AGENT_PROMPT = """你是一名专业的 SQL 数据分析师。
  16. ## 工作流程
  17. 1. 先用 sql_db_list_tables 查看数据库中有哪些表
  18. 2. 用 sql_db_schema 获取相关表的字段结构和类型
  19. 3. 生成 SQL 之前,用 sql_db_query_checker 检查语法
  20. 4. 确认无误后,用 sql_db_query 执行查询
  21. 5. 用中文总结查询结果,给出简洁的业务洞察
  22. ## 约束
  23. - 只使用数据库中实际存在的表和字段,不要凭空编造
  24. - 单次查询结果限制在 50 条以内
  25. - 如果查询出错,分析错误原因后重新生成 SQL
  26. - 回答要简洁专业,不要啰嗦
  27. """
  28. # ---- 创建 Agent ----
  29. sql_agent = create_agent(
  30. model=llm,
  31. tools=tools,
  32. system_prompt=SQL_AGENT_PROMPT,
  33. )
  34. if __name__ == "__main__":
  35. print(f"📊 数据库:{db}")
  36. print(f" 可用表:{db.get_usable_table_names()}")
  37. print(f"\nSQL 工具包已加载({len(tools)} 个工具):")
  38. for t in tools:
  39. print(f" - {t.name}")
  40. print("\n✅ NL2SQL Agent 创建完成,可以开始提问了!")