Langchain中SQLFilterTool为何收到意外参数arg1?
问题原因及解决方法
问题原因
LangChain的Agent在调用工具时,默认会将输入参数以arg1作为键名传递(这是多数Agent默认提示模板里的变量名),但你的SQLFilterTool的_run和_arun方法定义的参数名为query_input,参数名称不匹配,导致抛出unexpected keyword argument 'arg1'错误。
你之前尝试添加arg1到参数列表的方式有误(比如写成_run(self, query_input: str, arg1: str)),这会让方法同时期望query_input和arg1两个参数,但Agent只会传递arg1,所以仍会报错。
解决方法
有两种直接有效的解决方式:
方式一:修改方法参数名为arg1
直接把_run和_arun的参数名改成arg1,和Agent传递的键名匹配:
from langchain.tools import BaseTool class SQLFilterTool(BaseTool): name = "filter_user_query" description = SQL_FILTER_TOOL def __init__(self): super().__init__() def _run(self, arg1: str): return sql_filter_vanna(arg1) def _arun(self, arg1: str): return sql_filter_vanna(arg1)
方式二:通过args_schema指定自定义参数名
如果想保留query_input作为参数名,可以用Pydantic模型定义工具的参数schema,让Agent按照你指定的参数名传递:
from langchain.tools import BaseTool from pydantic import BaseModel, Field # 定义输入参数的schema class SQLFilterToolInput(BaseModel): query_input: str = Field(description="需要过滤的用户查询语句") class SQLFilterTool(BaseTool): name = "filter_user_query" description = SQL_FILTER_TOOL args_schema = SQLFilterToolInput # 绑定参数schema def __init__(self): super().__init__() def _run(self, query_input: str): return sql_filter_vanna(query_input) def _arun(self, query_input: str): return sql_filter_vanna(query_input)
这样Agent会自动识别query_input作为参数键名,不再传递arg1。
内容的提问来源于stack exchange,提问作者lauther27
相关产品推荐
相关产品推荐

