fix agent

This commit is contained in:
jxxghp
2025-11-01 10:39:08 +08:00
parent d523c7c916
commit 438d3210bc
18 changed files with 145 additions and 71 deletions
+21 -9
View File
@@ -3,8 +3,14 @@
from langchain.tools import BaseTool
from pydantic import PrivateAttr
from app.chain import ChainBase
from app.helper.message import MessageHelper
from app.log import logger
from app.schemas import Notification
class ToolChain(ChainBase):
pass
class MoviePilotTool(BaseTool):
@@ -14,11 +20,15 @@ class MoviePilotTool(BaseTool):
_user_id: str = PrivateAttr()
_message_helper: MessageHelper = PrivateAttr()
def __init__(self, session_id: str, user_id: str, message_helper: MessageHelper = None, **kwargs):
def __init__(self, session_id: str, user_id: str,
channel: str = None, source: str = None, username: str = None, **kwargs):
super().__init__(**kwargs)
self._session_id = session_id
self._user_id = user_id
self._message_helper = message_helper or MessageHelper()
self.channel = channel
self.source = source
self.username = username
self._message_helper = MessageHelper()
def _run(self, **kwargs) -> str:
raise NotImplementedError
@@ -28,11 +38,13 @@ class MoviePilotTool(BaseTool):
def _send_tool_message(self, message: str, title: str = None, **kwargs):
"""发送工具执行消息"""
try:
self._message_helper.put(
message=message,
role="system",
title=title or "工具执行"
ToolChain().post_message(
Notification(
channel=self.channel,
source=self.source,
userid=self.user_id,
username=self.username,
title=title,
text=message
)
except Exception as e:
logger.error(f"发送工具消息失败: {e}")
)
+15 -13
View File
@@ -2,26 +2,26 @@
from typing import List
from app.helper.message import MessageHelper
from app.agent.tools.impl.add_download import AddDownloadTool
from app.agent.tools.impl.add_subscribe import AddSubscribeTool
from app.agent.tools.impl.get_recommendations import GetRecommendationsTool
from app.agent.tools.impl.query_downloaders import QueryDownloadersTool
from app.agent.tools.impl.query_downloads import QueryDownloadsTool
from app.agent.tools.impl.query_media_library import QueryMediaLibraryTool
from app.agent.tools.impl.query_subscribes import QuerySubscribesTool
from app.agent.tools.impl.search_media import SearchMediaTool
from app.agent.tools.impl.search_torrents import SearchTorrentsTool
from app.agent.tools.impl.send_message import SendMessageTool
from app.log import logger
from .base import MoviePilotTool
from app.agent.tools.impl.search_media import SearchMediaTool
from app.agent.tools.impl.add_subscribe import AddSubscribeTool
from app.agent.tools.impl.search_torrents import SearchTorrentsTool
from app.agent.tools.impl.add_download import AddDownloadTool
from app.agent.tools.impl.query_subscribes import QuerySubscribesTool
from app.agent.tools.impl.query_downloads import QueryDownloadsTool
from app.agent.tools.impl.query_downloaders import QueryDownloadersTool
from app.agent.tools.impl.get_recommendations import GetRecommendationsTool
from app.agent.tools.impl.query_media_library import QueryMediaLibraryTool
from app.agent.tools.impl.send_message import SendMessageTool
class MoviePilotToolFactory:
"""MoviePilot工具工厂"""
@staticmethod
def create_tools(session_id: str, user_id: str, message_helper: MessageHelper = None) -> List[MoviePilotTool]:
def create_tools(session_id: str, user_id: str,
channel: str = None, source: str = None, username: str = None) -> List[MoviePilotTool]:
"""创建MoviePilot工具列表"""
tools = []
tool_definitions = [
@@ -40,7 +40,9 @@ class MoviePilotToolFactory:
tools.append(ToolClass(
session_id=session_id,
user_id=user_id,
message_helper=message_helper
channel=channel,
source=source,
username=username
))
logger.info(f"成功创建 {len(tools)} 个MoviePilot工具")
return tools
+1 -1
View File
@@ -16,7 +16,7 @@ class AddDownloadTool(MoviePilotTool):
async def _arun(self, torrent_title: str, torrent_url: str, explanation: str,
downloader: Optional[str] = None, save_path: Optional[str] = None,
labels: Optional[str] = None) -> str:
labels: Optional[str] = None, **kwargs) -> str:
logger.info(f"执行工具: {self.name}, 参数: torrent_title={torrent_title}, torrent_url={torrent_url}, downloader={downloader}, save_path={save_path}, labels={labels}")
# 发送工具执行说明
+1 -1
View File
@@ -13,7 +13,7 @@ class AddSubscribeTool(MoviePilotTool):
description: str = "添加媒体订阅,为用户感兴趣的媒体内容创建订阅规则。"
async def _arun(self, title: str, year: str, media_type: str, explanation: str,
season: Optional[int] = None, tmdb_id: Optional[str] = None) -> str:
season: Optional[int] = None, tmdb_id: Optional[str] = None, **kwargs) -> str:
logger.info(f"执行工具: {self.name}, 参数: title={title}, year={year}, media_type={media_type}, season={season}, tmdb_id={tmdb_id}")
# 发送工具执行说明
+1 -1
View File
@@ -13,7 +13,7 @@ class GetRecommendationsTool(MoviePilotTool):
description: str = "获取热门媒体推荐,包括电影、电视剧等热门内容。"
async def _arun(self, explanation: str, source: Optional[str] = "tmdb_trending",
media_type: Optional[str] = "all", limit: Optional[int] = 20) -> str:
media_type: Optional[str] = "all", limit: Optional[int] = 20, **kwargs) -> str:
logger.info(f"执行工具: {self.name}, 参数: source={source}, media_type={media_type}, limit={limit}")
try:
recommend_chain = RecommendChain()
+1 -1
View File
@@ -12,7 +12,7 @@ class QueryDownloadersTool(MoviePilotTool):
name: str = "query_downloaders"
description: str = "查询下载器配置,查看可用的下载器列表和配置信息。"
async def _arun(self, explanation: str) -> str:
async def _arun(self, explanation: str, **kwargs) -> str:
logger.info(f"执行工具: {self.name}")
try:
system_config_oper = SystemConfigOper()
+1 -1
View File
@@ -13,7 +13,7 @@ class QueryDownloadsTool(MoviePilotTool):
description: str = "查询下载状态,查看下载器的任务列表和进度。"
async def _arun(self, explanation: str, downloader: Optional[str] = None,
status: Optional[str] = "all") -> str:
status: Optional[str] = "all", **kwargs) -> str:
logger.info(f"执行工具: {self.name}, 参数: downloader={downloader}, status={status}")
try:
download_chain = DownloadChain()
+1 -1
View File
@@ -13,7 +13,7 @@ class QueryMediaLibraryTool(MoviePilotTool):
description: str = "查询媒体库状态,查看已入库的媒体文件情况。"
async def _arun(self, explanation: str, media_type: Optional[str] = "all",
title: Optional[str] = None) -> str:
title: Optional[str] = None, **kwargs) -> str:
logger.info(f"执行工具: {self.name}, 参数: media_type={media_type}, title={title}")
try:
media_server_oper = MediaServerOper()
+1 -1
View File
@@ -13,7 +13,7 @@ class QuerySubscribesTool(MoviePilotTool):
description: str = "查询订阅状态,查看用户的订阅列表和状态。"
async def _arun(self, explanation: str, status: Optional[str] = "all",
media_type: Optional[str] = "all") -> str:
media_type: Optional[str] = "all", **kwargs) -> str:
logger.info(f"执行工具: {self.name}, 参数: status={status}, media_type={media_type}")
try:
subscribe_oper = SubscribeOper()
+1 -1
View File
@@ -15,7 +15,7 @@ class SearchMediaTool(MoviePilotTool):
description: str = "搜索媒体资源,包括电影、电视剧、动漫等。可以根据标题、年份、类型等条件进行搜索。"
async def _arun(self, title: str, explanation: str, year: Optional[str] = None,
media_type: Optional[str] = None, season: Optional[int] = None) -> str:
media_type: Optional[str] = None, season: Optional[int] = None, **kwargs) -> str:
logger.info(f"执行工具: {self.name}, 参数: title={title}, year={year}, media_type={media_type}, season={season}")
# 发送工具执行说明
+1 -1
View File
@@ -15,7 +15,7 @@ class SearchTorrentsTool(MoviePilotTool):
async def _arun(self, title: str, explanation: str, year: Optional[str] = None,
media_type: Optional[str] = None, season: Optional[int] = None,
sites: Optional[List[int]] = None) -> str:
sites: Optional[List[int]] = None, **kwargs) -> str:
logger.info(f"执行工具: {self.name}, 参数: title={title}, year={year}, media_type={media_type}, season={season}, sites={sites}")
# 发送工具执行说明
+2 -2
View File
@@ -11,11 +11,11 @@ class SendMessageTool(MoviePilotTool):
name: str = "send_message"
description: str = "发送消息通知,向用户发送操作结果或重要信息。"
async def _arun(self, message: str, explanation: str, message_type: Optional[str] = "info") -> str:
async def _arun(self, message: str, explanation: str, message_type: Optional[str] = "info", **kwargs) -> str:
logger.info(f"执行工具: {self.name}, 参数: message={message}, message_type={message_type}")
try:
message_helper = MessageHelper()
message_helper.put(message=message, role="system", title=f"AI助手通知 ({message_type})")
message_helper.put(message=message, role="system", title=f"MoviePilot助手通知 ({message_type})")
return "消息已发送。"
except Exception as e:
logger.error(f"发送消息失败: {e}")