Agent学习记录-4

阅读学习时间:约30分钟

笔记

Agent的学习记录--4

来源Github的学习指南Agent Learning Hub的学习:

Stage 2—— 工具调用

在这里我们使用LangChain进行工作定义

import os
import re
import sys
import json
import sqlite3
from typing import List, Dict, Any, Optional
from html.parser import HTMLParser
from config import get_llm
from langchain_core.tools import tool
from langchain_core.messages import SystemMessage, HumanMessage, ToolMessage, AIMessage

我们在这里先模拟一个知识库

MOCK_SEARCH_KNOWLEDGE = {
    "transformer": "Transformer 由 Vaswani 等人在 2017 年论文《Attention Is All You Need》中提出,彻底摒弃了 RNN 循环结构,完全基于自注意力机制。",
    "deepseek": "DeepSeek (深度求索) 是一家专注于通用人工智能的中国科技公司,推出了 DeepSeek-V3 与 DeepSeek-R1 等开源大模型,以极高性价比和强悍推理能力闻名。",
    "agent loop": "智能体循环 (Agent Loop) 的经典范式包括 ReAct (Reasoning + Acting),其基本工作流为:观察 (Observe) -> 思考 (Think) -> 行动 (Act)。",
}

使用Langchain定义工具:

@tool
def search_academic(query: str) -> str:
    """当需要搜索最新的互联网信息、学术概念定义或未知的实体背景时,调用此搜索工具。
    参数 query: 搜索关键词,例如 'Transformer' 或 'DeepSeek'。
    """
    print(f"  [Tool] 正在执行网络搜索: '{query}' ...")
    query_lower = query.lower()
    results = []
    for key, val in MOCK_SEARCH_KNOWLEDGE.items():
        if key in query_lower or any(word in key for word in query_lower.split()):
            results.append(val)
    
    if results:
        return "\n".join(results)
    return f"未找到关于 '{query}' 的搜索结果,建议更换关键词重试。"

这是一个最基础的Search工具

Database

首先我们需要连接到一个数据库文件

DB_PATH = os.path.join(os.path.dirname(__file__), "sample_data", "research.db")

这里需要完成TODO:

【TODO (学员实操) 1.1】:

  1. 编写安全防御拦截:检查 sql_query,若包含 INSERT, UPDATE, DELETE, DROP, ALTER, TRUNCATE 等写入/破坏性危险关键字,立即返回拒绝执行的错误提示字符串(如: "安全拦截:仅允许 SELECT 只读查询!")。
  2. 连接 sqlite3.connect(DB_PATH),执行 SQL 并获取所有数据。
  3. 将查询结果转换为可读格式(如 JSON 字符串或易读列表文本)返回。
  4. 做好 try-except 异常处理,若 SQL 语法错误,返回错误提示。

现在开始建立工具

@tool
def query_sqlite(sql_query: str) -> str:
    """当需要查询本地科研数据库 (research.db) 时调用此工具。
    数据库包含两个表:
    1. papers (id, title, authors, year, citations, venue)
    2. experiments (exp_id, model, accuracy, status)
    注意:本工具只允许执行 SELECT 只读查询语句!严禁执行修改、删除等操作。
    """
    print(f"  [Tool] 正在执行 SQL 查询: {sql_query} ...")
    
    sql_query_low = sql_query.lower()
    if not sql_query_low.strip().startswith("select"):
        return f"Safety notice: only SELECT is allowed"
    try:
        conn = sqlite3.connect(DB_PATH)
        conn.row_factory = sqlite3.Row
        cur = conn.cursor()
        cur.execute(sql_query)
        rows = cur.fetchall()
        conn.close
        
        if not rows:
            return f"Success, but null result"
        
        result_list = [dict(r) for r in rows]
        return json.dumps(result_list, ensure_ascii=False, indent=2)
    expect Exception as e:
        return f"SQL error with {str(e)}"

这里的strip() 函数用于删除字符串两端的空格或指定字符

File

在这里我们使用**沙箱(Sandbox)**进行文件管理

首先我们一样的先定义沙箱地址

SANDBOX_DIR = os.path.join(os.path.dirname(__file__),"sample_data")

沙箱的核心目的就是划定安全边界只允许 Agent 在指定的目录内活动(比如sample_data),严禁越界访问操作系统的其他目录

我们开始定义工具

@tool
def manage_files(action: str, filepath: str, content: Optional[str]=None) -> str:
    """用于在本地安全沙箱目录内读写与浏览文件的工具。
    参数 action: 操作类型,可选 'read'(读取文件)、'write'(写入文件)、'list'(列出目录)。
    参数 filepath: 相对文件路径,例如 'papers/attention_mechanism.txt'。
    参数 content: 当 action='write' 时要写入的文本内容。
    """
    print(f"  [Tool] file management: action={action}, path={filepath} ...")

在这里我们要注意一些语法,关于python文件路径的一些表达:...是用来相对导航的特殊符号:

符号 含义
. 当前目录
.. 上一级目录(父目录)
../.. 上两级目录(祖父目录)

避免逃逸沙箱,所以我们不允许这种行为的发生,即在 filepath不允许非我们先定义的 SANDBOX_DIR 的开头的目录以外

    target_path = os.path.normpath(os.path.join(SANDBOX_DIR, filepath))
    if not target_path.startswith(os.path.abspath(SANDBOX_DIR)):
        return "[Error] Can't load the path out of the sandbox"

其中 os.path.normpath 为规范化文件路径字符串的工具。例如,它会将连续的多个斜杠(如//)合并为一个,将当前目录的符号(.)去除,将上一级目录的符号(..)与前面的路径部分相抵消。在Windows系统中,它还会将路径中的正斜杠(/)转换为反斜杠(\),以符合Windows的路径格式。需要注意的是,它不会检查路径是否真实存在,也不会创建新的文件或目录。

还有一点需要注意的是,这里需要用到绝对路径,而不能使用相对路径。因为相对路径取决于我们在哪里敲命令启动它。使用绝对路径放在一起才具有可比较性。另外还有一个好处在于其防止大模型直接传入“绝对路径”绕过沙箱。

然后我们就可以开始我们的工具部分,我们有三种参数 read, write, list

    try:
        if action == "read":
            if not os.path.exists(target_path):
                return f"Error, The file '{filepath}' does not exist。"
            with open(target_path, "r", encoding="utf-8") as f: #使用utf-8编码
                return f.read(3000) # 读取前 3000 字符防溢出
        elif action == "write":
            os.makedirs(os.path.dirname(target_path), exist_ok=True) #避免文件夹已经存在时还报错
            with open(target_path, "w", encoding="utf-8") as f:
                f.write(content or "")
            return f"成功写入文件 '{filepath}'。"
        elif action == "list":
            if not os.path.exists(target_path):
                return f"目录 '{filepath}' 不存在。"
            files = os.listdir(target_path)
            return f"目录内文件列表: {files}"
        else:
            return f"未知 action: {action},只支持 read/write/list。"
    except Exception as e:
        return f"文件操作异常: {str(e)}"

应注意,我们在这里是读不了除了非纯文字文档以外的文件,比如PDF, Word, Excel等,需要调用额外的库进行读取。

Browser

先做一个HTML文本提取器,我们定义一个子类SimpleHTMLTextExtractor,其继承父类 HTMLParser

class SimpleHTMLTextExtractor(HTMLParser):
    """一个轻量级的纯 Python HTML 文本提取器,自动过滤掉 script, style, head 等杂乱标签"""
    def __init__(self):
        super().__init__()
        self.ignored_tags = {"script", "style", "head", "title", "meta", "link", "nav", "footer"}
        self.void_tags = {"meta", "link", "img", "br", "hr", "input"}
        self.current_tag_stack = []
        self.text_pieces = []

    def handle_starttag(self, tag, attrs):
        # 过滤掉不需要闭合的单标签 (Void Tags),防止它们永久滞留在标签栈中
        if tag.lower() in self.void_tags:
            return
        self.current_tag_stack.append(tag.lower())

    def handle_endtag(self, tag):
        # 弹栈时匹配到对应标签即完成出栈
        tag_lower = tag.lower()
        if tag_lower in self.current_tag_stack:
            while self.current_tag_stack:
                top = self.current_tag_stack.pop()
                if top == tag_lower:
                    break

    def handle_data(self, data):
        # 如果当前处于需要忽略的标签(如 script/style/nav)内,则跳过
        if any(ignored in self.current_tag_stack for ignored in self.ignored_tags):
            return
        cleaned = data.strip()
        if cleaned:
            self.text_pieces.append(cleaned)

    def get_text(self) -> str:
        return "\n".join(self.text_pieces)

现在我们先不使用联网的html信息,而是使用本地的html文件:

这里的TODO为:

【TODO (学员实操) 1.2】:

  1. 判断 url_or_filepath:
  • 如果是以 'http://' 或 'https://' 开头,提示目前仅支持本地沙箱网页读取;
  • 否则读取沙箱目录 (SANDBOX_DIR) 下的本地 HTML 文件(例如 'web_cache/sample_article.html')。
  1. 使用上方提供的 SimpleHTMLTextExtractor 或正则表达式,过滤掉 HTML 标签、脚本及导航干扰, 提取出干净的文章正文纯文本。
  2. 将提取后的纯净文本(控制在 2000 字符内)返回。

现在开始建立工具:

@tool
def fetch_web_page(url_or_filepath: str) -> str:
    """当需要获取网页内容或分析 HTML 文档时调用此工具。
    参数 url_or_filepath: 可以是网页 URL,或者是本地测试 HTML 文件路径(如 'web_cache/sample_article.html')。
    """
    print(f"  [Tool] 正在抓取并清洗网页内容: {url_or_filepath} ...")
    if url_or_filepath.strip().startswith(("http://","https://")):
        return f"Now the tool only support the local file in sandbox"
    
    target_path = os.path.normpath(os.path.join(SANDBOX_DIR, url_or_filepath))
    if not os.path.exists(target_path):
        return f"Error, can not find the html file '{url_or_filepath}'"

    try:
        with open(target_path, "r", encoding="utf-8") as f:
            html_content = f.read()
        parser=SimpleHTMLTextExtractor()
        parser.feed(html_content)
        clean_text=parser.get_text()
        return clean_text[:2000]
    except Exception as e:
        return f"It is abnormal to read the html file with '{str(e)}'"

Code Execution

这个部分,我们要去执行代码,比如要去计算,要去数据统计,分析等等。但这里有一个很重要的点,我们需要限制内置函数,避免引入危险命令。

@tool
def execute_python_code(code: str) -> str:
    """当需要执行精确的数学计算、数据统计、字符串算法或验证代码逻辑时调用此工具。
    参数 code: 完整的 Python 代码字符串。
    注意:代码内请使用 print(...) 输出最终结果,工具将捕获标准输出 (stdout)。
    """
    print(f"  [Tool] 正在执行 Python 代码:\n{code}\n  ---")

    import io
    from contextlib import redirect_stdout, redirect_stderr
    
    stdout_buf = io.StringIO()
    stderr_buf = io.StringIO()
    
    # 限制内置函数,防止恶意 import os / sys 执行危险命令
    safe_globals = {
        "__builtins__": {
            "print": print, "range": range, "len": len, "sum": sum,
            "min": min, "max": max, "sorted": sorted, "abs": abs,
            "round": round, "int": int, "float": float, "str": str,
            "list": list, "dict": dict, "set": set, "tuple": tuple,
            "enumerate": enumerate, "zip": zip,
        }
    }
    
    try:
        with redirect_stdout(stdout_buf), redirect_stderr(stderr_buf):
            exec(code, safe_globals, {})
        output = stdout_buf.getvalue()
        errors = stderr_buf.getvalue()
        
        if errors:
            return f"代码执行产生警告/错误:\n{errors}\n输出内容:\n{output}"
        if not output:
            return "代码执行成功,但没有产生任何 print 输出。请确保在代码中使用 print(...) 输出结果。"
        return output.strip()
    except Exception as e:
        return f"代码执行异常: {type(e).__name__} - {str(e)}"

这里的 redirect_stdoutredirect_stderr为捕获代码执行的结果,并且保存至缓存区 stdout_bufstderr_buf

工具总装和调度闭环

TODO任务:

【TODO (学员实操) 1.3】: 实现工具调用的调度执行循环:

  1. 调用模型:response = model_with_tools.invoke(messages)
  2. 将 response 追加到 messages 中。
  3. 检查 response 是否有 tool_calls:
    • 如果没有 tool_calls,说明模型已得出最终结论,直接返回 response.content。
    • 如果有 tool_calls,遍历每一个 tool_call: a. 获取工具名称 tool_name = tool_call["name"] b. 获取实参 tool_args = tool_call["args"] c. 获取调用 ID tool_id = tool_call["id"] d. 从 TOOL_MAP 中获取对应工具并调用: result = TOOL_MAP[tool_name].invoke(tool_args) e. 构建 ToolMessage(content=str(result), tool_call_id=tool_id) 并追加到 messages。
  4. 再次调用模型,直到模型输出最终回复(设置最大轮数 max_steps=5 防止死循环)。

Langchain里,我们使用以下方式写入工具工具表

ALL_TOOLS = [search_academic, query_sqlite, manage_files, fetch_web_page, execute_python_code]
TOOL_MAP = {t.name: t for t in ALL_TOOLS}

另外还需要在config.py 里进行额外的配置:

from langchain_openai import ChatOpenAI
from langchain_core.embeddings import Embeddings

def get_llm(temperature: float = 0.0) -> ChatOpenAI:
    """获取 LangChain ChatOpenAI 实例,默认 temperature=0.0 以获得稳定输出"""
    if not API_KEY or API_KEY in ("your_api_key_here", "sk-xxxxxxxx"):
        print("\n" + "=" * 60)
        print("【错误提示】未检测到有效的 API Key!")
        print("请在 .env 文件中配置 OPENAI_API_KEY、OPENAI_BASE_URL 和 MODEL_NAME。")
        print("=" * 60 + "\n")
        sys.exit(1)
    
    return ChatOpenAI(
        model=MODEL_NAME,
        api_key=API_KEY,
        base_url=BASE_URL,
        temperature=temperature
    )

回到我们的工作py文件,因为这里我们用到了Langchian,所以跟之前用的OpenAI SDK不一样了,在这里我们会用更简单的写法,

def run_agent_with_tools(user_query: str) -> str:
    """
    接收用户问题,使用 LangChain bind_tools 赋予模型工具调用能力,
    并在模型请求调用工具时,自动分发执行并回填结果,直至得到最终答案。
    """
    llm = get_llm(temperature=0.0)
    
    model_with_tools = llm.bind_tools(ALL_TOOLS)
    messages = [
        SystemMessage(content=(
            "你是一个具备五大外部工具的研究助理。你可以使用搜索、SQL数据库、文件系统、"
            "网页浏览器和Python执行器来严谨地解答用户问题。\n"
            "【已知数据库表结构】:\n"
            "1. papers (id, title, authors, year, citations, venue)\n"
            "2. experiments (exp_id, model, accuracy, status)\n"
            "无需额外查询 sqlite_master 或 PRAGMA,请直接编写 SELECT 查询!遇到数据计算必须使用 Python 验证!"
        )),
        HumanMessage(content=user_query)
    ]
    print(f"\n[User Query]: {user_query}")
    print("=" * 60)

先打断一下,我们在messages里不需要再通过字典的格式 {"role": "system", "content": "XXX"}进行,而是直接通过SystemMessage(content=())HumanMessage(content=())的方式,更方便的是后面response:

    max_steps=5
    for step in range(max_steps):
        print(f"step {step} is using LLM")
        response = model_with_tools.invoke(messages)

        if not response.tool_calls:
            return f"\n {response.content}"

        messages.append(response)

        for tool_call in response.tool_calls:
            try:
                tool_name = tool_call["name"]
                tool_args = tool_call["args"]
                tool_id = tool_call["id"]
                if tool_name in TOOL_MAP:
                    try:
                        result = TOOL_MAP[tool_name].invoke(tool_args)
                    except Exception as e:
                        result = f"Error with {str(e)}"
                else: 
                    result = f"Do not have the tool name '{tool_name}'"
                messages.append(ToolMessage(content=str(result), tool_call_id = tool_id))
            except Exception as e:
                return f"Error with {str(e)}"
    return "Have reached the maximum of the step"

我们直接使用 response = model_with_tools.invoke(messages)messages.append(response)。不用再像以前要写

response = client.chat.completions.create(
        model="deepseek-v4-flash",
        messages=messages
    )
assistant_reply = response.choices[0].message.content

在这里我们就完成了一个小的工具调用

测试

def test():
	query = "请帮我查一下数据库里 2022 年及以后发表的论文有哪些,并用 Python 精确计算它们被引次数 (citations) 的平均值。"
	final_answer = run_agent_with_tools(query)
	print("\n【Agent 最终回答】:")
	print(final_answer)

if __name__ == "__main__":
    test()

运行结果为:

查询完成!以下是数据库里 **2022 年及以后发表**的论文:

| ID | 标题 | 作者 | 年份 | 被引次数 | 会议 |
|----|------|------|------|---------|------|
| 5 | Toolformer: Language Models Can Teach Themselves to Use Tools | Schick et al. | 2023 | 3,100 | NeurIPS |
| 4 | ReAct: Synergizing Reasoning and Acting in Language Models | Yao et al. | 2022 | 5,300 | ICLR |
| 6 | Chain-of-Thought Prompting Elicits Reasoning in Large Language Models | Wei et al. | 2022 | 8,900 | NeurIPS |

**被引次数平均值计算(Python 精确验证):**

- 论文数量:**3 篇**
- 被引次数总和:3100 + 5300 + 8900 = **17,300**
- 平均值:17300 ÷ 3 = **5766.666666...**

**平均被引次数 ≈ 5766.67 次**(保留两位小数)

📌 补充说明:这 3 篇都是大语言模型(LLM)推理与工具使用方向的经典高引论文,其中 Chain-of-Thought Prompting 被引最高(8,900 次),拉高了整体平均值。