ForcePilot/backend/package/yuxi/agents/toolkits/service.py
Kris 44dcb028b9 fix(toolkit): 修复搜索工具参数与工具信息展示问题
1. 为web_search和tavily_search添加runtime默认值None
2. 重构工具参数提取逻辑,跳过runtime参数并适配pydantic v2字段
3. 调整工具包导入顺序与格式化
4. 为外部工具调用添加上下文透传日志字段
2026-07-11 22:08:34 +08:00

186 lines
7.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

from typing import Any
from yuxi.utils import logger
# 工具元数据缓存
_metadata_cache: list[dict] = []
def _extract_tool_info(tool_obj) -> dict:
"""从 tool_obj 提取基础信息"""
metadata = getattr(tool_obj, "metadata", {}) or {}
info = {
"slug": tool_obj.name,
"name": metadata.get("name", tool_obj.name), # 显示名称优先从 metadata 获取
"description": tool_obj.description,
"metadata": metadata,
"args": [],
}
if hasattr(tool_obj, "args_schema") and tool_obj.args_schema:
schema = tool_obj.args_schema
# 优先从 pydantic v2 ``model_fields`` 提取参数信息,避免对
# ``ToolRuntime`` 等含 callable 字段(如 ``stream_writer``)的注入式
# 参数生成完整 JSON schema 时抛出 ``PydanticInvalidForJsonSchema``。
# ``runtime`` 字段由 langgraph 自动注入,不属于用户输入参数,跳过展示。
model_fields = getattr(schema, "model_fields", None)
if model_fields:
for arg_name, field_info in model_fields.items():
if arg_name == "runtime":
continue
annotation = field_info.annotation
info["args"].append(
{
"name": arg_name,
"type": getattr(annotation, "__name__", str(annotation)),
"description": field_info.description or "",
}
)
elif hasattr(schema, "schema"):
schema_dict = schema.schema()
for arg_name, arg_info in schema_dict.get("properties", {}).items():
if arg_name == "runtime":
continue
info["args"].append(
{
"name": arg_name,
"type": arg_info.get("type", ""),
"description": arg_info.get("description", ""),
}
)
return info
def _ensure_metadata_loaded():
"""延迟加载工具元数据(首次调用时自动触发)"""
global _metadata_cache
if _metadata_cache: # 已加载
return
from yuxi.agents.toolkits.registry import (
get_all_extra_metadata,
get_all_tool_instances,
)
# 获取所有工具实例
all_tools = get_all_tool_instances()
extra_meta = get_all_extra_metadata()
for tool in all_tools:
tool_name = tool.name
runtime_info = _extract_tool_info(tool)
# 合并附加元数据
if tool_name in extra_meta:
extra = extra_meta[tool_name]
runtime_info["category"] = extra.category
runtime_info["tags"] = extra.tags
runtime_info["config_guide"] = extra.config_guide
# display_name 优先级高于 tool.name
if extra.display_name:
runtime_info["name"] = extra.display_name
else:
# 未注册,设为默认分类
runtime_info["category"] = "buildin"
runtime_info["tags"] = []
runtime_info["config_guide"] = ""
_metadata_cache.append(runtime_info)
logger.info(f"Tool service loaded {len(_metadata_cache)} tools (lazy load)")
def get_tool_metadata(category: str = None) -> list[dict]:
"""获取工具元数据列表(延迟加载)"""
_ensure_metadata_loaded()
if category:
return [t for t in _metadata_cache if t.get("category") == category]
return _metadata_cache
def get_tool_instances_by_category(category: str) -> list[Any]:
from yuxi.agents.toolkits.registry import get_all_extra_metadata, get_all_tool_instances
extra_meta = get_all_extra_metadata()
tools = []
for tool in get_all_tool_instances():
tool_meta = extra_meta.get(tool.name)
tool_category = tool_meta.category if tool_meta else "buildin"
if tool_category == category:
tools.append(tool)
return tools
async def resolve_configured_runtime_tools(context) -> list[Any]:
from yuxi.agents.mcp.service import get_enabled_mcp_tools
from yuxi.external_systems.infrastructure.container import (
create_use_cases_from_db,
)
from yuxi.external_systems.use_cases.dto.tool import (
BuildRuntimeToolsInput,
)
from yuxi.storage.postgres.manager import pg_manager
selected_tools = []
selected_tool_names: set[str] = set()
buildin_tools = {tool.name: tool for tool in get_tool_instances_by_category("buildin")}
for tool_name in getattr(context, "tools", None) or []:
if not isinstance(tool_name, str) or tool_name in selected_tool_names:
continue
tool = buildin_tools.get(tool_name)
if tool is None:
logger.warning(f"Configured buildin tool not found, skip: {tool_name}")
continue
selected_tools.append(tool)
selected_tool_names.add(tool_name)
selected_mcp_servers: set[str] = set()
for server_name in getattr(context, "mcps", None) or []:
if not isinstance(server_name, str) or server_name in selected_mcp_servers:
continue
selected_mcp_servers.add(server_name)
try:
mcp_tools = await get_enabled_mcp_tools(server_name)
except Exception as e:
logger.warning(f"Failed to load configured MCP tools '{server_name}': {e}")
continue
if not mcp_tools:
logger.warning(f"Configured MCP unavailable, skip: {server_name}")
continue
for tool in mcp_tools:
if tool.name in selected_tool_names:
continue
selected_tools.append(tool)
selected_tool_names.add(tool.name)
# 追加加载外部系统工具(按 slug 匹配 context.tools 中的选择)
external_tool_names = [
name
for name in getattr(context, "tools", None) or []
if isinstance(name, str) and name not in selected_tool_names
]
if external_tool_names:
async with pg_manager.get_async_session_context() as db:
use_cases = create_use_cases_from_db(db)
# 透传 agent 上下文uid 作为 caller_idrun_id 作为 correlation_id
# 使 agent 调用外部工具时写入完整的可观测性字段
thread_id = getattr(context, "thread_id", None)
output = await use_cases.tool_service.build_runtime_tools(
BuildRuntimeToolsInput(
slugs=external_tool_names,
caller_id=getattr(context, "uid", None),
correlation_id=getattr(context, "run_id", None),
tags={"thread_id": thread_id} if thread_id else {},
),
)
selected_tools.extend(output.items)
selected_tool_names.update(tool.name for tool in output.items)
# 构建失败的工具 slug 必须显式记录,避免用户配置的工具被静默丢失
if output.failed_slugs:
logger.warning(f"Failed to build external runtime tools, skipped: {output.failed_slugs}")
return selected_tools