如何在自定义工具中用create_retrieval_chain替代RetrievalQA?
解决create_retrieval_chain替代RetrievalQA时的ValidationError问题
核心问题分析
create_retrieval_chain返回的是完整的链对象,而自定义工具通常要求传入可调用对象(如函数、带invoke方法的适配体)。直接将链实例传入工具会触发校验错误,因为工具无法直接识别链的调用格式。
正确实现步骤
1. 按官方规范构建检索链
先确保检索链本身构建正确,示例代码如下:
from langchain.chains import create_retrieval_chain from langchain.chains.combine_documents import create_stuff_documents_chain from langchain_core.prompts import ChatPromptTemplate from langchain_community.vectorstores import FAISS from langchain_openai import ChatOpenAI, OpenAIEmbeddings # 初始化基础组件 vectorstore = FAISS.from_texts(["测试文档内容"], embedding=OpenAIEmbeddings()) retriever = vectorstore.as_retriever() llm = ChatOpenAI() # 构建文档处理子链 prompt = ChatPromptTemplate.from_messages([ ("system", "仅根据提供的上下文回答问题:\n{context}"), ("human", "{input}") ]) document_chain = create_stuff_documents_chain(llm, prompt) # 构建完整检索链 retrieval_chain = create_retrieval_chain(retriever, document_chain)
2. 适配自定义工具的可调用要求
需要将检索链包装成工具能识别的格式,两种常用方式:
方式一:包装为函数
直接封装链的invoke方法,转换输入输出格式:
from langchain.tools import Tool def retrieval_tool(input_text): # 检索链要求输入为含"input"键的字典,提取返回结果中的"answer"字段 result = retrieval_chain.invoke({"input": input_text}) return result["answer"] # 定义自定义工具 custom_retrieval_tool = Tool( name="文档检索工具", func=retrieval_tool, description="检索相关文档并回答用户问题" )
方式二:用Runnable组件适配格式
借助LangChain的Runnable体系实现更灵活的格式转换:
from langchain.tools import Tool from langchain_core.runnables import RunnablePassthrough # 适配输入为字符串,输出直接返回answer字段 adapted_chain = ( RunnablePassthrough.assign(input=lambda x: x) | retrieval_chain | lambda output: output["answer"] ) # 传入工具 custom_retrieval_tool = Tool( name="文档检索工具", func=adapted_chain.invoke, description="检索相关文档并回答用户问题" )
常见错误排查点
- 确认工具的
func参数传入的是可调用对象(函数、带invoke方法的适配链),而非链实例本身 - 测试检索链独立运行:调用
retrieval_chain.invoke({"input": "测试问题"}),确认返回包含answer和context的字典 - 匹配输入输出格式:工具若要求字符串输入,必须将链的字典输入格式做转换
内容的提问来源于stack exchange,提问作者Skyward
相关产品推荐
相关产品推荐

