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