mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
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:
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
"""获取完整工具集
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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,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",
|
||||
]
|
||||
|
||||
@@ -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}"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user