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

126 lines
3.8 KiB
Python
Raw Normal View History

"""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:
"""获取限流 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}"
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)