fixx loop

This commit is contained in:
jxxghp
2025-11-20 08:15:37 +08:00
parent 5c983b64bc
commit 48da5c976c
8 changed files with 77 additions and 58 deletions
+8 -6
View File
@@ -8,7 +8,7 @@ from pydantic import BaseModel, Field
from app.agent.tools.base import MoviePilotTool from app.agent.tools.base import MoviePilotTool
from app.chain.media import MediaChain from app.chain.media import MediaChain
from app.core.config import GlobalVar from app.core.config import global_vars
from app.core.metainfo import MetaInfoPath from app.core.metainfo import MetaInfoPath
from app.log import logger from app.log import logger
from app.schemas import FileItem from app.schemas import FileItem
@@ -17,9 +17,12 @@ from app.schemas import FileItem
class ScrapeMetadataInput(BaseModel): class ScrapeMetadataInput(BaseModel):
"""刮削媒体元数据工具的输入参数模型""" """刮削媒体元数据工具的输入参数模型"""
explanation: str = Field(..., description="Clear explanation of why this tool is being used in the current context") explanation: str = Field(..., 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')") path: str = Field(...,
storage: Optional[str] = Field("local", description="Storage type: 'local' for local storage, 'smb', 'alist', etc. for remote storage (default: 'local')") description="Path to the file or directory to scrape metadata for (e.g., '/path/to/file.mkv' or '/path/to/directory')")
overwrite: Optional[bool] = Field(False, description="Whether to overwrite existing metadata files (default: False)") storage: Optional[str] = Field("local",
description="Storage type: 'local' for local storage, 'smb', 'alist', etc. for remote storage (default: 'local')")
overwrite: Optional[bool] = Field(False,
description="Whether to overwrite existing metadata files (default: False)")
class ScrapeMetadataTool(MoviePilotTool): class ScrapeMetadataTool(MoviePilotTool):
@@ -83,7 +86,7 @@ class ScrapeMetadataTool(MoviePilotTool):
}, ensure_ascii=False) }, ensure_ascii=False)
# 在线程池中执行同步的刮削操作 # 在线程池中执行同步的刮削操作
await GlobalVar.CURRENT_EVENT_LOOP.run_in_executor( await global_vars.loop.run_in_executor(
None, None,
lambda: media_chain.scrape_metadata( lambda: media_chain.scrape_metadata(
fileitem=fileitem, fileitem=fileitem,
@@ -114,4 +117,3 @@ class ScrapeMetadataTool(MoviePilotTool):
"message": error_message, "message": error_message,
"path": path "path": path
}, ensure_ascii=False) }, ensure_ascii=False)
+4 -4
View File
@@ -7,7 +7,7 @@ from pydantic import BaseModel, Field
from app.agent.tools.base import MoviePilotTool from app.agent.tools.base import MoviePilotTool
from app.chain.subscribe import SubscribeChain from app.chain.subscribe import SubscribeChain
from app.core.config import GlobalVar from app.core.config import global_vars
from app.db.subscribe_oper import SubscribeOper from app.db.subscribe_oper import SubscribeOper
from app.log import logger from app.log import logger
@@ -39,7 +39,8 @@ class SearchSubscribeTool(MoviePilotTool):
async def run(self, subscribe_id: int, manual: Optional[bool] = False, async def run(self, subscribe_id: int, manual: Optional[bool] = False,
filter_groups: Optional[List[str]] = None, **kwargs) -> str: filter_groups: Optional[List[str]] = None, **kwargs) -> str:
logger.info(f"执行工具: {self.name}, 参数: subscribe_id={subscribe_id}, manual={manual}, filter_groups={filter_groups}") logger.info(
f"执行工具: {self.name}, 参数: subscribe_id={subscribe_id}, manual={manual}, filter_groups={filter_groups}")
try: try:
# 先验证订阅是否存在 # 先验证订阅是否存在
@@ -85,7 +86,7 @@ class SearchSubscribeTool(MoviePilotTool):
# 在线程池中执行同步的搜索操作 # 在线程池中执行同步的搜索操作
# 当 sid 有值时,state 参数会被忽略,直接处理该订阅 # 当 sid 有值时,state 参数会被忽略,直接处理该订阅
await GlobalVar.CURRENT_EVENT_LOOP.run_in_executor( await global_vars.loop.run_in_executor(
None, None,
lambda: subscribe_chain.search( lambda: subscribe_chain.search(
sid=subscribe_id, sid=subscribe_id,
@@ -124,4 +125,3 @@ class SearchSubscribeTool(MoviePilotTool):
"message": error_message, "message": error_message,
"subscribe_id": subscribe_id "subscribe_id": subscribe_id
}, ensure_ascii=False) }, ensure_ascii=False)
+11 -10
View File
@@ -1,4 +1,3 @@
import asyncio
import re import re
import time import time
from datetime import datetime, timedelta from datetime import datetime, timedelta
@@ -10,7 +9,7 @@ from app.chain.download import DownloadChain
from app.chain.media import MediaChain from app.chain.media import MediaChain
from app.chain.search import SearchChain from app.chain.search import SearchChain
from app.chain.subscribe import SubscribeChain from app.chain.subscribe import SubscribeChain
from app.core.config import settings, GlobalVar from app.core.config import settings, global_vars
from app.core.context import MediaInfo, Context from app.core.context import MediaInfo, Context
from app.core.meta import MetaBase from app.core.meta import MetaBase
from app.db.user_oper import UserOper from app.db.user_oper import UserOper
@@ -174,7 +173,7 @@ class MessageChain(ChainBase):
elif text.startswith('/ai') or text.startswith('/AI'): elif text.startswith('/ai') or text.startswith('/AI'):
# AI智能体处理 # AI智能体处理
self._handle_ai_message(text=text, channel=channel, source=source, self._handle_ai_message(text=text, channel=channel, source=source,
userid=userid, username=username) userid=userid, username=username)
elif text.startswith('/'): elif text.startswith('/'):
# 执行命令 # 执行命令
self.eventmanager.send_event( self.eventmanager.send_event(
@@ -329,7 +328,8 @@ class MessageChain(ChainBase):
else: else:
best_version = True best_version = True
# 转换用户名 # 转换用户名
mp_name = UserOper().get_name(**{f"{channel.name.lower()}_userid": userid}) if channel else None mp_name = UserOper().get_name(
**{f"{channel.name.lower()}_userid": userid}) if channel else None
# 添加订阅,状态为N # 添加订阅,状态为N
SubscribeChain().add(title=mediainfo.title, SubscribeChain().add(title=mediainfo.title,
year=mediainfo.year, year=mediainfo.year,
@@ -505,7 +505,8 @@ class MessageChain(ChainBase):
# 开始搜索 # 开始搜索
if not medias: if not medias:
self.post_message(Notification( self.post_message(Notification(
channel=channel, source=source, title=f"{meta.name} 没有找到对应的媒体信息!", userid=userid)) channel=channel, source=source, title=f"{meta.name} 没有找到对应的媒体信息!",
userid=userid))
return return
logger.info(f"搜索到 {len(medias)} 条相关媒体信息") logger.info(f"搜索到 {len(medias)} 条相关媒体信息")
try: try:
@@ -847,7 +848,8 @@ class MessageChain(ChainBase):
if time_diff <= timedelta(minutes=MessageChain._session_timeout_minutes): if time_diff <= timedelta(minutes=MessageChain._session_timeout_minutes):
# 更新最后使用时间 # 更新最后使用时间
MessageChain._user_sessions[userid] = (session_id, current_time) MessageChain._user_sessions[userid] = (session_id, current_time)
logger.info(f"复用会话ID: {session_id}, 用户: {userid}, 距离上次会话: {time_diff.total_seconds() / 60:.1f}分钟") logger.info(
f"复用会话ID: {session_id}, 用户: {userid}, 距离上次会话: {time_diff.total_seconds() / 60:.1f}分钟")
return session_id return session_id
# 创建新的会话ID # 创建新的会话ID
@@ -881,7 +883,7 @@ class MessageChain(ChainBase):
# 如果有会话ID,同时清除智能体的会话记忆 # 如果有会话ID,同时清除智能体的会话记忆
if session_id: if session_id:
try: try:
GlobalVar.CURRENT_EVENT_LOOP.run_until_complete( global_vars.loop.run_until_complete(
agent_manager.clear_session( agent_manager.clear_session(
session_id=session_id, session_id=session_id,
user_id=str(userid) user_id=str(userid)
@@ -905,7 +907,7 @@ class MessageChain(ChainBase):
)) ))
def _handle_ai_message(self, text: str, channel: MessageChannel, source: str, def _handle_ai_message(self, text: str, channel: MessageChannel, source: str,
userid: Union[str, int], username: str) -> None: userid: Union[str, int], username: str) -> None:
""" """
处理AI智能体消息 处理AI智能体消息
""" """
@@ -948,7 +950,7 @@ class MessageChain(ChainBase):
session_id = self._get_or_create_session_id(userid) session_id = self._get_or_create_session_id(userid)
# 在事件循环中处理 # 在事件循环中处理
GlobalVar.CURRENT_EVENT_LOOP.run_until_complete( global_vars.loop.run_until_complete(
agent_manager.process_message( agent_manager.process_message(
session_id=session_id, session_id=session_id,
user_id=str(userid), user_id=str(userid),
@@ -962,4 +964,3 @@ class MessageChain(ChainBase):
except Exception as e: except Exception as e:
logger.error(f"处理AI智能体消息失败: {e}") logger.error(f"处理AI智能体消息失败: {e}")
self.messagehelper.put(f"AI智能体处理失败: {str(e)}", role="system", title="MoviePilot助手") self.messagehelper.put(f"AI智能体处理失败: {str(e)}", role="system", title="MoviePilot助手")
+13
View File
@@ -920,6 +920,19 @@ class GlobalVar(object):
return True return True
return False return False
@property
def loop(self) -> AbstractEventLoop:
"""
当前循环
"""
return self.CURRENT_EVENT_LOOP
def set_loop(self, loop: AbstractEventLoop):
"""
设置循环
"""
self.CURRENT_EVENT_LOOP = loop
# 全局标识 # 全局标识
global_vars = GlobalVar() global_vars = GlobalVar()
+2 -2
View File
@@ -11,7 +11,7 @@ from typing import Callable, Dict, List, Optional, Tuple, Union, Any
from fastapi.concurrency import run_in_threadpool from fastapi.concurrency import run_in_threadpool
from app.core.config import GlobalVar from app.core.config import global_vars
from app.helper.thread import ThreadHelper from app.helper.thread import ThreadHelper
from app.log import logger from app.log import logger
from app.schemas import ChainEventData from app.schemas import ChainEventData
@@ -453,7 +453,7 @@ class EventManager(metaclass=Singleton):
# 对于异步函数,直接在事件循环中运行 # 对于异步函数,直接在事件循环中运行
asyncio.run_coroutine_threadsafe( asyncio.run_coroutine_threadsafe(
self.__safe_invoke_handler_async(handler, isolated_event), self.__safe_invoke_handler_async(handler, isolated_event),
GlobalVar.CURRENT_EVENT_LOOP global_vars.loop
) )
else: else:
# 对于同步函数,在线程池中运行 # 对于同步函数,在线程池中运行
+2 -2
View File
@@ -21,7 +21,7 @@ from app.chain.site import SiteChain
from app.chain.subscribe import SubscribeChain from app.chain.subscribe import SubscribeChain
from app.chain.transfer import TransferChain from app.chain.transfer import TransferChain
from app.chain.workflow import WorkflowChain from app.chain.workflow import WorkflowChain
from app.core.config import settings, GlobalVar from app.core.config import settings, global_vars
from app.core.event import eventmanager, Event from app.core.event import eventmanager, Event
from app.core.plugin import PluginManager from app.core.plugin import PluginManager
from app.db.systemconfig_oper import SystemConfigOper from app.db.systemconfig_oper import SystemConfigOper
@@ -474,7 +474,7 @@ class Scheduler(metaclass=SingletonClass):
""" """
启动协程 启动协程
""" """
return asyncio.run_coroutine_threadsafe(coro, GlobalVar.CURRENT_EVENT_LOOP) return asyncio.run_coroutine_threadsafe(coro, global_vars.loop)
# 获取定时任务 # 获取定时任务
job = self.__prepare_job(job_id) job = self.__prepare_job(job_id)
+3
View File
@@ -4,6 +4,7 @@ from contextlib import asynccontextmanager
from fastapi import FastAPI from fastapi import FastAPI
from app.chain.system import SystemChain from app.chain.system import SystemChain
from app.core.config import global_vars
from app.helper.system import SystemHelper from app.helper.system import SystemHelper
from app.startup.command_initializer import init_command, stop_command, restart_command from app.startup.command_initializer import init_command, stop_command, restart_command
from app.startup.modules_initializer import init_modules, stop_modules from app.startup.modules_initializer import init_modules, stop_modules
@@ -35,6 +36,8 @@ async def lifespan(app: FastAPI):
定义应用的生命周期事件 定义应用的生命周期事件
""" """
print("Starting up...") print("Starting up...")
# 存储当前循环
global_vars.set_loop(asyncio.get_event_loop())
# 初始化路由 # 初始化路由
init_routers(app) init_routers(app)
# 初始化模块 # 初始化模块
+2 -2
View File
@@ -1,4 +1,4 @@
from app.core.config import GlobalVar from app.core.config import global_vars
from app.core.plugin import PluginManager from app.core.plugin import PluginManager
from app.log import logger from app.log import logger
@@ -8,7 +8,7 @@ async def sync_plugins() -> bool:
初始化安装插件,并动态注册后台任务及API 初始化安装插件,并动态注册后台任务及API
""" """
try: try:
loop = GlobalVar.CURRENT_EVENT_LOOP loop = global_vars.loop
plugin_manager = PluginManager() plugin_manager = PluginManager()
sync_result = await execute_task(loop, plugin_manager.sync, "插件同步到本地") sync_result = await execute_task(loop, plugin_manager.sync, "插件同步到本地")