126 lines
3.8 KiB
Python
126 lines
3.8 KiB
Python
|
|
"""web_search 与 tavily_search 工具注册。"""
|
|||
|
|
import json
|
|||
|
|
|
|||
|
|
from langgraph.prebuilt.tool_node import ToolRuntime
|
|||
|
|
|
|||
|
|
from yuxi.agents.toolkits.registry import tool
|
|||
|
|
|
|||
|
|
from .chain import build_provider_chain
|
|||
|
|
from .models import SearchResult
|
|||
|
|
from .rate_limiter import rate_limiter
|
|||
|
|
|
|||
|
|
|
|||
|
|
WEB_SEARCH_DESCRIPTION = """
|
|||
|
|
搜索互联网获取实时信息。
|
|||
|
|
|
|||
|
|
使用场景:
|
|||
|
|
1. 用户询问最新事件、新闻、实时数据
|
|||
|
|
2. 需要引用网页信息并给出原文链接
|
|||
|
|
3. 知识库中没有相关信息,需要联网查询
|
|||
|
|
|
|||
|
|
返回结果:
|
|||
|
|
- 每条结果包含 title(标题)、url(链接)、snippet(摘要)、published_at(发布时间)、source_provider(来源)
|
|||
|
|
- 结果按相关性排序,最多返回 max_results 条
|
|||
|
|
- snippet 已截断到 500 字符
|
|||
|
|
|
|||
|
|
后续动作:
|
|||
|
|
- 如需获取某条结果的全文,可调用 crawl_url(url) 抓取
|
|||
|
|
- 如需将全文入库到知识库,可调用 crawl_url 后再调用 crawl_to_kb 入库
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
|
|||
|
|
TAVILY_SEARCH_DESCRIPTION = """
|
|||
|
|
搜索互联网获取实时信息(兼容入口,新 agent 请使用 web_search)。
|
|||
|
|
|
|||
|
|
此工具为向后兼容保留,行为与 web_search 完全一致。
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _format_results(results: list[SearchResult], provider: str) -> str:
|
|||
|
|
"""将结构化结果格式化为 JSON 字符串供 LLM 消费。"""
|
|||
|
|
return json.dumps(
|
|||
|
|
{
|
|||
|
|
"results": [r.model_dump() for r in results],
|
|||
|
|
"provider": provider,
|
|||
|
|
"count": len(results),
|
|||
|
|
},
|
|||
|
|
ensure_ascii=False,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
|
|||
|
|
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}"
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def _execute_web_search(query: str, max_results: int, runtime: ToolRuntime) -> str:
|
|||
|
|
"""web_search 与 tavily_search 共享的核心逻辑。"""
|
|||
|
|
if not query or not query.strip():
|
|||
|
|
return json.dumps({"error": "查询词不能为空", "results": []}, ensure_ascii=False)
|
|||
|
|
|
|||
|
|
# clamp max_results 到 [1, 10]
|
|||
|
|
max_results = max(1, min(10, max_results))
|
|||
|
|
|
|||
|
|
# 限流检查
|
|||
|
|
key = _get_rate_limit_key(runtime)
|
|||
|
|
if not await rate_limiter.check(key):
|
|||
|
|
return json.dumps(
|
|||
|
|
{"error": "搜索调用过于频繁,请稍后再试", "results": []},
|
|||
|
|
ensure_ascii=False,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
chain = build_provider_chain()
|
|||
|
|
if chain.is_empty:
|
|||
|
|
return json.dumps(
|
|||
|
|
{"error": "搜索能力不可用,请稍后重试或联系管理员", "results": []},
|
|||
|
|
ensure_ascii=False,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
results, provider = await chain.search(query, max_results=max_results)
|
|||
|
|
if not results and not provider:
|
|||
|
|
return json.dumps(
|
|||
|
|
{"error": "搜索能力暂不可用,请稍后重试或联系管理员", "results": []},
|
|||
|
|
ensure_ascii=False,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
return _format_results(results, provider)
|
|||
|
|
|
|||
|
|
|
|||
|
|
@tool(
|
|||
|
|
category="buildin",
|
|||
|
|
tags=["搜索"],
|
|||
|
|
display_name="网页搜索",
|
|||
|
|
description=WEB_SEARCH_DESCRIPTION,
|
|||
|
|
)
|
|||
|
|
async def web_search(
|
|||
|
|
query: str,
|
|||
|
|
max_results: int = 5,
|
|||
|
|
runtime: ToolRuntime,
|
|||
|
|
) -> str:
|
|||
|
|
"""搜索互联网,返回结构化结果列表。"""
|
|||
|
|
return await _execute_web_search(query, max_results, runtime)
|
|||
|
|
|
|||
|
|
|
|||
|
|
@tool(
|
|||
|
|
category="buildin",
|
|||
|
|
tags=["搜索"],
|
|||
|
|
display_name="Tavily 网页搜索(兼容)",
|
|||
|
|
description=TAVILY_SEARCH_DESCRIPTION,
|
|||
|
|
)
|
|||
|
|
async def tavily_search(
|
|||
|
|
query: str,
|
|||
|
|
max_results: int = 5,
|
|||
|
|
runtime: ToolRuntime,
|
|||
|
|
) -> str:
|
|||
|
|
"""Tavily 网页搜索(兼容入口,内部委托 _execute_web_search)。"""
|
|||
|
|
return await _execute_web_search(query, max_results, runtime)
|