开头

大家好,这里是惬鹤频道!

那么,随着agent项目的学习进入了第三天,这次要给大家分享的是关于agent部分代码的编写,
今天的代码不算很难,很快就可以理解,那让我们开始吧!

如果你觉得这篇文章有用,请点点赞谢谢,可以的话关注我吧,我会努力更新的!

总会有的。。。项目结构图

在这里插入图片描述

代码文件:rag_service.py

"""
文件名:rag_service.py
描述:用于总结服务
"""
# 依赖文件导入
from rag.vector_store import VectorStoreService
from utils.prompt_loader import load_rag_prompts
from model.factory import chat_model

# 依赖导入
from langchain_core.prompts import PromptTemplate
from langchain_core.output_parsers import StrOutputParser
from langchain_core.documents import Document

# 打印提示词
def print_prompt(prompt):
    print("="*20)
    print(prompt.to_string())
    print("="*20)
    return prompt

# 定义RAG总结服务类
class RagSummarizeService(object):
    def __init__(self):
        self.vector_store = VectorStoreService()
        self.retriever = self.vector_store.get_retriever()
        self.prompt_text = load_rag_prompts()
        self.prompt_template = PromptTemplate.from_template(self.prompt_text)
        self.model = chat_model
        self.chain = self.__init__chain()

    # 初始化链
    def __init__chain(self):
        chain = self.prompt_template | print_prompt | self.model | StrOutputParser()
        return chain

    # 得到检索问题相关资料的结果(相关参数如k已经设置在配置文件中)
    def retriever_docs(self, query: str) -> list[Document]:
        return self.retriever.invoke(query)

    # 将用户问题转换为可插入提示词的一部分,并返回调用结果
    def rag_summarize(self, query: str) -> str:
        # 拿到Document类型的列表
        context_docs = self.retriever_docs(query)
        counter = 0
        context = ""
        for doc in context_docs:
            counter += 1
            context += f"[参考资料{counter}] 资料内容:{doc.page_content} | 参考源数据:{doc.metadata}\n"

        # 返回链调用的结果
        return self.chain.invoke(
            {
                "input": query,
                "context": context,
            }
        )

# 测试
if __name__ == '__main__':
    rag = RagSummarizeService()
    print(rag.rag_summarize("小户型适合哪种扫地机器人?"))

这是一个 RAG 总结服务类,它通过向量检索找到与问题相关的文档片段,将它们拼接成上下文,再交给大语言模型生成最终回答。

类内部封装了检索器、提示词模板和 LCEL 链,对外提供 rag_summarize(query) 方法即可完成从问题到答案的完整流程。

这段代码的结构分为以下几个层次:

模块导入部分

从项目内部模块导入

rag.vector_store.VectorStoreService:向量库服务,负责文档检索

utils.prompt_loader.load_rag_prompts:加载提示词模板

model.factory.chat_model:聊天模型实例

从 LangChain 导入

PromptTemplate:提示词模板

StrOutputParser:输出解析器(将模型输出转为字符串)

Document:文档类型注解

辅助函数

print_prompt(prompt):打印完整的提示词内容,用于调试,并原样返回 prompt

核心类 RagSummarizeService

内部有如下几个方法:

初始化 init
创建向量库服务实例 self.vector_store

获取检索器 self.retriever(用于根据问题检索相关文档)

加载提示词文本 self.prompt_text

创建提示词模板 self.prompt_template

引用全局聊天模型 self.model

调用私有方法 __init__chain() 初始化 LCEL 链

私有方法 __init__chain
构造链:prompt_template | print_prompt | model | StrOutputParser()

返回构建好的链(支持 invoke)

检索方法 retriever_docs(query)
调用检索器的 invoke 方法,返回相关文档列表(list[Document])

核心方法 rag_summarize(query)
调用 retriever_docs(query) 获取相关文档列表

遍历文档,格式化为带编号和元数据的上下文字符串 context
调用链的 invoke,传入 {“input”: query, “context”: context}
返回最终的回答字符串

整体的结构如下:
RagSummarizeService
├── init
│ ├── 初始化 VectorStoreService
│ ├── 获取 retriever
│ ├── 加载 prompt 文本
│ ├── 创建 PromptTemplate
│ ├── 引用 chat_model
│ └── 调用 __init__chain 创建 chain
├── retriever_docs(query) → list[Document]
├── rag_summarize(query) → str
│ ├── 调用 retriever_docs 获取相关文档
│ ├── 将文档格式化为 context 字符串
│ └── chain.invoke({“input”: query, “context”: context})
└── (内置 print_prompt 辅助函数)

代码文件:agent_tools.py

"""
文件名:agent_tools
描述:智能体需要使用的各种方法
"""
import os

# 依赖导入
from langchain_core.tools import tool
import random
# 依赖文件导入
from rag.rag_service import RagSummarizeService
from utils.config_handler import agent_conf
from utils.path_tool import get_abs_path
from utils.logger_handler import logger

rag = RagSummarizeService()

user_ids = ["1001", "1002", "1003", "1004", "1005", "1006", "1007", "1008", "1009", "1010",]
month_arr = ["2025-01", "2025-02", "2025-03", "2025-04", "2025-05", "2025-06",
             "2025-07", "2025-08", "2025-09", "2025-10", "2025-11", "2025-12", ]

external_data = {}

@tool(description="从向量存储中检索参考资料")
def rag_summarize(query: str) -> str:
    return rag.rag_summarize(query)

@tool(description="获取城市的天气,这里使用了固定值。")
def get_weather(city: str) -> str:
    return f"城市{city}天气为晴天,气温26摄氏度,空气湿度50%,南风1级,AQI21,最近六小时降雨概率极低。"

@tool(description="获取用户所在城市的名称")
def get_user_location() -> str:
    return random.choice(["北京", "厦门", "莆田"])

@tool(description="获取用户的ID")
def get_user_id() -> str:
    return random.choice(user_ids)

@tool(description="获取当前月份")
def get_current_month() -> str:
    return random.choice(month_arr)

def generate_external_data():
    if not external_data:
        external_data_path = get_abs_path(agent_conf["external_data_path"])

        if not os.path.exists(external_data_path):
            raise FileNotFoundError(f"外部数据文件{external_data_path}不存在")

        with open(external_data_path, "r", encoding="utf-8") as f:
            for line in f.readlines()[1:]:
                arr: list[arr] = line.strip().split(",")

                user_id: str = arr[0].replace('"', "")
                feature: str = arr[1].replace('"', "")
                efficiency: str = arr[2].replace('"', "")
                consumables: str = arr[3].replace('"', "")
                comparison: str = arr[4].replace('"', "")
                time: str = arr[5].replace('"', "")

                if user_id not in external_data:
                    external_data[user_id] = {}

                external_data[user_id][time] = {
                    "特征":feature,
                    "效率":efficiency,
                    "耗材":consumables,
                    "对比":comparison,
                }


@tool(description="从外部系统的获取用户的使用记录")
def fetch_external_data(user_id: str,month: str) -> str:
    generate_external_data()

    try:
        return external_data[user_id][month]
    except KeyError:
        logger.warning(f"[方法:fetch_external_data]未能检索到用户:{user_id}{month}的使用记录")
        return ""

# 测试
if __name__ == '__main__':
    print(fetch_external_data(user_id="1001", month="2025-01"))

代码中的注释比较少,下面进行介绍:
这段代码写了一些可供Agent调用的方法,有如下几种:

rag_summarize 从向量数据库检索知识
get_weather 查询城市天气
get_user_location 获取用户城市
get_user_id 获取用户ID
get_current_month 获取当前月份
fetch_external_data 获取用户某月使用记录

还包含一个辅助函数 generate_external_data,用来读取外部CSV文件,将数据存入 external_data 字典。

基本结构如下:
agent_tools.py
├── 导入依赖
├── 全局实例化:rag = RagSummarizeService()
├── 静态数据:user_ids, month_arr, external_data
├── 工具定义(@tool 装饰的6个函数)
│ ├── rag_summarize(query) → RAG检索
│ ├── get_weather(city) → 模拟天气
│ ├── get_user_location() → 模拟用户位置
│ ├── get_user_id() → 模拟用户ID
│ ├── get_current_month() → 模拟当前月份
│ └── fetch_external_data(user_id, month) → 从CSV查使用记录
├── 辅助函数:generate_external_data() → 加载CSV到external_data
└── 测试入口

这个文件对Agent来说很重要,它让Agent具备执行具体操作的能力。

同时,在实现功能的过程中,也借助了utils文件夹内的工具文件和RAG检索文件。

结尾

这期要分享的代码就是这些,同时,这次项目编写也差不多要接近尾声了,下一次大概就是最后一期。
老样子,在完结后,我会把所有代码放到GitHub上,需要研究的可以下载。
感谢大家的支持,我们下期再见!

Logo

小龙虾开发者社区是 CSDN 旗下专注 OpenClaw 生态的官方阵地,聚焦技能开发、插件实践与部署教程,为开发者提供可直接落地的方案、工具与交流平台,助力高效构建与落地 AI 应用

更多推荐