diff --git a/app/agent/memory/__init__.py b/app/agent/memory/__init__.py index d0d52c056..1f8dc4366 100644 --- a/app/agent/memory/__init__.py +++ b/app/agent/memory/__init__.py @@ -8,6 +8,7 @@ from langchain_core.messages import BaseMessage, messages_from_dict, messages_to from app.application.messaging.chat import ( AgentChatPersistenceService, + AgentChatRecord, AgentChatService, get_configured_agent_chat_persistence, get_configured_agent_chat_service, @@ -96,6 +97,40 @@ class MemoryManager: cache_key = self._get_memory_key(session_id, user_id) return self.memory_cache.get(cache_key) + @staticmethod + def _chat_lookup_params( + session_id: str, + user_id: str, + ) -> tuple[dict[str, str], ...]: + """返回用户专属会话优先、旧无用户会话兜底的查询顺序。""" + return ( + {"session_id": session_id, "user_id": user_id}, + {"session_id": session_id}, + ) + + def _restore_agent_messages( + self, + *, + session_id: str, + user_id: str, + chat: Optional[AgentChatRecord], + ) -> List[BaseMessage]: + """统一校验持久化快照、反序列化消息并回填内存缓存。""" + if not chat or not chat.agent_messages: + return [] + try: + messages = messages_from_dict(chat.agent_messages) + except Exception as e: + logger.debug(f"恢复持久化Agent消息失败: {e}") + return [] + memory = ConversationMemory( + session_id=session_id, + user_id=user_id, + messages=messages, + ) + self.save_memory(memory) + return memory.messages + def get_agent_messages( self, session_id: str, user_id: str ) -> List[BaseMessage]: @@ -109,31 +144,20 @@ class MemoryManager: return memory.messages try: - chat = self._chat_service().get_sync( - session_id=session_id, - user_id=user_id, - ) - if not chat: - chat = self._chat_service().get_sync(session_id=session_id) + service = self._chat_service() + chat = None + for lookup_params in self._chat_lookup_params(session_id, user_id): + chat = service.get_sync(**lookup_params) + if chat: + break except Exception as e: logger.debug(f"读取持久化Agent会话失败: {e}") return [] - if not chat or not chat.agent_messages: - return [] - - try: - messages = messages_from_dict(chat.agent_messages) - except Exception as e: - logger.debug(f"恢复持久化Agent消息失败: {e}") - return [] - - memory = ConversationMemory( + return self._restore_agent_messages( session_id=session_id, user_id=user_id, - messages=messages, + chat=chat, ) - self.save_memory(memory) - return memory.messages async def async_get_agent_messages( self, session_id: str, user_id: str @@ -145,31 +169,34 @@ class MemoryManager: try: service = self._chat_service() - chat = await service.get( - session_id=session_id, - user_id=user_id, - ) - if not chat: - chat = await service.get(session_id=session_id) + chat = None + for lookup_params in self._chat_lookup_params(session_id, user_id): + chat = await service.get(**lookup_params) + if chat: + break except Exception as e: logger.debug(f"读取持久化Agent会话失败: {e}") return [] - if not chat or not chat.agent_messages: - return [] - - try: - messages = messages_from_dict(chat.agent_messages) - except Exception as e: - logger.debug(f"恢复持久化Agent消息失败: {e}") - return [] - - memory = ConversationMemory( + return self._restore_agent_messages( session_id=session_id, user_id=user_id, - messages=messages, + chat=chat, ) + + def _update_agent_messages( + self, + *, + session_id: str, + user_id: str, + messages: List[BaseMessage], + ) -> None: + """统一更新同步和异步写路径共享的内存消息状态。""" + memory = self.get_memory(session_id, user_id) + if not memory: + memory = ConversationMemory(session_id=session_id, user_id=user_id) + memory.messages = messages + memory.updated_at = datetime.now() self.save_memory(memory) - return memory.messages def save_agent_messages( self, session_id: str, user_id: str, messages: List[BaseMessage] @@ -177,15 +204,11 @@ class MemoryManager: """ 保存Agent消息到内存缓存与持久化会话表。 """ - memory = self.get_memory(session_id, user_id) - if not memory: - memory = ConversationMemory(session_id=session_id, user_id=user_id) - - memory.messages = messages - memory.updated_at = datetime.now() - - # 更新内存缓存 - self.save_memory(memory) + self._update_agent_messages( + session_id=session_id, + user_id=user_id, + messages=messages, + ) try: self._chat_service().save_agent_messages( session_id=session_id, @@ -199,13 +222,11 @@ class MemoryManager: self, session_id: str, user_id: str, messages: List[BaseMessage] ) -> None: """异步保存 Agent 消息,持久化写入经有界数据库 worker 承接。""" - memory = self.get_memory(session_id, user_id) - if not memory: - memory = ConversationMemory(session_id=session_id, user_id=user_id) - - memory.messages = messages - memory.updated_at = datetime.now() - self.save_memory(memory) + self._update_agent_messages( + session_id=session_id, + user_id=user_id, + messages=messages, + ) try: persistence = self._chat_persistence() await persistence.async_save_agent_messages( diff --git a/app/application/server/report.py b/app/application/server/report.py index b60490aca..6cf88af90 100644 --- a/app/application/server/report.py +++ b/app/application/server/report.py @@ -53,7 +53,7 @@ class ServerReportService: reporter: Callable[[], bool], ) -> None: """首次成功上报后写入对应的完成标记。""" - if enabled and not self._config_reader(state_key) and reporter(): + if self._should_initialize_report(enabled=enabled, state_key=state_key) and reporter(): self._config_writer(state_key, "1") async def async_init_report( @@ -64,7 +64,7 @@ class ServerReportService: reporter: Callable[[], Awaitable[bool]], ) -> None: """异步完成首次上报,并通过异步配置端口持久化完成标记。""" - if not enabled or self._config_reader(state_key): + if not self._should_initialize_report(enabled=enabled, state_key=state_key): return if not await reporter(): return @@ -72,6 +72,10 @@ class ServerReportService: raise RuntimeError("中心服务上报未配置异步配置写入端口") await self._async_config_writer(state_key, "1") + def _should_initialize_report(self, *, enabled: bool, state_key: Any) -> bool: + """统一判断首次上报是否仍需执行。""" + return enabled and not self._config_reader(state_key) + def build_subscribe_payload(self, item: Optional[dict]) -> Optional[dict]: """构造中心服务订阅统计载荷并移除本地运行字段。""" if not isinstance(item, dict): @@ -108,23 +112,51 @@ class ServerReportService: if plugin_id ] - def report_subscribes(self, *, enabled: bool) -> bool: - """上报当前全部有效订阅的公开统计字段。""" - if not enabled: - return False - subscribes = self._subscribes_provider() + def _prepare_subscribe_report( + self, + subscribes: list[Any], + ) -> tuple[bool, Optional[list[dict[str, Any]]]]: + """把订阅读取结果投影为无需发送的终态或待发送载荷。""" if not subscribes: - return True + return True, None payloads = [ payload for subscribe in subscribes if (payload := self.build_subscribe_payload(subscribe.to_dict())) ] if not payloads: - return True - response = self._subscribe_report_sender(payloads) + return True, None + return False, payloads + + def _prepare_plugin_report( + self, + *, + enabled: bool, + items: Optional[list[tuple[str, Optional[str]]]], + ) -> Optional[list[dict[str, Any]]]: + """统一执行插件上报准入并构造脱敏载荷。""" + if not enabled: + return None + payload = self.build_plugin_payload(items) + return payload or None + + @staticmethod + def _report_succeeded(response: Any) -> bool: + """统一判定中心服务是否确认接收上报。""" return bool(response is not None and response.status_code == 200) + def report_subscribes(self, *, enabled: bool) -> bool: + """上报当前全部有效订阅的公开统计字段。""" + if not enabled: + return False + subscribes = self._subscribes_provider() + completed, payloads = self._prepare_subscribe_report(subscribes) + if completed: + return True + assert payloads is not None + response = self._subscribe_report_sender(payloads) + return self._report_succeeded(response) + def report_plugins( self, *, @@ -132,13 +164,11 @@ class ServerReportService: items: Optional[list[tuple[str, Optional[str]]]] = None, ) -> bool: """同步上报当前插件安装清单。""" - if not enabled: - return False - payload = self.build_plugin_payload(items) + payload = self._prepare_plugin_report(enabled=enabled, items=items) if not payload: return False response = self._plugin_report_sender(payload) - return bool(response is not None and response.status_code == 200) + return self._report_succeeded(response) async def async_report_plugins( self, @@ -147,13 +177,11 @@ class ServerReportService: items: Optional[list[tuple[str, Optional[str]]]] = None, ) -> bool: """异步上报当前插件安装清单。""" - if not enabled: - return False - payload = self.build_plugin_payload(items) + payload = self._prepare_plugin_report(enabled=enabled, items=items) if not payload: return False response = await self._async_plugin_report_sender(payload) - return bool(response is not None and response.status_code == 200) + return self._report_succeeded(response) async def async_report_subscribes(self, *, enabled: bool) -> bool: """异步上报当前全部有效订阅的公开统计字段。""" @@ -164,14 +192,9 @@ class ServerReportService: if self._async_subscribes_provider is None: raise RuntimeError("中心服务未配置异步订阅读取端口") subscribes = await self._async_subscribes_provider() - if not subscribes: - return True - payloads = [ - payload - for subscribe in subscribes - if (payload := self.build_subscribe_payload(subscribe.to_dict())) - ] - if not payloads: + completed, payloads = self._prepare_subscribe_report(subscribes) + if completed: return True + assert payloads is not None response = await self._async_subscribe_report_sender(payloads) - return bool(response is not None and response.status_code == 200) + return self._report_succeeded(response) diff --git a/app/chain/search/media.py b/app/chain/search/media.py index 634d81102..f6dff4ee5 100644 --- a/app/chain/search/media.py +++ b/app/chain/search/media.py @@ -14,6 +14,7 @@ from typing import ( List, Optional, Tuple, + Union, cast, ) @@ -71,6 +72,88 @@ def _build_missing_media_map( } +@dataclass(frozen=True, slots=True) +class _IdSearchCacheRequest: + """请求保存 ID 搜索参数并清理旧的 AI 推荐状态。""" + + params: Dict[str, Any] + + +@dataclass(frozen=True, slots=True) +class _IdSearchRecognizeRequest: + """请求通过媒体链识别一个来源原生 ID。""" + + params: Dict[str, Any] + + +@dataclass(frozen=True, slots=True) +class _IdSearchProcessRequest: + """请求执行已识别媒体的精确资源搜索。""" + + params: Dict[str, Any] + + +@dataclass(frozen=True, slots=True) +class _IdSearchSaveRequest: + """请求持久化成功的 ID 搜索结果。""" + + contexts: List[Context] + + +@dataclass(frozen=True, slots=True) +class _IdSearchResult: + """冻结 ID 搜索状态机结果与识别失败提示。""" + + contexts: List[Context] + warning: Optional[str] = None + + +_IdSearchRequest = Union[ + _IdSearchCacheRequest, + _IdSearchRecognizeRequest, + _IdSearchProcessRequest, + _IdSearchSaveRequest, +] +_IdSearchResolution = Generator[_IdSearchRequest, object, _IdSearchResult] + + +def _id_search_resolution( + recognition_params: Dict[str, Any], + cache_params: Dict[str, Any], + season: Optional[int], + sites: Optional[List[int]], + area: Optional[str], + cache_local: bool, + failure_keyword: str, +) -> _IdSearchResolution: + """统一 ID 搜索的缓存、识别、处理、失败短路与结果保存顺序。""" + if cache_local: + yield _IdSearchCacheRequest(params=cache_params) + mediainfo = cast( + Optional[MediaInfo], + (yield _IdSearchRecognizeRequest(params=recognition_params)), + ) + if not mediainfo: + return _IdSearchResult( + contexts=[], + warning=f"{failure_keyword} 媒体信息识别失败!", + ) + contexts = cast( + List[Context], + (yield _IdSearchProcessRequest( + params={ + "mediainfo": mediainfo, + "sites": sites, + "area": area, + "no_exists": _build_missing_media_map(mediainfo, season), + } + )), + ) + if cache_local: + yield _IdSearchSaveRequest(contexts=contexts) + return _IdSearchResult(contexts=contexts) + + def _normalize_media_search_input(mediainfo: MediaInfo) -> MediaInfo: """归一化非 TMDB 输入标题,并保留调用方对象由外层复制的所有权约束。""" if not mediainfo.tmdb_id: @@ -159,6 +242,168 @@ async def _run_keyword_search_async( return cast(_KeywordSearchResult, outcome.value) +@dataclass(frozen=True, slots=True) +class _MediaProcessPlan: + """冻结媒体资源处理入口的业务输入。""" + + mediainfo: MediaInfo + keyword: Optional[str] + no_exists: Optional[Dict[str, Dict[int, NotExistMediaInfo]]] + sites: Optional[List[int]] + rule_groups: Optional[List[str]] + area: Optional[str] + custom_words: Optional[List[str]] + filter_params: Optional[Dict[str, str]] + + +@dataclass(frozen=True, slots=True) +class _MediaMusicProcessRequest: + """请求执行音乐资源搜索外壳。""" + + params: Dict[str, Any] + + +@dataclass(frozen=True, slots=True) +class _MediaRecognizeRequest: + """请求补齐缺失名称的媒体信息。""" + + params: Dict[str, Any] + + +@dataclass(frozen=True, slots=True) +class _MediaSupplementRequest: + """请求聚合已启用媒体来源的附加信息。""" + + mediainfo: MediaInfo + + +@dataclass(frozen=True, slots=True) +class _MediaKeywordProcessRequest: + """请求按共享关键字计划执行 provider I/O。""" + + mediainfo: MediaInfo + keywords: List[str] + sites: Optional[List[int]] + area: Optional[str] + search_multiple_name: bool + + +@dataclass(frozen=True, slots=True) +class _MediaParseRequest: + """请求在同步或线程池 CPU 外壳中解析搜索结果。""" + + params: Dict[str, Any] + + +@dataclass(frozen=True, slots=True) +class _MediaLogRequest: + """请求记录共享状态机决定的运行日志。""" + + message: str + + +@dataclass(frozen=True, slots=True) +class _MediaProcessResult: + """冻结媒体处理结果与识别失败状态。""" + + contexts: List[Context] + recognition_failed: bool = False + + +_MediaProcessRequest = Union[ + _MediaMusicProcessRequest, + _MediaRecognizeRequest, + _MediaSupplementRequest, + _MediaKeywordProcessRequest, + _MediaParseRequest, + _MediaLogRequest, +] +_MediaProcessResolution = Generator[ + _MediaProcessRequest, object, _MediaProcessResult +] + + +def _media_process_resolution( + plan: _MediaProcessPlan, + search_multiple_name: Callable[[], bool], + copy_media: Callable[[MediaInfo], MediaInfo], + recognize_kwargs: Callable[[MediaInfo], Dict[str, Any]], + prepare_params: Callable[..., Tuple[Optional[Dict[int, List[int]]], List[str]]], +) -> _MediaProcessResolution: + """统一媒体处理的类型路由、识别、补充、搜索和解析状态。""" + if plan.mediainfo.type == MediaType.MUSIC: + contexts = cast( + List[Context], + (yield _MediaMusicProcessRequest( + params={ + "mediainfo": cast(MusicInfo, plan.mediainfo), + "keyword": plan.keyword, + "sites": plan.sites, + "rule_groups": plan.rule_groups, + "filter_params": plan.filter_params, + } + )), + ) + return _MediaProcessResult(contexts=contexts) + + mediainfo = _normalize_media_search_input(copy_media(plan.mediainfo)) + yield _MediaLogRequest( + message=f"开始搜索资源,关键词:{plan.keyword or mediainfo.title} ..." + ) + if not mediainfo.names: + recognized_media = cast( + Optional[MediaInfo], + (yield _MediaRecognizeRequest( + params={ + "mtype": mediainfo.type, + **recognize_kwargs(mediainfo), + } + )), + ) + if not recognized_media: + return _MediaProcessResult(contexts=[], recognition_failed=True) + mediainfo = recognized_media + + mediainfo = cast( + Optional[MediaInfo], + (yield _MediaSupplementRequest(mediainfo=mediainfo)), + ) or mediainfo + season_episodes, keywords = prepare_params( + mediainfo=mediainfo, + keyword=plan.keyword, + no_exists=plan.no_exists, + ) + outcome = cast( + _KeywordSearchResult, + (yield _MediaKeywordProcessRequest( + mediainfo=mediainfo, + keywords=keywords, + sites=plan.sites, + area=plan.area, + search_multiple_name=search_multiple_name(), + )), + ) + if outcome.stopped_early: + yield _MediaLogRequest( + message=f"共搜索到 {len(outcome.torrents)} 个资源,停止搜索" + ) + contexts = cast( + List[Context], + (yield _MediaParseRequest( + params=_build_result_params( + torrents=outcome.torrents, + mediainfo=mediainfo, + keyword=plan.keyword, + rule_groups=plan.rule_groups, + season_episodes=season_episodes, + custom_words=plan.custom_words, + filter_params=plan.filter_params, + ) + )), + ) + return _MediaProcessResult(contexts=contexts) + + def _build_result_params( torrents: List[TorrentInfo], mediainfo: MediaInfo, @@ -200,6 +445,50 @@ def _build_candidate_contexts( class SearchMediaOwner(_SearchOwnerBase): """精确媒体搜索与同步异步编排 owner。""" + def _run_id_search_sync( + self, resolution: _IdSearchResolution + ) -> _IdSearchResult: + """用同步 I/O 外壳驱动共享 ID 搜索状态机。""" + response: object = None + while True: + try: + request = resolution.send(response) + except StopIteration as completed: + return cast(_IdSearchResult, completed.value) + if isinstance(request, _IdSearchCacheRequest): + self.cancel_ai_recommend() + self.save_last_search_params(**request.params) + response = None + elif isinstance(request, _IdSearchRecognizeRequest): + response = MediaChain().recognize_media(**request.params) + elif isinstance(request, _IdSearchProcessRequest): + response = self.process(**request.params) + else: + self._save_results(request.contexts) + response = None + + async def _run_id_search_async( + self, resolution: _IdSearchResolution + ) -> _IdSearchResult: + """用异步 I/O 外壳驱动共享 ID 搜索状态机。""" + response: object = None + while True: + try: + request = resolution.send(response) + except StopIteration as completed: + return cast(_IdSearchResult, completed.value) + if isinstance(request, _IdSearchCacheRequest): + self.cancel_ai_recommend() + await self.async_save_last_search_params(**request.params) + response = None + elif isinstance(request, _IdSearchRecognizeRequest): + response = await MediaChain().async_recognize_media(**request.params) + elif isinstance(request, _IdSearchProcessRequest): + response = await self.async_process(**request.params) + else: + await self._async_save_results(request.contexts) + response = None + def search_by_id( self, media_source: MediaSource, @@ -231,20 +520,23 @@ class SearchMediaOwner(_SearchOwnerBase): sites=sites, music_type=music_type, ) - if cache_local: - self.cancel_ai_recommend() - self.save_last_search_params(**cache_params) - # 音乐统一在 MediaChain.recognize_media 内按固定来源路由 - mediainfo = MediaChain().recognize_media(**recognition_params) - if not mediainfo: - logger.error(f"{self._build_search_keyword(media_source, media_id)} 媒体信息识别失败!") - return [] - no_exists = _build_missing_media_map(mediainfo, season) - results = self.process(mediainfo=mediainfo, sites=sites, area=area, no_exists=no_exists) - # 保存到本地文件 - if cache_local: - self._save_results(results) - return results + result = SearchMediaOwner._run_id_search_sync( + self, + _id_search_resolution( + recognition_params=recognition_params, + cache_params=cache_params, + season=season, + sites=sites, + area=area, + cache_local=cache_local, + failure_keyword=self._build_search_keyword( + media_source, media_id + ), + ) + ) + if result.warning: + logger.error(result.warning) + return result.contexts async def async_search_by_id( self, @@ -277,20 +569,23 @@ class SearchMediaOwner(_SearchOwnerBase): sites=sites, music_type=music_type, ) - if cache_local: - self.cancel_ai_recommend() - await self.async_save_last_search_params(**cache_params) - # 音乐统一在 MediaChain.async_recognize_media 内按固定来源路由 - mediainfo = await MediaChain().async_recognize_media(**recognition_params) - if not mediainfo: - logger.error(f"{self._build_search_keyword(media_source, media_id)} 媒体信息识别失败!") - return [] - no_exists = _build_missing_media_map(mediainfo, season) - results = await self.async_process(mediainfo=mediainfo, sites=sites, area=area, no_exists=no_exists) - # 保存到本地文件 - if cache_local: - await self._async_save_results(results) - return results + result = await SearchMediaOwner._run_id_search_async( + self, + _id_search_resolution( + recognition_params=recognition_params, + cache_params=cache_params, + season=season, + sites=sites, + area=area, + cache_local=cache_local, + failure_keyword=self._build_search_keyword( + media_source, media_id + ), + ) + ) + if result.warning: + logger.error(result.warning) + return result.contexts async def async_search_by_id_stream( self, @@ -337,6 +632,112 @@ class SearchMediaOwner(_SearchOwnerBase): if cache_local: await self._async_save_results(contexts) + def _run_media_process_sync( + self, resolution: _MediaProcessResolution + ) -> _MediaProcessResult: + """用同步 provider 与 CPU 外壳驱动共享媒体处理状态机。""" + response: object = None + while True: + try: + request = resolution.send(response) + except StopIteration as completed: + return cast(_MediaProcessResult, completed.value) + if isinstance(request, _MediaMusicProcessRequest): + response = self._process_music(**request.params) + elif isinstance(request, _MediaRecognizeRequest): + response = MediaChain().recognize_media(**request.params) + elif isinstance(request, _MediaSupplementRequest): + response = MediaChain().supplement_media_info(request.mediainfo) + elif isinstance(request, _MediaKeywordProcessRequest): + + def execute_search( + keyword_request: _KeywordSearchRequest, + ) -> List[TorrentInfo]: + """执行共享关键字请求的同步站点 I/O。""" + if keyword_request.search_count > 0: + logger.info( + f"已搜索 {keyword_request.search_count} 次," + "强制休眠 1-10 秒 ..." + ) + time.sleep(random.randint(1, 10)) + return ( + self._SearchChain__search_all_sites( + mediainfo=request.mediainfo, + keyword=keyword_request.keyword, + sites=request.sites, + area=request.area, + ) + or [] + ) + + response = _run_keyword_search_sync( + _keyword_search_resolution( + request.keywords, request.search_multiple_name + ), + execute_search, + ) + elif isinstance(request, _MediaParseRequest): + response = self._parse_result(**request.params) + else: + logger.info(request.message) + response = None + + async def _run_media_process_async( + self, resolution: _MediaProcessResolution + ) -> _MediaProcessResult: + """用异步 provider 与线程池 CPU 外壳驱动共享媒体处理状态机。""" + response: object = None + while True: + try: + request = resolution.send(response) + except StopIteration as completed: + return cast(_MediaProcessResult, completed.value) + if isinstance(request, _MediaMusicProcessRequest): + response = await self._async_process_music(**request.params) + elif isinstance(request, _MediaRecognizeRequest): + response = await MediaChain().async_recognize_media( + **request.params + ) + elif isinstance(request, _MediaSupplementRequest): + response = await MediaChain().async_supplement_media_info( + request.mediainfo + ) + elif isinstance(request, _MediaKeywordProcessRequest): + + async def execute_search( + keyword_request: _KeywordSearchRequest, + ) -> List[TorrentInfo]: + """执行共享关键字请求的异步站点 I/O。""" + if keyword_request.search_count > 0: + logger.info( + f"已搜索 {keyword_request.search_count} 次," + "强制休眠 1-10 秒 ..." + ) + await asyncio.sleep(random.randint(1, 10)) + return ( + await self._SearchChain__async_search_all_sites( + mediainfo=request.mediainfo, + keyword=keyword_request.keyword, + sites=request.sites, + area=request.area, + ) + or [] + ) + + response = await _run_keyword_search_async( + _keyword_search_resolution( + request.keywords, request.search_multiple_name + ), + execute_search, + ) + elif isinstance(request, _MediaParseRequest): + response = await run_in_threadpool( + self._parse_result, **request.params + ) + else: + logger.info(request.message) + response = None + def process( self, mediainfo: MediaInfo, @@ -359,83 +760,30 @@ class SearchMediaOwner(_SearchOwnerBase): :param custom_words: 自定义识别词列表 :param filter_params: 过滤参数 """ - - if mediainfo.type == MediaType.MUSIC: - return cast( - List[Context], - self._process_music( - mediainfo=cast(MusicInfo, mediainfo), + result = SearchMediaOwner._run_media_process_sync( + self, + _media_process_resolution( + plan=_MediaProcessPlan( + mediainfo=mediainfo, keyword=keyword, + no_exists=no_exists, sites=sites, rule_groups=rule_groups, - filter_params=filter_params, - ), - ) - - mediainfo = _normalize_media_search_input(self._copy_media_input(mediainfo)) - logger.info(f"开始搜索资源,关键词:{keyword or mediainfo.title} ...") - - # 补充媒体信息 - if not mediainfo.names: - recognized_media = MediaChain().recognize_media( - mtype=mediainfo.type, - **self._media_recognize_kwargs(mediainfo), - ) - if not recognized_media: - logger.error("媒体信息识别失败!") - return [] - mediainfo = recognized_media - - # 搜索前按用户启用的数据源聚合别名;分类、风格与外部 ID 仅由 TMDB 补充。 - mediainfo = cast( - MediaInfo, - MediaChain().supplement_media_info(mediainfo) or mediainfo, - ) - - # 准备搜索参数 - season_episodes, keywords = self._prepare_params(mediainfo=mediainfo, keyword=keyword, no_exists=no_exists) - - def execute_search(request: _KeywordSearchRequest) -> List[TorrentInfo]: - """执行共享状态机请求的同步站点搜索。""" - if request.search_count > 0: - logger.info( - f"已搜索 {request.search_count} 次,强制休眠 1-10 秒 ..." - ) - time.sleep(random.randint(1, 10)) - return ( - self._SearchChain__search_all_sites( - mediainfo=mediainfo, - keyword=request.keyword, - sites=sites, area=area, - ) - or [] - ) - - outcome = _run_keyword_search_sync( - _keyword_search_resolution( - keywords, self.runtime_config.search_multiple_name - ), - execute_search, - ) - if outcome.stopped_early: - logger.info(f"共搜索到 {len(outcome.torrents)} 个资源,停止搜索") - - # 处理结果 - return cast( - List[Context], - self._parse_result( - **_build_result_params( - torrents=outcome.torrents, - mediainfo=mediainfo, - keyword=keyword, - rule_groups=rule_groups, - season_episodes=season_episodes, custom_words=custom_words, filter_params=filter_params, - ) - ), + ), + search_multiple_name=( + lambda: self.runtime_config.search_multiple_name + ), + copy_media=self._copy_media_input, + recognize_kwargs=self._media_recognize_kwargs, + prepare_params=self._prepare_params, + ) ) + if result.recognition_failed: + logger.error("媒体信息识别失败!") + return result.contexts async def async_process( self, @@ -459,86 +807,30 @@ class SearchMediaOwner(_SearchOwnerBase): :param custom_words: 自定义识别词列表 :param filter_params: 过滤参数 """ - - if mediainfo.type == MediaType.MUSIC: - return cast( - List[Context], - await self._async_process_music( - mediainfo=cast(MusicInfo, mediainfo), + result = await SearchMediaOwner._run_media_process_async( + self, + _media_process_resolution( + plan=_MediaProcessPlan( + mediainfo=mediainfo, keyword=keyword, + no_exists=no_exists, sites=sites, rule_groups=rule_groups, - filter_params=filter_params, - ), - ) - - mediainfo = _normalize_media_search_input(self._copy_media_input(mediainfo)) - logger.info(f"开始搜索资源,关键词:{keyword or mediainfo.title} ...") - - # 补充媒体信息 - if not mediainfo.names: - recognized_media = await MediaChain().async_recognize_media( - mtype=mediainfo.type, - **self._media_recognize_kwargs(mediainfo), - ) - if not recognized_media: - logger.error("媒体信息识别失败!") - return [] - mediainfo = recognized_media - - # 异步搜索与同步入口共享同一份多来源附加信息语义。 - mediainfo = cast( - MediaInfo, - await MediaChain().async_supplement_media_info(mediainfo) or mediainfo, - ) - - # 准备搜索参数 - season_episodes, keywords = self._prepare_params(mediainfo=mediainfo, keyword=keyword, no_exists=no_exists) - - async def execute_search( - request: _KeywordSearchRequest, - ) -> List[TorrentInfo]: - """执行共享状态机请求的异步站点搜索。""" - if request.search_count > 0: - logger.info( - f"已搜索 {request.search_count} 次,强制休眠 1-10 秒 ..." - ) - await asyncio.sleep(random.randint(1, 10)) - return ( - await self._SearchChain__async_search_all_sites( - mediainfo=mediainfo, - keyword=request.keyword, - sites=sites, area=area, - ) - or [] - ) - - outcome = await _run_keyword_search_async( - _keyword_search_resolution( - keywords, self.runtime_config.search_multiple_name - ), - execute_search, - ) - if outcome.stopped_early: - logger.info(f"共搜索到 {len(outcome.torrents)} 个资源,停止搜索") - - # 处理结果 - return cast( - List[Context], - await run_in_threadpool( - self._parse_result, - **_build_result_params( - torrents=outcome.torrents, - mediainfo=mediainfo, - keyword=keyword, - rule_groups=rule_groups, - season_episodes=season_episodes, custom_words=custom_words, filter_params=filter_params, ), - ), + search_multiple_name=( + lambda: self.runtime_config.search_multiple_name + ), + copy_media=self._copy_media_input, + recognize_kwargs=self._media_recognize_kwargs, + prepare_params=self._prepare_params, + ) ) + if result.recognition_failed: + logger.error("媒体信息识别失败!") + return result.contexts async def async_process_stream( self, diff --git a/app/chain/search/title.py b/app/chain/search/title.py index 2625ae26c..7dc90a3a9 100644 --- a/app/chain/search/title.py +++ b/app/chain/search/title.py @@ -1,7 +1,18 @@ """标题搜索入口与标题候选过滤 owner。""" from dataclasses import dataclass -from typing import Any, AsyncIterator, Callable, Dict, List, Optional, Tuple, cast +from typing import ( + Any, + AsyncIterator, + Callable, + Dict, + Generator, + List, + Optional, + Tuple, + Union, + cast, +) from app.application.configuration import ( get_configured_system_config, @@ -66,6 +77,79 @@ class _TitleSearchResult: warning: Optional[str] = None +@dataclass(frozen=True, slots=True) +class _TitleSearchCacheRequest: + """请求保存标题搜索参数并清理旧的 AI 推荐状态。""" + + params: Dict[str, Any] + + +@dataclass(frozen=True, slots=True) +class _TitleSearchProviderRequest: + """请求执行标题搜索 provider I/O。""" + + params: Dict[str, Any] + + +@dataclass(frozen=True, slots=True) +class _TitleSearchResolveRequest: + """请求过滤并投影标题搜索候选。""" + + title: str + torrents: List[TorrentInfo] + rule_groups: Optional[List[str]] + mtype: Optional[MediaType] + + +@dataclass(frozen=True, slots=True) +class _TitleSearchSaveRequest: + """请求持久化成功的标题搜索结果。""" + + contexts: List[Context] + + +_TitleSearchRequest = Union[ + _TitleSearchCacheRequest, + _TitleSearchProviderRequest, + _TitleSearchResolveRequest, + _TitleSearchSaveRequest, +] +_TitleSearchResolution = Generator[ + _TitleSearchRequest, object, _TitleSearchResult +] + + +def _title_search_resolution( + title: str, + search_params: Dict[str, Any], + cache_params: Dict[str, Any], + cache_local: bool, + mtype: Optional[MediaType], + rule_groups: Optional[List[str]], +) -> _TitleSearchResolution: + """统一标题搜索的缓存、provider、过滤、短路与结果保存顺序。""" + if cache_local: + yield _TitleSearchCacheRequest(params=cache_params) + torrents = cast( + List[TorrentInfo], + (yield _TitleSearchProviderRequest(params=search_params)), + ) + result = cast( + _TitleSearchResult, + (yield _TitleSearchResolveRequest( + title=title, + torrents=torrents, + rule_groups=rule_groups, + mtype=mtype, + )), + ) + if result.warning: + return result + if cache_local: + yield _TitleSearchSaveRequest(contexts=result.contexts) + return result + + def _resolve_title_search_result( title: str, torrents: List[TorrentInfo], @@ -101,6 +185,77 @@ def _resolve_title_search_result( class SearchTitleOwner(_SearchOwnerBase): """标题搜索入口与标题候选过滤 owner。""" + def _resolve_title_request( + self, request: _TitleSearchResolveRequest + ) -> _TitleSearchResult: + """通过可替换过滤与元数据投影点解析标题候选。""" + return _resolve_title_search_result( + title=request.title, + torrents=request.torrents, + rule_groups=request.rule_groups, + mtype=request.mtype, + filter_torrents=self._filter_title_search_torrents, + build_meta=self._build_title_search_meta, + ) + + def _run_title_search_sync( + self, resolution: _TitleSearchResolution + ) -> _TitleSearchResult: + """用同步 I/O 与过滤外壳驱动共享标题搜索状态机。""" + response: object = None + while True: + try: + request = resolution.send(response) + except StopIteration as completed: + return cast(_TitleSearchResult, completed.value) + if isinstance(request, _TitleSearchCacheRequest): + self.cancel_ai_recommend() + self.save_last_search_params(**request.params) + response = None + elif isinstance(request, _TitleSearchProviderRequest): + response = ( + self._SearchChain__search_all_sites(**request.params) or [] + ) + elif isinstance(request, _TitleSearchResolveRequest): + response = SearchTitleOwner._resolve_title_request(self, request) + else: + self._save_results(request.contexts) + response = None + + async def _run_title_search_async( + self, resolution: _TitleSearchResolution + ) -> _TitleSearchResult: + """用异步 I/O 与线程池过滤外壳驱动共享标题搜索状态机。""" + response: object = None + while True: + try: + request = resolution.send(response) + except StopIteration as completed: + return cast(_TitleSearchResult, completed.value) + if isinstance(request, _TitleSearchCacheRequest): + self.cancel_ai_recommend() + await self.async_save_last_search_params(**request.params) + response = None + elif isinstance(request, _TitleSearchProviderRequest): + response = ( + await self._SearchChain__async_search_all_sites( + **request.params + ) + or [] + ) + elif isinstance(request, _TitleSearchResolveRequest): + if request.torrents: + response = await run_in_threadpool( + SearchTitleOwner._resolve_title_request, self, request + ) + else: + response = SearchTitleOwner._resolve_title_request( + self, request + ) + else: + await self._async_save_results(request.contexts) + response = None + def search_by_title( self, title: str, @@ -125,28 +280,23 @@ class SearchTitleOwner(_SearchOwnerBase): sites=sites, mtype=mtype, ) - if cache_local: - self.cancel_ai_recommend() - self.save_last_search_params(**cache_params) if title: logger.info(f"开始搜索资源,关键词:{title} ...") else: logger.info(f"开始浏览资源,站点:{sites} ...") - # 搜索 - result = _resolve_title_search_result( - title=title, - torrents=self._SearchChain__search_all_sites(**search_params) or [], - rule_groups=rule_groups, - mtype=mtype, - filter_torrents=self._filter_title_search_torrents, - build_meta=self._build_title_search_meta, + result = SearchTitleOwner._run_title_search_sync( + self, + _title_search_resolution( + title=title, + search_params=search_params, + cache_params=cache_params, + cache_local=bool(cache_local), + mtype=mtype, + rule_groups=rule_groups, + ) ) if result.warning: logger.warning(result.warning) - return [] - # 保存到本地文件 - if cache_local: - self._save_results(result.contexts) return result.contexts async def async_search_by_title( @@ -173,45 +323,23 @@ class SearchTitleOwner(_SearchOwnerBase): sites=sites, mtype=mtype, ) - if cache_local: - self.cancel_ai_recommend() - await self.async_save_last_search_params(**cache_params) if title: logger.info(f"开始搜索资源,关键词:{title} ...") else: logger.info(f"开始浏览资源,站点:{sites} ...") - # 搜索 - torrents = ( - await self._SearchChain__async_search_all_sites(**search_params) or [] - ) - if not torrents: - result = _resolve_title_search_result( + result = await SearchTitleOwner._run_title_search_async( + self, + _title_search_resolution( title=title, - torrents=torrents, - rule_groups=rule_groups, + search_params=search_params, + cache_params=cache_params, + cache_local=bool(cache_local), mtype=mtype, - filter_torrents=self._filter_title_search_torrents, - build_meta=self._build_title_search_meta, - ) - else: - result = cast( - _TitleSearchResult, - await run_in_threadpool( - _resolve_title_search_result, - title=title, - torrents=torrents, - rule_groups=rule_groups, - mtype=mtype, - filter_torrents=self._filter_title_search_torrents, - build_meta=self._build_title_search_meta, - ), + rule_groups=rule_groups, ) + ) if result.warning: logger.warning(result.warning) - return [] - # 保存到本地文件 - if cache_local: - await self._async_save_results(result.contexts) return result.contexts async def async_search_by_title_stream( diff --git a/app/modules/acoustid/__init__.py b/app/modules/acoustid/__init__.py index 607b44985..ab30f6071 100644 --- a/app/modules/acoustid/__init__.py +++ b/app/modules/acoustid/__init__.py @@ -5,17 +5,25 @@ import subprocess import threading import time from collections import OrderedDict +from dataclasses import dataclass from pathlib import Path from typing import Any, Optional, Tuple, Union from uuid import UUID -from app.runtime.execution import run_in_threadpool -from app.runtime.settings import get_runtime_setting - -from app.runtime.log import logger -from app.modules import _ModuleBase -from app.schemas.types import ModuleType, OtherModulesType from app.adapters.network.http import AsyncRequestUtils, RequestUtils +from app.modules import _ModuleBase +from app.runtime.execution import run_in_threadpool +from app.runtime.log import logger +from app.runtime.settings import get_runtime_setting +from app.schemas.types import ModuleType, OtherModulesType + + +@dataclass(frozen=True, slots=True) +class _AcoustIdLookupPlan: + """冻结 AcoustID 请求载荷,避免同步与异步入口产生参数漂移。""" + + url: str + data: dict[str, Union[str, int]] class AcoustIdModule(_ModuleBase): @@ -291,14 +299,51 @@ class AcoustIdModule(_ModuleBase): if delay := cls._reserve_request_delay(): await asyncio.sleep(delay) + @classmethod + def _lookup_plan( + cls, + api_key: str, + duration: int, + fingerprint: str, + ) -> Optional[_AcoustIdLookupPlan]: + """校验 API Key 并构造同步与异步共用的指纹查询计划。""" + normalized_key = str(api_key or "").strip() + if not normalized_key: + return None + return _AcoustIdLookupPlan( + url=cls._base_url, + data={ + "client": normalized_key, + "duration": duration, + "fingerprint": fingerprint, + "meta": "recordingids", + "format": "json", + }, + ) + + @classmethod + def _project_lookup_response( + cls, + status_code: Optional[int], + payload: Any, + ) -> Optional[str]: + """按统一 HTTP 状态和 AcoustID 规则投影 Recording ID。""" + if status_code != 200: + return None + return cls._select_recording_id(payload) + def _lookup_recording_id( self, duration: int, fingerprint: str, ) -> Optional[str]: """查询 AcoustID 指纹库并提取 MusicBrainz Recording ID。""" - api_key = str(get_runtime_setting('ACOUSTID_API_KEY') or "").strip() - if not api_key: + plan = self._lookup_plan( + get_runtime_setting('ACOUSTID_API_KEY'), + duration, + fingerprint, + ) + if not plan: return None self._wait_for_rate_limit() response = RequestUtils( @@ -306,14 +351,8 @@ class AcoustIdModule(_ModuleBase): proxies=get_runtime_setting('PROXY'), timeout=30, ).post_res( - url=self._base_url, - data={ - "client": api_key, - "duration": duration, - "fingerprint": fingerprint, - "meta": "recordingids", - "format": "json", - }, + url=plan.url, + data=plan.data, ) if response is None: logger.warning("AcoustID 指纹查询失败:无响应") @@ -322,8 +361,9 @@ class AcoustIdModule(_ModuleBase): if response.status_code != 200: logger.warning(f"AcoustID 指纹查询失败:HTTP {response.status_code}") return None - payload = response.json() - return self._select_recording_id(payload) + return self._project_lookup_response( + response.status_code, response.json() + ) except (TypeError, ValueError) as err: logger.warning(f"AcoustID 响应解析失败:{err}") return None @@ -336,8 +376,12 @@ class AcoustIdModule(_ModuleBase): fingerprint: str, ) -> Optional[str]: """异步查询 AcoustID 指纹库并提取 MusicBrainz Recording ID。""" - api_key = str(get_runtime_setting('ACOUSTID_API_KEY') or "").strip() - if not api_key: + plan = self._lookup_plan( + get_runtime_setting('ACOUSTID_API_KEY'), + duration, + fingerprint, + ) + if not plan: return None await self._async_wait_for_rate_limit() response = await AsyncRequestUtils( @@ -345,14 +389,8 @@ class AcoustIdModule(_ModuleBase): proxies=get_runtime_setting('PROXY'), timeout=30, ).post_res( - url=self._base_url, - data={ - "client": api_key, - "duration": duration, - "fingerprint": fingerprint, - "meta": "recordingids", - "format": "json", - }, + url=plan.url, + data=plan.data, ) if response is None: logger.warning("AcoustID 指纹查询失败:无响应") @@ -361,7 +399,9 @@ class AcoustIdModule(_ModuleBase): if response.status_code != 200: logger.warning(f"AcoustID 指纹查询失败:HTTP {response.status_code}") return None - return self._select_recording_id(response.json()) + return self._project_lookup_response( + response.status_code, response.json() + ) except (TypeError, ValueError) as err: logger.warning(f"AcoustID 响应解析失败:{err}") return None diff --git a/app/modules/anilist/anilist.py b/app/modules/anilist/anilist.py index 947ecfe96..aa9517fa2 100644 --- a/app/modules/anilist/anilist.py +++ b/app/modules/anilist/anilist.py @@ -1,11 +1,19 @@ +from dataclasses import dataclass from datetime import date -from typing import Optional +from typing import Any, Optional, cast +from app.adapters.network.http import AsyncRequestUtils, RequestUtils from app.runtime.cache import cached +from app.runtime.log import logger from app.runtime.settings import get_runtime_setting -from app.runtime.log import logger -from app.adapters.network.http import AsyncRequestUtils, RequestUtils + +@dataclass(frozen=True, slots=True) +class _AniListRequestPlan: + """冻结 GraphQL 查询与变量,供同步和异步传输外壳共用。""" + + query: str + variables: dict[str, Any] class AniListApi: @@ -98,6 +106,55 @@ class AniListApi: }} }} """ + _search_query = f""" + query ($search: String!, $count: Int!) {{ + Page(page: 1, perPage: $count) {{ + media(search: $search, type: ANIME, sort: SEARCH_MATCH) {{ {_media_fields} }} + }} + }} + """ + _credits_query = """ + query ($id: Int!, $page: Int!, $count: Int!) { + Media(id: $id, type: ANIME) { + characters(page: $page, perPage: $count, sort: [ROLE, RELEVANCE]) { + edges { + role + node { id name { full native } } + voiceActors(language: JAPANESE, sort: [RELEVANCE]) { + id name { full native alternative } image { large medium } siteUrl + } + } + } + } + } + """ + _recommendations_query = """ + query ($id: Int!, $page: Int!, $count: Int!) { + Media(id: $id, type: ANIME) { + recommendations(page: $page, perPage: $count, sort: [RATING_DESC, ID]) { + nodes { mediaRecommendation { id } } + } + } + } + """ + _person_detail_query = """ + query ($id: Int!) { + Staff(id: $id) { + id name { full native alternative } image { large medium } + description(asHtml: false) dateOfBirth { year month day } + dateOfDeath { year month day } gender homeTown primaryOccupations siteUrl + } + } + """ + _person_credits_query = """ + query ($id: Int!, $page: Int!, $count: Int!) { + Staff(id: $id) { + characterMedia(page: $page, perPage: $count, sort: [POPULARITY_DESC]) { + nodes { id } + } + } + } + """ def __init__(self) -> None: """初始化同步与异步请求客户端""" @@ -137,6 +194,22 @@ class AniListApi: return None return result.get("data") + @staticmethod + def _request_plan( + query: str, variables: dict[str, Any] + ) -> _AniListRequestPlan: + """构造同步与异步 GraphQL 传输共同使用的请求计划。""" + return _AniListRequestPlan(query=query, variables=dict(variables)) + + @classmethod + def _project_translated( + cls, + result: Optional[dict[str, Any]], + translations: dict[int, dict[str, Any]], + ) -> Optional[dict[str, Any]]: + """把成功响应与中文数据集按同一规则合并。""" + return cls._inject_chinese(result, translations) if result else result + def _invoke(self, query: str, variables: dict) -> Optional[dict]: """ 执行同步 GraphQL 请求。 @@ -145,16 +218,17 @@ class AniListApi: :param variables: 查询变量 :return: GraphQL data 字段 """ - payload = {"query": query, "variables": variables} + plan = self._request_plan(query, variables) + payload = {"query": plan.query, "variables": plan.variables} if self._proxy_available: response = self._request.post_res(self._base_url, json=payload) result = self._extract_response(response) if result is not None: - return self._inject_chinese(result, self._translation_map()) + return self._project_translated(result, self._translation_map()) self._disable_proxy(response) response = self._request.post_res(self._official_url, json=payload) result = self._extract_response(response) - return self._inject_chinese(result, self._translation_map()) if result else result + return self._project_translated(result, self._translation_map()) async def _async_invoke(self, query: str, variables: dict) -> Optional[dict]: """ @@ -164,20 +238,21 @@ class AniListApi: :param variables: 查询变量 :return: GraphQL data 字段 """ - payload = {"query": query, "variables": variables} + plan = self._request_plan(query, variables) + payload = {"query": plan.query, "variables": plan.variables} if self._proxy_available: response = await self._async_request.post_res(self._base_url, json=payload) result = self._extract_response(response) if result is not None: translations = await self._async_translation_map() - return self._inject_chinese(result, translations) + return self._project_translated(result, translations) self._disable_proxy(response) response = await self._async_request.post_res(self._official_url, json=payload) result = self._extract_response(response) if not result: return result translations = await self._async_translation_map() - return self._inject_chinese(result, translations) + return self._project_translated(result, translations) def _disable_proxy(self, response) -> None: """ @@ -317,6 +392,80 @@ class AniListApi: media_map = {media.get("id"): media for media in medias if media.get("id")} return [media_map[media_id] for media_id in media_ids if media_id in media_map] + @staticmethod + def _media_id_plan( + media_ids: list[int] + ) -> tuple[list[int], dict[str, Any]]: + """去重批量媒体 ID,并构造根级 Page 查询变量。""" + unique_ids = list(dict.fromkeys(media_id for media_id in media_ids if media_id)) + return unique_ids, {"ids": unique_ids, "count": len(unique_ids)} + + @staticmethod + def _field( + result: Optional[dict[str, Any]], name: str + ) -> Optional[dict[str, Any]]: + """从 GraphQL data 中提取一个根字段。""" + return result.get(name) if result else None + + @staticmethod + def _credits_edges( + result: Optional[dict[str, Any]] + ) -> list[dict[str, Any]]: + """从动画人物响应中提取配音关系边。""" + if not result: + return [] + return result.get("Media", {}).get("characters", {}).get("edges") or [] + + @staticmethod + def _recommendation_ids( + result: Optional[dict[str, Any]] + ) -> list[int]: + """按推荐关系顺序提取有效 AniList 媒体 ID。""" + if not result: + return [] + nodes = result.get("Media", {}).get("recommendations", {}).get("nodes") or [] + return [ + media_id + for node in nodes + if isinstance(node, dict) + if isinstance( + (media_id := (node.get("mediaRecommendation") or {}).get("id")), + int, + ) + ] + + @staticmethod + def _person_media_ids( + result: Optional[dict[str, Any]] + ) -> list[int]: + """按人物作品关系顺序提取有效 AniList 媒体 ID。""" + if not result: + return [] + nodes = result.get("Staff", {}).get("characterMedia", {}).get("nodes") or [] + return [ + media_id + for node in nodes + if isinstance(node, dict) + if isinstance((media_id := node.get("id")), int) + ] + + @staticmethod + def _trend_filters(page: int, count: int) -> dict[str, Any]: + """构造当前趋势榜共同使用的探索参数。""" + return {"page": page, "count": count, "sort": "TRENDING_DESC"} + + @classmethod + def _season_filters(cls, page: int, count: int) -> dict[str, Any]: + """构造当前季度热门榜共同使用的探索参数。""" + season, season_year = cls._current_season() + return { + "page": page, + "count": count, + "season": season, + "season_year": season_year, + "sort": "POPULARITY_DESC", + } + def _medias_by_ids(self, media_ids: list[int]) -> list[dict]: """ 通过根级 Page.media 批量查询媒体,使中文代理能够注入标题。 @@ -324,12 +473,12 @@ class AniListApi: :param media_ids: AniList 媒体 ID 列表 :return: 按输入顺序排列的媒体列表 """ - unique_ids = list(dict.fromkeys(media_id for media_id in media_ids if media_id)) + unique_ids, variables = self._media_id_plan(media_ids) if not unique_ids: return [] result = self._invoke( self._media_by_ids_query, - {"ids": unique_ids, "count": len(unique_ids)}, + variables, ) return self._ordered_medias(media_ids, self._page_medias(result)) @@ -340,12 +489,12 @@ class AniListApi: :param media_ids: AniList 媒体 ID 列表 :return: 按输入顺序排列的媒体列表 """ - unique_ids = list(dict.fromkeys(media_id for media_id in media_ids if media_id)) + unique_ids, variables = self._media_id_plan(media_ids) if not unique_ids: return [] result = await self._async_invoke( self._media_by_ids_query, - {"ids": unique_ids, "count": len(unique_ids)}, + variables, ) return self._ordered_medias(media_ids, self._page_medias(result)) @@ -375,8 +524,7 @@ class AniListApi: :return: AniList 媒体详情 """ query = f"query ($id: Int!) {{ Media(id: $id, type: ANIME) {{ {self._media_fields} }} }}" - result = self._invoke(query, {"id": anilist_id}) - return result.get("Media") if result else None + return self._field(self._invoke(query, {"id": anilist_id}), "Media") @cached( maxsize=get_runtime_setting('CONF').anilist, @@ -393,7 +541,7 @@ class AniListApi: """ query = f"query ($id: Int!) {{ Media(id: $id, type: ANIME) {{ {self._media_fields} }} }}" result = await self._async_invoke(query, {"id": anilist_id}) - return result.get("Media") if result else None + return self._field(result, "Media") @cached( maxsize=get_runtime_setting('CONF').anilist, @@ -409,14 +557,9 @@ class AniListApi: :param count: 返回条数 :return: AniList 媒体列表 """ - query = f""" - query ($search: String!, $count: Int!) {{ - Page(page: 1, perPage: $count) {{ - media(search: $search, type: ANIME, sort: SEARCH_MATCH) {{ {self._media_fields} }} - }} - }} - """ - result = self._invoke(query, {"search": name, "count": count}) + result = self._invoke( + self._search_query, {"search": name, "count": count} + ) return self._page_medias(result) @cached( @@ -433,14 +576,9 @@ class AniListApi: :param count: 返回条数 :return: AniList 媒体列表 """ - query = f""" - query ($search: String!, $count: Int!) {{ - Page(page: 1, perPage: $count) {{ - media(search: $search, type: ANIME, sort: SEARCH_MATCH) {{ {self._media_fields} }} - }} - }} - """ - result = await self._async_invoke(query, {"search": name, "count": count}) + result = await self._async_invoke( + self._search_query, {"search": name, "count": count} + ) return self._page_medias(result) @cached( @@ -528,7 +666,10 @@ class AniListApi: :param count: 每页条数 :return: AniList 媒体列表 """ - return self.discover(page=page, count=count, sort="TRENDING_DESC") + return cast( + list[dict[str, Any]], + self.discover(**self._trend_filters(page, count)), + ) async def async_trending(self, page: int = 1, count: int = 20) -> list[dict]: """ @@ -538,7 +679,10 @@ class AniListApi: :param count: 每页条数 :return: AniList 媒体列表 """ - return await self.async_discover(page=page, count=count, sort="TRENDING_DESC") + return cast( + list[dict[str, Any]], + await self.async_discover(**self._trend_filters(page, count)), + ) def popular_this_season(self, page: int = 1, count: int = 20) -> list[dict]: """ @@ -548,13 +692,9 @@ class AniListApi: :param count: 每页条数 :return: AniList 媒体列表 """ - season, season_year = self._current_season() - return self.discover( - page=page, - count=count, - season=season, - season_year=season_year, - sort="POPULARITY_DESC", + return cast( + list[dict[str, Any]], + self.discover(**self._season_filters(page, count)), ) async def async_popular_this_season(self, page: int = 1, count: int = 20) -> list[dict]: @@ -565,13 +705,9 @@ class AniListApi: :param count: 每页条数 :return: AniList 媒体列表 """ - season, season_year = self._current_season() - return await self.async_discover( - page=page, - count=count, - season=season, - season_year=season_year, - sort="POPULARITY_DESC", + return cast( + list[dict[str, Any]], + await self.async_discover(**self._season_filters(page, count)), ) @cached( @@ -586,23 +722,11 @@ class AniListApi: :return: AniList 人物边列表 """ - query = """ - query ($id: Int!, $page: Int!, $count: Int!) { - Media(id: $id, type: ANIME) { - characters(page: $page, perPage: $count, sort: [ROLE, RELEVANCE]) { - edges { - role - node { id name { full native } } - voiceActors(language: JAPANESE, sort: [RELEVANCE]) { - id name { full native alternative } image { large medium } siteUrl - } - } - } - } - } - """ - result = self._invoke(query, {"id": anilist_id, "page": page, "count": count}) - return result.get("Media", {}).get("characters", {}).get("edges") or [] if result else [] + result = self._invoke( + self._credits_query, + {"id": anilist_id, "page": page, "count": count}, + ) + return self._credits_edges(result) @cached( maxsize=get_runtime_setting('CONF').anilist, @@ -616,23 +740,11 @@ class AniListApi: :return: AniList 人物边列表 """ - query = """ - query ($id: Int!, $page: Int!, $count: Int!) { - Media(id: $id, type: ANIME) { - characters(page: $page, perPage: $count, sort: [ROLE, RELEVANCE]) { - edges { - role - node { id name { full native } } - voiceActors(language: JAPANESE, sort: [RELEVANCE]) { - id name { full native alternative } image { large medium } siteUrl - } - } - } - } - } - """ - result = await self._async_invoke(query, {"id": anilist_id, "page": page, "count": count}) - return result.get("Media", {}).get("characters", {}).get("edges") or [] if result else [] + result = await self._async_invoke( + self._credits_query, + {"id": anilist_id, "page": page, "count": count}, + ) + return self._credits_edges(result) @cached( maxsize=get_runtime_setting('CONF').anilist, @@ -646,19 +758,11 @@ class AniListApi: :return: AniList 媒体列表 """ - query = """ - query ($id: Int!, $page: Int!, $count: Int!) { - Media(id: $id, type: ANIME) { - recommendations(page: $page, perPage: $count, sort: [RATING_DESC, ID]) { - nodes { mediaRecommendation { id } } - } - } - } - """ - result = self._invoke(query, {"id": anilist_id, "page": page, "count": count}) - nodes = result.get("Media", {}).get("recommendations", {}).get("nodes") or [] if result else [] - media_ids = [node.get("mediaRecommendation", {}).get("id") for node in nodes] - return self._medias_by_ids(media_ids) + result = self._invoke( + self._recommendations_query, + {"id": anilist_id, "page": page, "count": count}, + ) + return self._medias_by_ids(self._recommendation_ids(result)) @cached( maxsize=get_runtime_setting('CONF').anilist, @@ -672,19 +776,11 @@ class AniListApi: :return: AniList 媒体列表 """ - query = """ - query ($id: Int!, $page: Int!, $count: Int!) { - Media(id: $id, type: ANIME) { - recommendations(page: $page, perPage: $count, sort: [RATING_DESC, ID]) { - nodes { mediaRecommendation { id } } - } - } - } - """ - result = await self._async_invoke(query, {"id": anilist_id, "page": page, "count": count}) - nodes = result.get("Media", {}).get("recommendations", {}).get("nodes") or [] if result else [] - media_ids = [node.get("mediaRecommendation", {}).get("id") for node in nodes] - return await self._async_medias_by_ids(media_ids) + result = await self._async_invoke( + self._recommendations_query, + {"id": anilist_id, "page": page, "count": count}, + ) + return await self._async_medias_by_ids(self._recommendation_ids(result)) @cached( maxsize=get_runtime_setting('CONF').anilist, @@ -699,17 +795,10 @@ class AniListApi: :param person_id: AniList 人物 ID :return: AniList 人物详情 """ - query = """ - query ($id: Int!) { - Staff(id: $id) { - id name { full native alternative } image { large medium } - description(asHtml: false) dateOfBirth { year month day } - dateOfDeath { year month day } gender homeTown primaryOccupations siteUrl - } - } - """ - result = self._invoke(query, {"id": person_id}) - return result.get("Staff") if result else None + return self._field( + self._invoke(self._person_detail_query, {"id": person_id}), + "Staff", + ) @cached( maxsize=get_runtime_setting('CONF').anilist, @@ -724,17 +813,10 @@ class AniListApi: :param person_id: AniList 人物 ID :return: AniList 人物详情 """ - query = """ - query ($id: Int!) { - Staff(id: $id) { - id name { full native alternative } image { large medium } - description(asHtml: false) dateOfBirth { year month day } - dateOfDeath { year month day } gender homeTown primaryOccupations siteUrl - } - } - """ - result = await self._async_invoke(query, {"id": person_id}) - return result.get("Staff") if result else None + result = await self._async_invoke( + self._person_detail_query, {"id": person_id} + ) + return self._field(result, "Staff") @cached( maxsize=get_runtime_setting('CONF').anilist, @@ -748,18 +830,11 @@ class AniListApi: :return: AniList 媒体列表 """ - query = """ - query ($id: Int!, $page: Int!, $count: Int!) { - Staff(id: $id) { - characterMedia(page: $page, perPage: $count, sort: [POPULARITY_DESC]) { - nodes { id } - } - } - } - """ - result = self._invoke(query, {"id": person_id, "page": page, "count": count}) - nodes = result.get("Staff", {}).get("characterMedia", {}).get("nodes") or [] if result else [] - return self._medias_by_ids([node.get("id") for node in nodes]) + result = self._invoke( + self._person_credits_query, + {"id": person_id, "page": page, "count": count}, + ) + return self._medias_by_ids(self._person_media_ids(result)) @cached( maxsize=get_runtime_setting('CONF').anilist, @@ -773,18 +848,11 @@ class AniListApi: :return: AniList 媒体列表 """ - query = """ - query ($id: Int!, $page: Int!, $count: Int!) { - Staff(id: $id) { - characterMedia(page: $page, perPage: $count, sort: [POPULARITY_DESC]) { - nodes { id } - } - } - } - """ - result = await self._async_invoke(query, {"id": person_id, "page": page, "count": count}) - nodes = result.get("Staff", {}).get("characterMedia", {}).get("nodes") or [] if result else [] - return await self._async_medias_by_ids([node.get("id") for node in nodes]) + result = await self._async_invoke( + self._person_credits_query, + {"id": person_id, "page": page, "count": count}, + ) + return await self._async_medias_by_ids(self._person_media_ids(result)) def clear_cache(self) -> None: """清理 AniList 接口缓存""" diff --git a/app/modules/bangumi/bangumi.py b/app/modules/bangumi/bangumi.py index 9493b1ace..6a714e918 100644 --- a/app/modules/bangumi/bangumi.py +++ b/app/modules/bangumi/bangumi.py @@ -1,17 +1,24 @@ +from dataclasses import dataclass from datetime import datetime -from typing import Optional +from typing import Any, Optional, cast from app.adapters.network.http import AsyncRequestUtils, RequestUtils from app.runtime.cache import cached +from app.runtime.log import logger from app.runtime.settings import get_runtime_setting -class BangumiApi(object): - """ - Bangumi API客户端。 +@dataclass(frozen=True, slots=True) +class _BangumiRequestPlan: + """冻结 Bangumi 请求地址、参数和结果字段,供同步与异步传输共用。""" - 接口文档:https://bangumi.github.io/api/ - """ + url: str + params: dict[str, Any] + key: Optional[str] = None + + +class BangumiApi: + """Bangumi API 客户端,统一同步与异步请求决策和结果投影。""" _urls = { "discover": "v0/subjects", @@ -26,7 +33,8 @@ class BangumiApi(object): } _base_url = "https://api.bgm.tv/" - def __init__(self): + def __init__(self) -> None: + """初始化同步与异步 Bangumi 请求客户端。""" self._req = RequestUtils( ua=get_runtime_setting('NORMAL_USER_AGENT'), proxies=get_runtime_setting('PROXY'), @@ -37,283 +45,236 @@ class BangumiApi(object): proxies=get_runtime_setting('PROXY'), ) - @cached(maxsize=get_runtime_setting('CONF').bangumi, ttl=get_runtime_setting('CONF').meta, shared_key="get") - def __invoke(self, url, key: Optional[str] = None, **kwargs): - req_url = self._base_url + url - params = {} - if kwargs: - params.update(kwargs) - resp = self._req.get_res(url=req_url, params=params) - try: - if resp is None or resp.status_code != 200: - return None - result = resp.json() - return result.get(key) if key else result - except Exception as e: - print(e) + @classmethod + def _request_plan( + cls, + path: str, + key: Optional[str] = None, + **params: Any, + ) -> _BangumiRequestPlan: + """构造同步与异步请求共同使用的不可变调用计划。""" + return _BangumiRequestPlan( + url=f"{cls._base_url}{path}", + params=dict(params), + key=key, + ) + + @staticmethod + def _project_response( + status_code: Optional[int], + payload: Any, + key: Optional[str], + ) -> Any: + """按统一状态码与字段规则投影 Bangumi 响应。""" + if status_code != 200: return None + if key: + return payload.get(key) if isinstance(payload, dict) else None + return payload - @cached(maxsize=get_runtime_setting('CONF').bangumi, ttl=get_runtime_setting('CONF').meta, shared_key="get") - async def __async_invoke(self, url, key: Optional[str] = None, **kwargs): - req_url = self._base_url + url - params = {} - if kwargs: - params.update(kwargs) - resp = await self._async_req.get_res(url=req_url, params=params) - try: - if resp is None or resp.status_code != 200: - return None - result = resp.json() - return result.get(key) if key else result - except Exception as e: - print(e) + @classmethod + def _decode_response(cls, response: Any, key: Optional[str]) -> Any: + """解析 HTTP 响应,并把格式错误统一映射为空结果。""" + if response is None: return None + try: + payload = response.json() if response.status_code == 200 else None + except (TypeError, ValueError) as err: + logger.warning(f"Bangumi 响应解析失败:{str(err)}") + return None + return cls._project_response(response.status_code, payload, key) - def search(self, name): - """ - 搜索媒体信息 - """ - result = self.__invoke("search/subject/%s" % name) - if result: - return result.get("list") - return [] + @cached( + maxsize=get_runtime_setting('CONF').bangumi, + ttl=get_runtime_setting('CONF').meta, + shared_key="get", + ) + def __invoke(self, url, key=None, **kwargs): + """执行同步 HTTP 请求,业务计划与响应规则由共享 helper 决定。""" + plan = self._request_plan(url, key=key, **kwargs) + response = self._req.get_res(url=plan.url, params=plan.params) + return self._decode_response(response, plan.key) - async def async_search(self, name): - """ - 搜索媒体信息(异步版本) - """ - result = await self.__async_invoke("search/subject/%s" % name) - if result: - return result.get("list") - return [] + @cached( + maxsize=get_runtime_setting('CONF').bangumi, + ttl=get_runtime_setting('CONF').meta, + shared_key="get", + ) + async def __async_invoke(self, url, key=None, **kwargs): + """执行异步 HTTP 请求,业务计划与响应规则由共享 helper 决定。""" + plan = self._request_plan(url, key=key, **kwargs) + response = await self._async_req.get_res(url=plan.url, params=plan.params) + return self._decode_response(response, plan.key) - def calendar(self): - """ - 获取每日放送,返回items - """ - """ - [ - { - "weekday": { - "en": "Mon", - "cn": "星期一", - "ja": "月耀日", - "id": 1 - }, - "items": [ - { - "id": 350235, - "url": "http://bgm.tv/subject/350235", - "type": 2, - "name": "月が導く異世界道中 第二幕", - "name_cn": "月光下的异世界之旅 第二幕", - "summary": "", - "air_date": "2024-01-08", - "air_weekday": 1, - "rating": { - "total": 257, - "count": { - "1": 1, - "2": 1, - "3": 4, - "4": 15, - "5": 51, - "6": 111, - "7": 49, - "8": 13, - "9": 5, - "10": 7 - }, - "score": 6.1 - }, - "rank": 6125, - "images": { - "large": "http://lain.bgm.tv/pic/cover/l/3c/a5/350235_A0USf.jpg", - "common": "http://lain.bgm.tv/pic/cover/c/3c/a5/350235_A0USf.jpg", - "medium": "http://lain.bgm.tv/pic/cover/m/3c/a5/350235_A0USf.jpg", - "small": "http://lain.bgm.tv/pic/cover/s/3c/a5/350235_A0USf.jpg", - "grid": "http://lain.bgm.tv/pic/cover/g/3c/a5/350235_A0USf.jpg" - }, - "collection": { - "doing": 920 - } - }, - { - "id": 358561, - "url": "http://bgm.tv/subject/358561", - "type": 2, - "name": "大宇宙时代", - "name_cn": "大宇宙时代", - "summary": "", - "air_date": "2024-01-22", - "air_weekday": 1, - "rating": { - "total": 2, - "count": { - "1": 0, - "2": 0, - "3": 0, - "4": 0, - "5": 1, - "6": 1, - "7": 0, - "8": 0, - "9": 0, - "10": 0 - }, - "score": 5.5 - }, - "images": { - "large": "http://lain.bgm.tv/pic/cover/l/71/66/358561_UzsLu.jpg", - "common": "http://lain.bgm.tv/pic/cover/c/71/66/358561_UzsLu.jpg", - "medium": "http://lain.bgm.tv/pic/cover/m/71/66/358561_UzsLu.jpg", - "small": "http://lain.bgm.tv/pic/cover/s/71/66/358561_UzsLu.jpg", - "grid": "http://lain.bgm.tv/pic/cover/g/71/66/358561_UzsLu.jpg" - }, - "collection": { - "doing": 9 - } - } - ] - } + @staticmethod + def _dated_params(**params: Any) -> dict[str, Any]: + """为缓存键稳定附加当天日期,并保留调用方筛选参数。""" + return { + "_ts": datetime.strftime(datetime.now(), '%Y%m%d'), + **params, + } + + @staticmethod + def _search_results(result: Any) -> list[dict[str, Any]]: + """从旧版搜索响应中提取条目列表。""" + return result.get("list") or [] if isinstance(result, dict) else [] + + @staticmethod + def _calendar_items(result: Any) -> list[dict[str, Any]]: + """按星期顺序展开每日放送条目。""" + return [ + item + for weekday in result or [] + if isinstance(weekday, dict) + for item in weekday.get("items") or [] ] - """ - ret_list = [] - result = self.__invoke(self._urls["calendar"], _ts=datetime.strftime(datetime.now(), '%Y%m%d')) - if result: - for item in result: - ret_list.extend(item.get("items") or []) - return ret_list - async def async_calendar(self): - """ - 获取每日放送,返回items(异步版本) - """ - ret_list = [] - result = await self.__async_invoke(self._urls["calendar"], _ts=datetime.strftime(datetime.now(), '%Y%m%d')) - if result: - for item in result: - ret_list.extend(item.get("items") or []) - return ret_list + @staticmethod + def _credit_people(result: Any) -> list[dict[str, Any]]: + """把角色配音关系投影为带角色职业信息的人物列表。""" + people: list[dict[str, Any]] = [] + for item in result or []: + if not isinstance(item, dict) or not item.get("id"): + continue + actors = item.get("actors") or [] + if not actors or not isinstance(actors[0], dict): + continue + actor = actors[0] + actor.update({"career": [item.get("name")]}) + people.append(actor) + return people - def detail(self, bid: int): - """ - 获取番剧详情 - """ - return self.__invoke(self._urls["detail"] % bid, _ts=datetime.strftime(datetime.now(), '%Y%m%d')) + @staticmethod + def _list_result(result: Any) -> list[dict[str, Any]]: + """把列表响应统一投影为空安全列表。""" + return list(result) if isinstance(result, list) else [] - async def async_detail(self, bid: int): - """ - 获取番剧详情(异步版本) - """ - return await self.__async_invoke(self._urls["detail"] % bid, _ts=datetime.strftime(datetime.now(), '%Y%m%d')) + def search(self, name: str) -> list[dict[str, Any]]: + """搜索媒体信息。""" + return self._search_results(self.__invoke(f"search/subject/{name}")) - def credits(self, bid: int): - """ - 获取番剧人物 - """ - ret_list = [] - result = self.__invoke(self._urls["characters"] % bid, _ts=datetime.strftime(datetime.now(), '%Y%m%d')) - if result: - for item in result: - character_id = item.get("id") - actors = item.get("actors") - if character_id and actors and actors[0]: - actor_info = actors[0] - actor_info.update({'career': [item.get('name')]}) - ret_list.append(actor_info) - return ret_list + async def async_search(self, name: str) -> list[dict[str, Any]]: + """异步搜索媒体信息。""" + return self._search_results(await self.__async_invoke(f"search/subject/{name}")) - async def async_credits(self, bid: int): - """ - 获取番剧人物(异步版本) - """ - ret_list = [] - result = await self.__async_invoke(self._urls["characters"] % bid, - _ts=datetime.strftime(datetime.now(), '%Y%m%d')) - if result: - for item in result: - character_id = item.get("id") - actors = item.get("actors") - if character_id and actors and actors[0]: - actor_info = actors[0] - actor_info.update({'career': [item.get('name')]}) - ret_list.append(actor_info) - return ret_list + def calendar(self) -> list[dict[str, Any]]: + """获取每日放送条目。""" + result = self.__invoke(self._urls["calendar"], **self._dated_params()) + return self._calendar_items(result) - def subjects(self, bid: int): - """ - 获取关联条目信息 - """ - return self.__invoke(self._urls["subjects"] % bid, _ts=datetime.strftime(datetime.now(), '%Y%m%d')) + async def async_calendar(self) -> list[dict[str, Any]]: + """异步获取每日放送条目。""" + result = await self.__async_invoke( + self._urls["calendar"], **self._dated_params() + ) + return self._calendar_items(result) - async def async_subjects(self, bid: int): - """ - 获取关联条目信息(异步版本) - """ - return await self.__async_invoke(self._urls["subjects"] % bid, _ts=datetime.strftime(datetime.now(), '%Y%m%d')) + def detail(self, bid: int) -> Optional[dict[str, Any]]: + """获取番剧详情。""" + return cast( + Optional[dict[str, Any]], + self.__invoke(self._urls["detail"] % bid, **self._dated_params()), + ) - def person_detail(self, person_id: int): - """ - 获取人物详细信息 - """ - return self.__invoke(self._urls["person_detail"] % person_id, _ts=datetime.strftime(datetime.now(), '%Y%m%d')) + async def async_detail(self, bid: int) -> Optional[dict[str, Any]]: + """异步获取番剧详情。""" + return cast( + Optional[dict[str, Any]], + await self.__async_invoke( + self._urls["detail"] % bid, **self._dated_params() + ), + ) - async def async_person_detail(self, person_id: int): - """ - 获取人物详细信息(异步版本) - """ - return await self.__async_invoke(self._urls["person_detail"] % person_id, - _ts=datetime.strftime(datetime.now(), '%Y%m%d')) + def credits(self, bid: int) -> list[dict[str, Any]]: + """获取番剧配音人物。""" + result = self.__invoke(self._urls["characters"] % bid, **self._dated_params()) + return self._credit_people(result) - def person_credits(self, person_id: int): - """ - 获取人物参演作品 - """ - ret_list = [] - result = self.__invoke(self._urls["person_credits"] % person_id, - _ts=datetime.strftime(datetime.now(), '%Y%m%d')) - if result: - for item in result: - ret_list.append(item) - return ret_list + async def async_credits(self, bid: int) -> list[dict[str, Any]]: + """异步获取番剧配音人物。""" + result = await self.__async_invoke( + self._urls["characters"] % bid, **self._dated_params() + ) + return self._credit_people(result) - async def async_person_credits(self, person_id: int): - """ - 获取人物参演作品(异步版本) - """ - ret_list = [] - result = await self.__async_invoke(self._urls["person_credits"] % person_id, - _ts=datetime.strftime(datetime.now(), '%Y%m%d')) - if result: - for item in result: - ret_list.append(item) - return ret_list + def subjects(self, bid: int) -> Optional[list[dict[str, Any]]]: + """获取关联条目信息。""" + return cast( + Optional[list[dict[str, Any]]], + self.__invoke(self._urls["subjects"] % bid, **self._dated_params()), + ) - def discover(self, **kwargs): - """ - 发现 - """ - return self.__invoke(self._urls["discover"], - key="data", - _ts=datetime.strftime(datetime.now(), '%Y%m%d'), **kwargs) + async def async_subjects(self, bid: int) -> Optional[list[dict[str, Any]]]: + """异步获取关联条目信息。""" + return cast( + Optional[list[dict[str, Any]]], + await self.__async_invoke( + self._urls["subjects"] % bid, **self._dated_params() + ), + ) - async def async_discover(self, **kwargs): - """ - 发现(异步版本) - """ - return await self.__async_invoke(self._urls["discover"], - key="data", - _ts=datetime.strftime(datetime.now(), '%Y%m%d'), **kwargs) + def person_detail(self, person_id: int) -> Optional[dict[str, Any]]: + """获取人物详细信息。""" + return cast( + Optional[dict[str, Any]], + self.__invoke( + self._urls["person_detail"] % person_id, **self._dated_params() + ), + ) - def clear_cache(self): - """ - 清除缓存 - """ + async def async_person_detail( + self, person_id: int + ) -> Optional[dict[str, Any]]: + """异步获取人物详细信息。""" + return cast( + Optional[dict[str, Any]], + await self.__async_invoke( + self._urls["person_detail"] % person_id, **self._dated_params() + ), + ) + + def person_credits(self, person_id: int) -> list[dict[str, Any]]: + """获取人物参演作品。""" + result = self.__invoke( + self._urls["person_credits"] % person_id, **self._dated_params() + ) + return self._list_result(result) + + async def async_person_credits(self, person_id: int) -> list[dict[str, Any]]: + """异步获取人物参演作品。""" + result = await self.__async_invoke( + self._urls["person_credits"] % person_id, **self._dated_params() + ) + return self._list_result(result) + + def discover(self, **kwargs: Any) -> Optional[list[dict[str, Any]]]: + """按筛选条件发现番剧。""" + return cast( + Optional[list[dict[str, Any]]], + self.__invoke( + self._urls["discover"], + key="data", + **self._dated_params(**kwargs), + ), + ) + + async def async_discover( + self, **kwargs: Any + ) -> Optional[list[dict[str, Any]]]: + """异步按筛选条件发现番剧。""" + return cast( + Optional[list[dict[str, Any]]], + await self.__async_invoke( + self._urls["discover"], + key="data", + **self._dated_params(**kwargs), + ), + ) + + def clear_cache(self) -> None: + """清除 Bangumi 请求缓存。""" self.__invoke.cache_clear() def close(self) -> None: - """ - 关闭Bangumi会话 - """ + """关闭 Bangumi 同步会话。""" self._req.close() diff --git a/app/modules/indexer/spider/yema.py b/app/modules/indexer/spider/yema.py index b44fec625..b6d735859 100644 --- a/app/modules/indexer/spider/yema.py +++ b/app/modules/indexer/spider/yema.py @@ -1,13 +1,12 @@ import base64 import json -from typing import List, Optional, Tuple +from typing import Any, List, Optional, Tuple -from app.runtime.settings import get_runtime_setting - -from app.runtime.log import logger -from app.schemas.types import MediaType from app.adapters.network.http import AsyncRequestUtils, RequestUtils from app.foundation import temporal as time_tools +from app.runtime.log import logger +from app.runtime.settings import get_runtime_setting +from app.schemas.types import MediaType class YemaSpider: @@ -81,7 +80,7 @@ class YemaSpider: keyword: Optional[str], page: Optional[int], category_id: Optional[int] = None, - ) -> dict: + ) -> dict[str, Any]: """ 构造公开种子列表查询参数 @@ -90,7 +89,7 @@ class YemaSpider: :param category_id: 可选的站点分类 ID :return: YemaPT 开放 API 请求体 """ - params = { + params: dict[str, Any] = { "pageParam": { "current": int(page or 0) + 1, "pageSize": self._size, @@ -159,6 +158,8 @@ class YemaSpider: if label_id in self._labels ] torrent_id = result.get("id") + if not isinstance(torrent_id, int): + continue torrents.append({ "title": result.get("showName"), "description": result.get("shortDesc"), @@ -168,8 +169,12 @@ class YemaSpider: "seeders": result.get("seedNum"), "peers": result.get("leechNum"), "grabs": result.get("completedNum"), - "downloadvolumefactor": self._download_factor(result.get("downloadPromotion")), - "uploadvolumefactor": self._upload_factor(result.get("uploadPromotion")), + "downloadvolumefactor": self._download_factor( + str(result.get("downloadPromotion") or "") + ), + "uploadvolumefactor": self._upload_factor( + str(result.get("uploadPromotion") or "") + ), "freedate": time_tools.normalize_datetime(result.get("downloadPromotionEndTime")), "page_url": f"{self._site_url}/#/torrent/detail/{torrent_id}/", "labels": labels, @@ -208,6 +213,36 @@ class YemaSpider: return True, [] return False, self._parse_result(results) + def _prepare_search_requests( + self, + *, + keyword: Optional[str], + mtype: Optional[MediaType], + page: Optional[int], + ) -> Optional[List[dict[str, Any]]]: + """统一执行认证准入并生成各分类的搜索请求体。""" + if not self._api_key: + logger.warning(f"{self._name} 未配置 API AuthKey") + return None + return [ + self._build_params(keyword, page, category_id) + for category_id in self._search_category_ids(mtype) + ] + + def _finalize_search_responses( + self, + responses: List[Any], + ) -> Tuple[bool, List[dict[str, Any]]]: + """统一投影各分类响应并合并去重后的搜索结果。""" + result_groups = [] + errors = [] + for response in responses: + error, results = self._process_search_response(response) + errors.append(error) + if not error: + result_groups.append(results) + return all(errors), self._merge_search_results(result_groups) + def search( self, keyword: Optional[str], @@ -222,26 +257,26 @@ class YemaSpider: :param page: MoviePilot 从 0 开始的页码 :return: 是否失败及标准种子列表 """ - if not self._api_key: - logger.warning(f"{self._name} 未配置 API AuthKey") + payloads = self._prepare_search_requests( + keyword=keyword, + mtype=mtype, + page=page, + ) + if payloads is None: return True, [] request = RequestUtils( headers=self._request_headers(), proxies=self._proxy, timeout=self._timeout, ) - result_groups = [] - errors = [] - for category_id in self._search_category_ids(mtype): - response = request.post_res( + responses = [ + request.post_res( url=self._search_url, - json=self._build_params(keyword, page, category_id), + json=payload, ) - error, results = self._process_search_response(response) - errors.append(error) - if not error: - result_groups.append(results) - return all(errors), self._merge_search_results(result_groups) + for payload in payloads + ] + return self._finalize_search_responses(responses) async def async_search( self, @@ -257,26 +292,25 @@ class YemaSpider: :param page: MoviePilot 从 0 开始的页码 :return: 是否失败及标准种子列表 """ - if not self._api_key: - logger.warning(f"{self._name} 未配置 API AuthKey") + payloads = self._prepare_search_requests( + keyword=keyword, + mtype=mtype, + page=page, + ) + if payloads is None: return True, [] request = AsyncRequestUtils( headers=self._request_headers(), proxies=self._proxy, timeout=self._timeout, ) - result_groups = [] - errors = [] - for category_id in self._search_category_ids(mtype): - response = await request.post_res( + responses = [] + for payload in payloads: + responses.append(await request.post_res( url=self._search_url, - json=self._build_params(keyword, page, category_id), - ) - error, results = self._process_search_response(response) - errors.append(error) - if not error: - result_groups.append(results) - return all(errors), self._merge_search_results(result_groups) + json=payload, + )) + return self._finalize_search_responses(responses) @staticmethod def _download_factor(promotion: str) -> float: diff --git a/app/modules/musicbrainz/__init__.py b/app/modules/musicbrainz/__init__.py index 999a41842..742fb8eff 100644 --- a/app/modules/musicbrainz/__init__.py +++ b/app/modules/musicbrainz/__init__.py @@ -65,6 +65,30 @@ class _MusicBrainzRecognitionPlan: raise RuntimeError("MusicBrainz 详情识别计划缺少原生 ID") return self.media_id + def detail_kwargs(self) -> dict[str, str]: + """生成详情入口兼容旧签名所需的可选实体参数。""" + return ( + {"music_type": self.music_type} + if self.music_type is not None + else {} + ) + + +@dataclass(frozen=True, slots=True) +class _MusicBrainzResponseDecision: + """描述一次 MusicBrainz 响应的投影结果与退避决策。""" + + payload: Optional[dict[str, Any]] = None + retry_delay: Optional[float] = None + + +@dataclass(frozen=True, slots=True) +class _MusicBrainzRequestPlan: + """冻结 MusicBrainz 请求路径与参数,供同步异步 I/O 外壳共用。""" + + path: str + params: dict[str, Any] + class MusicBrainzModule(_ModuleBase): """通过 MusicBrainz 提供音乐元数据搜索和详情识别。""" @@ -497,14 +521,9 @@ class MusicBrainzModule(_ModuleBase): if not tracks: return None details: list[dict[str, Any]] = [] - for release in self._search_release_candidates(meta, tracks, limit=limit): - release_id = release.get("id") - if not release_id: - continue - detail = self._request_json( - f"/release/{release_id}", - params={"inc": "recordings+media+artist-credits", "fmt": "json"}, - ) + releases = self._search_release_candidates(meta, tracks, limit=limit) + for request in self._release_detail_requests(releases): + detail = self._request_json(request.path, params=request.params) if not detail: continue details.append(detail) @@ -525,13 +544,9 @@ class MusicBrainzModule(_ModuleBase): tracks, limit=limit, ) - for release in releases: - release_id = release.get("id") - if not release_id: - continue + for request in self._release_detail_requests(releases): detail = await self._async_request_json( - f"/release/{release_id}", - params={"inc": "recordings+media+artist-credits", "fmt": "json"}, + request.path, params=request.params ) if not detail: continue @@ -568,10 +583,9 @@ class MusicBrainzModule(_ModuleBase): """按专辑名和曲名线索搜索候选发行版本,多个查询按命中顺序去重。""" releases: list[dict[str, Any]] = [] seen: set[str] = set() - for query in self._release_queries(meta, tracks): + for request in self._release_search_requests(meta, tracks, limit): payload = self._request_json( - "/release", - params={"query": query, "limit": max(1, min(limit, 25)), "fmt": "json"}, + request.path, params=request.params ) self._merge_release_candidates(releases, seen, payload) if len(releases) >= limit: @@ -587,14 +601,9 @@ class MusicBrainzModule(_ModuleBase): """异步按专辑名和曲名线索搜索并去重候选发行版本。""" releases: list[dict[str, Any]] = [] seen: set[str] = set() - for query in self._release_queries(meta, tracks): + for request in self._release_search_requests(meta, tracks, limit): payload = await self._async_request_json( - "/release", - params={ - "query": query, - "limit": max(1, min(limit, 25)), - "fmt": "json", - }, + request.path, params=request.params ) self._merge_release_candidates(releases, seen, payload) if len(releases) >= limit: @@ -614,6 +623,44 @@ class MusicBrainzModule(_ModuleBase): seen.add(release_id) releases.append(item) + @classmethod + def _release_search_requests( + cls, + meta: MetaMusic, + tracks: list[MetaMusic], + limit: int, + ) -> list[_MusicBrainzRequestPlan]: + """构造发行候选查询计划,统一同步与异步的限额和参数。""" + normalized_limit = max(1, min(limit, 25)) + return [ + _MusicBrainzRequestPlan( + path="/release", + params={ + "query": query, + "limit": normalized_limit, + "fmt": "json", + }, + ) + for query in cls._release_queries(meta, tracks) + ] + + @staticmethod + def _release_detail_requests( + releases: Iterable[dict[str, Any]], + ) -> list[_MusicBrainzRequestPlan]: + """按候选顺序构造发行详情请求计划并跳过无 ID 条目。""" + return [ + _MusicBrainzRequestPlan( + path=f"/release/{release_id}", + params={ + "inc": "recordings+media+artist-credits", + "fmt": "json", + }, + ) + for release in releases + if (release_id := release.get("id")) + ] + @classmethod def _release_queries(cls, meta: MetaMusic, tracks: list[MetaMusic]) -> list[str]: """构造专辑搜索表达式:优先专辑名+歌手,无专辑线索时用曲名兜底。""" @@ -828,19 +875,12 @@ class MusicBrainzModule(_ModuleBase): if not plan: return None if plan.media_id: - detail_kwargs = ( - {"music_type": plan.music_type} - if plan.music_type is not None - else {} - ) info = self.recognize_music( plan.media_source, plan.require_media_id(), - **detail_kwargs, + **plan.detail_kwargs(), ) - if info and plan.meta: - self._update_recognize_cache(plan.meta, info) - return info + return self._finalize_detail_recognition(plan, info) return self._recognize_from_candidates_sync(plan) def _update_recognize_cache(self, meta: MetaMusic, info: Optional[MusicInfo]) -> None: @@ -891,19 +931,12 @@ class MusicBrainzModule(_ModuleBase): if not plan: return None if plan.media_id: - detail_kwargs = ( - {"music_type": plan.music_type} - if plan.music_type is not None - else {} - ) info = await self.async_recognize_music( plan.media_source, plan.require_media_id(), - **detail_kwargs, + **plan.detail_kwargs(), ) - if info and plan.meta: - self._update_recognize_cache(plan.meta, info) - return info + return self._finalize_detail_recognition(plan, info) return await self._recognize_from_candidates_async(plan) @classmethod @@ -962,6 +995,16 @@ class MusicBrainzModule(_ModuleBase): cached_info.recognize_cache_hit = True return cached_info + def _finalize_detail_recognition( + self, + plan: _MusicBrainzRecognitionPlan, + info: Optional[MusicInfo], + ) -> Optional[MusicInfo]: + """统一完成显式详情识别后的缓存回填。""" + if info and plan.meta: + self._update_recognize_cache(plan.meta, info) + return info + @classmethod def _select_recognition_candidate( cls, @@ -980,6 +1023,26 @@ class MusicBrainzModule(_ModuleBase): return matched return cls._select_album_candidate(meta, albums) + @staticmethod + def _should_search_albums( + plan: _MusicBrainzRecognitionPlan, + preliminary: Optional[MusicInfo], + ) -> bool: + """统一决定候选识别是否需要继续查询专辑。""" + meta = plan.require_meta() + return bool( + plan.music_type == MUSIC_ENTITY_ALBUM + or (not preliminary and plan.search_album and meta.artists) + ) + + @staticmethod + def _should_probe_album( + plan: _MusicBrainzRecognitionPlan, + recording: Optional[MusicInfo], + ) -> bool: + """统一决定 Recording 详情未命中后是否继续探测专辑。""" + return bool(not recording and plan.search_album) + def _finalize_recognition( self, plan: _MusicBrainzRecognitionPlan, @@ -1007,10 +1070,7 @@ class MusicBrainzModule(_ModuleBase): preliminary = self._select_recognition_candidate(plan, recordings) albums = ( self._search_albums(meta, limit=10) - if ( - plan.music_type == MUSIC_ENTITY_ALBUM - or (not preliminary and plan.search_album and meta.artists) - ) + if self._should_search_albums(plan, preliminary) else [] ) matched = self._select_recognition_candidate(plan, recordings, albums) @@ -1034,10 +1094,7 @@ class MusicBrainzModule(_ModuleBase): preliminary = self._select_recognition_candidate(plan, recordings) albums = ( await self._async_search_albums(meta, limit=10) - if ( - plan.music_type == MUSIC_ENTITY_ALBUM - or (not preliminary and plan.search_album and meta.artists) - ) + if self._should_search_albums(plan, preliminary) else [] ) matched = self._select_recognition_candidate(plan, recordings, albums) @@ -1301,6 +1358,7 @@ class MusicBrainzModule(_ModuleBase): plan = self._detail_plan(media_source, media_id, music_type) if not plan: return None + result: Optional[MusicInfo] = None if plan.search_recording: payload = self._request_json( f"/recording/{plan.require_media_id()}", @@ -1312,8 +1370,8 @@ class MusicBrainzModule(_ModuleBase): result = self._project_recording_detail(payload) if result: return result - if not plan.search_album: - return None + if not self._should_probe_album(plan, result): + return None # MusicBrainz 各实体共用 UUID 形式,统一详情入口在 Recording 未命中后继续探测专辑。 album = self.music_album(self._source, plan.require_media_id()) return self._project_album_result(album) @@ -1328,6 +1386,7 @@ class MusicBrainzModule(_ModuleBase): plan = self._detail_plan(media_source, media_id, music_type) if not plan: return None + result: Optional[MusicInfo] = None if plan.search_recording: payload = await self._async_request_json( f"/recording/{plan.require_media_id()}", @@ -1339,8 +1398,8 @@ class MusicBrainzModule(_ModuleBase): result = self._project_recording_detail(payload) if result: return result - if not plan.search_album: - return None + if not self._should_probe_album(plan, result): + return None album = await self._async_music_album( self._source, plan.require_media_id() ) @@ -1976,6 +2035,42 @@ class MusicBrainzModule(_ModuleBase): if delay := cls._reserve_request_delay(): await asyncio.sleep(delay) + @classmethod + def _response_decision( + cls, + response: Any, + path: str, + attempt: int, + attempts: int, + ) -> _MusicBrainzResponseDecision: + """统一分类响应、解析 JSON,并决定是否执行下一次退避重试。""" + status_code = response.status_code + if status_code == 404: + logger.debug(f"MusicBrainz 资源不存在:{path}") + return _MusicBrainzResponseDecision(payload={}) + if status_code == 429 or status_code >= 500: + logger.warning( + f"MusicBrainz 服务繁忙:{status_code} {response.text[:200]}" + ) + if attempt < attempts - 1: + return _MusicBrainzResponseDecision( + retry_delay=cls._busy_backoff * (2 ** attempt) + ) + return _MusicBrainzResponseDecision() + if status_code != 200: + logger.warning( + f"MusicBrainz 请求失败:{status_code} {response.text[:200]}" + ) + return _MusicBrainzResponseDecision() + try: + payload = response.json() + except (TypeError, ValueError) as err: + logger.warning(f"MusicBrainz 响应解析失败:{err}") + return _MusicBrainzResponseDecision() + return _MusicBrainzResponseDecision( + payload=payload if isinstance(payload, dict) else None + ) + @classmethod @cached(maxsize=get_runtime_setting('CONF').musicbrainz, ttl=get_runtime_setting('CONF').meta, skip_none=True) def _request_json( @@ -1996,32 +2091,16 @@ class MusicBrainzModule(_ModuleBase): ) if response is None: return None - status_code = response.status_code try: - if status_code == 404: - # 单曲与专辑共用同一套 ID 入口,404 属于正常的探测结果 - logger.debug(f"MusicBrainz 资源不存在:{path}") - # 使用空对象区分稳定的不存在与瞬时请求失败,使有界缓存能够复用探测结果。 - return {} - if status_code == 429 or status_code >= 500: - logger.warning( - f"MusicBrainz 服务繁忙:{status_code} {response.text[:200]}" - ) - if attempt < attempts - 1: - time.sleep(cls._busy_backoff * (2 ** attempt)) - continue - return None - if status_code != 200: - logger.warning( - f"MusicBrainz 请求失败:{status_code} {response.text[:200]}" - ) - return None - return response.json() - except (TypeError, ValueError) as err: - logger.warning(f"MusicBrainz 响应解析失败:{err}") - return None + decision = cls._response_decision( + response, path, attempt, attempts + ) finally: response.close() + if decision.retry_delay is not None: + time.sleep(decision.retry_delay) + continue + return decision.payload return None @classmethod @@ -2050,29 +2129,14 @@ class MusicBrainzModule(_ModuleBase): ).get_res(f"{cls._base_url}{path}", params=params) if response is None: return None - status_code = response.status_code try: - if status_code == 404: - logger.debug(f"MusicBrainz 资源不存在:{path}") - return {} - if status_code == 429 or status_code >= 500: - logger.warning( - f"MusicBrainz 服务繁忙:{status_code} {response.text[:200]}" - ) - if attempt < attempts - 1: - await asyncio.sleep(cls._busy_backoff * (2 ** attempt)) - continue - return None - if status_code != 200: - logger.warning( - f"MusicBrainz 请求失败:{status_code} {response.text[:200]}" - ) - return None - payload = response.json() - return payload if isinstance(payload, dict) else None - except (TypeError, ValueError) as err: - logger.warning(f"MusicBrainz 响应解析失败:{err}") - return None + decision = cls._response_decision( + response, path, attempt, attempts + ) finally: await response.aclose() + if decision.retry_delay is not None: + await asyncio.sleep(decision.retry_delay) + continue + return decision.payload return None diff --git a/app/modules/theaudiodb/__init__.py b/app/modules/theaudiodb/__init__.py index 1811dd29f..a912232dd 100644 --- a/app/modules/theaudiodb/__init__.py +++ b/app/modules/theaudiodb/__init__.py @@ -56,6 +56,14 @@ class _TheAudioDbRecognitionPlan: return self.media_id +@dataclass(frozen=True, slots=True) +class _TheAudioDbRequestPlan: + """冻结 TheAudioDB 请求地址和参数,供同步异步传输共用。""" + + url: str + params: dict[str, Any] + + class TheAudioDbModule(_ModuleBase): """通过 TheAudioDB V1 API 提供音乐搜索、详情和手动识别能力。""" @@ -146,12 +154,13 @@ class TheAudioDbModule(_ModuleBase): self._source, plan.require_media_id(), music_type=plan.music_type ) plan_meta = plan.require_meta() + matched: Optional[MusicInfo] = None if plan.search_recording: matched = self._select_track(plan_meta, self._search_tracks(plan_meta)) if matched: return matched - if not plan.search_album: - return None + if not self._should_search_album(plan, matched): + return None album = self._select_album(plan_meta, self._search_albums(plan_meta)) return self._project_album_result(album) @@ -180,6 +189,7 @@ class TheAudioDbModule(_ModuleBase): music_type=plan.music_type, ) plan_meta = plan.require_meta() + matched: Optional[MusicInfo] = None if plan.search_recording: matched = self._select_track( plan_meta, @@ -187,8 +197,8 @@ class TheAudioDbModule(_ModuleBase): ) if matched: return matched - if not plan.search_album: - return None + if not self._should_search_album(plan, matched): + return None album = self._select_album( plan_meta, await self._async_search_albums(plan_meta), @@ -219,6 +229,14 @@ class TheAudioDbModule(_ModuleBase): music_type=music_type, ) + @staticmethod + def _should_search_album( + plan: _TheAudioDbRecognitionPlan, + matched: Optional[MusicInfo], + ) -> bool: + """统一决定单曲未命中后是否继续查询专辑。""" + return bool(not matched and plan.search_album) + def recognize_music( self, media_source: MediaSource, @@ -229,6 +247,7 @@ class TheAudioDbModule(_ModuleBase): plan = self._detail_plan(media_source, media_id, music_type) if not plan: return None + result: Optional[MusicInfo] = None if plan.search_recording: payload = self._request_json( "track.php", {"h": plan.require_media_id()} @@ -236,8 +255,8 @@ class TheAudioDbModule(_ModuleBase): result = self._project_track_detail(payload) if result: return result - if not plan.search_album: - return None + if not self._should_search_album(plan, result): + return None album = self.music_album(self._source, plan.require_media_id()) return self._project_album_result(album) @@ -251,6 +270,7 @@ class TheAudioDbModule(_ModuleBase): plan = self._detail_plan(media_source, media_id, music_type) if not plan: return None + result: Optional[MusicInfo] = None if plan.search_recording: payload = await self._async_request_json( "track.php", {"h": plan.require_media_id()} @@ -258,8 +278,8 @@ class TheAudioDbModule(_ModuleBase): result = self._project_track_detail(payload) if result: return result - if not plan.search_album: - return None + if not self._should_search_album(plan, result): + return None album = await self._async_music_album( self._source, plan.require_media_id() ) @@ -665,6 +685,42 @@ class TheAudioDbModule(_ModuleBase): return results return results + @classmethod + def _response_payload( + cls, response: Any, endpoint: str + ) -> Optional[dict[str, Any]]: + """统一校验并解析 TheAudioDB 响应,避免同步异步错误语义漂移。""" + if response.status_code != 200: + return None + diagnostic = cls._response_diagnostic(response, endpoint) + if getattr(response, "content", None) in (b"", ""): + logger.warning(f"TheAudioDB 返回空响应:{diagnostic}") + return None + try: + payload = response.json() + except (TypeError, ValueError) as err: + logger.warning( + f"TheAudioDB 响应解析失败:{diagnostic},错误:{str(err)}" + ) + return None + return payload if isinstance(payload, dict) else None + + @classmethod + def _request_plan( + cls, + api_key: str, + endpoint: str, + params: Optional[dict[str, Any]], + ) -> Optional[_TheAudioDbRequestPlan]: + """校验 API Key 并构造同步异步共用的请求计划。""" + normalized_key = str(api_key or "").strip() + if not normalized_key: + return None + return _TheAudioDbRequestPlan( + url=f"{cls._base_url}/{normalized_key}/{endpoint}", + params=dict(params or {}), + ) + @classmethod @cached(maxsize=get_runtime_setting('CONF').theaudiodb, ttl=get_runtime_setting('CONF').meta, skip_none=True) def _request_json( @@ -673,8 +729,10 @@ class TheAudioDbModule(_ModuleBase): params: Optional[dict[str, Any]] = None, ) -> Optional[dict[str, Any]]: """请求 TheAudioDB V1 JSON 接口并统一处理错误响应。""" - api_key = str(get_runtime_setting('THEAUDIODB_API_KEY') or "").strip() - if not api_key: + plan = cls._request_plan( + get_runtime_setting('THEAUDIODB_API_KEY'), endpoint, params + ) + if not plan: logger.warning("TheAudioDB API Key 未配置,跳过请求") return None response = RequestUtils( @@ -682,26 +740,13 @@ class TheAudioDbModule(_ModuleBase): proxies=get_runtime_setting('PROXY'), timeout=30, ).get_res( - url=f"{cls._base_url}/{api_key}/{endpoint}", - params=params or {}, + url=plan.url, + params=plan.params, ) if response is None: return None try: - if response.status_code != 200: - return None - diagnostic = cls._response_diagnostic(response, endpoint) - if getattr(response, "content", None) in (b"", ""): - logger.warning(f"TheAudioDB 返回空响应:{diagnostic}") - return None - try: - payload = response.json() - except (TypeError, ValueError) as err: - logger.warning( - f"TheAudioDB 响应解析失败:{diagnostic},错误:{str(err)}" - ) - return None - return payload if isinstance(payload, dict) else None + return cls._response_payload(response, endpoint) finally: response.close() @@ -718,8 +763,10 @@ class TheAudioDbModule(_ModuleBase): params: Optional[dict[str, Any]] = None, ) -> Optional[dict[str, Any]]: """异步请求 TheAudioDB V1 JSON 接口并统一处理错误响应。""" - api_key = str(get_runtime_setting('THEAUDIODB_API_KEY') or "").strip() - if not api_key: + plan = cls._request_plan( + get_runtime_setting('THEAUDIODB_API_KEY'), endpoint, params + ) + if not plan: logger.warning("TheAudioDB API Key 未配置,跳过请求") return None response = await AsyncRequestUtils( @@ -727,26 +774,13 @@ class TheAudioDbModule(_ModuleBase): proxies=get_runtime_setting('PROXY'), timeout=30, ).get_res( - url=f"{cls._base_url}/{api_key}/{endpoint}", - params=params or {}, + url=plan.url, + params=plan.params, ) if response is None: return None try: - if response.status_code != 200: - return None - diagnostic = cls._response_diagnostic(response, endpoint) - if getattr(response, "content", None) in (b"", ""): - logger.warning(f"TheAudioDB 返回空响应:{diagnostic}") - return None - try: - payload = response.json() - except (TypeError, ValueError) as err: - logger.warning( - f"TheAudioDB 响应解析失败:{diagnostic},错误:{str(err)}" - ) - return None - return payload if isinstance(payload, dict) else None + return cls._response_payload(response, endpoint) finally: await response.aclose() diff --git a/app/modules/thetvdb/__init__.py b/app/modules/thetvdb/__init__.py index 734ac564e..36f65d3f8 100644 --- a/app/modules/thetvdb/__init__.py +++ b/app/modules/thetvdb/__init__.py @@ -1,15 +1,15 @@ +from dataclasses import dataclass from threading import Lock -from typing import Optional, Tuple, Union, cast +from typing import Optional, Tuple, Union from app.domain.context import MediaInfo from app.domain.media import is_media_source_enabled from app.domain.meta.metabase import MetaBase -from app.runtime.execution import run_in_threadpool -from app.runtime.settings import get_runtime_setting - from app.modules import _ModuleBase from app.modules.thetvdb import client +from app.runtime.execution import run_in_threadpool from app.runtime.log import logger +from app.runtime.settings import get_runtime_setting from app.schemas.types import ( MediaRecognizeType, MediaSource, @@ -19,6 +19,14 @@ from app.schemas.types import ( ) +@dataclass(frozen=True, slots=True) +class _TvdbAuxiliaryLookup: + """描述附加信息查询需要执行的单次 TVDB I/O。""" + + method_name: str + argument: Union[int, str] + + class TheTvDbModule(_ModuleBase): """ TVDB媒体信息匹配 @@ -230,21 +238,68 @@ class TheTvDbModule(_ModuleBase): metainfo: Optional[MetaBase] = None, ) -> list[MediaInfo]: """从 TVDB 补充电视剧别名,不向主媒体写入 TVDB 专用字段。""" + lookup = self._build_auxiliary_lookup( + mediainfo=mediainfo, + media_source=media_source, + metainfo=metainfo, + ) + if not lookup: + return [] + candidates = self._load_auxiliary_candidates(lookup) + return self._resolve_auxiliary_candidates(mediainfo, candidates) + + @staticmethod + def _build_auxiliary_lookup( + mediainfo: MediaInfo, + media_source: Optional[MediaSourceSelection], + metainfo: Optional[MetaBase], + ) -> Optional[_TvdbAuxiliaryLookup]: + """校验 TVDB 附加信息请求并选择原生 ID 或标题查询。""" if ( not mediainfo or mediainfo.type != MediaType.TV or not is_media_source_enabled(media_source, MediaSource.TVDB) ): - return [] + return None del metainfo if ( mediainfo.media_source == MediaSource.TVDB and str(mediainfo.media_id or "").isdigit() ): - info = self.tvdb_info(int(mediainfo.media_id)) - candidates = [info] if info else [] - else: - candidates = self.search_tvdb(mediainfo.title) + return _TvdbAuxiliaryLookup( + method_name="tvdb_info", + argument=int(mediainfo.media_id), + ) + return _TvdbAuxiliaryLookup( + method_name="search_tvdb", + argument=mediainfo.title, + ) + + def _load_auxiliary_candidates( + self, lookup: _TvdbAuxiliaryLookup + ) -> list[dict[str, object]]: + """通过同步 TVDB I/O 获取候选详情。""" + result = getattr(self, lookup.method_name)(lookup.argument) + if lookup.method_name == "tvdb_info": + return [result] if result else [] + return result or [] + + async def _async_load_auxiliary_candidates( + self, lookup: _TvdbAuxiliaryLookup + ) -> list[dict[str, object]]: + """仅在线程池中执行 TVDB 客户端的阻塞网络调用。""" + result = await run_in_threadpool( + getattr(self, lookup.method_name), + lookup.argument, + ) + if lookup.method_name == "tvdb_info": + return [result] if result else [] + return result or [] + + def _resolve_auxiliary_candidates( + self, mediainfo: MediaInfo, candidates: list[dict[str, object]] + ) -> list[MediaInfo]: + """按年份、名称和来源 ID 统一解析 TVDB 候选。""" target_names = { " ".join(str(name).casefold().split()) for name in [mediainfo.title, *(mediainfo.names or [])] @@ -282,16 +337,16 @@ class TheTvDbModule(_ModuleBase): media_source: Optional[MediaSourceSelection] = None, metainfo: Optional[MetaBase] = None, ) -> list[MediaInfo]: - """在线程池中执行 TVDB 同步附加信息查询。""" - return cast( - list[MediaInfo], - await run_in_threadpool( - self.get_media_auxiliary_info, - mediainfo=mediainfo, - media_source=media_source, - metainfo=metainfo, - ), + """异步获取 TVDB 候选,并复用同步入口的纯解析决策。""" + lookup = self._build_auxiliary_lookup( + mediainfo=mediainfo, + media_source=media_source, + metainfo=metainfo, ) + if not lookup: + return [] + candidates = await self._async_load_auxiliary_candidates(lookup) + return self._resolve_auxiliary_candidates(mediainfo, candidates) def clear_cache(self): """ diff --git a/tests/test_agent_chat_history.py b/tests/test_agent_chat_history.py index ccce5f586..cc595ed7c 100644 --- a/tests/test_agent_chat_history.py +++ b/tests/test_agent_chat_history.py @@ -1,7 +1,7 @@ import asyncio import json from types import SimpleNamespace -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, MagicMock, call from langchain_core.messages import AIMessage, HumanMessage @@ -327,3 +327,82 @@ def test_async_memory_manager_restores_through_native_async_service(monkeypatch) session_id=session_id, user_id=user_id, ) + + +def test_memory_restore_entries_share_lookup_and_projection(monkeypatch): + """记忆同步、异步恢复只保留查询 I/O 差异并共享回退与投影决策。""" + snapshot = SimpleNamespace(agent_messages=[ + { + "type": "human", + "data": { + "content": "共享恢复", + "additional_kwargs": {}, + "response_metadata": {}, + "type": "human", + "name": None, + "id": None, + "example": False, + }, + } + ]) + chat = MagicMock() + chat.get_sync.side_effect = [None, snapshot] + chat.get = AsyncMock(side_effect=[None, snapshot]) + manager = MemoryManager(chat=chat, persistence=MagicMock()) + lookup = MagicMock(wraps=manager._chat_lookup_params) + restore = MagicMock(wraps=manager._restore_agent_messages) + monkeypatch.setattr(manager, "_chat_lookup_params", lookup) + monkeypatch.setattr(manager, "_restore_agent_messages", restore) + + sync_messages = manager.get_agent_messages("session-parity", "user-parity") + manager.clear_memory("session-parity", "user-parity") + async_messages = asyncio.run( + manager.async_get_agent_messages("session-parity", "user-parity") + ) + + assert [message.content for message in sync_messages] == ["共享恢复"] + assert [message.content for message in async_messages] == ["共享恢复"] + assert lookup.call_args_list == [ + call("session-parity", "user-parity"), + call("session-parity", "user-parity"), + ] + assert restore.call_count == 2 + expected_calls = [ + call(session_id="session-parity", user_id="user-parity", chat=snapshot), + call(session_id="session-parity", user_id="user-parity", chat=snapshot), + ] + assert restore.call_args_list == expected_calls + assert chat.get_sync.call_args_list == [ + call(session_id="session-parity", user_id="user-parity"), + call(session_id="session-parity"), + ] + assert chat.get.await_args_list == [ + call(session_id="session-parity", user_id="user-parity"), + call(session_id="session-parity"), + ] + + +def test_memory_save_entries_share_cache_state_projection(monkeypatch): + """记忆同步、异步保存必须共享内存状态更新和消息序列化结果。""" + chat = MagicMock() + persistence = MagicMock() + persistence.async_save_agent_messages = AsyncMock() + manager = MemoryManager(chat=chat, persistence=persistence) + update = MagicMock(wraps=manager._update_agent_messages) + monkeypatch.setattr(manager, "_update_agent_messages", update) + messages = [HumanMessage(content="共享保存")] + + manager.save_agent_messages("session-save", "user-save", messages) + asyncio.run( + manager.async_save_agent_messages("session-save", "user-save", messages) + ) + + expected = { + "session_id": "session-save", + "user_id": "user-save", + "messages": messages, + } + assert update.call_args_list == [call(**expected), call(**expected)] + assert chat.save_agent_messages.call_args.kwargs == ( + persistence.async_save_agent_messages.await_args.kwargs + ) diff --git a/tests/test_chain_mixin_contracts.py b/tests/test_chain_mixin_contracts.py index 701794241..2b6679273 100644 --- a/tests/test_chain_mixin_contracts.py +++ b/tests/test_chain_mixin_contracts.py @@ -86,7 +86,7 @@ def test_domain_mixins_keep_concrete_imports_explicit_until_next_migration() -> assert set(violations) <= { "app/chain/_music.py:app.chain.download", "app/chain/_music.py:app.chain.media", - "app/chain/_music.py:app.chain.search", + "app/chain/_music.py:app.chain.search.facade", "app/chain/transfer/filter.py:app.chain.media", "app/chain/transfer/filter.py:app.chain.storage", "app/chain/transfer/format.py:app.chain.storage", diff --git a/tests/test_downloader_path_mapping.py b/tests/test_downloader_path_mapping.py index f2edf878b..1e760bcc9 100644 --- a/tests/test_downloader_path_mapping.py +++ b/tests/test_downloader_path_mapping.py @@ -141,6 +141,8 @@ def _load_transmission_module(): cache_module = types.ModuleType("app.runtime.cache") runtime_settings_module = types.ModuleType("app.runtime.settings") base_module = types.ModuleType("app.modules._base") + base_module.__path__ = [] + base_downloader_module = types.ModuleType("app.modules._base.downloader") modules_module = types.ModuleType("app.modules") modules_module.__path__ = [] transmission_package_module = types.ModuleType("app.modules.transmission") @@ -295,7 +297,8 @@ def _load_transmission_module(): modules_module._ModuleBase = _ModuleBase modules_module._DownloaderBase = _DownloaderBase modules_module._base = base_module - base_module._DownloaderModuleBase = _DownloaderModuleBase + base_module.downloader = base_downloader_module + base_downloader_module._DownloaderModuleBase = _DownloaderModuleBase torrent_rules_module.is_magnet_link = _is_magnet_link size_tools_module.format_compact_size = _format_size temporal_tools_module.format_duration = _format_duration @@ -343,6 +346,7 @@ def _load_transmission_module(): "app.runtime.settings": runtime_settings_module, "app.modules": modules_module, "app.modules._base": base_module, + "app.modules._base.downloader": base_downloader_module, "app.modules.transmission": transmission_package_module, "app.modules.transmission.transmission": transmission_client_module, "app.schemas": schemas_module, diff --git a/tests/test_indexer_spider_response_parity.py b/tests/test_indexer_spider_response_parity.py index a1cde3f63..b958ad6bb 100644 --- a/tests/test_indexer_spider_response_parity.py +++ b/tests/test_indexer_spider_response_parity.py @@ -312,6 +312,52 @@ def test_missing_credentials_reject_sync_and_async_search_without_network( assert asyncio.run(spider.async_search(**search_kwargs)) == (True, []) +def test_yema_sync_async_entries_share_request_plan_and_result_finalizer( + monkeypatch, +): + """Yema 分类搜索的两种入口仅替换 HTTP,规划和归并决策必须同源。""" + spider = _build_api_spider("yema", monkeypatch) + prepare = [] + finalized = [] + original_prepare = spider._prepare_search_requests + + def prepare_search_requests(**kwargs): + """记录两种入口交给共享请求规划器的参数。""" + prepare.append(kwargs) + return original_prepare(**kwargs) + + def finalize_search_responses(responses): + """记录两种入口交给共享结果归并器的响应序列。""" + finalized.append(responses) + return False, [{"title": "projected"}] + + def sync_request(*_args, **kwargs): + """返回可按分类请求体识别的同步响应。""" + return ("sync", kwargs["json"]["categoryId"]) + + async def async_request(*_args, **kwargs): + """返回可按分类请求体识别的异步响应。""" + return ("async", kwargs["json"]["categoryId"]) + + monkeypatch.setattr(spider, "_prepare_search_requests", prepare_search_requests) + monkeypatch.setattr(spider, "_finalize_search_responses", finalize_search_responses) + monkeypatch.setattr(RequestUtils, "post_res", sync_request) + monkeypatch.setattr(AsyncRequestUtils, "post_res", async_request) + + sync_result = spider.search(keyword="Movie", mtype=MediaType.MUSIC, page=2) + async_result = asyncio.run( + spider.async_search(keyword="Movie", mtype=MediaType.MUSIC, page=2) + ) + + expected_plan = {"keyword": "Movie", "mtype": MediaType.MUSIC, "page": 2} + assert sync_result == async_result == (False, [{"title": "projected"}]) + assert prepare == [expected_plan, expected_plan] + assert finalized == [ + [("sync", 8), ("sync", 16)], + [("async", 8), ("async", 16)], + ] + + def test_tnode_and_torrentleech_preflight_failures_skip_network(monkeypatch): """令牌缺失和不支持中文的入口都应在同步、异步 HTTP 前终止。""" tnode = _build_api_spider("tnode", monkeypatch) diff --git a/tests/test_module_client_parity.py b/tests/test_module_client_parity.py new file mode 100644 index 000000000..35c41e7b2 --- /dev/null +++ b/tests/test_module_client_parity.py @@ -0,0 +1,167 @@ +import asyncio +from unittest.mock import AsyncMock, Mock + +import pytest + +from app.domain.context import MusicInfo +from app.domain.meta.metamusic import MetaMusic +from app.modules.acoustid import AcoustIdModule +from app.modules.anilist.anilist import AniListApi +from app.modules.bangumi.bangumi import BangumiApi +from app.modules.musicbrainz import MusicBrainzModule +from app.modules.theaudiodb import TheAudioDbModule +from app.schemas.types import MediaSource + + +def test_bangumi_sync_async_share_request_and_projection( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Bangumi 同步异步搜索与人物投影应使用同一请求和结果规则。""" + api = BangumiApi() + sync_invoke = Mock( + side_effect=[ + {"list": [{"id": 1}]}, + [{"id": 2, "name": "角色", "actors": [{"id": 3}]}], + ] + ) + async_invoke = AsyncMock( + side_effect=[ + {"list": [{"id": 1}]}, + [{"id": 2, "name": "角色", "actors": [{"id": 3}]}], + ] + ) + monkeypatch.setattr(api, "_BangumiApi__invoke", sync_invoke) + monkeypatch.setattr(api, "_BangumiApi__async_invoke", async_invoke) + + sync_search = api.search("葬送的芙莉莲") + async_search = asyncio.run(api.async_search("葬送的芙莉莲")) + sync_credits = api.credits(154587) + async_credits = asyncio.run(api.async_credits(154587)) + + assert sync_search == async_search == [{"id": 1}] + assert sync_credits == async_credits == [ + {"id": 3, "career": ["角色"]} + ] + assert sync_invoke.call_args_list[0].args == async_invoke.await_args_list[0].args + assert sync_invoke.call_args_list[1].args == async_invoke.await_args_list[1].args + assert sync_invoke.call_args_list[1].kwargs.keys() == async_invoke.await_args_list[1].kwargs.keys() + + +def test_anilist_sync_async_share_queries_and_projection( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """AniList 同步异步搜索与人物入口应共享查询文本、变量和投影。""" + api = AniListApi() + search_payload = {"Page": {"media": [{"id": 154587}]}} + person_payload = {"Staff": {"id": 95075, "name": {"native": "种崎敦美"}}} + sync_invoke = Mock(side_effect=[search_payload, person_payload]) + async_invoke = AsyncMock(side_effect=[search_payload, person_payload]) + monkeypatch.setattr(api, "_invoke", sync_invoke) + monkeypatch.setattr(api, "_async_invoke", async_invoke) + + sync_search = AniListApi.search.__wrapped__(api, "Frieren", 12) + async_search = asyncio.run( + AniListApi.async_search.__wrapped__(api, "Frieren", 12) + ) + sync_person = AniListApi.person_detail.__wrapped__(api, 95075) + async_person = asyncio.run( + AniListApi.async_person_detail.__wrapped__(api, 95075) + ) + + assert sync_search == async_search == [{"id": 154587}] + assert sync_person == async_person == person_payload["Staff"] + assert sync_invoke.call_args_list[0] == async_invoke.await_args_list[0] + assert sync_invoke.call_args_list[1] == async_invoke.await_args_list[1] + + +def test_musicbrainz_sync_async_share_detail_decision( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """MusicBrainz 显式详情识别应共享来源准入、参数和缓存收尾。""" + module = MusicBrainzModule() + module.cache = Mock() + meta = MetaMusic( + title="Yellow", + media_source=MediaSource.MusicBrainz, + media_id="recording-1", + ) + expected = MusicInfo( + media_source=MediaSource.MusicBrainz, + media_id="recording-1", + title="Yellow", + ) + sync_detail = Mock(return_value=expected) + async_detail = AsyncMock(return_value=expected) + monkeypatch.setattr(module, "recognize_music", sync_detail) + monkeypatch.setattr(module, "async_recognize_music", async_detail) + + sync_result = module.recognize_media(meta=meta, cache=False) + async_result = asyncio.run( + module.async_recognize_media(meta=meta, cache=False) + ) + + assert sync_result == async_result == expected + assert async_detail.await_args is not None + assert sync_detail.call_args.args == async_detail.await_args.args + assert sync_detail.call_args.kwargs == async_detail.await_args.kwargs == {} + assert module.cache.update.call_count == 2 + + +def test_theaudiodb_sync_async_share_request_plan_and_response( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """TheAudioDB 同步异步详情应共享请求计划、响应解析和实体选择。""" + module = TheAudioDbModule() + payload = { + "track": [ + { + "idTrack": "32793500", + "strTrack": "Yellow", + "strArtist": "Coldplay", + } + ] + } + sync_request = Mock(return_value=payload) + async_request = AsyncMock(return_value=payload) + monkeypatch.setattr(module, "_request_json", sync_request) + monkeypatch.setattr(module, "_async_request_json", async_request) + + sync_result = module.recognize_music( + MediaSource.TheAudioDB, "32793500", music_type="recording" + ) + async_result = asyncio.run( + module.async_recognize_music( + MediaSource.TheAudioDB, "32793500", music_type="recording" + ) + ) + + assert sync_result == async_result + assert sync_result is not None and sync_result.media_id == "32793500" + assert sync_request.call_args == async_request.await_args + plan = module._request_plan("client-key", "track.php", {"h": "32793500"}) + assert plan is not None + assert plan.url.endswith("/client-key/track.php") + assert plan.params == {"h": "32793500"} + + +def test_acoustid_lookup_plan_and_response_projection_are_transport_neutral() -> None: + """AcoustID 请求载荷和响应筛选不得依赖同步或异步传输实现。""" + recording_id = "38035858-f990-4fbb-b3b2-f2f8b958eeba" + plan = AcoustIdModule._lookup_plan(" client-key ", 243, "AQADtM...") + payload = { + "status": "ok", + "results": [ + {"score": 0.98, "recordings": [{"id": recording_id}]} + ], + } + + assert plan is not None + assert plan.data == { + "client": "client-key", + "duration": 243, + "fingerprint": "AQADtM...", + "meta": "recordingids", + "format": "json", + } + assert AcoustIdModule._project_lookup_response(200, payload) == recording_id + assert AcoustIdModule._project_lookup_response(503, payload) is None diff --git a/tests/test_search_sync_async_parity.py b/tests/test_search_sync_async_parity.py index 0fdb872bf..fb7219357 100644 --- a/tests/test_search_sync_async_parity.py +++ b/tests/test_search_sync_async_parity.py @@ -5,6 +5,8 @@ from copy import deepcopy from types import SimpleNamespace from typing import Any +import pytest + from app.chain.search import media as media_module from app.chain.search import title as title_module from app.chain.search.facade import SearchChain @@ -561,3 +563,95 @@ def test_title_stream_preserves_provider_and_result_order(): item["torrent_info"]["title"] for item in events[-1]["items"] ] == [first.title, second.title] assert events[-1]["candidate_items"] == 2 + + +def test_id_search_resolution_freezes_success_and_failure_effect_order(): + """ID 搜索状态机应唯一决定成功保存顺序与识别失败短路。""" + recognition_params = {"media_source": MediaSource.TMDB, "media_id": "100"} + cache_params = {**recognition_params, "sites": [1]} + contexts = [SimpleNamespace(name="context")] + + success = media_module._id_search_resolution( + recognition_params=recognition_params, + cache_params=cache_params, + season=2, + sites=[1], + area="title", + cache_local=True, + failure_keyword="tmdb:100", + ) + assert isinstance(next(success), media_module._IdSearchCacheRequest) + assert isinstance(success.send(None), media_module._IdSearchRecognizeRequest) + process_request = success.send(_media()) + assert isinstance(process_request, media_module._IdSearchProcessRequest) + assert list(process_request.params["no_exists"]["tmdb:100"]) == [2] + save_request = success.send(contexts) + assert isinstance(save_request, media_module._IdSearchSaveRequest) + assert save_request.contexts is contexts + with pytest.raises(StopIteration) as success_completed: + success.send(None) + assert success_completed.value.value.contexts is contexts + + failure = media_module._id_search_resolution( + recognition_params=recognition_params, + cache_params=cache_params, + season=None, + sites=None, + area="title", + cache_local=False, + failure_keyword="tmdb:100", + ) + assert isinstance(next(failure), media_module._IdSearchRecognizeRequest) + with pytest.raises(StopIteration) as failure_completed: + failure.send(None) + assert failure_completed.value.value.contexts == [] + assert failure_completed.value.value.warning == "tmdb:100 媒体信息识别失败!" + + +def test_title_search_resolution_owns_filter_short_circuit_and_save_order(): + """标题状态机应在过滤失败时短路,并仅保存成功投影。""" + search_params = {"keyword": "Movie", "sites": [1], "page": 0} + cache_params = {"keyword": "Movie", "area": "title"} + torrent = _torrent("Movie 2026") + contexts = [SimpleNamespace(name="context")] + + success = title_module._title_search_resolution( + title="Movie", + search_params=search_params, + cache_params=cache_params, + cache_local=True, + mtype=MediaType.MOVIE, + rule_groups=["quality"], + ) + assert isinstance(next(success), title_module._TitleSearchCacheRequest) + assert isinstance(success.send(None), title_module._TitleSearchProviderRequest) + resolve_request = success.send([torrent]) + assert isinstance(resolve_request, title_module._TitleSearchResolveRequest) + save_request = success.send( + title_module._TitleSearchResult(contexts=contexts) + ) + assert isinstance(save_request, title_module._TitleSearchSaveRequest) + assert save_request.contexts is contexts + with pytest.raises(StopIteration) as success_completed: + success.send(None) + assert success_completed.value.value.contexts is contexts + + failure = title_module._title_search_resolution( + title="Movie", + search_params=search_params, + cache_params=cache_params, + cache_local=True, + mtype=MediaType.MOVIE, + rule_groups=["quality"], + ) + next(failure) + failure.send(None) + failure.send([torrent]) + with pytest.raises(StopIteration) as completed: + failure.send( + title_module._TitleSearchResult( + contexts=[], + warning="Movie 没有符合过滤规则的资源", + ) + ) + assert completed.value.value.contexts == [] diff --git a/tests/test_server_report_service.py b/tests/test_server_report_service.py index 49bfb01f1..386488f22 100644 --- a/tests/test_server_report_service.py +++ b/tests/test_server_report_service.py @@ -1,5 +1,5 @@ from types import SimpleNamespace -from unittest.mock import AsyncMock, Mock +from unittest.mock import AsyncMock, Mock, call import pytest @@ -121,3 +121,56 @@ def test_plugin_report_sanitizes_explicit_sources_before_transport(): "plugin_id": "Demo", "repo_url": "local://Demo", }]) + + +@pytest.mark.asyncio +async def test_plugin_report_entries_share_payload_and_response_decisions(monkeypatch): + """插件同步、异步入口只替换发送端口,准入和结果判定必须同源。""" + sync_sender = Mock(return_value=SimpleNamespace(status_code=200)) + async_sender = AsyncMock(return_value=SimpleNamespace(status_code=200)) + service = _service( + plugin_report_sender=sync_sender, + async_plugin_report_sender=async_sender, + ) + prepare = Mock(wraps=service._prepare_plugin_report) + succeeded = Mock(wraps=service._report_succeeded) + monkeypatch.setattr(service, "_prepare_plugin_report", prepare) + monkeypatch.setattr(service, "_report_succeeded", succeeded) + items = [("Demo", "https://repo.example/Demo")] + + assert service.report_plugins(enabled=True, items=items) is True + assert await service.async_report_plugins(enabled=True, items=items) is True + + assert prepare.call_args_list == [ + call(enabled=True, items=items), + call(enabled=True, items=items), + ] + assert sync_sender.call_args == async_sender.await_args + assert succeeded.call_count == 2 + + +@pytest.mark.asyncio +async def test_subscribe_report_entries_share_payload_projection(monkeypatch): + """订阅读取后的过滤与终态决策必须由同步、异步入口共享。""" + subscribe = SimpleNamespace(to_dict=lambda: { + "name": "Demo", + "media_source": "themoviedb", + "media_id": "123", + }) + service = _service( + subscribes_provider=Mock(return_value=[subscribe]), + async_subscribes_provider=AsyncMock(return_value=[subscribe]), + async_subscribe_report_sender=AsyncMock( + return_value=SimpleNamespace(status_code=200) + ), + ) + prepare = Mock(wraps=service._prepare_subscribe_report) + monkeypatch.setattr(service, "_prepare_subscribe_report", prepare) + + assert service.report_subscribes(enabled=True) is True + assert await service.async_report_subscribes(enabled=True) is True + + assert prepare.call_args_list == [(([subscribe],), {}), (([subscribe],), {})] + assert service._subscribe_report_sender.call_args == ( + service._async_subscribe_report_sender.await_args + ) diff --git a/tests/test_tvdb_sync_async_parity.py b/tests/test_tvdb_sync_async_parity.py new file mode 100644 index 000000000..a77239b98 --- /dev/null +++ b/tests/test_tvdb_sync_async_parity.py @@ -0,0 +1,151 @@ +"""TVDB 附加信息同步异步共享决策回归测试。""" + +import asyncio +from typing import Any + +from app.domain.context import MediaInfo +from app.modules import thetvdb as tvdb_module +from app.modules.thetvdb import TheTvDbModule +from app.schemas.types import MediaSource, MediaType + + +def _module() -> TheTvDbModule: + """绕过模块运行时构造可局部注入的 TVDB provider。""" + return object.__new__(TheTvDbModule) + + +def _media( + media_source: MediaSource = MediaSource.Douban, + media_id: str = "100", +) -> MediaInfo: + """构造用于 TVDB 别名匹配的电视剧媒体。""" + return MediaInfo( + media_source=media_source, + media_id=media_id, + type=MediaType.TV, + title="测试剧集", + year="2026", + names=["Test Series"], + ) + + +def test_tvdb_sync_async_share_title_lookup_and_candidate_resolution( + monkeypatch, +) -> None: + """标题查询双入口应只切换 I/O 外壳并复用同一候选解析。""" + module = _module() + calls: list[tuple[str, object]] = [] + threaded: list[str] = [] + candidates = [ + {"type": "series", "tvdb_id": "series-404", "name": "Other"}, + { + "type": "series", + "tvdb_id": "series-123", + "name": "测试剧集", + "year": "2026", + "translations": [{"name": "Test Series"}], + }, + ] + + def search_tvdb(title: str) -> list[dict[str, Any]]: + """记录阻塞标题查询并返回稳定候选。""" + calls.append(("search_tvdb", title)) + return candidates + + async def run_blocking_call(function, *args, **kwargs): + """记录异步入口实际提交到线程池的函数边界。""" + threaded.append(function.__name__) + return function(*args, **kwargs) + + module.search_tvdb = search_tvdb + monkeypatch.setattr(tvdb_module, "run_in_threadpool", run_blocking_call) + + sync_result = module.get_media_auxiliary_info( + _media(), media_source=(MediaSource.TVDB,) + ) + async_result = asyncio.run( + module.async_get_media_auxiliary_info( + _media(), media_source=(MediaSource.TVDB,) + ) + ) + + assert [item.to_dict() for item in sync_result] == [ + item.to_dict() for item in async_result + ] + assert calls == [ + ("search_tvdb", "测试剧集"), + ("search_tvdb", "测试剧集"), + ] + assert threaded == ["search_tvdb"] + assert sync_result[0].media_id == "123" + assert sync_result[0].names == ["测试剧集", "Test Series"] + + +def test_tvdb_sync_async_share_native_id_lookup(monkeypatch) -> None: + """TVDB 原生身份双入口应选择相同详情查询,线程池不包业务方法。""" + module = _module() + calls: list[int] = [] + threaded: list[str] = [] + info = { + "id": 123, + "name": "测试剧集", + "year": "2026", + "aliases": ["Test Series"], + } + + def tvdb_info(tvdb_id: int) -> dict[str, Any]: + """记录阻塞详情查询。""" + calls.append(tvdb_id) + return info + + async def run_blocking_call(function, *args, **kwargs): + """验证异步入口只提交详情 I/O。""" + threaded.append(function.__name__) + return function(*args, **kwargs) + + module.tvdb_info = tvdb_info + monkeypatch.setattr(tvdb_module, "run_in_threadpool", run_blocking_call) + media = _media(media_source=MediaSource.TVDB, media_id="123") + + sync_result = module.get_media_auxiliary_info( + media, media_source=(MediaSource.TVDB,) + ) + async_result = asyncio.run( + module.async_get_media_auxiliary_info( + media, media_source=(MediaSource.TVDB,) + ) + ) + + assert [item.to_dict() for item in sync_result] == [ + item.to_dict() for item in async_result + ] + assert calls == [123, 123] + assert threaded == ["tvdb_info"] + + +def test_tvdb_disabled_source_short_circuits_before_sync_and_async_io( + monkeypatch, +) -> None: + """未启用 TVDB 时双入口都不得触发网络或线程池。""" + module = _module() + + def unexpected_io(*_args, **_kwargs): + """标记禁用来源错误触发了同步 I/O。""" + raise AssertionError("禁用 TVDB 时不应触发 I/O") + + async def unexpected_threadpool(*_args, **_kwargs): + """标记禁用来源错误触发了线程池。""" + raise AssertionError("禁用 TVDB 时不应触发线程池") + + module.search_tvdb = unexpected_io + module.tvdb_info = unexpected_io + monkeypatch.setattr(tvdb_module, "run_in_threadpool", unexpected_threadpool) + + assert module.get_media_auxiliary_info( + _media(), media_source=(MediaSource.TMDB,) + ) == [] + assert asyncio.run( + module.async_get_media_auxiliary_info( + _media(), media_source=(MediaSource.TMDB,) + ) + ) == []