实现了完整的网页搜索与抓取内置工具集: 1. 新增搜索工具链:支持SearXNG与Tavily多Provider降级、滑动窗口限流、统一结果格式 2. 新增网页抓取工具:单页抓取、站点递归抓取、内容入库知识库功能 3. 完善安全防护:SSRF校验、robots.txt合规、请求限流 4. 自动注册工具到系统注册表,无需额外配置即可使用
406 lines
13 KiB
Python
406 lines
13 KiB
Python
"""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_code:HTTP 状态码
|
||
- 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_pattern:URL 通配符过滤(如 /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. 需要将网页内容持久化到知识库供后续检索
|
||
|
||
输入参数:
|
||
- content:Markdown 内容(通常来自 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:
|
||
"""获取限流 key:thread_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,
|
||
)
|