1、概述

本次项目是在以下链接基础上增加自定义工具相关代码,进行演示。其中自定义工具是参考百度地图开放平台中的天气情况API创作中心-CSDN https://mp.csdn.net/mp_blog/creation/editor/150424650

2、完整代码

from typing import Optional, Type
import streamlit as st
import tempfile
import os, sys, csv, requests
from pydantic import BaseModel, Field
from langchain.memory import ConversationBufferMemory
from langchain_community.chat_message_histories import StreamlitChatMessageHistory
from langchain_community.document_loaders import (
    PyPDFLoader,
    TextLoader,
    Docx2txtLoader,
)

# from langchain_openai import OpenAIEmbeddings
from langchain_community.embeddings import DashScopeEmbeddings
from langchain_chroma import Chroma
from langchain_core.prompts import PromptTemplate
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain.agents import create_react_agent, AgentExecutor
from langchain_community.callbacks.streamlit import StreamlitCallbackHandler
from langchain_deepseek import ChatDeepSeek
from dotenv import load_dotenv
from langchain_core.tools import BaseTool
from langchain_core.callbacks import CallbackManagerForToolRun

sys.stdout.reconfigure(encoding="utf-8")
# 加载环境变量
load_dotenv(override=True)
VECTOR_STORE = "vector_store"
# 设置Steamlit应用的页面标题和布局
st.set_page_config(page_title="文档助手", layout="wide")

# 设置应用的标题
st.title("文档助手")

# st.markdown("📁 **请上传文档(支持 pdf / docx / txt):**")
# 上传文件
upload_files = st.sidebar.file_uploader(
    label="📁 **请上传文档:**",
    type=["pdf", "docx", "txt"],
    accept_multiple_files=True,
    help="可以一次选择多个文件",  # 鼠标悬停时的中文提示
)

# 如果没有上传文件,提示用户上传文件并停止运行
if not upload_files:
    st.warning("请上传文件!")
    st.stop()


# 实现检索器
@st.cache_resource(ttl="1h")
def configure_retriever(upload_files):

    # 读取上传的文件,并写入一个临时目录

    docs = []
    temp_dir = tempfile.TemporaryDirectory(
        dir=r"D:\\PythonProject\\langgraph\\langchain_txt_rag\\"
    )
    for file in upload_files:
        temp_filepath = os.path.join(temp_dir.name, file.name)
        with open(temp_filepath, "wb") as f:
            f.write(file.getvalue())
            # 根据扩展名选择 loader
            suffix = file.name.lower().rsplit(".", 1)[-1]
            if suffix == "pdf":
                loader = PyPDFLoader(temp_filepath)
            elif suffix == "docx":
                loader = Docx2txtLoader(temp_filepath)
            elif suffix in {"txt", "md"}:
                loader = TextLoader(temp_filepath, encoding="utf-8")
            else:
                st.warning(f"跳过不支持的格式:{file.name}")
                continue
        # 使用TextLoader加载文件
        # loader = TextLoader(temp_filepath, encoding="utf-8")
        docs.extend(loader.load())
    # 添加检查逻辑
    if not docs:
        raise ValueError("文档列表为空,请检查 loader 或过滤逻辑")
    text_splitter = RecursiveCharacterTextSplitter(chunk_size=1000, chunk_overlap=200)
    chunks = text_splitter.split_documents(docs)

    # 使用openAI的向量模型生成文档的向量表示
    # embeddings = OpenAIEmbeddings()
    embeddings = DashScopeEmbeddings(
        model="text-embedding-v1", dashscope_api_key=os.environ["DASHSCOPE_API_KEY"]
    )
    if not embeddings:
        raise ValueError("向量初始化失败,请检查 API 密钥或模型初始化逻辑")
    vectordb = None
    try:
        vectordb = Chroma.from_documents(
            chunks, embeddings, persist_directory=VECTOR_STORE
        )
    except Exception as e:
        st.error(f"向量数据库创建失败: {str(e)}")
        st.write("请检查文档内容和API密钥")
        raise
    if not vectordb:
        raise ValueError("文档向量化持久化失败,请检查持久化目录或向量初始化逻辑")
    # 加载磁盘中的向量数据库
    # vector_store = Chroma(persist_directory=VECTOR_STORE,embedding_function=embeddings)

    # 创建文档检索器
    retriever = vectordb.as_retriever()

    return retriever


# 配置检索器
retriever = configure_retriever(upload_files)

# 如果seesion_state中没有消息记录或者用户点击了清空聊天记录按钮,则初始化消息记录
if "messages" not in st.session_state or st.sidebar.button("清空聊天记录"):
    st.session_state["messages"] = [
        {"role": "assistant", "content": "你好,我是文档助手."}
    ]

# 加载历史聊天记录
for message in st.session_state.messages:
    st.chat_message(message["role"]).write(message["content"])

# 创建检索工具
from langchain.tools.retriever import create_retriever_tool

# 创建用于文档检索工具
file_tool = create_retriever_tool(
    retriever, "文档检索", "用于检索用户的问题,并基于检索到的文档内容进行回复"
)


def find_district(csv_file_path, district_name) -> str:
    """
    根据区域的名称返回该对应区域的编码
    Args:
        csv_file_path (_type_): 文件路径
        district_name (_type_): 位置

    Returns:
        str: 位置编码
    """
    district_map = {}
    if not os.path.exists(csv_file_path):
        print(f"CSV文件不存在: {csv_file_path}")
        return None
    with open(csv_file_path, "r", encoding="utf-8") as f:
        csv_reader = csv.DictReader(f)
        for row in csv_reader:
            district_code = row["district_geocode"].strip()
            district = row["district"].strip()
            if district not in district_map:
                district_map[district] = district_code
    code = district_map.get(district_name.strip(), None)
    print(f"【{district_name}】查询结果code= {code}")
    return code


class WeatherSearchResults(BaseModel):
    """用于位置查询天气的信息

    Args:
        BaseModel (_type_): _description_
    """

    location: str = Field(description="用于位置查询天气的信息")


class WeatherSearchTool(BaseTool):
    """查询实时天气情况工具"""

    name: str = "weather_search"
    description: str = "可以查询任意位置置的实时天气情况"
    args_schema: Type[WeatherSearchResults] = WeatherSearchResults

    def _run(
        self,
        location: str,
        run_manager: Optional[CallbackManagerForToolRun] = None,
    ) -> str:
        """调用工具的时候,自动执行函数"""
        # 确保使用正确的CSV文件路径
        current_dir = os.path.dirname(os.path.abspath(__file__))
        csv_file_path = os.path.join(current_dir, "weather_district_id.csv")
        code = find_district(csv_file_path, location)
        if not code:
            error_msg = f"{location}的编码不存在"
            st.write(error_msg)
            raise ValueError(error_msg)

        print(f"需要查询的{location}的地区编码是{code}")
        CHINA_WEATHER_API_KEY = os.environ.get("CHINA_WEATHER_API_KEY")
        url = f"https://api.map.baidu.com/weather/v1/?district_id={code}&data_type=now&ak={CHINA_WEATHER_API_KEY}"
        response = requests.get(url)
        data = response.json()
        text = data["result"]["now"]["text"]
        temp = data["result"]["now"]["temp"]
        rh = data["result"]["now"]["rh"]
        feels_like = data["result"]["now"]["feels_like"]
        wind_dir = data["result"]["now"]["wind_dir"]
        wind_class = data["result"]["now"]["wind_class"]

        return f"位置:{location}。当前天气:{text},温度为{temp}度,相对湿度为{rh},体感温度为{feels_like}度,风向为{wind_dir},风力等级为{wind_class}"


weather_tool = WeatherSearchTool()

tools = [file_tool, weather_tool]

# 创建聊天消息历史记录
msgs = StreamlitChatMessageHistory()

# 创建对话缓存区内存
memory = ConversationBufferMemory(
    chat_memory=msgs,
    return_messages=True,
    memory_key="chat_history",
    output_key="output",
)

# 指令模版
instructions = """
你是一个全能的查询助手。
1、如果涉及到查询文档,你可以使用文档检索助手,并基于检索内容来回答问题。
2、如果涉及到查询天气,你需要准确提取地理位置名称,然后使用天气查询助手,并基于查询结果来回答问题。
3、如果涉及到查询数据库,你可以使用数据库查询助手,并基于查询结果来回答问题。
4、如果涉及到查询图片,你可以使用图片查询助手,并基于查询结果来回答问题。
5、如果涉及到查询代码,你可以使用代码查询助手,并基于查询结果来回答问题。
6、如果涉及到查询图表,你可以使用图表查询助手,并基于查询结果来回答问题。
你可能不查询就知道答案,但是你仍然应用查询来获得答案。
如果你从中找不到任何信息用于回答问题,则只需返回'抱歉,这个问题我不知道'。

重要提示:
- 当使用文档检索工具时,如果检索到相关信息,必须将完整的信息整合到最终回答中。
- 如果检索到人员信息(如姓名、手机号、身份证号、入委时间等),请完整地展示这些信息。
- 不要因为信息敏感而省略,要完整展示检索到的所有相关内容。
- 你可能不查询就知道答案,但是你仍然应该使用查询工具来获得准确答案。
- 当使用天气查询工具时,请确保只传递地理位置名称,不要包含其他词汇。例如:传递"北京"而不是"北京天气"或"北京市"。
回答格式要求:
- 对于人员信息,请按照清晰的格式列出每个人的详细信息。
- 保持信息的完整性和准确性。
"""


# 基础提示模版
base_prompt_template = """
{instructions}
TOOLS:
------
You have access to the following tools:
{tools}
To use a tool, please use the following format:
```
Thought: Do I need to use a tool? Yes
Action: the action to take, should be one of [{tool_names}]
Action Input: {input}
Observation: the result of the action
```
When you have a response to say to the human, or don't need to use a tool,
you MUST use the format:
```
Thought: Do I need to use a tool? No
Final Answer: [ your response here ]
```
Begin!
Previous conversation history:
{chat_history}

New input: {input}
{agent_scratchpad}
"""

# 创建基础提示模版
base_prompt = PromptTemplate.from_template(base_prompt_template)

# 创建部分填充的提示模版
prompt = base_prompt.partial(instructions=instructions)

# 创建llm
llm = ChatDeepSeek(model="deepseek-chat", api_key=os.getenv("DEEPSEEK_API_KEY"))
# 创建agent
agent = create_react_agent(llm, tools, prompt)

# 创建agent执行器
agent_executor = AgentExecutor(
    agent=agent,
    tools=tools,
    memory=memory,
    verbose=True,
    max_iterations=3,
    handle_parsing_errors=True,
    early_stopping_method="generate",
)


# 创建聊天输入框
user_query = st.chat_input(placeholder="请输入问题...")

# 如果有用户输入的查询
if user_query:
    # 添加用户消息到session_state
    st.session_state.messages.append({"role": "user", "content": user_query})
    # 显示用户信息
    st.chat_message("user").write(user_query)

    with st.chat_message("assistant"):
        # 创建Streamlit回调处理器
        st_cb = StreamlitCallbackHandler(st.container(), expand_new_thoughts=False)
        config = {"callbacks": [st_cb]}
        try:
            # 执行agent并获取响应
            response = agent_executor.invoke({"input": user_query}, config=config)
            st.write("DEBUG: Agent response:", response)  # 调试信息
            # 确保响应中包含output字段
            answer = ""
            if "output" in response and response["output"]:
                answer = response["output"]
            else:
                answer = "抱歉,我没有找到相关答案。"
            # 添加助手消息到session_state
            st.session_state.messages.append({"role": "assistant", "content": answer})
            # 显示助手响应
            st.write(answer)
        except Exception as e:
            error_msg = f"处理您的问题时出现错误: {str(e)}"
            st.session_state.messages.append(
                {"role": "assistant", "content": "抱歉,处理您的问题时出现错误。"}
            )
            st.write("抱歉,处理您的问题时出现错误:", {str(e)})
            import traceback

            st.text(traceback.format_exc())
            print(error_msg)

3、运行

steamlit run xx.py

4、浏览器访问

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐