ForcePilot/backend/package/yuxi/agents/toolkits/buildin/crawl/tools.py
Kris 077de8e6d6 feat(agent-toolkit): add web search and crawl built-in toolkits
实现了完整的网页搜索与抓取内置工具集:
1. 新增搜索工具链:支持SearXNG与Tavily多Provider降级、滑动窗口限流、统一结果格式
2. 新增网页抓取工具:单页抓取、站点递归抓取、内容入库知识库功能
3. 完善安全防护:SSRF校验、robots.txt合规、请求限流
4. 自动注册工具到系统注册表,无需额外配置即可使用
2026-06-22 21:20:02 +08:00

406 lines
13 KiB
Python
Raw 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.

"""Web 抓取 buildin 工具注册crawl_url、crawl_site、crawl_to_kb。
工具注册通过 @tool 装饰器完成,模块加载即注册。
- crawl_url单页抓取返回 Markdown
- crawl_site站点递归抓取异步任务立即返回 task_id
- crawl_to_kb将爬取内容入库到知识库
"""
import json
from langgraph.prebuilt.tool_node import ToolRuntime
from yuxi.agents.toolkits.registry import tool
from yuxi.config import config
from yuxi.services.task_service import TaskContext, tasker
from yuxi.utils.logging_config import logger
from yuxi.utils.paths import VIRTUAL_PATH_OUTPUTS
from .client import crawl4ai_client
from .ingester import check_kb_permission, ingest_crawl_result, resolve_user_dict
from .rate_limiter import crawl_rate_limiter
from .security import should_respect_robots, validate_crawl_url
CRAWL_URL_DESCRIPTION = """
抓取单个网页 URL返回干净的 Markdown 内容。
使用场景:
1. 从 web_search 结果中选取感兴趣的 URL抓取全文
2. 用户指定某个网页 URL需要获取完整内容
3. 需要将网页内容入库到知识库前的抓取步骤
返回结果:
- markdown网页正文 Markdown已剥离导航栏、广告、页脚等噪声
- url最终 URL可能经过重定向
- title页面标题
- status_codeHTTP 状态码
- crawled_at抓取时间
注意事项:
- 结果长度限制 50000 字符(约 12k token超出截断并提示
- 遵守 robots.txt被禁止的 URL 返回错误
- SSRF 防护:内网地址、私有 IP 段会被拒绝
"""
CRAWL_SITE_DESCRIPTION = """
递归抓取一个站点的多个页面,返回每页 Markdown 摘要。
使用场景:
1. 把一个技术文档站(如 FastAPI 官网)整站入库
2. 批量抓取同域下的多个页面
输入参数:
- start_url起始 URL
- max_depth递归深度1-3默认 2
- max_pages最大抓取页数1-100默认 20
- url_patternURL 通配符过滤(如 /docs/*
返回结果:
- 摘要:成功页数、失败页数、总字符数
- 每页含 URL + Markdown 摘要(前 500 字符)
- 完整内容存到临时文件,通过 file_path 字段引用,可按需读取
执行方式:
- 异步任务,立即返回 task_id
- 通过 GET /api/tasks/{task_id} 查询进度
- 通过 POST /api/tasks/{task_id}/cancel 取消任务
"""
CRAWL_TO_KB_DESCRIPTION = """
将爬取的 Markdown 内容入库到指定知识库。
使用场景:
1. agent 已通过 crawl_url 抓取到全文,用户确认入库到某个知识库
2. 需要将网页内容持久化到知识库供后续检索
输入参数:
- contentMarkdown 内容(通常来自 crawl_url 的返回)
- kb_id目标知识库 ID
- source_url来源 URL记入文件元数据
- filename可选文件名默认从 URL 推导
返回结果:
- file_id入库后的文件 ID
- kb_id知识库 ID
- status入库状态indexed 表示成功)
注意事项:
- 用户必须有目标知识库的写入权限
- 基于 content_hash 去重,重复内容返回提示
- 入库后自动触发解析与索引
"""
def _get_rate_limit_key(runtime: ToolRuntime) -> str:
"""获取限流 keythread_id + uid。
BaseContext 仅暴露 thread_id 与 uid无 agent_id 字段),
故限流维度为 (thread_id, uid)。
"""
ctx = getattr(runtime, "context", None)
if ctx is None:
return "::"
thread_id = getattr(ctx, "thread_id", "") or ""
uid = getattr(ctx, "uid", "") or ""
return f"{thread_id}:{uid}"
def _extract_runtime_identity(runtime: ToolRuntime) -> tuple[str, str]:
"""从 runtime 提取 (thread_id, uid),供闭包捕获。
必须在工具调用同步阶段提取,不能在异步任务协程内访问 runtime——
工具调用返回后 runtime 可能失效。
"""
ctx = getattr(runtime, "context", None)
if ctx is None:
return "", ""
thread_id = getattr(ctx, "thread_id", "") or ""
uid = getattr(ctx, "uid", "") or ""
return thread_id, uid
@tool(
category="buildin",
tags=["抓取"],
display_name="网页抓取",
description=CRAWL_URL_DESCRIPTION,
)
async def crawl_url(
url: str,
output_format: str = "markdown",
runtime: ToolRuntime = None,
) -> str:
"""抓取单个 URL返回干净 Markdown。"""
# 限流检查
key = _get_rate_limit_key(runtime)
if not await crawl_rate_limiter.check(key):
return json.dumps(
{"error": "抓取调用过于频繁,请稍后再试", "url": url},
ensure_ascii=False,
)
# SSRF 校验
ok, reason = await validate_crawl_url(url)
if not ok:
return json.dumps(
{"error": f"URL 被 SSRF 防护拒绝: {reason}", "url": url},
ensure_ascii=False,
)
respect_robots = should_respect_robots()
try:
result = await crawl4ai_client.crawl_single(
url,
output_format=output_format,
check_robots_txt=respect_robots,
)
except Exception as exc:
logger.warning(f"crawl_url_failed url={url} error={exc}")
return json.dumps(
{"error": f"抓取失败: {exc}", "url": url},
ensure_ascii=False,
)
if not result.get("success", False):
error_msg = result.get("error_message", "未知错误")
status_code = result.get("status_code")
return json.dumps(
{
"error": f"抓取失败: {error_msg}",
"url": url,
"status_code": status_code,
},
ensure_ascii=False,
)
markdown = result.get("markdown", "")
# 截断到 50000 字符
truncated = False
if len(markdown) > 50000:
markdown = markdown[:50000]
truncated = True
return json.dumps(
{
"markdown": markdown,
"url": result.get("url", url),
"title": result.get("metadata", {}).get("title", ""),
"status_code": result.get("status_code"),
"crawled_at": result.get("metadata", {}).get("crawl_date", ""),
"truncated": truncated,
},
ensure_ascii=False,
)
@tool(
category="buildin",
tags=["抓取"],
display_name="站点递归抓取",
description=CRAWL_SITE_DESCRIPTION,
)
async def crawl_site(
start_url: str,
max_depth: int = 2,
max_pages: int = 20,
url_pattern: str = "",
runtime: ToolRuntime = None,
) -> str:
"""递归抓取站点,异步执行,立即返回 task_id。"""
# 限流检查
key = _get_rate_limit_key(runtime)
if not await crawl_rate_limiter.check(key):
return json.dumps(
{"error": "抓取调用过于频繁,请稍后再试", "url": start_url},
ensure_ascii=False,
)
# 参数 clamp
max_depth = max(1, min(3, max_depth))
max_pages = max(1, min(100, max_pages))
# SSRF 校验起始 URL
ok, reason = await validate_crawl_url(start_url)
if not ok:
return json.dumps(
{"error": f"起始 URL 被 SSRF 防护拒绝: {reason}", "url": start_url},
ensure_ascii=False,
)
# 在工具调用同步阶段提取 (thread_id, uid),供闭包捕获。
# 不能在异步任务协程内访问 runtime——工具调用返回后 runtime 可能失效。
thread_id, uid = _extract_runtime_identity(runtime)
# 闭包捕获参数TaskContext 仅持有 task_id不持有 Task 引用。
# 闭包内不访问 runtime仅使用已提取的 thread_id / uid。
async def _run_crawl_site(context: TaskContext) -> dict:
"""crawl_site 异步任务协程,通过闭包捕获 start_url 等参数。"""
respect_robots = should_respect_robots()
await context.set_progress(5.0, f"开始抓取 {start_url}")
await context.raise_if_cancelled()
try:
results = await crawl4ai_client.crawl_many(
[start_url],
max_depth=max_depth,
max_pages=max_pages,
url_pattern=url_pattern,
check_robots_txt=respect_robots,
semaphore_count=config.web_crawl_default_concurrency,
)
except Exception as exc:
await context.set_result({"error": str(exc), "start_url": start_url})
raise
await context.set_progress(80.0, f"抓取完成,共 {len(results)}")
await context.raise_if_cancelled()
# 聚合结果:每页含 URL + 摘要(前 500 字符)
pages = []
full_content_parts = []
total_chars = 0
success_count = 0
failed_count = 0
for item in results:
if item.get("success", False):
markdown = item.get("markdown", "")
total_chars += len(markdown)
success_count += 1
pages.append(
{
"url": item.get("url", ""),
"title": item.get("metadata", {}).get("title", ""),
"snippet": markdown[:500],
"status_code": item.get("status_code"),
}
)
full_content_parts.append(f"# {item.get('url', '')}\n\n{markdown}\n\n---\n")
else:
failed_count += 1
pages.append(
{
"url": item.get("url", ""),
"error": item.get("error_message", "未知错误"),
"status_code": item.get("status_code"),
}
)
# 完整内容存到临时文件agent 可按需读取
file_path = ""
if full_content_parts and thread_id and uid:
from yuxi.agents.backends.sandbox.paths import (
ensure_thread_dirs,
sandbox_outputs_dir,
)
ensure_thread_dirs(thread_id, uid)
outputs_dir = sandbox_outputs_dir(thread_id)
file_name = f"crawl_site_{start_url.replace('/', '_').replace(':', '_')[:50]}_{success_count}pages.md"
full_path = outputs_dir / file_name
full_path.write_text("\n".join(full_content_parts), encoding="utf-8")
file_path = f"{VIRTUAL_PATH_OUTPUTS}/{file_name}"
summary = {
"start_url": start_url,
"success_count": success_count,
"failed_count": failed_count,
"total_chars": total_chars,
"pages": pages,
"file_path": file_path, # 完整内容文件路径agent 可按需读取
}
await context.set_result(summary)
await context.set_progress(100.0, f"抓取完成:成功 {success_count} 页,失败 {failed_count}")
return summary
# 异步任务去重:基于 start_url 去重
task, created = await tasker.enqueue_unique_by_payload(
name=f"站点抓取 ({start_url})",
task_type="web_crawl_site",
payload={
"start_url": start_url,
"max_depth": max_depth,
"max_pages": max_pages,
"url_pattern": url_pattern,
},
coroutine=_run_crawl_site,
payload_match={"start_url": start_url},
statuses={"pending", "running"},
)
if not created:
return json.dumps(
{
"message": "已有相同起始 URL 的抓取任务正在执行",
"task_id": task.id,
"status": task.status,
},
ensure_ascii=False,
)
return json.dumps(
{
"message": "抓取任务已提交",
"task_id": task.id,
"status": task.status,
"progress_query": f"GET /api/tasks/{task.id}",
"cancel_query": f"POST /api/tasks/{task.id}/cancel",
},
ensure_ascii=False,
)
@tool(
category="buildin",
tags=["抓取", "知识库"],
display_name="爬取内容入库",
description=CRAWL_TO_KB_DESCRIPTION,
)
async def crawl_to_kb(
content: str,
kb_id: str,
source_url: str,
filename: str | None = None,
runtime: ToolRuntime = None,
) -> str:
"""将爬取的 Markdown 内容入库到指定知识库。"""
if not content or not content.strip():
return json.dumps(
{"error": "内容不能为空", "kb_id": kb_id},
ensure_ascii=False,
)
# 权限校验BaseContext 仅暴露 uid需查询 UserRepository 补全 role/department_id
# 复用 KnowledgeBaseManager.check_accessible(user, kb_id)。
_, uid = _extract_runtime_identity(runtime)
user = await resolve_user_dict(uid)
has_permission = await check_kb_permission(kb_id, user)
if not has_permission:
return json.dumps(
{"error": f"无知识库 {kb_id} 的写入权限", "kb_id": kb_id},
ensure_ascii=False,
)
try:
result = await ingest_crawl_result(
content=content,
kb_id=kb_id,
source_url=source_url,
filename=filename,
operator_id=uid,
)
return json.dumps(result, ensure_ascii=False)
except ValueError as exc:
return json.dumps(
{"error": str(exc), "kb_id": kb_id, "source_url": source_url},
ensure_ascii=False,
)
except Exception as exc:
logger.error(f"crawl_to_kb_failed kb_id={kb_id} source_url={source_url} error={exc}")
return json.dumps(
{"error": f"入库失败: {exc}", "kb_id": kb_id, "source_url": source_url},
ensure_ascii=False,
)