ForcePilot/backend/package/yuxi/agents/toolkits/buildin/crawl/tools.py

406 lines
13 KiB
Python
Raw Normal View History

"""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,
)