mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-05 07:27:15 +08:00
feat: enhance agent download task controls
This commit is contained in:
@@ -53,7 +53,7 @@ from app.agent.tools.impl.update_site_cookie import UpdateSiteCookieTool
|
||||
from app.agent.tools.impl.delete_download import DeleteDownloadTool
|
||||
from app.agent.tools.impl.delete_download_history import DeleteDownloadHistoryTool
|
||||
from app.agent.tools.impl.delete_transfer_history import DeleteTransferHistoryTool
|
||||
from app.agent.tools.impl.modify_download import ModifyDownloadTool
|
||||
from app.agent.tools.impl.update_download_tasks import UpdateDownloadTasksTool
|
||||
from app.agent.tools.impl.query_directory_settings import QueryDirectorySettingsTool
|
||||
from app.agent.tools.impl.list_directory import ListDirectoryTool
|
||||
from app.agent.tools.impl.query_transfer_history import QueryTransferHistoryTool
|
||||
@@ -186,7 +186,7 @@ class MoviePilotToolFactory:
|
||||
DeleteDownloadTool,
|
||||
DeleteDownloadHistoryTool,
|
||||
DeleteTransferHistoryTool,
|
||||
ModifyDownloadTool,
|
||||
UpdateDownloadTasksTool,
|
||||
QueryDownloadersTool,
|
||||
QuerySitesTool,
|
||||
UpdateSiteTool,
|
||||
|
||||
@@ -1,143 +0,0 @@
|
||||
"""修改下载任务工具"""
|
||||
|
||||
from typing import Optional, Type, List
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.chain.download import DownloadChain
|
||||
from app.log import logger
|
||||
|
||||
|
||||
class ModifyDownloadInput(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)"
|
||||
)
|
||||
action: Optional[str] = Field(
|
||||
None,
|
||||
description="Action to perform on the task: 'start' to resume downloading, 'stop' to pause downloading. "
|
||||
"If not provided, no start/stop action will be performed.",
|
||||
)
|
||||
tags: Optional[List[str]] = Field(
|
||||
None,
|
||||
description="List of tags to set on the download task. If provided, these tags will be added to the task. "
|
||||
"Example: ['movie', 'hd']",
|
||||
)
|
||||
downloader: Optional[str] = Field(
|
||||
None,
|
||||
description="Name of specific downloader (optional, if not provided will search all downloaders)",
|
||||
)
|
||||
|
||||
|
||||
class ModifyDownloadTool(MoviePilotTool):
|
||||
"""修改下载任务工具"""
|
||||
|
||||
name: str = "modify_download"
|
||||
tags: list[str] = [
|
||||
ToolTag.Write,
|
||||
ToolTag.Download,
|
||||
ToolTag.Admin,
|
||||
]
|
||||
description: str = (
|
||||
"Modify a download task in the downloader by task hash. "
|
||||
"Supports: 1) Setting tags on a download task, "
|
||||
"2) Starting (resuming) a paused download task, "
|
||||
"3) Stopping (pausing) a downloading task. "
|
||||
"Multiple operations can be performed in a single call."
|
||||
)
|
||||
args_schema: Type[BaseModel] = ModifyDownloadInput
|
||||
require_admin: bool = True
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
hash_value = kwargs.get("hash", "")
|
||||
action = kwargs.get("action")
|
||||
tags = kwargs.get("tags")
|
||||
downloader = kwargs.get("downloader")
|
||||
|
||||
parts = [f"修改下载任务: {hash_value}"]
|
||||
if action == "start":
|
||||
parts.append("操作: 开始下载")
|
||||
elif action == "stop":
|
||||
parts.append("操作: 暂停下载")
|
||||
if tags:
|
||||
parts.append(f"标签: {', '.join(tags)}")
|
||||
if downloader:
|
||||
parts.append(f"下载器: {downloader}")
|
||||
return " | ".join(parts)
|
||||
|
||||
@staticmethod
|
||||
def _modify_download_sync(
|
||||
hash_value: str,
|
||||
action: Optional[str] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
downloader: Optional[str] = None,
|
||||
) -> List[str]:
|
||||
"""同步修改下载任务状态和标签,避免下载器 SDK 阻塞事件循环。"""
|
||||
download_chain = DownloadChain()
|
||||
results = []
|
||||
|
||||
if tags:
|
||||
tag_result = download_chain.set_torrents_tag(
|
||||
hashs=[hash_value], tags=tags, downloader=downloader
|
||||
)
|
||||
if tag_result:
|
||||
results.append(f"成功设置标签:{', '.join(tags)}")
|
||||
else:
|
||||
results.append("设置标签失败,请检查任务是否存在或下载器是否可用")
|
||||
|
||||
if action:
|
||||
action_result = download_chain.set_downloading(
|
||||
hash_str=hash_value, oper=action, name=downloader
|
||||
)
|
||||
action_desc = "开始" if action == "start" else "暂停"
|
||||
if action_result:
|
||||
results.append(f"成功{action_desc}下载任务")
|
||||
else:
|
||||
results.append(f"{action_desc}下载任务失败,请检查任务是否存在或下载器是否可用")
|
||||
|
||||
return results
|
||||
|
||||
async def run(
|
||||
self,
|
||||
hash: str,
|
||||
action: Optional[str] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
downloader: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
logger.info(
|
||||
f"执行工具: {self.name}, 参数: hash={hash}, action={action}, tags={tags}, downloader={downloader}"
|
||||
)
|
||||
|
||||
try:
|
||||
# 校验 hash 格式
|
||||
if len(hash) != 40 or not all(c in "0123456789abcdefABCDEF" for c in hash):
|
||||
return "参数错误:hash 格式无效,请先使用 query_download_tasks 工具获取正确的 hash。"
|
||||
|
||||
# 校验参数:至少需要一个操作
|
||||
if not action and not tags:
|
||||
return "参数错误:至少需要指定 action(start/stop)或 tags 中的一个。"
|
||||
|
||||
# 校验 action 参数
|
||||
if action and action not in ("start", "stop"):
|
||||
return f"参数错误:action 只支持 'start'(开始下载)或 'stop'(暂停下载),收到: '{action}'。"
|
||||
|
||||
results = await self.run_blocking(
|
||||
"downloader",
|
||||
self._modify_download_sync,
|
||||
hash,
|
||||
action,
|
||||
tags,
|
||||
downloader,
|
||||
)
|
||||
|
||||
return f"下载任务 {hash}:" + ";".join(results)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"修改下载任务失败: {e}", exc_info=True)
|
||||
return f"修改下载任务时发生错误: {str(e)}"
|
||||
@@ -25,6 +25,10 @@ class QueryDownloadTasksInput(BaseModel):
|
||||
False,
|
||||
description="Include tasks without the MoviePilot built-in tag. Default false keeps the normal MoviePilot task scope.",
|
||||
)
|
||||
include_trackers: Optional[bool] = Field(
|
||||
False,
|
||||
description="Include tracker URLs when supported. Hash queries always include trackers.",
|
||||
)
|
||||
hash: Optional[str] = Field(None, description="Query specific download task by hash (optional, if provided will search for this specific task regardless of status)")
|
||||
title: Optional[str] = Field(None, description="Query download tasks by title/name (optional, supports partial match, searches all tasks if provided)")
|
||||
tag: Optional[str] = Field(None, description="Filter download tasks by tag (optional, supports partial match, e.g. 'movie' will match tasks with tag 'movie' or 'movie_2024')")
|
||||
@@ -131,6 +135,7 @@ class QueryDownloadTasksTool(MoviePilotTool):
|
||||
title: Optional[str] = None,
|
||||
tag: Optional[str] = None,
|
||||
include_all_tags: bool = False,
|
||||
include_trackers: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
同步查询下载器和下载历史,整个链路放在线程池中执行。
|
||||
@@ -214,6 +219,16 @@ class QueryDownloadTasksTool(MoviePilotTool):
|
||||
if not filtered_downloads:
|
||||
return {"message": "未找到相关下载任务"}
|
||||
|
||||
if hash_value or include_trackers:
|
||||
for torrent in filtered_downloads:
|
||||
if not getattr(torrent, "hash", None):
|
||||
continue
|
||||
tracker_map = download_chain.get_torrent_trackers(
|
||||
hash_string=torrent.hash,
|
||||
downloader=getattr(torrent, "downloader", None) or downloader,
|
||||
) or {}
|
||||
torrent.trackers = tracker_map.get(getattr(torrent, "downloader", None)) or []
|
||||
|
||||
return {"downloads": filtered_downloads}
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
@@ -245,6 +260,8 @@ class QueryDownloadTasksTool(MoviePilotTool):
|
||||
parts.append(f"标签: {tag}")
|
||||
if include_all_tags:
|
||||
parts.append("范围: 全部标签")
|
||||
if kwargs.get("include_trackers"):
|
||||
parts.append("包含Tracker")
|
||||
|
||||
return " | ".join(parts) if len(parts) > 1 else parts[0]
|
||||
|
||||
@@ -254,10 +271,12 @@ class QueryDownloadTasksTool(MoviePilotTool):
|
||||
title: Optional[str] = None,
|
||||
tag: Optional[str] = None,
|
||||
include_all_tags: Optional[bool] = False,
|
||||
include_trackers: Optional[bool] = False,
|
||||
**kwargs) -> str:
|
||||
logger.info(
|
||||
f"执行工具: {self.name}, 参数: downloader={downloader}, status={status}, "
|
||||
f"hash={hash}, title={title}, tag={tag}, include_all_tags={include_all_tags}"
|
||||
f"hash={hash}, title={title}, tag={tag}, include_all_tags={include_all_tags}, "
|
||||
f"include_trackers={include_trackers}"
|
||||
)
|
||||
try:
|
||||
payload = await self.run_blocking(
|
||||
@@ -269,6 +288,7 @@ class QueryDownloadTasksTool(MoviePilotTool):
|
||||
title,
|
||||
tag,
|
||||
self._normalize_include_all_tags(include_all_tags),
|
||||
self._normalize_include_all_tags(include_trackers),
|
||||
)
|
||||
if payload.get("message"):
|
||||
return payload["message"]
|
||||
@@ -294,6 +314,16 @@ class QueryDownloadTasksTool(MoviePilotTool):
|
||||
"upspeed": getattr(d, "upspeed", None),
|
||||
"dlspeed": getattr(d, "dlspeed", None),
|
||||
"tags": d.tags,
|
||||
"save_path": getattr(d, "save_path", None),
|
||||
"content_path": getattr(d, "content_path", None) or (
|
||||
d.path.as_posix() if getattr(d, "path", None) else None
|
||||
),
|
||||
"category": getattr(d, "category", None),
|
||||
"download_limit": getattr(d, "download_limit", None),
|
||||
"upload_limit": getattr(d, "upload_limit", None),
|
||||
"ratio_limit": getattr(d, "ratio_limit", None),
|
||||
"seeding_time_limit": getattr(d, "seeding_time_limit", None),
|
||||
"trackers": getattr(d, "trackers", None) or [],
|
||||
"left_time": getattr(d, "left_time", None)
|
||||
}
|
||||
# 精简 media 字段
|
||||
|
||||
@@ -0,0 +1,310 @@
|
||||
"""更新下载任务工具"""
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional, Type
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.chain.download import DownloadChain
|
||||
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)"
|
||||
)
|
||||
action: Optional[str] = Field(
|
||||
None,
|
||||
description="Action to perform on the task: 'start' to resume downloading, 'stop' to pause downloading.",
|
||||
)
|
||||
tags: Optional[List[str]] = Field(
|
||||
None,
|
||||
description="List of tags to add to the download task. Example: ['movie', 'hd']",
|
||||
)
|
||||
downloader: Optional[str] = Field(
|
||||
None,
|
||||
description="Name of specific downloader. If omitted, the tool resolves it from the task hash.",
|
||||
)
|
||||
download_limit: Optional[float] = Field(
|
||||
None,
|
||||
description="Per-task download speed limit in KB/s. Use 0 to disable the limit when supported.",
|
||||
)
|
||||
upload_limit: Optional[float] = Field(
|
||||
None,
|
||||
description="Per-task upload speed limit in KB/s. Use 0 to disable the limit when supported.",
|
||||
)
|
||||
trackers: Optional[List[str]] = Field(
|
||||
None,
|
||||
description="Tracker URL list to add or set, depending on downloader support.",
|
||||
)
|
||||
save_path: Optional[str] = Field(
|
||||
None,
|
||||
description="New save/download directory for the task, when supported.",
|
||||
)
|
||||
category: Optional[str] = Field(
|
||||
None,
|
||||
description="Downloader category to set, when supported.",
|
||||
)
|
||||
ratio_limit: Optional[float] = Field(
|
||||
None,
|
||||
description="Per-task share ratio limit, when supported.",
|
||||
)
|
||||
seeding_time_limit: Optional[int] = Field(
|
||||
None,
|
||||
description="Per-task seeding time limit in minutes, when supported.",
|
||||
)
|
||||
|
||||
|
||||
class UpdateDownloadTasksTool(MoviePilotTool):
|
||||
"""更新下载任务工具"""
|
||||
|
||||
name: str = "update_download_tasks"
|
||||
tags: list[str] = [
|
||||
ToolTag.Write,
|
||||
ToolTag.Download,
|
||||
ToolTag.Admin,
|
||||
]
|
||||
description: str = (
|
||||
"Update a download task by hash. Supports start/stop, adding tags, per-task "
|
||||
"upload/download speed limits, trackers, save directory, category, share ratio, "
|
||||
"and seeding time where the configured downloader supports them. "
|
||||
"Use query_download_tasks first to get the hash and current downloader."
|
||||
)
|
||||
args_schema: Type[BaseModel] = UpdateDownloadTasksInput
|
||||
require_admin: bool = True
|
||||
|
||||
@staticmethod
|
||||
def _is_valid_hash(hash_value: str) -> bool:
|
||||
"""校验下载任务Hash格式。"""
|
||||
return len(hash_value) == 40 and all(c in "0123456789abcdefABCDEF" for c in hash_value)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_non_empty_list(values: Optional[List[str]]) -> Optional[List[str]]:
|
||||
"""清理字符串列表中的空值。"""
|
||||
if values is None:
|
||||
return None
|
||||
return [str(value).strip() for value in values if str(value).strip()]
|
||||
|
||||
@staticmethod
|
||||
def _has_update_params(**kwargs) -> bool:
|
||||
"""判断是否传入至少一个修改参数。"""
|
||||
return any(value is not None and value != [] for value in kwargs.values())
|
||||
|
||||
@staticmethod
|
||||
def _build_result(operation: str, success: bool, message: str) -> Dict[str, Any]:
|
||||
"""构造单项操作结果。"""
|
||||
return {
|
||||
"operation": operation,
|
||||
"success": success,
|
||||
"message": message,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def _resolve_downloader(
|
||||
cls,
|
||||
download_chain: DownloadChain,
|
||||
hash_value: str,
|
||||
downloader: Optional[str],
|
||||
) -> Optional[str]:
|
||||
"""根据Hash解析下载任务所在下载器。"""
|
||||
if downloader:
|
||||
return downloader
|
||||
torrents = download_chain.list_torrents(
|
||||
hashs=[hash_value],
|
||||
include_all_tags=True,
|
||||
) or []
|
||||
return getattr(torrents[0], "downloader", None) if torrents else None
|
||||
|
||||
@classmethod
|
||||
def _update_download_sync(
|
||||
cls,
|
||||
hash_value: str,
|
||||
action: Optional[str] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
downloader: Optional[str] = None,
|
||||
download_limit: Optional[float] = None,
|
||||
upload_limit: Optional[float] = None,
|
||||
trackers: Optional[List[str]] = None,
|
||||
save_path: Optional[str] = None,
|
||||
category: Optional[str] = None,
|
||||
ratio_limit: Optional[float] = None,
|
||||
seeding_time_limit: Optional[int] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""同步更新下载任务,避免下载器 SDK 阻塞事件循环。"""
|
||||
download_chain = DownloadChain()
|
||||
resolved_downloader = cls._resolve_downloader(
|
||||
download_chain=download_chain,
|
||||
hash_value=hash_value,
|
||||
downloader=downloader,
|
||||
)
|
||||
if not resolved_downloader:
|
||||
return {
|
||||
"hash": hash_value,
|
||||
"downloader": downloader,
|
||||
"results": [
|
||||
cls._build_result("resolve_downloader", False, "未找到下载任务或下载器不可用")
|
||||
],
|
||||
}
|
||||
|
||||
results = []
|
||||
if tags:
|
||||
tag_result = download_chain.set_torrents_tag(
|
||||
hashs=[hash_value], tags=tags, downloader=resolved_downloader
|
||||
)
|
||||
results.append(
|
||||
cls._build_result(
|
||||
"tags",
|
||||
bool(tag_result),
|
||||
f"成功设置标签:{', '.join(tags)}" if tag_result else "设置标签失败",
|
||||
)
|
||||
)
|
||||
|
||||
if action:
|
||||
action_result = download_chain.set_downloading(
|
||||
hash_str=hash_value, oper=action, name=resolved_downloader
|
||||
)
|
||||
action_desc = "开始" if action == "start" else "暂停"
|
||||
results.append(
|
||||
cls._build_result(
|
||||
action,
|
||||
bool(action_result),
|
||||
f"成功{action_desc}下载任务" if action_result else f"{action_desc}下载任务失败",
|
||||
)
|
||||
)
|
||||
|
||||
update_result = {}
|
||||
if cls._has_update_params(
|
||||
download_limit=download_limit,
|
||||
upload_limit=upload_limit,
|
||||
trackers=trackers,
|
||||
save_path=save_path,
|
||||
category=category,
|
||||
ratio_limit=ratio_limit,
|
||||
seeding_time_limit=seeding_time_limit,
|
||||
):
|
||||
update_result = download_chain.update_torrent(
|
||||
hash_string=hash_value,
|
||||
downloader=resolved_downloader,
|
||||
download_limit=download_limit,
|
||||
upload_limit=upload_limit,
|
||||
tracker_list=trackers,
|
||||
save_path=save_path,
|
||||
category=category,
|
||||
ratio_limit=ratio_limit,
|
||||
seeding_time_limit=seeding_time_limit,
|
||||
)
|
||||
operation_messages = {
|
||||
"limits": "限速/做种策略",
|
||||
"trackers": "Tracker",
|
||||
"save_path": "保存目录",
|
||||
"category": "分类",
|
||||
}
|
||||
for operation, success in (update_result or {}).items():
|
||||
label = operation_messages.get(operation, operation)
|
||||
results.append(
|
||||
cls._build_result(
|
||||
operation,
|
||||
bool(success),
|
||||
f"{label}修改成功" if success else f"{label}修改失败或下载器不支持",
|
||||
)
|
||||
)
|
||||
|
||||
return {
|
||||
"hash": hash_value,
|
||||
"downloader": resolved_downloader,
|
||||
"results": results,
|
||||
}
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
"""根据更新参数生成友好的提示消息。"""
|
||||
hash_value = kwargs.get("hash", "")
|
||||
parts = [f"更新下载任务: {hash_value}"]
|
||||
action = kwargs.get("action")
|
||||
if action == "start":
|
||||
parts.append("操作: 开始下载")
|
||||
elif action == "stop":
|
||||
parts.append("操作: 暂停下载")
|
||||
if kwargs.get("tags"):
|
||||
parts.append(f"标签: {', '.join(kwargs.get('tags'))}")
|
||||
if kwargs.get("download_limit") is not None or kwargs.get("upload_limit") is not None:
|
||||
parts.append("限速")
|
||||
if kwargs.get("trackers") is not None:
|
||||
parts.append("Tracker")
|
||||
if kwargs.get("save_path"):
|
||||
parts.append("保存目录")
|
||||
if kwargs.get("category") is not None:
|
||||
parts.append("分类")
|
||||
if kwargs.get("downloader"):
|
||||
parts.append(f"下载器: {kwargs.get('downloader')}")
|
||||
return " | ".join(parts)
|
||||
|
||||
async def run(
|
||||
self,
|
||||
hash: str,
|
||||
action: Optional[str] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
downloader: Optional[str] = None,
|
||||
download_limit: Optional[float] = None,
|
||||
upload_limit: Optional[float] = None,
|
||||
trackers: Optional[List[str]] = None,
|
||||
save_path: Optional[str] = None,
|
||||
category: Optional[str] = None,
|
||||
ratio_limit: Optional[float] = None,
|
||||
seeding_time_limit: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""执行下载任务更新。"""
|
||||
logger.info(
|
||||
f"执行工具: {self.name}, 参数: hash={hash}, action={action}, tags={tags}, "
|
||||
f"downloader={downloader}, download_limit={download_limit}, upload_limit={upload_limit}, "
|
||||
f"trackers={trackers}, save_path={save_path}, category={category}, "
|
||||
f"ratio_limit={ratio_limit}, seeding_time_limit={seeding_time_limit}"
|
||||
)
|
||||
try:
|
||||
if not self._is_valid_hash(hash):
|
||||
return "参数错误:hash 格式无效,请先使用 query_download_tasks 工具获取正确的 hash。"
|
||||
|
||||
tags = self._normalize_non_empty_list(tags)
|
||||
trackers = self._normalize_non_empty_list(trackers)
|
||||
if action and action not in ("start", "stop"):
|
||||
return f"参数错误:action 只支持 'start'(开始下载)或 'stop'(暂停下载),收到: '{action}'。"
|
||||
if not self._has_update_params(
|
||||
action=action,
|
||||
tags=tags,
|
||||
download_limit=download_limit,
|
||||
upload_limit=upload_limit,
|
||||
trackers=trackers,
|
||||
save_path=save_path,
|
||||
category=category,
|
||||
ratio_limit=ratio_limit,
|
||||
seeding_time_limit=seeding_time_limit,
|
||||
):
|
||||
return "参数错误:至少需要指定一个要更新的字段。"
|
||||
|
||||
result = await self.run_blocking(
|
||||
"downloader",
|
||||
self._update_download_sync,
|
||||
hash,
|
||||
action,
|
||||
tags,
|
||||
downloader,
|
||||
download_limit,
|
||||
upload_limit,
|
||||
trackers,
|
||||
save_path,
|
||||
category,
|
||||
ratio_limit,
|
||||
seeding_time_limit,
|
||||
)
|
||||
return json.dumps(result, ensure_ascii=False, indent=2)
|
||||
except Exception as e:
|
||||
logger.error(f"更新下载任务失败: {e}", exc_info=True)
|
||||
return f"更新下载任务时发生错误: {str(e)}"
|
||||
Reference in New Issue
Block a user