From 6504cb36f5446b522326512f0bfd169f94fec0ab Mon Sep 17 00:00:00 2001 From: jxxghp Date: Sun, 16 Aug 2026 05:36:40 +0800 Subject: [PATCH] =?UTF-8?q?fix(recognize):=20=E4=BF=AE=E5=A4=8D=E8=AF=86?= =?UTF-8?q?=E5=88=AB=E6=B5=8B=E8=AF=95=E4=BB=85=E4=BC=A0=E6=95=B0=E6=8D=AE?= =?UTF-8?q?=E6=BA=90=E6=97=B6=E8=A2=AB=E6=88=90=E5=AF=B9=E8=BA=AB=E4=BB=BD?= =?UTF-8?q?=E6=A0=A1=E9=AA=8C=E8=AF=AF=E6=8B=92=EF=BC=8C=E5=90=8D=E7=A7=B0?= =?UTF-8?q?=E8=AF=86=E5=88=AB=E4=B8=8E=E4=B8=B4=E6=97=B6=E8=AF=86=E5=88=AB?= =?UTF-8?q?=E8=AF=8D=E6=81=A2=E5=A4=8D=E7=94=9F=E6=95=88?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 8bf2f601 引入的显式身份成对校验把'仅传 media_source'也视为显式身份,识别测试恒传数据源导致直接返回未识别 - 成对约束仅在显式传 media_id 时生效;仅传数据源时保留为请求级识别源约束,按名称在指定数据源内识别 - 字符串来源经 normalize_media_source 规范,meta 自带同源身份(如 {tmdbid=})时优先按身份识别 - 新增 tests/test_recognize_source_selection.py 覆盖 source-only、同源身份采纳与成对拒绝场景 --- app/chain/__init__.py | 22 +++- tests/test_recognize_source_selection.py | 146 +++++++++++++++++++++++ 2 files changed, 165 insertions(+), 3 deletions(-) create mode 100644 tests/test_recognize_source_selection.py diff --git a/app/chain/__init__.py b/app/chain/__init__.py index b4411aebe..5cb69f08b 100644 --- a/app/chain/__init__.py +++ b/app/chain/__init__.py @@ -40,7 +40,7 @@ from app.schemas import ( MessageResponse, ) from app.foundation.identity import normalize_internal_user_id -from app.schemas.media import resolve_media_identity +from app.schemas.media import normalize_media_source, resolve_media_identity from app.schemas.message import ChannelCapability, ChannelCapabilityManager from app.schemas.category import CategoryConfig from app.schemas.types import ( @@ -626,7 +626,9 @@ class ChainBase(metaclass=ABCMeta): :param music_type: 音乐实体类型,显式音乐 ID 必须据此区分单曲与专辑 :return: 识别的媒体信息,包括剧集信息 """ - explicit_identity = media_source is not None or media_id is not None + # 仅传数据源是请求级识别源约束(按名称识别限定数据源),显式 media_id 才要求来源成对 + explicit_identity = media_id is not None + requested_source = normalize_media_source(media_source) or media_source media_source, media_id = resolve_media_identity( media=meta, media_source=media_source, @@ -635,6 +637,12 @@ class ChainBase(metaclass=ABCMeta): if explicit_identity and (not media_source or not media_id): logger.warning("媒体识别需要同时提供有效的 media_source 和 media_id") return None + if not media_id and requested_source is not None: + media_source = requested_source + # meta 自带同源身份(如 {tmdbid=} 标题)时直接按身份识别,避免退化为名称搜索 + meta_source, meta_id = resolve_media_identity(media=meta) + if meta_id and meta_source == requested_source: + media_source, media_id = meta_source, meta_id if not episode_group and hasattr(meta, "episode_group"): episode_group = meta.episode_group if not mtype and not (media_source and media_id) and meta and meta.type in [ @@ -736,7 +744,9 @@ class ChainBase(metaclass=ABCMeta): :param music_type: 音乐实体类型,显式音乐 ID 必须据此区分单曲与专辑 :return: 识别的媒体信息,包括剧集信息 """ - explicit_identity = media_source is not None or media_id is not None + # 仅传数据源是请求级识别源约束(按名称识别限定数据源),显式 media_id 才要求来源成对 + explicit_identity = media_id is not None + requested_source = normalize_media_source(media_source) or media_source media_source, media_id = resolve_media_identity( media=meta, media_source=media_source, @@ -745,6 +755,12 @@ class ChainBase(metaclass=ABCMeta): if explicit_identity and (not media_source or not media_id): logger.warning("媒体识别需要同时提供有效的 media_source 和 media_id") return None + if not media_id and requested_source is not None: + media_source = requested_source + # meta 自带同源身份(如 {tmdbid=} 标题)时直接按身份识别,避免退化为名称搜索 + meta_source, meta_id = resolve_media_identity(media=meta) + if meta_id and meta_source == requested_source: + media_source, media_id = meta_source, meta_id if not episode_group and hasattr(meta, "episode_group"): episode_group = meta.episode_group if not mtype and not (media_source and media_id) and meta and meta.type in [ diff --git a/tests/test_recognize_source_selection.py b/tests/test_recognize_source_selection.py new file mode 100644 index 000000000..3c0db6375 --- /dev/null +++ b/tests/test_recognize_source_selection.py @@ -0,0 +1,146 @@ +# -*- coding: utf-8 -*- +"""识别测试类请求的数据源选择回归测试。 + +识别测试(名称测试)前端只传请求级数据源 media_source,不携带 media_id; +该路径必须按名称在指定数据源内识别,而不是被"显式身份必须成对"规则拦截。 +""" +import asyncio +from unittest.mock import AsyncMock, Mock, patch + +from app.chain import ChainBase +from app.domain.context import MediaInfo +from app.domain.metainfo import MetaInfo +from app.schemas.types import MediaSource, MediaType + +# 用户反馈的识别测试失败样例 +FAILING_TITLE = "[ANI] 關於我轉生變成史萊姆這檔事 第四季 - 90 [1080P][Baha][WEB-DL][AAC AVC][CHT].mp4" + + +def _tmdb_media() -> MediaInfo: + """构造带 TMDB 身份的识别结果。""" + return MediaInfo( + media_source=MediaSource.TMDB, + media_id="120089", + tmdb_id=120089, + title="关于我转生变成史莱姆这档事", + type=MediaType.TV, + ) + + +def test_recognize_media_with_source_only_uses_name_search(): + """仅传 media_source 时应按名称在指定数据源识别,不得直接拒绝。""" + chain = ChainBase() + meta = MetaInfo(FAILING_TITLE) + captured = {} + + def fake_native(module_kwargs, cache): + captured.update(module_kwargs) + return _tmdb_media() + + with patch.object(chain, "_run_native_media_recognize", side_effect=fake_native), \ + patch.object(chain, "_supplement_media_recognize", side_effect=lambda **kw: kw["mediainfo"]), \ + patch("app.chain.MoviePilotServerHelper"): + result = chain.recognize_media(meta=meta, media_source=MediaSource.TMDB, cache=False) + + assert result is not None + assert result.tmdb_id == 120089 + # 数据源约束透传到模块层,且无显式 ID + assert captured["media_source"] == MediaSource.TMDB + assert captured["media_id"] is None + # 名称识别应沿用 meta 推断的类型 + assert captured["mtype"] == MediaType.TV + + +def test_async_recognize_media_with_source_only_uses_name_search(): + """异步入口与同步入口保持同一语义。""" + chain = ChainBase() + meta = MetaInfo(FAILING_TITLE) + captured = {} + + async def fake_native(module_kwargs, cache): + captured.update(module_kwargs) + return _tmdb_media() + + async def fake_supplement(**kwargs): + return kwargs["mediainfo"] + + # 异步路径上报共享识别为协程,需要可 await 的桩 + helper = Mock() + helper.async_report_recognize_share = AsyncMock() + + with patch.object(chain, "_async_run_native_media_recognize", side_effect=fake_native), \ + patch.object(chain, "_async_supplement_media_recognize", side_effect=fake_supplement), \ + patch("app.chain.MoviePilotServerHelper", helper): + result = asyncio.run( + chain.async_recognize_media(meta=meta, media_source=MediaSource.TMDB, cache=False) + ) + + assert result is not None + assert result.tmdb_id == 120089 + assert captured["media_source"] == MediaSource.TMDB + assert captured["media_id"] is None + + +def test_recognize_media_accepts_string_source_only(): + """第三方客户端可能传字符串数据源,应规范为枚举后按名称识别。""" + chain = ChainBase() + meta = MetaInfo(FAILING_TITLE) + captured = {} + + def fake_native(module_kwargs, cache): + captured.update(module_kwargs) + return _tmdb_media() + + with patch.object(chain, "_run_native_media_recognize", side_effect=fake_native), \ + patch.object(chain, "_supplement_media_recognize", side_effect=lambda **kw: kw["mediainfo"]), \ + patch("app.chain.MoviePilotServerHelper"): + result = chain.recognize_media(meta=meta, media_source="themoviedb", cache=False) + + assert result is not None + assert captured["media_source"] == MediaSource.TMDB + + +def test_recognize_media_meta_identity_same_source_uses_id(): + """仅传数据源且 meta 自带同源身份(如 {tmdbid=})时应直接按身份识别。""" + chain = ChainBase() + meta = MetaInfo("空之境界 第五章 矛盾螺旋 (2008) {tmdbid=23155}") + captured = {} + + def fake_native(module_kwargs, cache): + captured.update(module_kwargs) + return _tmdb_media() + + with patch.object(chain, "_run_native_media_recognize", side_effect=fake_native), \ + patch.object(chain, "_supplement_media_recognize", side_effect=lambda **kw: kw["mediainfo"]), \ + patch("app.chain.MoviePilotServerHelper"): + result = chain.recognize_media(meta=meta, media_source=MediaSource.TMDB, cache=False) + + assert result is not None + assert captured["media_source"] == MediaSource.TMDB + assert captured["media_id"] == "23155" + + +def test_recognize_media_media_id_without_source_still_rejected(): + """显式 media_id 缺少有效来源时仍应拒绝,成对约束不放宽。""" + chain = ChainBase() + native = Mock() + + with patch.object(chain, "_run_native_media_recognize", native): + result = chain.recognize_media(meta=MetaInfo("任意标题"), media_id="12345") + + assert result is None + native.assert_not_called() + + +def test_recognize_media_invalid_pair_still_rejected(): + """media_id 为 0 等无效值与来源组合仍应拒绝。""" + chain = ChainBase() + native = Mock() + + with patch.object(chain, "_run_native_media_recognize", native): + result = chain.recognize_media( + meta=MetaInfo("任意标题"), media_source=MediaSource.TMDB, media_id="0" + ) + + assert result is None + native.assert_not_called()