如何从Langchain的create_sql_agent中提取生成的SQL查询?
提取LangChain SQL Agent生成的SQL查询方法
针对你使用的openai-tools类型SQL Agent,以下是两种可靠的SQL提取方案:
方法一:自定义回调函数捕获SQL(推荐)
通过自定义回调处理器,在工具调用阶段直接拦截生成的SQL查询:
from langchain.callbacks.base import BaseCallbackHandler class SQLCaptureCallback(BaseCallbackHandler): def __init__(self): self.generated_sql = None def on_tool_start(self, serialized: dict, input_str: str, **kwargs) -> None: # 匹配SQL数据库查询工具的调用 if serialized.get("name") == "sql_db_query": self.generated_sql = input_str # 初始化回调实例 sql_callback = SQLCaptureCallback() # 创建Agent时传入回调 agent = create_sql_agent( llm=llm, db=db, verbose=True, agent_type="openai-tools", callbacks=[sql_callback] ) # 执行查询 response = agent.invoke({"input": "How many resources are there in XYZ location?"}) # 获取生成的SQL print("生成的SQL查询:", sql_callback.generated_sql)
方法二:解析Verbose输出(调试用)
如果你开启了verbose=True,可以通过捕获标准输出来提取SQL:
import io import sys # 捕获Agent的verbose输出 output_capture = io.StringIO() sys.stdout = output_capture # 执行查询 response = agent.invoke({"input": "How many resources are there in XYZ location?"}) # 恢复标准输出 sys.stdout = sys.__stdout__ # 从输出中提取SQL(根据实际输出格式调整匹配逻辑) full_output = output_capture.getvalue() for line in full_output.split("\n"): if "SQLQuery:" in line or "Executing SQL:" in line: extracted_sql = line.split(":", 1)[1].strip() print("生成的SQL查询:", extracted_sql) break
注意:方法二依赖于verbose输出的格式,LangChain版本更新可能会导致匹配逻辑失效,因此方法一的回调方式更稳定。
内容的提问来源于stack exchange,提问作者AlfiyaFaisy
相关产品推荐
相关产品推荐

