refactor: improve astrbot builtin tool management (#7418)

* feat: Refactor astrbot builtin tool management

- Introduced a new registry for builtin tools to streamline their management.
- Added `SendMessageToUserTool`, `KnowledgeBaseQueryTool`, and various web search tools as builtin tools.
- Updated `FunctionToolManager` to cache and retrieve builtin tools efficiently.
- Modified `CronJobManager` to utilize the new `SendMessageToUserTool`.
- Enhanced the dashboard to display readonly status for builtin tools and prevent toggling their state.
- Added tests for builtin tool injection and retrieval to ensure proper functionality.

* fix: escape file path in shell command to prevent injection vulnerabilities
This commit is contained in:
Soulter
2026-04-08 14:49:58 +08:00
committed by GitHub
parent 94a529d3fd
commit 38b1b4d4ea
20 changed files with 648 additions and 446 deletions
+4 -2
View File
@@ -25,7 +25,6 @@ from astrbot.core.astr_main_agent_resources import (
LOCAL_EXECUTE_SHELL_TOOL,
LOCAL_PYTHON_TOOL,
PYTHON_TOOL,
SEND_MESSAGE_TO_USER_TOOL,
)
from astrbot.core.cron.events import CronMessageEvent
from astrbot.core.message.components import Image
@@ -37,6 +36,7 @@ from astrbot.core.message.message_event_result import (
from astrbot.core.platform.message_session import MessageSession
from astrbot.core.provider.entites import ProviderRequest
from astrbot.core.provider.register import llm_tools
from astrbot.core.tools.message_tools import SendMessageToUserTool
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from astrbot.core.utils.history_saver import persist_agent_history
from astrbot.core.utils.image_ref_utils import is_supported_image_ref
@@ -515,7 +515,9 @@ class FunctionToolExecutor(BaseFunctionToolExecutor[AstrAgentContext]):
)
if not req.func_tool:
req.func_tool = ToolSet()
req.func_tool.add_tool(SEND_MESSAGE_TO_USER_TOOL)
req.func_tool.add_tool(
ctx.get_llm_tool_manager().get_builtin_tool(SendMessageToUserTool)
)
result = await build_main_agent(
event=cron_event, plugin_context=ctx, config=config, req=req
+35 -23
View File
@@ -32,7 +32,6 @@ from astrbot.core.astr_main_agent_resources import (
FILE_UPLOAD_TOOL,
GET_EXECUTION_HISTORY_TOOL,
GET_SKILL_PAYLOAD_TOOL,
KNOWLEDGE_BASE_QUERY_TOOL,
LIST_SKILL_CANDIDATES_TOOL,
LIST_SKILL_RELEASES_TOOL,
LIVE_MODE_SYSTEM_PROMPT,
@@ -44,11 +43,9 @@ from astrbot.core.astr_main_agent_resources import (
ROLLBACK_SKILL_RELEASE_TOOL,
RUN_BROWSER_SKILL_TOOL,
SANDBOX_MODE_PROMPT,
SEND_MESSAGE_TO_USER_TOOL,
SYNC_SKILL_RELEASE_TOOL,
TOOL_CALL_PROMPT,
TOOL_CALL_PROMPT_SKILLS_LIKE_MODE,
retrieve_knowledge_base,
)
from astrbot.core.conversation_mgr import Conversation
from astrbot.core.message.components import File, Image, Record, Reply
@@ -63,16 +60,21 @@ from astrbot.core.skills.skill_manager import SkillManager, build_skills_prompt
from astrbot.core.star.context import Context
from astrbot.core.star.star_handler import star_map
from astrbot.core.tools.cron_tools import (
CREATE_CRON_JOB_TOOL,
DELETE_CRON_JOB_TOOL,
LIST_CRON_JOBS_TOOL,
CreateActiveCronTool,
DeleteCronJobTool,
ListCronJobsTool,
)
from astrbot.core.tools.knowledge_base_tools import (
KnowledgeBaseQueryTool,
retrieve_knowledge_base,
)
from astrbot.core.tools.message_tools import SendMessageToUserTool
from astrbot.core.tools.web_search_tools import (
TAVILY_EXTRACT_WEB_PAGE_TOOL,
WEB_SEARCH_BAIDU_TOOL,
WEB_SEARCH_BOCHA_TOOL,
WEB_SEARCH_BRAVE_TOOL,
WEB_SEARCH_TAVILY_TOOL,
BaiduWebSearchTool,
BochaWebSearchTool,
BraveWebSearchTool,
TavilyExtractWebPageTool,
TavilyWebSearchTool,
normalize_legacy_web_search_config,
)
from astrbot.core.utils.file_extract import extract_file_moonshotai
@@ -226,7 +228,11 @@ async def _apply_kb(
else:
if req.func_tool is None:
req.func_tool = ToolSet()
req.func_tool.add_tool(KNOWLEDGE_BASE_QUERY_TOOL)
req.func_tool.add_tool(
plugin_context.get_llm_tool_manager().get_builtin_tool(
KnowledgeBaseQueryTool
)
)
async def _apply_file_extract(
@@ -1054,12 +1060,13 @@ def _apply_sandbox_tools(
req.system_prompt = f"{req.system_prompt or ''}\n{SANDBOX_MODE_PROMPT}\n"
def _proactive_cron_job_tools(req: ProviderRequest) -> None:
def _proactive_cron_job_tools(req: ProviderRequest, plugin_context: Context) -> None:
if req.func_tool is None:
req.func_tool = ToolSet()
req.func_tool.add_tool(CREATE_CRON_JOB_TOOL)
req.func_tool.add_tool(DELETE_CRON_JOB_TOOL)
req.func_tool.add_tool(LIST_CRON_JOBS_TOOL)
tool_mgr = plugin_context.get_llm_tool_manager()
req.func_tool.add_tool(tool_mgr.get_builtin_tool(CreateActiveCronTool))
req.func_tool.add_tool(tool_mgr.get_builtin_tool(DeleteCronJobTool))
req.func_tool.add_tool(tool_mgr.get_builtin_tool(ListCronJobsTool))
async def _apply_web_search_tools(
@@ -1077,16 +1084,17 @@ async def _apply_web_search_tools(
if req.func_tool is None:
req.func_tool = ToolSet()
tool_mgr = plugin_context.get_llm_tool_manager()
provider = prov_settings.get("websearch_provider", "tavily")
if provider == "tavily":
req.func_tool.add_tool(WEB_SEARCH_TAVILY_TOOL)
req.func_tool.add_tool(TAVILY_EXTRACT_WEB_PAGE_TOOL)
req.func_tool.add_tool(tool_mgr.get_builtin_tool(TavilyWebSearchTool))
req.func_tool.add_tool(tool_mgr.get_builtin_tool(TavilyExtractWebPageTool))
elif provider == "bocha":
req.func_tool.add_tool(WEB_SEARCH_BOCHA_TOOL)
req.func_tool.add_tool(tool_mgr.get_builtin_tool(BochaWebSearchTool))
elif provider == "brave":
req.func_tool.add_tool(WEB_SEARCH_BRAVE_TOOL)
req.func_tool.add_tool(tool_mgr.get_builtin_tool(BraveWebSearchTool))
elif provider == "baidu_ai_search":
req.func_tool.add_tool(WEB_SEARCH_BAIDU_TOOL)
req.func_tool.add_tool(tool_mgr.get_builtin_tool(BaiduWebSearchTool))
def _get_compress_provider(
@@ -1348,12 +1356,16 @@ async def build_main_agent(
)
if config.add_cron_tools:
_proactive_cron_job_tools(req)
_proactive_cron_job_tools(req, plugin_context)
if event.platform_meta.support_proactive_message:
if req.func_tool is None:
req.func_tool = ToolSet()
req.func_tool.add_tool(SEND_MESSAGE_TO_USER_TOOL)
req.func_tool.add_tool(
plugin_context.get_llm_tool_manager().get_builtin_tool(
SendMessageToUserTool
)
)
if provider.provider_config.get("max_context_tokens", 0) <= 0:
model = provider.get_model()
-363
View File
@@ -1,17 +1,5 @@
import base64
import json
import os
import uuid
from pydantic import Field
from pydantic.dataclasses import dataclass
import astrbot.core.message.components as Comp
from astrbot.api import logger, sp
from astrbot.core.agent.run_context import ContextWrapper
from astrbot.core.agent.tool import FunctionTool, ToolExecResult
from astrbot.core.astr_agent_context import AstrAgentContext
from astrbot.core.computer.computer_client import get_booter
from astrbot.core.computer.tools import (
AnnotateExecutionTool,
BrowserBatchExecTool,
@@ -33,11 +21,6 @@ from astrbot.core.computer.tools import (
RunBrowserSkillTool,
SyncSkillReleaseTool,
)
from astrbot.core.knowledge_base.kb_helper import KBHelper
from astrbot.core.message.message_event_result import MessageChain
from astrbot.core.platform.message_session import MessageSession
from astrbot.core.star.context import Context
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
LLM_SAFETY_MODE_SYSTEM_PROMPT = """You are running in Safe Mode.
@@ -148,352 +131,6 @@ BACKGROUND_TASK_RESULT_WOKE_SYSTEM_PROMPT = (
)
@dataclass
class KnowledgeBaseQueryTool(FunctionTool[AstrAgentContext]):
name: str = "astr_kb_search"
description: str = (
"Query the knowledge base for facts or relevant context. "
"Use this tool when the user's question requires factual information, "
"definitions, background knowledge, or previously indexed content. "
"Only send short keywords or a concise question as the query."
)
parameters: dict = Field(
default_factory=lambda: {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "A concise keyword query for the knowledge base.",
},
},
"required": ["query"],
}
)
async def call(
self, context: ContextWrapper[AstrAgentContext], **kwargs
) -> ToolExecResult:
query = kwargs.get("query", "")
if not query:
return "error: Query parameter is empty."
result = await retrieve_knowledge_base(
query=kwargs.get("query", ""),
umo=context.context.event.unified_msg_origin,
context=context.context.context,
)
if not result:
return "No relevant knowledge found."
return result
@dataclass
class SendMessageToUserTool(FunctionTool[AstrAgentContext]):
name: str = "send_message_to_user"
description: str = (
"Send message to the user. "
"Supports various message types including `plain`, `image`, `record`, `video`, `file`, and `mention_user`. "
"Use this tool to send media files (`image`, `record`, `video`, `file`), "
"or when you need to proactively message the user(such as cron job). For normal text replies, you can output directly."
)
parameters: dict = Field(
default_factory=lambda: {
"type": "object",
"properties": {
"messages": {
"type": "array",
"description": "An ordered list of message components to send. `mention_user` type can be used to mention the user.",
"items": {
"type": "object",
"properties": {
"type": {
"type": "string",
"description": (
"Component type. One of: "
"plain, image, record, video, file, mention_user. Record is voice message."
),
},
"text": {
"type": "string",
"description": "Text content for `plain` type.",
},
"path": {
"type": "string",
"description": "File path for `image`, `record`, or `file` types. Both local path and sandbox path are supported.",
},
"url": {
"type": "string",
"description": "URL for `image`, `record`, or `file` types.",
},
"mention_user_id": {
"type": "string",
"description": "User ID to mention for `mention_user` type.",
},
},
"required": ["type"],
},
},
},
"required": ["messages"],
}
)
async def _resolve_path_from_sandbox(
self, context: ContextWrapper[AstrAgentContext], path: str
) -> tuple[str, bool]:
"""
If the path exists locally, return it directly.
Otherwise, check if it exists in the sandbox and download it.
bool: indicates whether the file was downloaded from sandbox.
"""
if os.path.exists(path):
return path, False
# Try to check if the file exists in the sandbox
try:
sb = await get_booter(
context.context.context,
context.context.event.unified_msg_origin,
)
# Use shell to check if the file exists in sandbox
result = await sb.shell.exec(f"test -f {path} && echo '_&exists_'")
if "_&exists_" in json.dumps(result):
# Download the file from sandbox
name = os.path.basename(path)
local_path = os.path.join(
get_astrbot_temp_path(), f"sandbox_{uuid.uuid4().hex[:4]}_{name}"
)
await sb.download_file(path, local_path)
logger.info(f"Downloaded file from sandbox: {path} -> {local_path}")
return local_path, True
except Exception as e:
logger.warning(f"Failed to check/download file from sandbox: {e}")
# Return the original path (will likely fail later, but that's expected)
return path, False
async def call(
self, context: ContextWrapper[AstrAgentContext], **kwargs
) -> ToolExecResult:
session = kwargs.get("session") or context.context.event.unified_msg_origin
messages = kwargs.get("messages")
if not isinstance(messages, list) or not messages:
return "error: messages parameter is empty or invalid."
components: list[Comp.BaseMessageComponent] = []
for idx, msg in enumerate(messages):
if not isinstance(msg, dict):
return f"error: messages[{idx}] should be an object."
msg_type = str(msg.get("type", "")).lower()
if not msg_type:
return f"error: messages[{idx}].type is required."
file_from_sandbox = False
try:
if msg_type == "plain":
text = str(msg.get("text", "")).strip()
if not text:
return f"error: messages[{idx}].text is required for plain component."
components.append(Comp.Plain(text=text))
elif msg_type == "image":
path = msg.get("path")
url = msg.get("url")
if path:
(
local_path,
file_from_sandbox,
) = await self._resolve_path_from_sandbox(context, path)
components.append(Comp.Image.fromFileSystem(path=local_path))
elif url:
components.append(Comp.Image.fromURL(url=url))
else:
return f"error: messages[{idx}] must include path or url for image component."
elif msg_type == "record":
path = msg.get("path")
url = msg.get("url")
if path:
(
local_path,
file_from_sandbox,
) = await self._resolve_path_from_sandbox(context, path)
components.append(Comp.Record.fromFileSystem(path=local_path))
elif url:
components.append(Comp.Record.fromURL(url=url))
else:
return f"error: messages[{idx}] must include path or url for record component."
elif msg_type == "video":
path = msg.get("path")
url = msg.get("url")
if path:
(
local_path,
file_from_sandbox,
) = await self._resolve_path_from_sandbox(context, path)
components.append(Comp.Video.fromFileSystem(path=local_path))
elif url:
components.append(Comp.Video.fromURL(url=url))
else:
return f"error: messages[{idx}] must include path or url for video component."
elif msg_type == "file":
path = msg.get("path")
url = msg.get("url")
name = (
msg.get("text")
or (os.path.basename(path) if path else "")
or (os.path.basename(url) if url else "")
or "file"
)
if path:
(
local_path,
file_from_sandbox,
) = await self._resolve_path_from_sandbox(context, path)
components.append(Comp.File(name=name, file=local_path))
elif url:
components.append(Comp.File(name=name, url=url))
else:
return f"error: messages[{idx}] must include path or url for file component."
elif msg_type == "mention_user":
mention_user_id = msg.get("mention_user_id")
if not mention_user_id:
return f"error: messages[{idx}].mention_user_id is required for mention_user component."
components.append(
Comp.At(
qq=mention_user_id,
),
)
else:
return (
f"error: unsupported message type '{msg_type}' at index {idx}."
)
except Exception as exc: # 捕获组件构造异常,避免直接抛出
return f"error: failed to build messages[{idx}] component: {exc}"
try:
target_session = (
MessageSession.from_str(session)
if isinstance(session, str)
else session
)
except Exception as e:
return f"error: invalid session: {e}"
await context.context.context.send_message(
target_session,
MessageChain(chain=components),
)
# if file_from_sandbox:
# try:
# os.remove(local_path)
# except Exception as e:
# logger.error(f"Error removing temp file {local_path}: {e}")
return f"Message sent to session {target_session}"
def check_all_kb(kb_list: list[KBHelper | None]) -> bool:
"""检查是否所有的知识库都为空
Args:
kb_list: 所选的知识库
Returns:
bool: 是否全为空
"""
return not any(
kb and (kb.kb.doc_count != 0 or kb.kb.chunk_count != 0) for kb in kb_list
)
async def retrieve_knowledge_base(
query: str,
umo: str,
context: Context,
) -> str | None:
"""Inject knowledge base context into the provider request
Args:
umo: Unique message object (session ID)
p_ctx: Pipeline context
"""
kb_mgr = context.kb_manager
config = context.get_config(umo=umo)
# 1. 优先读取会话级配置
session_config = await sp.session_get(umo, "kb_config", default={})
if session_config and "kb_ids" in session_config:
# 会话级配置
kb_ids = session_config.get("kb_ids", [])
# 如果配置为空列表,明确表示不使用知识库
if not kb_ids:
logger.info(f"[知识库] 会话 {umo} 已被配置为不使用知识库")
return
top_k = session_config.get("top_k", 5)
# 将 kb_ids 转换为 kb_names
kb_names = []
invalid_kb_ids = []
for kb_id in kb_ids:
kb_helper = await kb_mgr.get_kb(kb_id)
if kb_helper:
kb_names.append(kb_helper.kb.kb_name)
else:
logger.warning(f"[知识库] 知识库不存在或未加载: {kb_id}")
invalid_kb_ids.append(kb_id)
if invalid_kb_ids:
logger.warning(
f"[知识库] 会话 {umo} 配置的以下知识库无效: {invalid_kb_ids}",
)
if not kb_names:
return
logger.debug(f"[知识库] 使用会话级配置,知识库数量: {len(kb_names)}")
else:
kb_names = config.get("kb_names", [])
top_k = config.get("kb_final_top_k", 5)
logger.debug(f"[知识库] 使用全局配置,知识库数量: {len(kb_names)}")
top_k_fusion = config.get("kb_fusion_top_k", 20)
if not kb_names:
return
all_kbs = [await kb_mgr.get_kb_by_name(kb) for kb in kb_names]
if check_all_kb(all_kbs):
logger.debug("所配置的所有知识库全为空,跳过检索过程")
return
logger.debug(f"[知识库] 开始检索知识库,数量: {len(kb_names)}, top_k={top_k}")
kb_context = await kb_mgr.retrieve(
query=query,
kb_names=kb_names,
top_k_fusion=top_k_fusion,
top_m_final=top_k,
)
if not kb_context:
return
formatted = kb_context.get("context_text", "")
if formatted:
results = kb_context.get("results", [])
logger.debug(f"[知识库] 为会话 {umo} 注入了 {len(results)} 条相关知识块")
return formatted
KNOWLEDGE_BASE_QUERY_TOOL = KnowledgeBaseQueryTool()
SEND_MESSAGE_TO_USER_TOOL = SendMessageToUserTool()
EXECUTE_SHELL_TOOL = ExecuteShellTool()
LOCAL_EXECUTE_SHELL_TOOL = ExecuteShellTool(is_local=True)
PYTHON_TOOL = PythonTool()
+4 -2
View File
@@ -275,8 +275,8 @@ class CronJobManager:
)
from astrbot.core.astr_main_agent_resources import (
PROACTIVE_AGENT_CRON_WOKE_SYSTEM_PROMPT,
SEND_MESSAGE_TO_USER_TOOL,
)
from astrbot.core.tools.message_tools import SendMessageToUserTool
try:
session = (
@@ -342,7 +342,9 @@ class CronJobManager:
)
if not req.func_tool:
req.func_tool = ToolSet()
req.func_tool.add_tool(SEND_MESSAGE_TO_USER_TOOL)
req.func_tool.add_tool(
self.ctx.get_llm_tool_manager().get_builtin_tool(SendMessageToUserTool)
)
result = await build_main_agent(
event=cron_event, plugin_context=self.ctx, config=config, req=req
+53 -1
View File
@@ -17,6 +17,12 @@ from astrbot import logger
from astrbot.core import sp
from astrbot.core.agent.mcp_client import MCPClient, MCPTool
from astrbot.core.agent.tool import FunctionTool, ToolSet
from astrbot.core.tools.registry import (
ensure_builtin_tools_loaded,
get_builtin_tool_class,
get_builtin_tool_name,
iter_builtin_tool_classes,
)
from astrbot.core.utils.astrbot_path import get_astrbot_data_path
DEFAULT_MCP_CONFIG = {"mcpServers": {}}
@@ -207,8 +213,12 @@ async def _quick_test_mcp_connection(config: dict) -> tuple[bool, str]:
class FunctionToolManager:
def __init__(self) -> None:
self.func_list: list[FuncTool] = []
"""All tools include mcp tools and plugin tools, except astrbot builtin tools."""
self.builtin_func_list: dict[type[FuncTool], FuncTool] = {}
"""All astrbot builtin tools, keyed by their class. Values are instantiated tool objects, created on demand."""
self._mcp_server_runtime: dict[str, _MCPServerRuntime] = {}
"""MCP 服务运行时状态(唯一事实来源)"""
"""MCP runtime metadata, keyed by server name. Updated atomically on MCP lifecycle changes."""
self._mcp_server_runtime_view = MappingProxyType(self._mcp_server_runtime)
self._mcp_client_dict_view = _MCPClientDictView(self._mcp_server_runtime)
self._timeout_mismatch_warned = False
@@ -320,8 +330,50 @@ class FunctionToolManager:
for f in reversed(self.func_list):
if f.name == name:
return f
if isinstance(name, str):
try:
builtin_tool = self.get_builtin_tool(name)
except KeyError:
return None
if getattr(builtin_tool, "active", True):
return builtin_tool
return builtin_tool
return None
def get_builtin_tool(self, tool: str | type[FuncTool]) -> FuncTool:
ensure_builtin_tools_loaded()
if isinstance(tool, str):
tool_cls = get_builtin_tool_class(tool)
if tool_cls is None:
raise KeyError(f"Builtin tool {tool} is not registered.")
elif isinstance(tool, type) and issubclass(tool, FunctionTool):
tool_cls = tool
if get_builtin_tool_name(tool_cls) is None:
raise KeyError(
f"Builtin tool class {tool_cls.__module__}.{tool_cls.__name__} is not registered.",
)
else:
raise TypeError("tool must be a builtin tool name or FunctionTool class.")
cached_tool = self.builtin_func_list.get(tool_cls)
if cached_tool is not None:
return cached_tool
builtin_tool = tool_cls() # type: ignore
self.builtin_func_list[tool_cls] = builtin_tool
return builtin_tool
def iter_builtin_tools(self) -> list[FuncTool]:
ensure_builtin_tools_loaded()
return [
self.get_builtin_tool(tool_cls) for tool_cls in iter_builtin_tool_classes()
]
def is_builtin_tool(self, name: str) -> bool:
ensure_builtin_tools_loaded()
return get_builtin_tool_class(name) is not None
def get_full_tool_set(self) -> ToolSet:
"""获取完整工具集
+4 -7
View File
@@ -7,6 +7,7 @@ from pydantic.dataclasses import dataclass
from astrbot.core.agent.run_context import ContextWrapper
from astrbot.core.agent.tool import FunctionTool, ToolExecResult
from astrbot.core.astr_agent_context import AstrAgentContext
from astrbot.core.tools.registry import builtin_tool
def _extract_job_session(job: Any) -> str | None:
@@ -17,6 +18,7 @@ def _extract_job_session(job: Any) -> str | None:
return str(session) if session is not None else None
@builtin_tool
@dataclass
class CreateActiveCronTool(FunctionTool[AstrAgentContext]):
name: str = "create_future_task"
@@ -105,6 +107,7 @@ class CreateActiveCronTool(FunctionTool[AstrAgentContext]):
return f"Scheduled future task {job.job_id} ({job.name}) {suffix}."
@builtin_tool
@dataclass
class DeleteCronJobTool(FunctionTool[AstrAgentContext]):
name: str = "delete_future_task"
@@ -141,6 +144,7 @@ class DeleteCronJobTool(FunctionTool[AstrAgentContext]):
return f"Deleted cron job {job_id}."
@builtin_tool
@dataclass
class ListCronJobsTool(FunctionTool[AstrAgentContext]):
name: str = "list_future_tasks"
@@ -180,14 +184,7 @@ class ListCronJobsTool(FunctionTool[AstrAgentContext]):
return "\n".join(lines)
CREATE_CRON_JOB_TOOL = CreateActiveCronTool()
DELETE_CRON_JOB_TOOL = DeleteCronJobTool()
LIST_CRON_JOBS_TOOL = ListCronJobsTool()
__all__ = [
"CREATE_CRON_JOB_TOOL",
"DELETE_CRON_JOB_TOOL",
"LIST_CRON_JOBS_TOOL",
"CreateActiveCronTool",
"DeleteCronJobTool",
"ListCronJobsTool",
+129
View File
@@ -0,0 +1,129 @@
from pydantic import Field
from pydantic.dataclasses import dataclass
from astrbot.api import logger, sp
from astrbot.core.agent.run_context import ContextWrapper
from astrbot.core.agent.tool import FunctionTool, ToolExecResult
from astrbot.core.astr_agent_context import AstrAgentContext
from astrbot.core.knowledge_base.kb_helper import KBHelper
from astrbot.core.star.context import Context
from astrbot.core.tools.registry import builtin_tool
def check_all_kb(kb_list: list[KBHelper | None]) -> bool:
"""检查是否所有的知识库都为空"""
return not any(
kb and (kb.kb.doc_count != 0 or kb.kb.chunk_count != 0) for kb in kb_list
)
async def retrieve_knowledge_base(
query: str,
umo: str,
context: Context,
) -> str | None:
"""Retrieve knowledge base context for the given query."""
kb_mgr = context.kb_manager
config = context.get_config(umo=umo)
session_config = await sp.session_get(umo, "kb_config", default={})
if session_config and "kb_ids" in session_config:
kb_ids = session_config.get("kb_ids", [])
if not kb_ids:
logger.info(f"[知识库] 会话 {umo} 已被配置为不使用知识库")
return None
top_k = session_config.get("top_k", 5)
kb_names = []
invalid_kb_ids = []
for kb_id in kb_ids:
kb_helper = await kb_mgr.get_kb(kb_id)
if kb_helper:
kb_names.append(kb_helper.kb.kb_name)
else:
logger.warning(f"[知识库] 知识库不存在或未加载: {kb_id}")
invalid_kb_ids.append(kb_id)
if invalid_kb_ids:
logger.warning(
f"[知识库] 会话 {umo} 配置的以下知识库无效: {invalid_kb_ids}",
)
if not kb_names:
return None
logger.debug(f"[知识库] 使用会话级配置,知识库数量: {len(kb_names)}")
else:
kb_names = config.get("kb_names", [])
top_k = config.get("kb_final_top_k", 5)
logger.debug(f"[知识库] 使用全局配置,知识库数量: {len(kb_names)}")
top_k_fusion = config.get("kb_fusion_top_k", 20)
if not kb_names:
return None
all_kbs = [await kb_mgr.get_kb_by_name(kb) for kb in kb_names]
if check_all_kb(all_kbs):
logger.debug("所配置的所有知识库全为空,跳过检索过程")
return None
logger.debug(f"[知识库] 开始检索知识库,数量: {len(kb_names)}, top_k={top_k}")
kb_context = await kb_mgr.retrieve(
query=query,
kb_names=kb_names,
top_k_fusion=top_k_fusion,
top_m_final=top_k,
)
if not kb_context:
return None
formatted = kb_context.get("context_text", "")
if formatted:
results = kb_context.get("results", [])
logger.debug(f"[知识库] 为会话 {umo} 注入了 {len(results)} 条相关知识块")
return formatted
return None
@builtin_tool
@dataclass
class KnowledgeBaseQueryTool(FunctionTool[AstrAgentContext]):
name: str = "astr_kb_search"
description: str = (
"Query the knowledge base for facts or relevant context. "
"Use this tool when the user's question requires factual information, "
"definitions, background knowledge, or previously indexed content. "
"Only send short keywords or a concise question as the query."
)
parameters: dict = Field(
default_factory=lambda: {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "A concise keyword query for the knowledge base.",
},
},
"required": ["query"],
}
)
async def call(
self, context: ContextWrapper[AstrAgentContext], **kwargs
) -> ToolExecResult:
query = kwargs.get("query", "")
if not query:
return "error: Query parameter is empty."
result = await retrieve_knowledge_base(
query=query,
umo=context.context.event.unified_msg_origin,
context=context.context.context,
)
if not result:
return "No relevant knowledge found."
return result
__all__ = [
"KnowledgeBaseQueryTool",
"check_all_kb",
"retrieve_knowledge_base",
]
+210
View File
@@ -0,0 +1,210 @@
import json
import os
import shlex
import uuid
from pydantic import Field
from pydantic.dataclasses import dataclass
import astrbot.core.message.components as Comp
from astrbot.api import logger
from astrbot.core.agent.run_context import ContextWrapper
from astrbot.core.agent.tool import FunctionTool, ToolExecResult
from astrbot.core.astr_agent_context import AstrAgentContext
from astrbot.core.computer.computer_client import get_booter
from astrbot.core.message.message_event_result import MessageChain
from astrbot.core.platform.message_session import MessageSession
from astrbot.core.tools.registry import builtin_tool
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
@builtin_tool
@dataclass
class SendMessageToUserTool(FunctionTool[AstrAgentContext]):
name: str = "send_message_to_user"
description: str = (
"Send message to the user. "
"Supports various message types including `plain`, `image`, `record`, `video`, `file`, and `mention_user`. "
"Use this tool to send media files (`image`, `record`, `video`, `file`), "
"or when you need to proactively message the user(such as cron job). For normal text replies, you can output directly."
)
parameters: dict = Field(
default_factory=lambda: {
"type": "object",
"properties": {
"messages": {
"type": "array",
"description": "An ordered list of message components to send. `mention_user` type can be used to mention the user.",
"items": {
"type": "object",
"properties": {
"type": {
"type": "string",
"description": (
"Component type. One of: "
"plain, image, record, video, file, mention_user. Record is voice message."
),
},
"text": {
"type": "string",
"description": "Text content for `plain` type.",
},
"path": {
"type": "string",
"description": "File path for `image`, `record`, `video`, or `file` types. Both local path and sandbox path are supported.",
},
"url": {
"type": "string",
"description": "URL for `image`, `record`, `video`, or `file` types.",
},
"mention_user_id": {
"type": "string",
"description": "User ID to mention for `mention_user` type.",
},
},
"required": ["type"],
},
},
"session": {
"type": "string",
"description": "Optional. Target session string. Defaults to current session.",
},
},
"required": ["messages"],
}
)
async def _resolve_path_from_sandbox(
self, context: ContextWrapper[AstrAgentContext], path: str
) -> tuple[str, bool]:
if os.path.exists(path):
return path, False
try:
sb = await get_booter(
context.context.context,
context.context.event.unified_msg_origin,
)
quoted_path = shlex.quote(path)
result = await sb.shell.exec(f"test -f {quoted_path} && echo '_&exists_'")
if "_&exists_" in json.dumps(result):
name = os.path.basename(path)
local_path = os.path.join(
get_astrbot_temp_path(), f"sandbox_{uuid.uuid4().hex[:4]}_{name}"
)
await sb.download_file(path, local_path)
logger.info(f"Downloaded file from sandbox: {path} -> {local_path}")
return local_path, True
except Exception as exc:
logger.warning(f"Failed to check/download file from sandbox: {exc}")
return path, False
async def call(
self, context: ContextWrapper[AstrAgentContext], **kwargs
) -> ToolExecResult:
session = kwargs.get("session") or context.context.event.unified_msg_origin
messages = kwargs.get("messages")
if not isinstance(messages, list) or not messages:
return "error: messages parameter is empty or invalid."
components: list[Comp.BaseMessageComponent] = []
for idx, msg in enumerate(messages):
if not isinstance(msg, dict):
return f"error: messages[{idx}] should be an object."
msg_type = str(msg.get("type", "")).lower()
if not msg_type:
return f"error: messages[{idx}].type is required."
try:
if msg_type == "plain":
text = str(msg.get("text", "")).strip()
if not text:
return f"error: messages[{idx}].text is required for plain component."
components.append(Comp.Plain(text=text))
elif msg_type == "image":
path = msg.get("path")
url = msg.get("url")
if path:
local_path, _ = await self._resolve_path_from_sandbox(
context, path
)
components.append(Comp.Image.fromFileSystem(path=local_path))
elif url:
components.append(Comp.Image.fromURL(url=url))
else:
return f"error: messages[{idx}] must include path or url for image component."
elif msg_type == "record":
path = msg.get("path")
url = msg.get("url")
if path:
local_path, _ = await self._resolve_path_from_sandbox(
context, path
)
components.append(Comp.Record.fromFileSystem(path=local_path))
elif url:
components.append(Comp.Record.fromURL(url=url))
else:
return f"error: messages[{idx}] must include path or url for record component."
elif msg_type == "video":
path = msg.get("path")
url = msg.get("url")
if path:
local_path, _ = await self._resolve_path_from_sandbox(
context, path
)
components.append(Comp.Video.fromFileSystem(path=local_path))
elif url:
components.append(Comp.Video.fromURL(url=url))
else:
return f"error: messages[{idx}] must include path or url for video component."
elif msg_type == "file":
path = msg.get("path")
url = msg.get("url")
name = (
msg.get("text")
or (os.path.basename(path) if path else "")
or (os.path.basename(url) if url else "")
or "file"
)
if path:
local_path, _ = await self._resolve_path_from_sandbox(
context, path
)
components.append(Comp.File(name=name, file=local_path))
elif url:
components.append(Comp.File(name=name, url=url))
else:
return f"error: messages[{idx}] must include path or url for file component."
elif msg_type == "mention_user":
mention_user_id = msg.get("mention_user_id")
if not mention_user_id:
return f"error: messages[{idx}].mention_user_id is required for mention_user component."
components.append(Comp.At(qq=mention_user_id))
else:
return (
f"error: unsupported message type '{msg_type}' at index {idx}."
)
except Exception as exc:
return f"error: failed to build messages[{idx}] component: {exc}"
try:
target_session = (
MessageSession.from_str(session)
if isinstance(session, str)
else session
)
except Exception as exc:
return f"error: invalid session: {exc}"
await context.context.context.send_message(
target_session,
MessageChain(chain=components),
)
return f"Message sent to session {target_session}"
__all__ = [
"SendMessageToUserTool",
]
+83
View File
@@ -0,0 +1,83 @@
from __future__ import annotations
from importlib import import_module
from typing import TypeVar
from astrbot.core.agent.tool import FunctionTool
TFunctionTool = TypeVar("TFunctionTool", bound=type[FunctionTool])
_BUILTIN_TOOL_MODULES = (
"astrbot.core.tools.cron_tools",
"astrbot.core.tools.knowledge_base_tools",
"astrbot.core.tools.message_tools",
"astrbot.core.tools.web_search_tools",
)
_builtin_tool_classes_by_name: dict[str, type[FunctionTool]] = {}
_builtin_tool_names_by_class: dict[type[FunctionTool], str] = {}
_builtin_tools_loaded = False
def _resolve_builtin_tool_name(tool_cls: type[FunctionTool]) -> str:
tool_name = getattr(tool_cls, "name", None)
if isinstance(tool_name, str) and tool_name:
return tool_name
dataclass_fields = getattr(tool_cls, "__dataclass_fields__", {})
name_field = dataclass_fields.get("name")
if name_field is not None and isinstance(name_field.default, str):
return name_field.default
raise ValueError(
f"Builtin tool class {tool_cls.__module__}.{tool_cls.__name__} does not define a valid name.",
)
def builtin_tool(tool_cls: TFunctionTool) -> TFunctionTool:
tool_name = _resolve_builtin_tool_name(tool_cls)
existing = _builtin_tool_classes_by_name.get(tool_name)
if existing is not None and existing is not tool_cls:
raise ValueError(
f"Builtin tool name conflict detected: {tool_name} is already registered by "
f"{existing.__module__}.{existing.__name__}.",
)
_builtin_tool_classes_by_name[tool_name] = tool_cls
_builtin_tool_names_by_class[tool_cls] = tool_name
return tool_cls
def ensure_builtin_tools_loaded() -> None:
global _builtin_tools_loaded
if _builtin_tools_loaded:
return
for module_name in _BUILTIN_TOOL_MODULES:
import_module(module_name)
_builtin_tools_loaded = True
def get_builtin_tool_class(name: str) -> type[FunctionTool] | None:
ensure_builtin_tools_loaded()
return _builtin_tool_classes_by_name.get(name)
def get_builtin_tool_name(tool_cls: type[FunctionTool]) -> str | None:
ensure_builtin_tools_loaded()
return _builtin_tool_names_by_class.get(tool_cls)
def iter_builtin_tool_classes() -> tuple[type[FunctionTool], ...]:
ensure_builtin_tools_loaded()
return tuple(_builtin_tool_classes_by_name.values())
__all__ = [
"builtin_tool",
"ensure_builtin_tools_loaded",
"get_builtin_tool_class",
"get_builtin_tool_name",
"iter_builtin_tool_classes",
]
+11 -11
View File
@@ -11,6 +11,7 @@ from pydantic.dataclasses import dataclass as pydantic_dataclass
from astrbot.core import logger, sp
from astrbot.core.agent.tool import FunctionTool, ToolExecResult
from astrbot.core.astr_agent_context import AstrAgentContext
from astrbot.core.tools.registry import builtin_tool
WEB_SEARCH_TOOL_NAMES = [
"web_search_baidu",
@@ -275,6 +276,7 @@ async def _baidu_search(
]
@builtin_tool
@pydantic_dataclass
class TavilyWebSearchTool(FunctionTool[AstrAgentContext]):
name: str = "web_search_tavily"
@@ -357,6 +359,7 @@ class TavilyWebSearchTool(FunctionTool[AstrAgentContext]):
return _search_result_payload(results)
@builtin_tool
@pydantic_dataclass
class TavilyExtractWebPageTool(FunctionTool[AstrAgentContext]):
name: str = "tavily_extract_web_page"
@@ -403,6 +406,7 @@ class TavilyExtractWebPageTool(FunctionTool[AstrAgentContext]):
return ret or "Error: Tavily web searcher does not return any results."
@builtin_tool
@pydantic_dataclass
class BochaWebSearchTool(FunctionTool[AstrAgentContext]):
name: str = "web_search_bocha"
@@ -466,6 +470,7 @@ class BochaWebSearchTool(FunctionTool[AstrAgentContext]):
return _search_result_payload(results)
@builtin_tool
@pydantic_dataclass
class BraveWebSearchTool(FunctionTool[AstrAgentContext]):
name: str = "web_search_brave"
@@ -523,6 +528,7 @@ class BraveWebSearchTool(FunctionTool[AstrAgentContext]):
return _search_result_payload(results)
@builtin_tool
@pydantic_dataclass
class BaiduWebSearchTool(FunctionTool[AstrAgentContext]):
name: str = "web_search_baidu"
@@ -585,18 +591,12 @@ class BaiduWebSearchTool(FunctionTool[AstrAgentContext]):
return _search_result_payload(results)
WEB_SEARCH_TAVILY_TOOL = TavilyWebSearchTool()
TAVILY_EXTRACT_WEB_PAGE_TOOL = TavilyExtractWebPageTool()
WEB_SEARCH_BOCHA_TOOL = BochaWebSearchTool()
WEB_SEARCH_BRAVE_TOOL = BraveWebSearchTool()
WEB_SEARCH_BAIDU_TOOL = BaiduWebSearchTool()
__all__ = [
"WEB_SEARCH_BAIDU_TOOL",
"WEB_SEARCH_BOCHA_TOOL",
"WEB_SEARCH_BRAVE_TOOL",
"WEB_SEARCH_TAVILY_TOOL",
"TAVILY_EXTRACT_WEB_PAGE_TOOL",
"BaiduWebSearchTool",
"BochaWebSearchTool",
"BraveWebSearchTool",
"TavilyExtractWebPageTool",
"TavilyWebSearchTool",
"WEB_SEARCH_TOOL_NAMES",
"normalize_legacy_web_search_config",
]
+20 -2
View File
@@ -428,10 +428,20 @@ class ToolsRoute(Route):
async def get_tool_list(self):
"""Get all registered tools."""
try:
tools = self.tool_mgr.func_list
tools = list(self.tool_mgr.func_list)
existing_names = {tool.name for tool in tools}
for tool in self.tool_mgr.iter_builtin_tools():
if tool.name not in existing_names:
tools.append(tool)
tools_dict = []
for tool in tools:
if isinstance(tool, MCPTool):
readonly = False
if self.tool_mgr.is_builtin_tool(tool.name):
origin = "builtin"
origin_name = "AstrBot Core"
readonly = True
elif isinstance(tool, MCPTool):
origin = "mcp"
origin_name = tool.mcp_server_name
elif tool.handler_module_path and star_map.get(
@@ -451,6 +461,7 @@ class ToolsRoute(Route):
"active": tool.active,
"origin": origin,
"origin_name": origin_name,
"readonly": readonly,
}
tools_dict.append(tool_info)
return Response().ok(data=tools_dict).__dict__
@@ -472,6 +483,13 @@ class ToolsRoute(Route):
.__dict__
)
if self.tool_mgr.is_builtin_tool(tool_name):
return (
Response()
.error("Builtin tools are read-only and cannot be toggled.")
.__dict__
)
if action:
try:
ok = self.tool_mgr.activate_llm_tool(tool_name, star_map=star_map)
@@ -4,7 +4,6 @@ import { useModuleI18n } from '@/i18n/composables';
import type { ToolItem } from '../types';
const { tm: tmTool } = useModuleI18n('features/tooluse');
const { tm: tmCommand } = useModuleI18n('features/command');
const props = defineProps<{
items: ToolItem[];
@@ -16,11 +15,10 @@ const emit = defineEmits<{
}>();
const toolHeaders = computed(() => [
{ title: tmTool('functionTools.title'), key: 'name', minWidth: '160px' },
{ title: tmTool('functionTools.title'), key: 'name', minWidth: '240px' },
{ title: tmTool('functionTools.description'), key: 'description' },
{ title: tmTool('functionTools.table.origin'), key: 'origin', sortable: false, width: '120px' },
{ title: tmTool('functionTools.table.originName'), key: 'origin_name', sortable: false, width: '160px' },
{ title: tmCommand('status.enabled'), key: 'active', sortable: false, width: '120px' },
{ title: tmTool('functionTools.table.actions'), key: 'actions', sortable: false, width: '120px' }
]);
@@ -39,13 +37,8 @@ const parameterEntries = (tool: ToolItem) => Object.entries(tool.parameters?.pro
:loading="props.loading"
>
<template #item.name="{ item }">
<div class="d-flex align-center py-2">
<v-icon color="primary" class="mr-2" size="18">
{{ item.name.includes(':') ? 'mdi-server-network' : 'mdi-function-variant' }}
</v-icon>
<div>
<div class="text-subtitle-1 font-weight-medium">{{ item.name }}</div>
</div>
<div class="py-2">
<div class="tool-name text-body-2 font-weight-medium">{{ item.name }}</div>
</div>
</template>
@@ -56,7 +49,7 @@ const parameterEntries = (tool: ToolItem) => Object.entries(tool.parameters?.pro
</template>
<template #item.origin="{ item }">
<v-chip size="small" variant="tonal" color="info" class="text-caption font-weight-medium">
<v-chip size="x-small" variant="tonal" color="info" class="text-caption font-weight-medium">
{{ item.origin || '-' }}
</v-chip>
</template>
@@ -67,14 +60,10 @@ const parameterEntries = (tool: ToolItem) => Object.entries(tool.parameters?.pro
</div>
</template>
<template #item.active="{ item }">
<v-chip :color="item.active ? 'success' : 'error'" size="small" class="font-weight-medium" :variant="item.active ? 'flat' : 'outlined'">
{{ item.active ? tmCommand('status.enabled') : tmCommand('status.disabled') }}
</v-chip>
</template>
<template #item.actions="{ item }">
<span v-if="item.readonly" class="text-medium-emphasis">-</span>
<v-switch
v-else
:model-value="item.active"
color="primary"
density="compact"
@@ -141,4 +130,9 @@ const parameterEntries = (tool: ToolItem) => Object.entries(tool.parameters?.pro
.tool-table :deep(.v-data-table__td) {
vertical-align: middle;
}
.tool-name {
font-size: 0.9rem;
line-height: 1.35;
}
</style>
@@ -80,4 +80,3 @@ export function useComponentData() {
fetchTools
};
}
@@ -102,6 +102,10 @@ const handleUpdatePermission = async (cmd: CommandItem, permission: 'admin' | 'm
};
const handleToggleTool = async (tool: ToolItem) => {
if (tool.readonly) {
toast(tmTool('messages.toggleToolReadonly'), 'info');
return;
}
const previous = tool.active;
tool.active = !tool.active;
try {
@@ -264,19 +268,6 @@ watch(viewMode, async (mode) => {
clearable
/>
</div>
<div class="d-flex align-center ga-2">
<div class="d-flex align-center">
<v-icon size="18" color="primary" class="mr-1">mdi-function-variant</v-icon>
<span class="text-body-2 text-medium-emphasis mr-1">{{ tm('summary.total') }}:</span>
<span class="text-body-1 font-weight-bold text-primary">{{ filteredTools.length }}</span>
</div>
<v-divider vertical class="mx-1" style="height: 20px;" />
<div class="d-flex align-center">
<v-icon size="18" color="success" class="mr-1">mdi-check-circle-outline</v-icon>
<span class="text-body-2 text-medium-emphasis mr-1">{{ tm('status.enabled') }}:</span>
<span class="text-body-1 font-weight-bold text-success">{{ filteredTools.filter(t => t.active).length }}</span>
</div>
</div>
</div>
<ToolTable
@@ -94,10 +94,10 @@ export interface ToolItem {
name: string;
description: string;
active: boolean;
readonly?: boolean;
parameters?: {
properties?: Record<string, ToolParameter>;
};
origin?: string;
origin_name?: string;
}
@@ -45,6 +45,7 @@
"required": "Required",
"origin": "Origin",
"originName": "Origin Name",
"readonly": "Read-only",
"actions": "Actions"
}
},
@@ -153,6 +154,7 @@
"configParseError": "Configuration parse error: {error}",
"noAvailableConfig": "No available configuration",
"toggleToolSuccess": "Tool status toggled successfully!",
"toggleToolReadonly": "Builtin tools are read-only and cannot be enabled or disabled.",
"toggleToolError": "Failed to toggle tool status: {error}",
"testError": "Test connection failed: {error}"
}
@@ -45,6 +45,7 @@
"required": "Обяз.",
"origin": "Источник",
"originName": "Имя источника",
"readonly": "Только чтение",
"actions": "Действия"
}
},
@@ -153,6 +154,7 @@
"configParseError": "Ошибка разбора конфигурации: {error}",
"noAvailableConfig": "Конфигурация отсутствует",
"toggleToolSuccess": "Статус инструмента изменен!",
"toggleToolReadonly": "Встроенные инструменты доступны только для чтения и не могут быть включены или выключены.",
"toggleToolError": "Не удалось изменить статус: {error}",
"testError": "Ошибка теста связи: {error}"
},
@@ -192,4 +194,4 @@
"tokenHelp": "Как получить токен доступа ModelScope? Нажмите кнопку справа для получения инструкций"
}
}
}
}
@@ -45,6 +45,7 @@
"required": "必填",
"origin": "来源",
"originName": "来源名称",
"readonly": "只读",
"actions": "操作"
}
},
@@ -153,7 +154,8 @@
"configParseError": "配置解析错误: {error}",
"noAvailableConfig": "无可用配置",
"toggleToolSuccess": "工具状态切换成功!",
"toggleToolReadonly": "内置工具为只读,无法进行启用或停用操作。",
"toggleToolError": "工具状态切换失败: {error}",
"testError": "测试连接失败: {error}"
}
}
}
+57 -1
View File
@@ -1,7 +1,7 @@
"""Tests for astr_main_agent module."""
import os
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import AsyncMock, MagicMock, call, patch
import pytest
@@ -368,6 +368,62 @@ class TestApplyKb:
assert req.func_tool is not None
class TestBuiltinToolInjection:
"""Tests for builtin tool injection paths."""
@pytest.mark.asyncio
async def test_apply_web_search_tools_uses_builtin_tool_manager(
self, mock_event, mock_context
):
"""Test web search tool injection through the builtin tool manager."""
module = ama
req = ProviderRequest()
mock_context.get_config.return_value = {
"provider_settings": {
"web_search": True,
"websearch_provider": "baidu_ai_search",
}
}
builtin_tool = MagicMock(spec=FunctionTool)
builtin_tool.name = "web_search_baidu"
tool_mgr = MagicMock()
tool_mgr.get_builtin_tool.return_value = builtin_tool
mock_context.get_llm_tool_manager.return_value = tool_mgr
await module._apply_web_search_tools(mock_event, req, mock_context)
tool_mgr.get_builtin_tool.assert_called_once_with(module.BaiduWebSearchTool)
assert req.func_tool is not None
assert req.func_tool.get_tool("web_search_baidu") is builtin_tool
def test_proactive_cron_job_tools_uses_builtin_tool_manager(self, mock_context):
"""Test cron tool injection through the builtin tool manager."""
module = ama
req = ProviderRequest()
tool_mgr = MagicMock()
create_tool = MagicMock(spec=FunctionTool)
create_tool.name = "create_future_task"
delete_tool = MagicMock(spec=FunctionTool)
delete_tool.name = "delete_future_task"
list_tool = MagicMock(spec=FunctionTool)
list_tool.name = "list_future_tasks"
tool_mgr.get_builtin_tool.side_effect = [create_tool, delete_tool, list_tool]
mock_context.get_llm_tool_manager.return_value = tool_mgr
module._proactive_cron_job_tools(req, mock_context)
assert tool_mgr.get_builtin_tool.call_args_list == [
call(module.CreateActiveCronTool),
call(module.DeleteCronJobTool),
call(module.ListCronJobsTool),
]
assert req.func_tool is not None
assert req.func_tool.get_tool("create_future_task") is create_tool
assert req.func_tool.get_tool("delete_future_task") is delete_tool
assert req.func_tool.get_tool("list_future_tasks") is list_tool
class TestApplyFileExtract:
"""Tests for _apply_file_extract function."""
+14
View File
@@ -0,0 +1,14 @@
from astrbot.core import sp
from astrbot.core.provider.func_tool_manager import FunctionToolManager
from astrbot.core.tools.message_tools import SendMessageToUserTool
def test_get_builtin_tool_by_class_returns_cached_instance():
manager = FunctionToolManager()
tool_by_class = manager.get_builtin_tool(SendMessageToUserTool)
tool_by_name = manager.get_builtin_tool("send_message_to_user")
assert tool_by_class is tool_by_name
assert manager.get_func("send_message_to_user") is tool_by_class
assert tool_by_class.name == "send_message_to_user"