笔记
Agent的学习记录--4
来源Github的学习指南Agent Learning Hub的学习:
Stage 2—— 工具调用
Search
在这里我们使用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】:
- 编写安全防御拦截:检查 sql_query,若包含 INSERT, UPDATE, DELETE, DROP, ALTER, TRUNCATE 等写入/破坏性危险关键字,立即返回拒绝执行的错误提示字符串(如: "安全拦截:仅允许 SELECT 只读查询!")。
- 连接 sqlite3.connect(DB_PATH),执行 SQL 并获取所有数据。
- 将查询结果转换为可读格式(如 JSON 字符串或易读列表文本)返回。
- 做好 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】:
- 判断 url_or_filepath:
- 如果是以 'http://' 或 'https://' 开头,提示目前仅支持本地沙箱网页读取;
- 否则读取沙箱目录 (SANDBOX_DIR) 下的本地 HTML 文件(例如 'web_cache/sample_article.html')。
- 使用上方提供的 SimpleHTMLTextExtractor 或正则表达式,过滤掉 HTML 标签、脚本及导航干扰, 提取出干净的文章正文纯文本。
- 将提取后的纯净文本(控制在 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_stdout和redirect_stderr为捕获代码执行的结果,并且保存至缓存区 stdout_buf和stderr_buf
工具总装和调度闭环
TODO任务:
【TODO (学员实操) 1.3】: 实现工具调用的调度执行循环:
- 调用模型:response = model_with_tools.invoke(messages)
- 将 response 追加到 messages 中。
- 检查 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。
- 再次调用模型,直到模型输出最终回复(设置最大轮数 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 次),拉高了整体平均值。
