Refactor agent tool inputs and background activity logging

This commit is contained in:
jxxghp
2026-06-27 07:53:09 +08:00
parent 544ed6d84d
commit 318cfe68e9
88 changed files with 223 additions and 238 deletions
+43 -25
View File
@@ -6,6 +6,7 @@
并在每次 Agent 启动时注入轻量索引,完整日志由工具按需查询。
"""
import asyncio
import json
import os
import re
@@ -88,10 +89,6 @@ ACTIVITY_ENTRY_PATTERN = re.compile(r"^-\s+\*\*(?P<time>\d{2}:\d{2})\*\*\s+(?P<s
class QueryActivityLogInput(BaseModel):
"""查询活动日志工具的输入参数模型。"""
explanation: Optional[str] = Field(
None,
description="Clear explanation of why this tool is being used in the current context",
)
keyword: Optional[str] = Field(
None,
description=(
@@ -288,17 +285,15 @@ class _ActivityLogToolProvider:
date: Optional[str] = None,
days: Optional[int] = DEFAULT_QUERY_DAYS,
limit: Optional[int] = DEFAULT_QUERY_LIMIT,
explanation: Optional[str] = None,
) -> str:
"""查询活动日志并返回 JSON 字符串。"""
logger.info(
"查询活动日志: keyword=%s, use_regex=%s, date=%s, days=%s, limit=%s, explanation=%s",
"查询活动日志: keyword=%s, use_regex=%s, date=%s, days=%s, limit=%s",
keyword,
use_regex,
date,
days,
limit,
explanation or "-",
)
try:
payload = query_activity_logs(
@@ -505,6 +500,7 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
self.retention_days = retention_days
self.prompt_load_days = prompt_load_days
self.stream_handler = stream_handler
self._background_tasks: set[asyncio.Task[None]] = set()
self._tool_provider = _ActivityLogToolProvider(activity_dir=activity_dir)
self.tools = [
StructuredTool.from_function(
@@ -631,6 +627,44 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
except Exception as e:
logger.warning(f"Failed to cleanup old activity logs: {e}")
def _schedule_activity_recording(self, messages: list) -> None:
"""提交后台活动记录任务,不阻塞当前 Agent 会话结束。"""
task = asyncio.create_task(self._record_activity(messages))
self._background_tasks.add(task)
task.add_done_callback(self._on_activity_recording_done)
def _on_activity_recording_done(self, task: asyncio.Task[None]) -> None:
"""清理已完成的后台任务并记录未捕获异常。"""
self._background_tasks.discard(task)
try:
task.result()
except asyncio.CancelledError:
logger.debug("活动日志后台记录任务已取消")
except Exception as err:
logger.warning(f"活动日志后台记录任务失败: {err}")
async def _record_activity(self, messages: list) -> None:
"""在后台生成本轮活动摘要并写入活动日志。"""
try:
# 提取本轮交互
round_messages = _extract_last_round(messages)
if not round_messages:
return
if _should_skip_activity_summary(round_messages):
return
# 格式化对话文本
conversation_text = _format_conversation_for_summary(round_messages)
if not conversation_text:
return
# 调用 LLM 生成摘要
summary = await _summarize_with_llm(conversation_text)
if summary:
await self._append_activity(summary)
except Exception as e:
logger.warning(f"Failed to record activity: {e}")
async def abefore_agent(
self, state: ActivityLogState, runtime: Runtime
) -> Optional[ActivityLogStateUpdate]:
@@ -699,28 +733,12 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
async def aafter_agent(
self, state: ActivityLogState, runtime: Runtime
) -> Optional[dict[str, Any]]:
"""Agent 执行完毕后,调用 LLM 对本轮对话生成摘要并追加到当日活动日志"""
"""Agent 执行完毕后,异步提交活动日志记录任务"""
try:
messages = state.get("messages", [])
if not messages:
return None
# 提取本轮交互
round_messages = _extract_last_round(messages)
if not round_messages:
return None
if _should_skip_activity_summary(round_messages):
return None
# 格式化对话文本
conversation_text = _format_conversation_for_summary(round_messages)
if not conversation_text:
return None
# 调用 LLM 生成摘要
summary = await _summarize_with_llm(conversation_text)
if summary:
await self._append_activity(summary)
self._schedule_activity_recording(list(messages))
except Exception as e:
logger.warning(f"Failed to record activity: {e}")
+3 -8
View File
@@ -92,10 +92,6 @@ class SkillsStateUpdate(TypedDict):
class SkillToolInput(BaseModel):
"""Skill 加载工具的输入参数模型。"""
explanation: Optional[str] = Field(
None,
description="Clear explanation of why this skill is needed in the current context",
)
name: str = Field(
...,
description="Skill name or id from the available skills list.",
@@ -460,9 +456,9 @@ class _SkillToolProvider:
raw_content = await handle.read(MAX_SKILL_FILE_SIZE)
return raw_content.decode("utf-8", errors="replace"), truncated
async def load_skill(self, name: str, explanation: Optional[str] = None) -> str:
async def load_skill(self, name: str) -> str:
"""加载指定 Skill 的完整说明并返回 JSON 字符串。"""
logger.info(f"加载 Skill: name={name}, explanation={explanation or '-'}")
logger.info(f"加载 Skill: name={name}")
try:
skill = await self._find_skill(name)
if not skill:
@@ -674,8 +670,7 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
if not isinstance(tool_args, dict):
tool_args = {}
logger.info(
f"开始执行 Skill 工具: name={tool_args.get('name') or '-'}, "
f"explanation={tool_args.get('explanation') or '-'}"
f"开始执行 Skill 工具: name={tool_args.get('name') or '-'}"
)
if self.stream_handler and getattr(self.stream_handler, "is_streaming", False):
self.stream_handler.record_tool_call(
+3 -10
View File
@@ -238,10 +238,6 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
# 获取工具执行提示消息
tool_message = self.get_tool_message(**kwargs)
if not tool_message:
explanation = kwargs.get("explanation")
if explanation:
tool_message = explanation
# 发送工具执行过程消息(流式传输且非最后终结工具时)
if self._stream_handler and self._stream_handler.is_streaming and not self.return_direct:
@@ -325,16 +321,13 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
获取工具执行时的友好提示消息。
子类可以重写此方法,根据实际参数生成个性化的提示消息。
如果返回 None 或空字符串,将回退使用 explanation 参数。
Args:
**kwargs: 工具的所有参数(包括 explanation
**kwargs: 工具的所有参数
Returns:
str: 友好的提示消息,如果返回 None 或空字符串则使用 explanation
str: 友好的提示消息
"""
explanation = kwargs.get("explanation")
return str(explanation) if explanation else None
return None
@abstractmethod
async def run(self, **kwargs) -> str:
@@ -20,8 +20,6 @@ from app.schemas.types import SystemConfigKey
class AddCustomFilterRuleInput(BaseModel):
"""新增自定义过滤规则工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
rule_id: str = Field(
...,
description="Unique custom rule ID. Only letters and numbers are allowed.",
@@ -24,7 +24,6 @@ from app.utils.crypto import HashUtils
class AddDownloadTasksInput(BaseModel):
"""添加下载任务工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
torrent_url: List[str] = Field(
...,
description="One or more torrent_url values. Supports refs from get_search_results (`hash:id`) and magnet links."
-2
View File
@@ -24,8 +24,6 @@ from app.schemas.types import SystemConfigKey
class AddRuleGroupInput(BaseModel):
"""新增过滤规则组工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
name: str = Field(..., description="New rule group name.")
rule_string: str = Field(
...,
-2
View File
@@ -15,8 +15,6 @@ from app.schemas.types import MediaType, MessageChannel
class AddSubscribeInput(BaseModel):
"""添加订阅工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
title: str = Field(
...,
description="The title of the media to subscribe to (e.g., 'The Matrix', 'Breaking Bad')",
-8
View File
@@ -24,10 +24,6 @@ class UserChoiceOptionInput(BaseModel):
...,
description="The exact content that will be sent back to the agent after the user clicks this button",
)
description: Optional[str] = Field(
None,
description="Optional user-facing description stored in chat history after this option is selected",
)
@model_validator(mode="after")
def validate_option(self):
@@ -44,8 +40,6 @@ class UserChoiceOptionInput(BaseModel):
class AskUserChoiceInput(BaseModel):
"""按钮选择工具输入。"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why the agent needs the user to choose from buttons",)
message: str = Field(
...,
description="Question or prompt shown to the user together with the buttons",
@@ -166,7 +160,6 @@ class AskUserChoiceTool(MoviePilotTool):
AgentInteractionOption(
label=option.label.strip(),
value=option.value.strip(),
description=(option.description.strip() if option.description else None),
)
for option in options
]
@@ -190,7 +183,6 @@ class AskUserChoiceTool(MoviePilotTool):
"callback_data": (
f"agent_interaction:choice:{request.request_id}:{index}"
),
"description": option.description or option.label,
}
)
if len(current_row) >= max_per_row:
-2
View File
@@ -47,8 +47,6 @@ class BrowserAction(str, Enum):
class BrowseWebpageInput(BaseModel):
"""浏览器操作工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this browser action is being performed",)
action: str = Field(
...,
description=(
@@ -20,8 +20,6 @@ from app.schemas.types import SystemConfigKey
class DeleteCustomFilterRuleInput(BaseModel):
"""删除自定义过滤规则工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
rule_id: str = Field(..., description="Custom rule ID to delete.")
@@ -13,8 +13,6 @@ from app.log import logger
class DeleteDownloadHistoryInput(BaseModel):
"""删除下载历史记录工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
history_id: int = Field(
..., description="The ID of the download history record to delete"
)
@@ -13,8 +13,6 @@ from app.log import logger
class DeleteDownloadTasksInput(BaseModel):
"""删除下载任务工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
hash: str = Field(
..., description="Task hash (can be obtained from query_download_tasks tool)"
)
@@ -19,8 +19,6 @@ from app.schemas.types import SystemConfigKey
class DeleteRuleGroupInput(BaseModel):
"""删除过滤规则组工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
name: str = Field(..., description="Rule group name to delete.")
-2
View File
@@ -16,8 +16,6 @@ from app.schemas.types import EventType
class DeleteSubscribeInput(BaseModel):
"""删除订阅工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
subscribe_id: int = Field(
...,
description="The ID of the subscription to delete (can be obtained from query_subscribes tool)",
@@ -15,8 +15,6 @@ from app.schemas import FileItem
class DeleteTransferHistoryInput(BaseModel):
"""删除整理历史记录工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
history_id: int = Field(
..., description="The ID of the transfer history record to delete"
)
-1
View File
@@ -135,7 +135,6 @@ class _CommandOutput:
class ExecuteCommandInput(BaseModel):
"""执行 Shell 命令工具的输入参数模型。"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this command action is needed")
action: Optional[Literal["start", "read", "wait", "write", "kill", "run"]] = Field(
"start",
description=(
@@ -15,8 +15,6 @@ from app.schemas.types import MediaType, media_type_to_agent
class GetRecommendationsInput(BaseModel):
"""获取推荐工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
source: Optional[str] = Field(
"tmdb_trending",
description="Recommendation source: "
@@ -21,8 +21,6 @@ from ._torrent_search_utils import (
class GetSearchResultsInput(BaseModel):
"""获取搜索结果工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
site: Optional[List[str]] = Field(None, description="Site name filters")
season: Optional[List[str]] = Field(None, description="Season or episode filters")
free_state: Optional[List[str]] = Field(None, description="Promotion state filters")
-2
View File
@@ -19,8 +19,6 @@ from app.log import logger
class InstallPluginInput(BaseModel):
"""安装插件工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
plugin_id: str = Field(
...,
description="Exact plugin ID to install. Use query_market_plugins first to find the correct plugin_id.",
-1
View File
@@ -17,7 +17,6 @@ from app.utils.string import StringUtils
class ListDirectoryInput(BaseModel):
"""查询文件系统目录内容工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
path: str = Field(..., description="Directory path to list contents (e.g., '/home/user/downloads' or 'C:/Downloads')")
storage: Optional[str] = Field("local", description="Storage type (default: 'local' for local file system, can be 'smb', 'alist', etc.)")
sort_by: Optional[str] = Field("name", description="Sort order: 'name' for alphabetical sorting, 'time' for modification time sorting (default: 'name')")
@@ -13,8 +13,6 @@ from app.log import logger
class ListSlashCommandsInput(BaseModel):
"""查询所有可用斜杠命令工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
class ListSlashCommandsTool(MoviePilotTool):
@@ -18,8 +18,6 @@ from app.log import logger
class QueryBuiltinFilterRulesInput(BaseModel):
"""查询内置过滤规则工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
rule_ids: Optional[List[str]] = Field(
None,
description="Optional list of built-in rule IDs to query. If omitted, return all built-in rules.",
@@ -19,8 +19,6 @@ from app.log import logger
class QueryCustomFilterRulesInput(BaseModel):
"""查询自定义过滤规则工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
rule_ids: Optional[List[str]] = Field(
None,
description="Optional list of custom rule IDs to query. If omitted, return all custom rules.",
@@ -15,8 +15,6 @@ from app.schemas.types import SystemConfigKey
class QueryCustomIdentifiersInput(BaseModel):
"""查询自定义识别词工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
class QueryCustomIdentifiersTool(MoviePilotTool):
@@ -13,7 +13,6 @@ from app.log import logger
class QueryDirectorySettingsInput(BaseModel):
"""查询系统目录设置工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
directory_type: Optional[str] = Field("all",
description="Filter directories by type: 'download' for download directories, 'library' for media library directories, 'all' for all directories")
storage_type: Optional[str] = Field("all",
@@ -14,10 +14,6 @@ from app.log import logger
class QueryDoctorReportInput(BaseModel):
"""查询 Doctor 诊断报告工具的输入参数模型。"""
explanation: Optional[str] = Field(
None,
description="Clear explanation of why this tool is being used in the current context",
)
deep: Optional[bool] = Field(
False,
description=(
@@ -16,7 +16,6 @@ from app.schemas.types import TorrentQueryStatus, media_type_to_agent
class QueryDownloadTasksInput(BaseModel):
"""查询下载工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
downloader: Optional[str] = Field(None,
description="Name of specific downloader to query (optional, if not provided queries all configured downloaders)")
status: Optional[str] = Field("all",
@@ -14,9 +14,6 @@ from app.schemas.types import SystemConfigKey
class QueryDownloadersInput(BaseModel):
"""查询下载器工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
class QueryDownloadersTool(MoviePilotTool):
name: str = "query_downloaders"
tags: list[str] = [
@@ -13,7 +13,6 @@ from app.log import logger
class QueryEpisodeScheduleInput(BaseModel):
"""查询剧集上映时间工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
tmdb_id: int = Field(..., description="TMDB ID of the TV series (can be obtained from search_media tool)")
season: int = Field(..., description="Season number to query")
episode_group: Optional[str] = Field(None, description="Episode group ID (optional)")
@@ -22,8 +22,6 @@ from app.log import logger
class QueryInstalledPluginsInput(BaseModel):
"""查询已安装插件工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
query: Optional[str] = Field(
None,
description="Optional keyword to filter installed plugins by plugin ID, name, description, or author.",
@@ -77,7 +77,6 @@ def _build_tv_server_result(existing_seasons: OrderedDict, total_seasons: Ordere
class QueryLibraryExistsInput(BaseModel):
"""查询媒体库工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
tmdb_id: Optional[int] = Field(None, description="TMDB ID (can be obtained from search_media tool). Either tmdb_id or douban_id must be provided.")
douban_id: Optional[str] = Field(None, description="Douban ID (can be obtained from search_media tool). Either tmdb_id or douban_id must be provided.")
media_type: Optional[str] = Field(None, description="Allowed values: movie, tv")
@@ -18,8 +18,6 @@ PAGE_SIZE = 20
class QueryLibraryLatestInput(BaseModel):
"""查询媒体服务器最近入库影片工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
server: Optional[str] = Field(
None,
description="Media server name (optional, if not specified queries all enabled media servers)",
@@ -21,8 +21,6 @@ from app.log import logger
class QueryMarketPluginsInput(BaseModel):
"""查询插件市场工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
query: Optional[str] = Field(
None,
description="Optional keyword to filter plugin market results by plugin ID, name, description, or author.",
@@ -18,7 +18,6 @@ SEASON_PREVIEW_LIMIT = 100
class QueryMediaDetailInput(BaseModel):
"""查询媒体详情工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
tmdb_id: Optional[int] = Field(None, description="TMDB ID of the media (movie or TV series, can be obtained from search_media tool)")
douban_id: Optional[str] = Field(None, description="Douban ID of the media (alternative to tmdb_id)")
media_type: str = Field(..., description="Allowed values: movie, tv")
-2
View File
@@ -14,8 +14,6 @@ from app.log import logger
class QueryPersonasInput(BaseModel):
"""查询人格工具的输入参数模型。"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
query: Optional[str] = Field(
None,
description=(
@@ -14,8 +14,6 @@ from app.log import logger
class QueryPluginCapabilitiesInput(BaseModel):
"""查询插件能力工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
plugin_id: Optional[str] = Field(
None,
description="Optional plugin ID to query capabilities for a specific plugin. "
@@ -15,8 +15,6 @@ from app.log import logger
class QueryPluginConfigInput(BaseModel):
"""查询插件配置工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
plugin_id: str = Field(
...,
description="The plugin ID to query. Use query_installed_plugins first to discover valid plugin IDs.",
@@ -19,8 +19,6 @@ from app.log import logger
class QueryPluginDataInput(BaseModel):
"""查询插件数据工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
plugin_id: str = Field(
...,
description="The plugin ID to query. Use query_installed_plugins first to discover valid plugin IDs.",
@@ -18,7 +18,6 @@ MAX_PAGE_SIZE = 50
class QueryPopularSubscribesInput(BaseModel):
"""查询热门订阅工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
media_type: str = Field(..., description="Allowed values: movie, tv")
page: Optional[int] = Field(1, description="Page number for pagination (default: 1)")
count: Optional[int] = Field(30, description="Number of items per page (default: 30, max: 50)")
@@ -19,8 +19,6 @@ from app.log import logger
class QueryRuleGroupsInput(BaseModel):
"""查询规则组工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
group_names: Optional[List[str]] = Field(
None,
description="Optional list of rule group names to query. If omitted, return all rule groups.",
-3
View File
@@ -12,9 +12,6 @@ from app.log import logger
class QuerySchedulersInput(BaseModel):
"""查询定时服务工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
class QuerySchedulersTool(MoviePilotTool):
name: str = "query_schedulers"
tags: list[str] = [
@@ -22,8 +22,6 @@ def _preview_list(value, limit: int = SITE_USERDATA_DETAIL_PREVIEW_LIMIT) -> tup
class QuerySiteUserdataInput(BaseModel):
"""查询站点用户数据工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
site_id: int = Field(
...,
description="The ID of the site to query user data for (can be obtained from query_sites tool)",
-2
View File
@@ -14,8 +14,6 @@ from app.log import logger
class QuerySitesInput(BaseModel):
"""查询站点工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
status: Optional[str] = Field(
"all",
description="Filter sites by status: 'active' for enabled sites, 'inactive' for disabled sites, 'all' for all sites",
@@ -17,8 +17,6 @@ PAGE_SIZE = 20
class QuerySubscribeHistoryInput(BaseModel):
"""查询订阅历史工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
media_type: Optional[str] = Field(
"all", description="Allowed values: movie, tv, all"
)
@@ -15,7 +15,6 @@ MAX_PAGE_SIZE = 50
class QuerySubscribeSharesInput(BaseModel):
"""查询订阅分享工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
name: Optional[str] = Field(None, description="Filter shares by media name (partial match, optional)")
page: Optional[int] = Field(1, description="Page number for pagination (default: 1)")
count: Optional[int] = Field(30, description="Number of items per page (default: 30, max: 50)")
-2
View File
@@ -48,8 +48,6 @@ QUERY_SUBSCRIBE_OUTPUT_FIELDS = [
class QuerySubscribesInput(BaseModel):
"""查询订阅工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
status: Optional[str] = Field(
"all",
description="Filter subscriptions by status: 'R' for enabled subscriptions, 'S' for paused ones, 'all' for all subscriptions",
@@ -23,8 +23,6 @@ from app.log import logger
class QuerySystemSettingsInput(BaseModel):
"""查询系统设置工具的输入参数模型。"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
setting_key: Optional[str] = Field(
None,
description=(
@@ -15,7 +15,6 @@ from app.utils.jieba import cut as jieba_cut
class QueryTransferHistoryInput(BaseModel):
"""查询整理历史记录工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
title: Optional[str] = Field(None, description="Search by title (optional, supports partial match)")
status: Optional[str] = Field("all",
description="Filter by status: 'success' for successful transfers, 'failed' for failed transfers, 'all' for all records (default: 'all')")
-1
View File
@@ -13,7 +13,6 @@ from app.log import logger
class QueryWorkflowsInput(BaseModel):
"""查询工作流工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
state: Optional[str] = Field("all", description="Filter workflows by state: 'W' for waiting, 'R' for running, 'P' for paused, 'S' for success, 'F' for failed, 'all' for all workflows (default: 'all')")
name: Optional[str] = Field(None, description="Filter workflows by name (partial match, optional)")
trigger_type: Optional[str] = Field("all", description="Filter workflows by trigger type: 'timer' for scheduled, 'event' for event-triggered, 'manual' for manual, 'all' for all types (default: 'all')")
@@ -15,10 +15,6 @@ from app.log import logger
class RecognizeCaptchaInput(BaseModel):
"""识别图形验证码工具的输入参数模型。"""
explanation: Optional[str] = Field(
None,
description="Clear explanation of why this captcha image needs to be recognized",
)
image_url: str = Field(
...,
description=(
-1
View File
@@ -16,7 +16,6 @@ from app.schemas.types import media_type_to_agent
class RecognizeMediaInput(BaseModel):
"""识别媒体信息工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
title: Optional[str] = Field(None, description="The title of the torrent/media to recognize (required for torrent recognition)")
subtitle: Optional[str] = Field(None, description="The subtitle or description of the torrent (optional, helps improve recognition accuracy)")
path: Optional[str] = Field(None, description="The file path to recognize (required for file recognition, mutually exclusive with title)")
-2
View File
@@ -17,8 +17,6 @@ from app.log import logger
class ReloadPluginInput(BaseModel):
"""重载插件工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
plugin_id: str = Field(
...,
description="The plugin ID to reload so the latest saved config takes effect.",
-2
View File
@@ -12,8 +12,6 @@ from app.log import logger
class RunSchedulerInput(BaseModel):
"""运行定时服务工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
job_id: str = Field(
...,
description="The ID of the scheduled job to run (can be obtained from query_schedulers tool)",
@@ -15,8 +15,6 @@ from app.schemas.types import EventType, MessageChannel
class RunSlashCommandInput(BaseModel):
"""运行斜杠命令工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
command: str = Field(
...,
description="The slash command to execute, e.g. '/cookiecloud'. "
-2
View File
@@ -14,8 +14,6 @@ from app.log import logger
class RunWorkflowInput(BaseModel):
"""执行工作流工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
workflow_id: int = Field(
..., description="Workflow ID (can be obtained from query_workflows tool)"
)
-2
View File
@@ -16,8 +16,6 @@ from app.schemas import FileItem
class ScrapeMetadataInput(BaseModel):
"""刮削媒体元数据工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
path: str = Field(
...,
description="Path to the file or directory to scrape metadata for (e.g., '/path/to/file.mkv' or '/path/to/directory')",
-1
View File
@@ -14,7 +14,6 @@ from app.schemas.types import MediaType, media_type_to_agent
class SearchMediaInput(BaseModel):
"""搜索媒体工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
title: str = Field(..., description="The title of the media to search for (e.g., 'The Matrix', 'Breaking Bad')")
year: Optional[str] = Field(None, description="Release year of the media (optional, helps narrow down results)")
media_type: Optional[str] = Field(None,
-1
View File
@@ -13,7 +13,6 @@ from app.log import logger
class SearchPersonInput(BaseModel):
"""搜索人物工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
name: str = Field(..., description="The name of the person to search for (e.g., 'Tom Hanks', '周杰伦')")
@@ -15,7 +15,6 @@ from app.log import logger
class SearchPersonCreditsInput(BaseModel):
"""搜索演员参演作品工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
person_id: int = Field(..., description="The ID of the person/actor to search for credits (e.g., 31 for Tom Hanks in TMDB)")
source: str = Field(..., description="The data source: 'tmdb' for TheMovieDB, 'douban' for Douban, 'bangumi' for Bangumi")
page: Optional[int] = Field(1, description="Page number for pagination (default: 1)")
-1
View File
@@ -15,7 +15,6 @@ from app.schemas.types import media_type_to_agent
class SearchSubscribeInput(BaseModel):
"""搜索订阅缺失剧集工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
subscribe_id: int = Field(..., description="The ID of the subscription to search for missing episodes (can be obtained from query_subscribes tool)")
manual: Optional[bool] = Field(False, description="Whether this is a manual search (default: False)")
filter_groups: Optional[List[str]] = Field(None,
-1
View File
@@ -20,7 +20,6 @@ from ._torrent_search_utils import (
class SearchTorrentsInput(BaseModel):
"""搜索种子工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
tmdb_id: Optional[int] = Field(None, description="TMDB ID (can be obtained from search_media tool). Either tmdb_id or douban_id must be provided.")
douban_id: Optional[str] = Field(None, description="Douban ID (can be obtained from search_media tool). Either tmdb_id or douban_id must be provided.")
media_type: Optional[str] = Field(None, description="Allowed values: movie, tv")
-4
View File
@@ -48,10 +48,6 @@ class _SearchSiteFilter:
class SearchWebInput(BaseModel):
"""搜索网络内容工具的输入参数模型"""
explanation: Optional[str] = Field(
None,
description="Clear explanation of why this tool is being used in the current context",
)
query: str = Field(
..., description="The search query string to search for on the web"
)
-2
View File
@@ -16,8 +16,6 @@ from app.schemas.types import MessageChannel
class SendLocalFileInput(BaseModel):
"""发送本地附件工具输入。"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why sending this local file helps the user",)
file_path: str = Field(
...,
description="Absolute path to the local image or file to send to the user",
-4
View File
@@ -14,10 +14,6 @@ from app.schemas.types import NotificationType
class SendMessageInput(BaseModel):
"""发送消息工具的输入参数模型"""
explanation: Optional[str] = Field(
None,
description="Clear explanation of why this tool is being used in the current context",
)
message: Optional[str] = Field(
None,
description="The message content to send to the user (should be clear and informative)",
@@ -14,10 +14,6 @@ from app.schemas import Notification, NotificationType
class SendVoiceMessageInput(BaseModel):
"""发送语音消息工具输入。"""
explanation: Optional[str] = Field(
None,
description="Clear explanation of why a voice reply is the best fit in the current context",
)
message: str = Field(
...,
description="The spoken content to send back to the user",
-2
View File
@@ -14,8 +14,6 @@ from app.log import logger
class SwitchPersonaInput(BaseModel):
"""切换人格工具的输入参数模型。"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
persona_id: str = Field(
...,
description=(
-1
View File
@@ -13,7 +13,6 @@ from app.log import logger
class TestSiteInput(BaseModel):
"""测试站点连通性工具的输入参数模型"""
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
site_identifier: int = Field(..., description="Site ID to test (can be obtained from query_sites tool)")
-2
View File
@@ -14,8 +14,6 @@ from app.schemas import FileItem, MediaType
class TransferFileInput(BaseModel):
"""整理文件或目录工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
file_path: str = Field(
...,
description="Path to the file or directory to transfer (e.g., '/path/to/file.mkv' or '/path/to/directory')",
-2
View File
@@ -18,8 +18,6 @@ from app.log import logger
class UninstallPluginInput(BaseModel):
"""卸载插件工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
plugin_id: str = Field(
...,
description="Exact plugin ID to uninstall. Use query_installed_plugins first to find the correct plugin_id.",
@@ -23,8 +23,6 @@ from app.schemas.types import SystemConfigKey
class UpdateCustomFilterRuleInput(BaseModel):
"""更新自定义过滤规则工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
current_rule_id: str = Field(
..., description="Existing custom rule ID to update."
)
@@ -15,8 +15,6 @@ from app.schemas.types import SystemConfigKey
class UpdateCustomIdentifiersInput(BaseModel):
"""更新自定义识别词工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
identifiers: List[str] = Field(
...,
description=(
@@ -14,10 +14,6 @@ from app.log import logger
class UpdateDownloadTasksInput(BaseModel):
"""更新下载任务工具的输入参数模型"""
explanation: Optional[str] = Field(
None,
description="Clear explanation of why this tool is being used in the current context",
)
hash: str = Field(
..., description="Task hash (can be obtained from query_download_tasks tool)"
)
@@ -14,8 +14,6 @@ from app.log import logger
class UpdatePersonaDefinitionInput(BaseModel):
"""更新人格定义工具的输入参数模型。"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
persona_id: str = Field(
...,
description=(
@@ -15,8 +15,6 @@ from app.log import logger
class UpdatePluginConfigInput(BaseModel):
"""修改插件配置工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
plugin_id: str = Field(
...,
description="The plugin ID to update. Use query_plugin_config first to inspect the current config.",
@@ -25,8 +25,6 @@ from app.schemas.types import SystemConfigKey
class UpdateRuleGroupInput(BaseModel):
"""更新过滤规则组工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
current_name: str = Field(..., description="Existing rule group name to update.")
new_name: Optional[str] = Field(
None,
-2
View File
@@ -17,8 +17,6 @@ from app.utils.string import StringUtils
class UpdateSiteInput(BaseModel):
"""更新站点工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
site_id: int = Field(
...,
description="The ID of the site to update (can be obtained from query_sites tool)",
@@ -14,8 +14,6 @@ from app.log import logger
class UpdateSiteCookieInput(BaseModel):
"""更新站点Cookie和UA工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
site_identifier: int = Field(
...,
description="Site ID to update Cookie and User-Agent for (can be obtained from query_sites tool)",
-2
View File
@@ -16,8 +16,6 @@ from app.schemas.types import EventType
class UpdateSubscribeInput(BaseModel):
"""更新订阅工具的输入参数模型"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
subscribe_id: int = Field(
...,
description="The ID of the subscription to update (can be obtained from query_subscribes tool)",
@@ -29,8 +29,6 @@ SettingValue = Optional[Union[list, dict, bool, int, float, str]]
class UpdateSystemSettingsInput(BaseModel):
"""更新系统设置工具的输入参数模型。"""
explanation: Optional[str] = Field(None,
description="Clear explanation of why this tool is being used in the current context",)
setting_key: str = Field(
...,
description=(
+3 -9
View File
@@ -559,8 +559,6 @@ def _format_tool_detail(tool: Dict[str, Any]) -> None:
required = set((tool.get("inputSchema") or {}).get("required") or [])
fields = []
for name, schema in properties.items():
if name == "explanation":
continue
fields.append(
(
f"{name}*" if name in required else name,
@@ -1133,8 +1131,7 @@ def tool_show(tool_name: str) -> None:
@click.argument("args", nargs=-1, type=click.UNPROCESSED)
def tool_run(tool_name: str, args: tuple[str, ...]) -> None:
"""运行指定工具"""
arguments = {"explanation": "CLI invocation"}
arguments.update(_parse_key_value_pairs(args))
arguments = _parse_key_value_pairs(args)
result = _call_tool(tool_name, arguments, runtime=_backend_runtime())
if isinstance(result, (dict, list)):
_print_json(result)
@@ -1152,7 +1149,7 @@ def scheduler_list() -> None:
"""列出调度任务"""
result = _call_tool(
"query_schedulers",
{"explanation": "List scheduler jobs from local CLI"},
{},
runtime=_backend_runtime(),
)
if isinstance(result, list):
@@ -1168,10 +1165,7 @@ def scheduler_run(job_id: str) -> None:
"""立即执行某个调度任务"""
result = _call_tool(
"run_scheduler",
{
"explanation": "Run a scheduler job from local CLI",
"job_id": job_id,
},
{"job_id": job_id},
runtime=_backend_runtime(),
)
if isinstance(result, (dict, list)):
+2
View File
@@ -128,6 +128,8 @@ MoviePilot 也提供普通 REST API 给前端和自动化客户端使用。所
获取所有可用的MCP工具列表。
工具的 `inputSchema` 只包含实际执行业务所需的参数,不包含用于解释调用原因的通用 `explanation` 参数,以减少 Agent 上下文消耗。
**认证**: 需要API KEY,在请求头中添加 `X-API-KEY: <api_key>` 或在查询参数中添加 `apikey=<api_key>`
**响应示例**:
+2 -3
View File
@@ -150,7 +150,7 @@ Write the rule using the appropriate format. Ensure:
Use the `query_custom_identifiers` tool to get all current rules:
```
query_custom_identifiers(explanation="Checking existing identifiers before adding new rules to avoid duplicates")
query_custom_identifiers()
```
### Step 4: Check for Duplicates
@@ -166,7 +166,6 @@ Merge new non-duplicate rules into the existing list, then use `update_custom_id
```
update_custom_identifiers(
explanation="Adding new identifier rules for [description]",
identifiers=["existing rule 1", "existing rule 2", "# new comment", "new rule"]
)
```
@@ -178,7 +177,7 @@ update_custom_identifiers(
If the user wants to verify the rule works, use `recognize_media` to test:
```
recognize_media(explanation="Testing recognition after adding identifier", title="the torrent title to test")
recognize_media(title="the torrent title to test")
```
### Step 7: Report
+89 -21
View File
@@ -28,6 +28,13 @@ def _write_activity_log(activity_dir, date_str: str, lines: list[str]) -> None:
)
async def _wait_activity_log_tasks(middleware: ActivityLogMiddleware) -> None:
"""等待活动日志后台任务完成,避免测试与后台写入竞态。"""
tasks = list(middleware._background_tasks)
if tasks:
await asyncio.gather(*tasks)
def test_activity_log_index_counts_entries_without_body(tmp_path):
"""活动日志索引只应包含条目数量,不暴露完整摘要正文。"""
date_str = datetime.now().strftime("%Y-%m-%d")
@@ -91,19 +98,19 @@ def test_activity_log_abefore_agent_refreshes_existing_state(tmp_path):
def test_activity_log_skips_trivial_greeting_without_llm(tmp_path):
"""无实际任务的寒暄不应调用 LLM,也不应写入活动日志。"""
middleware = ActivityLogMiddleware(activity_dir=str(tmp_path))
summarize_mock = AsyncMock(return_value="不应写入")
append_mock = AsyncMock()
async def _run_test():
middleware = ActivityLogMiddleware(activity_dir=str(tmp_path))
summarize_mock = AsyncMock(return_value="不应写入")
append_mock = AsyncMock()
with (
patch(
"app.agent.middleware.activity_log._summarize_with_llm",
new=summarize_mock,
),
patch.object(middleware, "_append_activity", new=append_mock),
):
asyncio.run(
middleware.aafter_agent(
with (
patch(
"app.agent.middleware.activity_log._summarize_with_llm",
new=summarize_mock,
),
patch.object(middleware, "_append_activity", new=append_mock),
):
await middleware.aafter_agent(
{
"messages": [
HumanMessage(content="你好"),
@@ -112,7 +119,11 @@ def test_activity_log_skips_trivial_greeting_without_llm(tmp_path):
},
runtime=None,
)
)
await _wait_activity_log_tasks(middleware)
return summarize_mock, append_mock
summarize_mock, append_mock = asyncio.run(_run_test())
summarize_mock.assert_not_awaited()
append_mock.assert_not_awaited()
@@ -137,18 +148,18 @@ def test_summarize_with_llm_ignores_skip_marker():
def test_activity_log_records_detailed_summary(tmp_path):
"""有实际工具动作的交互应写入较完整的活动摘要。"""
middleware = ActivityLogMiddleware(activity_dir=str(tmp_path))
summary = (
"用户要求整理 `/downloads/Show`,助手调用 transfer_file 识别并转移剧集,"
"结果成功写入目标媒体库。"
)
with patch(
"app.agent.middleware.activity_log._summarize_with_llm",
new=AsyncMock(return_value=summary),
):
asyncio.run(
middleware.aafter_agent(
async def _run_test():
middleware = ActivityLogMiddleware(activity_dir=str(tmp_path))
with patch(
"app.agent.middleware.activity_log._summarize_with_llm",
new=AsyncMock(return_value=summary),
):
await middleware.aafter_agent(
{
"messages": [
HumanMessage(content="帮我整理 /downloads/Show"),
@@ -170,7 +181,9 @@ def test_activity_log_records_detailed_summary(tmp_path):
},
runtime=None,
)
)
await _wait_activity_log_tasks(middleware)
asyncio.run(_run_test())
log_files = list(tmp_path.glob("*.md"))
assert len(log_files) == 1
@@ -179,6 +192,61 @@ def test_activity_log_records_detailed_summary(tmp_path):
assert "- **" in content
def test_activity_log_after_agent_does_not_wait_for_summary(tmp_path):
"""活动日志摘要生成应在后台执行,不阻塞当前 Agent 会话结束。"""
async def _slow_summarize(_conversation_text: str) -> str:
"""模拟较慢的活动摘要生成。"""
await asyncio.sleep(0.05)
return "用户要求检查下载任务,助手调用工具完成检查。"
async def _run_test():
middleware = ActivityLogMiddleware(activity_dir=str(tmp_path))
append_mock = AsyncMock()
with (
patch(
"app.agent.middleware.activity_log._summarize_with_llm",
side_effect=_slow_summarize,
) as summarize_mock,
patch.object(middleware, "_append_activity", new=append_mock),
):
await middleware.aafter_agent(
{
"messages": [
HumanMessage(content="帮我检查下载任务"),
AIMessage(
content="",
tool_calls=[
{
"name": "query_download_tasks",
"args": {},
"id": "call_1",
}
],
),
ToolMessage(
content='{"success": true}',
tool_call_id="call_1",
),
],
},
runtime=None,
)
called_before_wait = summarize_mock.await_count
pending_before_wait = len(middleware._background_tasks)
await _wait_activity_log_tasks(middleware)
return called_before_wait, pending_before_wait, summarize_mock, append_mock
called_before_wait, pending_before_wait, summarize_mock, append_mock = asyncio.run(
_run_test()
)
assert called_before_wait == 0
assert pending_before_wait == 1
summarize_mock.assert_awaited_once()
append_mock.assert_awaited_once_with("用户要求检查下载任务,助手调用工具完成检查。")
def test_query_activity_logs_filters_by_keyword_and_date(tmp_path):
"""活动日志查询应支持日期和关键词过滤。"""
_write_activity_log(
-2
View File
@@ -563,7 +563,6 @@ class AgentImageSupportTest(unittest.TestCase):
def test_send_message_input_accepts_image_only_payload(self):
payload = SendMessageInput(
explanation="send poster image",
image_url="https://example.com/poster.png",
)
@@ -659,7 +658,6 @@ class AgentImageSupportTest(unittest.TestCase):
def test_send_local_file_input_accepts_file_payload(self):
payload = SendLocalFileInput(
explanation="send generated report",
file_path="/tmp/report.txt",
message="请下载查看",
)
+1
View File
@@ -125,6 +125,7 @@ class TestAgentInteraction(unittest.TestCase):
notification = async_post_message.await_args.args[0]
self.assertEqual(notification.text, "请选择要执行的操作")
self.assertEqual(sum(len(row) for row in notification.buttons), 2)
self.assertNotIn("description", notification.buttons[0][0])
callback_data = notification.buttons[0][0]["callback_data"]
_, _, request_id, option_index = callback_data.split(":")
-2
View File
@@ -132,7 +132,6 @@ async def test_skill_tool_call_records_streaming_summary(tmp_path):
tool_call={
"args": {
"name": "moviepilot-cli",
"explanation": "测试加载技能",
}
},
)
@@ -150,7 +149,6 @@ async def test_skill_tool_call_records_streaming_summary(tmp_path):
"tool_message": "Skill loaded",
"tool_kwargs": {
"name": "moviepilot-cli",
"explanation": "测试加载技能",
},
}
]
+69 -1
View File
@@ -1,10 +1,18 @@
import importlib.util
from types import SimpleNamespace
from typing import Iterator, Optional
from pathlib import Path
from typing import Iterator, Optional, Type
import pytest
from pydantic import BaseModel
from app.agent.middleware.activity_log import QueryActivityLogInput
from app.agent.middleware.skills import SkillToolInput
from app.agent.tools.base import MoviePilotTool
from app.agent.tools.factory import MoviePilotToolFactory
from app.agent.tools.impl.ask_user_choice import AskUserChoiceInput, AskUserChoiceTool
from app.agent.tools.impl.send_local_file import SendLocalFileTool
from app.agent.tools.impl.send_voice_message import SendVoiceMessageTool
from app.core.plugin import PluginManager
from app.utils.singleton import Singleton
@@ -56,6 +64,66 @@ def _build_plugin(
)
def _schema_properties(args_schema: Type[BaseModel]) -> dict:
"""返回工具输入模型的 JSON Schema 属性。"""
return args_schema.model_json_schema().get("properties", {})
def _load_lexiannot_tool_schemas() -> list[Type[BaseModel]]:
"""只加载 LexiAnnot schema 文件,避免触发插件包可选依赖。"""
schema_path = (
Path(__file__).resolve().parents[1]
/ "app"
/ "plugins"
/ "lexiannot"
/ "schemas.py"
)
spec = importlib.util.spec_from_file_location(
"_test_lexiannot_schemas",
schema_path,
)
assert spec and spec.loader
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return [
module.VocabularyAnnotatingToolInput,
module.QueryAnnotationTasksToolInput,
]
def test_agent_tool_schemas_do_not_expose_explanation_parameter() -> None:
"""所有 Agent 工具输入模型都不应暴露 explanation 参数。"""
tool_classes = [
*MoviePilotToolFactory.BUILTIN_TOOL_CLASSES,
AskUserChoiceTool,
SendLocalFileTool,
SendVoiceMessageTool,
]
middleware_schemas = [
SkillToolInput,
QueryActivityLogInput,
]
plugin_schemas = _load_lexiannot_tool_schemas()
for tool_class in tool_classes:
args_schema = getattr(tool_class, "args_schema", None)
if args_schema is None:
continue
assert "explanation" not in _schema_properties(args_schema), tool_class.name
for args_schema in middleware_schemas + plugin_schemas:
assert "explanation" not in _schema_properties(args_schema), args_schema.__name__
def test_ask_user_choice_option_schema_does_not_expose_description() -> None:
"""询问用户意图工具的选项参数不应暴露 description 字段。"""
schema = AskUserChoiceInput.model_json_schema()
option_schema = schema["$defs"]["UserChoiceOptionInput"]
assert "description" not in option_schema["properties"]
assert option_schema["required"] == ["label", "value"]
def test_plugin_agent_tools_are_cached(plugin_manager: PluginManager) -> None:
"""插件智能体工具注册表应缓存,避免同一轮启动反复询问插件实例。"""
calls: list[int] = []
+8 -4
View File
@@ -49,6 +49,10 @@ class DummyTool(MoviePilotTool):
name: str = "dummy_tool"
description: str = "Dummy tool for streaming tests."
def get_tool_message(self, **kwargs) -> str:
"""返回固定工具执行提示。"""
return "run test tool"
async def run(self, **kwargs) -> str:
"""返回固定工具执行结果。"""
return "ok"
@@ -67,7 +71,7 @@ class TestAgentToolStreaming:
tool.set_stream_handler(handler)
with patch.object(settings, "AI_AGENT_VERBOSE", False):
result = await tool._arun(explanation="run test tool")
result = await tool._arun()
buffered_message = await handler.take()
return result, buffered_message
@@ -103,7 +107,7 @@ class TestAgentToolStreaming:
tool.set_stream_handler(handler)
with patch.object(settings, "AI_AGENT_VERBOSE", False):
await tool._arun(explanation="run test tool")
await tool._arun()
handler.emit("已经拿到结果")
return await handler.take()
@@ -470,7 +474,7 @@ class TestAgentToolStreaming:
DummyTool, "send_tool_message", new_callable=AsyncMock
) as send_tool_message,
):
result = await tool._arun(explanation="run test tool")
result = await tool._arun()
buffered_message = await handler.take()
return result, buffered_message, send_tool_message
@@ -497,7 +501,7 @@ class TestAgentToolStreaming:
DummyTool, "send_tool_message", new_callable=AsyncMock
) as send_tool_message,
):
result = await tool._arun(explanation="run test tool")
result = await tool._arun()
buffered_message = await handler.take()
return result, buffered_message, send_tool_message