1. 为web_search和tavily_search添加runtime默认值None 2. 重构工具参数提取逻辑,跳过runtime参数并适配pydantic v2字段 3. 调整工具包导入顺序与格式化 4. 为外部工具调用添加上下文透传日志字段
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 = None,
|
||
) -> 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 = None,
|
||
) -> str:
|
||
"""Tavily 网页搜索(兼容入口,内部委托 _execute_web_search)。"""
|
||
return await _execute_web_search(query, max_results, runtime)
|