Files
liqiang b119135836
Build latest book artifacts / build (push) Canceled after 0s
dependency resolution / resolve (3.11) (push) Canceled after 0s
dependency resolution / resolve (3.13) (push) Canceled after 0s
deploy-pages / build (push) Canceled after 0s
deploy-pages / deploy (push) Canceled after 0s
i18n consistency check / check (push) Canceled after 0s
provider adoption tests / test (chapter2/context-compression) (push) Canceled after 0s
provider adoption tests / test (chapter2/prompt-injection) (push) Canceled after 0s
provider adoption tests / test (chapter2/system-hint) (push) Canceled after 0s
provider adoption tests / test (chapter3/log-sanitization) (push) Canceled after 0s
web-search-agent tests / test (push) Canceled after 0s
web-search-agent tests / agentbook (push) Canceled after 0s
ai-agent-book 精选快照(<2MB 代码与文档,来自 github.com/bojieli/ai-agent-book)
2026-08-20 13:12:50 +00:00

298 lines
11 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
tools.py —— 各专业角色的专属工具实现 + OpenAI function-calling schema。
设计原则(配合实验 10-1):
- 所有被实验场景实际调用的工具都执行真实工作,不用预置答案冒充检索。
- research.web_searchTavily 真实联网检索,并返回可追溯 URL 与摘录。
- coding.execute_python:真实执行 Python 代码并捕获标准输出(子进程 + 超时)。
- data_analysis.calculate / descriptive_stats:真实的安全计算。
- writing.count_characters:真实的中英文字数统计。
每个工具函数签名为 func(**kwargs) -> str(统一返回字符串,方便塞回对话历史)。
"""
from __future__ import annotations
import ast
import json
import operator
import os
import subprocess
import sys
import tempfile
import time
import urllib.error
import urllib.request
from typing import Callable, Dict, List, Optional
# Keep live campaigns bounded when a provider stalls. The value is configurable
# for readers running in a slower network, while the default is short enough that
# one unavailable search cannot consume the whole paired comparison.
TAVILY_TIMEOUT_SECONDS = float(os.environ.get("TAVILY_TIMEOUT_SECONDS", "20"))
TAVILY_MAX_RESULTS = int(os.environ.get("TAVILY_MAX_RESULTS", "5"))
TAVILY_CONTENT_CHARS = int(os.environ.get("TAVILY_CONTENT_CHARS", "1400"))
# ---------------------------------------------------------------------------
# research 角色:web_search —— 真实 Tavily 搜索
# ---------------------------------------------------------------------------
def web_search(query: str, receipt_sink: Optional[Callable[[dict], None]] = None) -> str:
"""Run a real Tavily web search and return attributable source excerpts."""
api_key = os.environ.get("TAVILY_API_KEY", "").strip()
if not api_key:
raise RuntimeError("web_search requires TAVILY_API_KEY; no mock fallback is allowed")
body = {
"api_key": api_key,
"query": query,
"search_depth": "advanced",
"max_results": TAVILY_MAX_RESULTS,
"include_answer": True,
"include_raw_content": False,
}
request = urllib.request.Request(
"https://api.tavily.com/search",
data=json.dumps(body).encode("utf-8"),
headers={"Content-Type": "application/json"},
method="POST",
)
started = time.monotonic()
try:
with urllib.request.urlopen(request, timeout=TAVILY_TIMEOUT_SECONDS) as response:
status = response.status
raw_response = response.read().decode("utf-8", "replace")
payload = json.loads(raw_response)
except urllib.error.HTTPError as exc:
detail = exc.read().decode("utf-8", "replace")[:1000]
raise RuntimeError(f"Tavily HTTP {exc.code}: {detail}") from None
except (urllib.error.URLError, TimeoutError, json.JSONDecodeError) as exc:
raise RuntimeError(f"Tavily 请求失败:{exc}") from None
if receipt_sink:
receipt_sink({
"kind": "tavily_search",
"request": {
"method": "POST",
"url": "https://api.tavily.com/search",
"headers": {"Content-Type": "application/json"},
"body": {key: value for key, value in body.items() if key != "api_key"},
},
"response": {
"http_status": status,
"raw_body": raw_response,
},
"duration_seconds": round(time.monotonic() - started, 3),
})
results = []
for item in payload.get("results") or []:
if not isinstance(item, dict):
continue
results.append({
"title": item.get("title"),
"url": item.get("url"),
# Search snippets are evidence pointers, not a second context
# window. Bound their size so repeated role transitions do not
# make later API requests quadratic in prompt length.
"content": str(item.get("content") or "")[:TAVILY_CONTENT_CHARS],
"score": item.get("score"),
})
if not results:
return json.dumps({
"provider": "tavily",
"query": query,
"answer": payload.get("answer"),
"results": [],
}, ensure_ascii=False)
return json.dumps({
"provider": "tavily",
"query": query,
"answer": payload.get("answer"),
"results": results,
}, ensure_ascii=False)
# ---------------------------------------------------------------------------
# coding 角色:execute_python —— 真实执行代码并捕获 stdout(带超时)
# ---------------------------------------------------------------------------
def execute_python(code: str, timeout: int = 10) -> str:
"""把源码写到临时文件并用子进程执行,返回 stdout(带超时)。"""
with tempfile.TemporaryDirectory() as tmp:
script = os.path.join(tmp, "snippet.py")
with open(script, "w", encoding="utf-8") as fh:
fh.write(code)
try:
proc = subprocess.run(
[sys.executable, script],
capture_output=True,
text=True,
timeout=timeout,
cwd=tmp,
)
except subprocess.TimeoutExpired:
return f"执行超时(>{timeout}s"
out = (proc.stdout or "").strip()
err = (proc.stderr or "").strip()
if proc.returncode != 0:
return (
f"代码执行出错:退出码 {proc.returncode}\n"
f"stderr\n{err}\n"
f"已捕获输出:\n{out}"
)
return out if out else "(代码已执行,但没有任何 print 输出)"
# ---------------------------------------------------------------------------
# data_analysis 角色:calculate(安全表达式求值)+ descriptive_stats
# ---------------------------------------------------------------------------
_ALLOWED_OPERATORS = {
ast.Add: operator.add,
ast.Sub: operator.sub,
ast.Mult: operator.mul,
ast.Div: operator.truediv,
ast.Pow: operator.pow,
ast.Mod: operator.mod,
ast.USub: operator.neg,
ast.UAdd: operator.pos,
}
def _safe_eval(node: ast.AST) -> float:
"""只支持四则运算/幂/取模的安全表达式求值(不走 Python 内置 eval)。"""
if isinstance(node, ast.Constant) and isinstance(node.value, (int, float)):
return float(node.value)
if isinstance(node, ast.BinOp) and type(node.op) in _ALLOWED_OPERATORS:
return _ALLOWED_OPERATORS[type(node.op)](_safe_eval(node.left), _safe_eval(node.right))
if isinstance(node, ast.UnaryOp) and type(node.op) in _ALLOWED_OPERATORS:
return _ALLOWED_OPERATORS[type(node.op)](_safe_eval(node.operand))
raise ValueError("表达式包含不被支持的运算,只允许 + - * / ** % 与括号。")
def calculate(expression: str) -> str:
"""安全地计算一个纯数学表达式,例如 (949.5/352.1)**(1/2)-1 。"""
try:
tree = ast.parse(expression, mode="eval")
result = _safe_eval(tree.body)
except Exception as exc: # noqa: BLE001
return f"计算失败:{exc}"
return f"{expression} = {result}"
def descriptive_stats(numbers: List[float]) -> str:
"""给一组数值返回基本描述统计(均值/最大/最小/极差)。"""
if not numbers:
return "输入为空,无法统计。"
nums = [float(x) for x in numbers]
n = len(nums)
mean = sum(nums) / n
return (
f"样本量={n}, 均值={mean:.4f}, 最小={min(nums)}, "
f"最大={max(nums)}, 极差={max(nums) - min(nums)}"
)
# ---------------------------------------------------------------------------
# writing 角色:count_characters —— 中英文字数统计
# ---------------------------------------------------------------------------
def count_characters(text: str) -> str:
"""统计文本的字符数与中文字符数,帮助控制篇幅。"""
if text is None:
text = ""
total = len(text)
chinese = sum(1 for ch in text if "一" <= ch <= "鿿")
return f"总字符数={total}, 其中中文字符={chinese}"
# ---------------------------------------------------------------------------
# 工具注册表:名称 -> (实现函数, OpenAI schema)
# ---------------------------------------------------------------------------
# 每个工具的 OpenAI function-calling schema。
TOOL_SCHEMAS: Dict[str, dict] = {
"web_search": {
"type": "function",
"function": {
"name": "web_search",
"description": "通过 Tavily 真实联网检索信息,返回带 URL 的来源摘录。用于查数据、事实、资料。",
"parameters": {
"type": "object",
"properties": {
"query": {"type": "string", "description": "检索关键词或问题"},
},
"required": ["query"],
},
},
},
"execute_python": {
"type": "function",
"function": {
"name": "execute_python",
"description": "执行一段 Python 代码并返回其 print 输出。适合写脚本、跑逻辑。",
"parameters": {
"type": "object",
"properties": {
"code": {"type": "string", "description": "要执行的 Python 源码,用 print 输出结果"},
},
"required": ["code"],
},
},
},
"calculate": {
"type": "function",
"function": {
"name": "calculate",
"description": "安全计算一个数学表达式,支持 + - * / ** % 和括号。",
"parameters": {
"type": "object",
"properties": {
"expression": {"type": "string", "description": "数学表达式,如 (949.5/352.1)**(1/2)-1"},
},
"required": ["expression"],
},
},
},
"descriptive_stats": {
"type": "function",
"function": {
"name": "descriptive_stats",
"description": "对一组数值做基本描述统计(均值/最大/最小/极差)。",
"parameters": {
"type": "object",
"properties": {
"numbers": {
"type": "array",
"items": {"type": "number"},
"description": "数值数组",
},
},
"required": ["numbers"],
},
},
},
"count_characters": {
"type": "function",
"function": {
"name": "count_characters",
"description": "统计文本字符数与中文字符数,帮助控制篇幅。",
"parameters": {
"type": "object",
"properties": {
"text": {"type": "string", "description": "要统计的文本"},
},
"required": ["text"],
},
},
},
}
# 工具名 -> 实现函数
TOOL_IMPLEMENTATIONS: Dict[str, Callable[..., str]] = {
"web_search": web_search,
"execute_python": execute_python,
"calculate": calculate,
"descriptive_stats": descriptive_stats,
"count_characters": count_characters,
}