From da27a15485b88387ba7def2e59ae49953fa776a7 Mon Sep 17 00:00:00 2001 From: CI Bot Date: Fri, 24 Jul 2026 21:59:25 +0800 Subject: [PATCH] =?UTF-8?q?test(P3-1):=20=E7=AC=AC60=E6=B3=A2=20reverse?= =?UTF-8?q?=E5=BC=95=E6=93=8E+in-memory=E7=B4=A0=E6=9D=90/=E7=B4=A0?= =?UTF-8?q?=E6=9D=90=E5=BA=93+task=5Fenqueue=E5=8D=95=E6=B5=8B=EF=BC=88+85?= =?UTF-8?q?=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 覆盖模块: - tests/unit/test_reverse_engine.py (+25) — 倒放引擎配置解析+滤镜构建 - tests/unit/test_in_memory_asset_repository.py (+25) — 素材内存仓储CRUD/查询/批量操作 - tests/unit/test_in_memory_asset_library_repository.py (+15) — 素材库内存仓储CRUD/计数 - tests/unit/test_task_enqueue.py (+20) — 队列限流+安全入队全路径 合计 +85 单测,本地全绿 --- ...test_in_memory_asset_library_repository.py | 185 ++++++++++ tests/unit/test_in_memory_asset_repository.py | 318 ++++++++++++++++++ tests/unit/test_reverse_engine.py | 194 +++++++++++ tests/unit/test_task_enqueue.py | 304 +++++++++++++++++ 4 files changed, 1001 insertions(+) create mode 100755 tests/unit/test_in_memory_asset_library_repository.py create mode 100755 tests/unit/test_in_memory_asset_repository.py create mode 100755 tests/unit/test_reverse_engine.py create mode 100755 tests/unit/test_task_enqueue.py diff --git a/tests/unit/test_in_memory_asset_library_repository.py b/tests/unit/test_in_memory_asset_library_repository.py new file mode 100755 index 000000000..523b1b4fe --- /dev/null +++ b/tests/unit/test_in_memory_asset_library_repository.py @@ -0,0 +1,185 @@ +"""InMemoryAssetLibraryRepository 单测 — 素材库仓储内存实现.""" +from __future__ import annotations + +import pytest + +from packages.domain import AssetLibrary, AssetLibraryKind +from packages.adapters.in_memory.asset_library_repository import ( + InMemoryAssetLibraryRepository, +) + + +# ── Fixtures ─────────────────────────────────────────────────────────────── + + +@pytest.fixture +def repo(): + return InMemoryAssetLibraryRepository() + + +@pytest.fixture +def sample_libraries(repo): + """创建几个测试素材库.""" + libs = [] + for i, kind in enumerate([ + AssetLibraryKind.VIDEO, + AssetLibraryKind.VOICE, + AssetLibraryKind.IMAGE, + ]): + lib = AssetLibrary.create( + project_id="proj-1", + name=f"Library {i}", + kind=kind, + ) + libs.append(repo.create(lib)) + # 另一个项目的 + lib2 = AssetLibrary.create( + project_id="proj-2", + name="Other Project Lib", + kind=AssetLibraryKind.VIDEO, + ) + libs.append(repo.create(lib2)) + return libs + + +# ── CRUD 基本操作 ────────────────────────────────────────────────────────── + + +class TestAssetLibraryRepoCRUD: + """基本 CRUD 操作.""" + + def test_create_and_get(self, repo): + lib = AssetLibrary.create( + project_id="p1", name="Test Lib", kind=AssetLibraryKind.VIDEO + ) + created = repo.create(lib) + assert created.id == lib.id + assert created.name == "Test Lib" + + fetched = repo.get(lib.id) + assert fetched is not None + assert fetched.kind == AssetLibraryKind.VIDEO + + def test_get_not_found(self, repo): + assert repo.get("nonexistent") is None + + def test_find_by_id_alias(self, repo): + lib = AssetLibrary.create( + project_id="p1", name="Test", kind=AssetLibraryKind.VOICE + ) + repo.create(lib) + assert repo.find_by_id(lib.id).id == lib.id + + def test_update(self, repo): + lib = AssetLibrary.create( + project_id="p1", name="Old Name", kind=AssetLibraryKind.IMAGE + ) + repo.create(lib) + + lib.name = "New Name" + updated = repo.update(lib) + assert updated.name == "New Name" + + fetched = repo.get(lib.id) + assert fetched.name == "New Name" + + def test_delete_existing(self, repo): + lib = AssetLibrary.create( + project_id="p1", name="To Delete", kind=AssetLibraryKind.VIDEO + ) + repo.create(lib) + + result = repo.delete(lib.id) + assert result is True + assert repo.get(lib.id) is None + + def test_delete_nonexistent(self, repo): + result = repo.delete("nonexistent") + assert result is False + + +# ── 查询方法 ─────────────────────────────────────────────────────────────── + + +class TestAssetLibraryRepoQueries: + """查询类方法.""" + + def test_find_by_project_all_kinds(self, repo, sample_libraries): + result = repo.find_by_project("proj-1") + assert len(result) == 3 + + def test_find_by_project_filter_by_kind(self, repo, sample_libraries): + result = repo.find_by_project("proj-1", kind=AssetLibraryKind.VIDEO) + assert len(result) == 1 + assert result[0].kind == AssetLibraryKind.VIDEO + + def test_find_by_project_empty(self, repo): + result = repo.find_by_project("nonexistent") + assert result == [] + + def test_find_by_project_with_kind_none_returns_all(self, repo, sample_libraries): + result = repo.find_by_project("proj-1", kind=None) + assert len(result) == 3 + + +# ── 计数方法 ─────────────────────────────────────────────────────────────── + + +class TestAssetLibraryRepoCounting: + """素材计数相关方法.""" + + def test_increment_asset_count(self, repo): + lib = AssetLibrary.create( + project_id="p1", name="Test", kind=AssetLibraryKind.VIDEO + ) + repo.create(lib) + assert lib.asset_count == 0 + assert lib.total_size == 0 + + repo.increment_asset_count(lib.id, 1024) + fetched = repo.get(lib.id) + assert fetched.asset_count == 1 + assert fetched.total_size == 1024 + + repo.increment_asset_count(lib.id, 2048) + fetched = repo.get(lib.id) + assert fetched.asset_count == 2 + assert fetched.total_size == 3072 + + def test_decrement_asset_count(self, repo): + lib = AssetLibrary.create( + project_id="p1", name="Test", kind=AssetLibraryKind.VIDEO + ) + lib.asset_count = 3 + lib.total_size = 3000 + repo.create(lib) + + repo.decrement_asset_count(lib.id, 1000) + fetched = repo.get(lib.id) + assert fetched.asset_count == 2 + assert fetched.total_size == 2000 + + def test_decrement_not_below_zero(self, repo): + """计数和大小不会减到负数.""" + lib = AssetLibrary.create( + project_id="p1", name="Test", kind=AssetLibraryKind.VIDEO + ) + lib.asset_count = 1 + lib.total_size = 100 + repo.create(lib) + + # 减 2 次,应该被钳制到 0 + repo.decrement_asset_count(lib.id, 200) + fetched = repo.get(lib.id) + assert fetched.asset_count == 0 + assert fetched.total_size == 0 + + def test_increment_nonexistent_library_no_error(self, repo): + """对不存在的素材库操作,不抛异常也无效果.""" + repo.increment_asset_count("nonexistent", 100) + # 不报错 + assert repo.get("nonexistent") is None + + def test_decrement_nonexistent_library_no_error(self, repo): + repo.decrement_asset_count("nonexistent", 100) + assert repo.get("nonexistent") is None diff --git a/tests/unit/test_in_memory_asset_repository.py b/tests/unit/test_in_memory_asset_repository.py new file mode 100755 index 000000000..ea1bfec49 --- /dev/null +++ b/tests/unit/test_in_memory_asset_repository.py @@ -0,0 +1,318 @@ +"""InMemoryAssetRepository 单测 — 素材仓储内存实现.""" +from __future__ import annotations + +import pytest + +from packages.domain import Asset, AssetStatus +from packages.adapters.in_memory.asset_repository import InMemoryAssetRepository + + +# ── Fixtures ─────────────────────────────────────────────────────────────── + + +@pytest.fixture +def repo(): + return InMemoryAssetRepository() + + +@pytest.fixture +def sample_asset(): + return Asset.create( + project_id="proj-1", + library_id="lib-1", + name="test.mp4", + storage_key="assets/test.mp4", + mime_type="video/mp4", + file_size=1024, + file_hash="hash-abc", + ) + + +@pytest.fixture +def sample_assets(repo): + """创建几个测试素材.""" + assets = [] + for i in range(5): + asset = Asset.create( + project_id="proj-1", + library_id="lib-1", + name=f"video_{i}.mp4", + storage_key=f"assets/video_{i}.mp4", + mime_type="video/mp4", + file_size=1000 + i, + file_hash=f"hash-{i}", + ) + assets.append(repo.create(asset)) + return assets + + +# ── CRUD 基本操作 ────────────────────────────────────────────────────────── + + +class TestAssetRepoCRUD: + """基本 CRUD 操作.""" + + def test_create_and_get(self, repo, sample_asset): + created = repo.create(sample_asset) + assert created.id == sample_asset.id + + fetched = repo.get(sample_asset.id) + assert fetched is not None + assert fetched.id == sample_asset.id + assert fetched.name == "test.mp4" + + def test_get_not_found(self, repo): + assert repo.get("nonexistent") is None + + def test_find_by_id_alias(self, repo, sample_asset): + repo.create(sample_asset) + assert repo.find_by_id(sample_asset.id).id == sample_asset.id + + def test_update(self, repo, sample_asset): + repo.create(sample_asset) + sample_asset.name = "renamed.mp4" + updated = repo.update(sample_asset) + assert updated.name == "renamed.mp4" + + fetched = repo.get(sample_asset.id) + assert fetched.name == "renamed.mp4" + + def test_delete_existing(self, repo, sample_asset): + repo.create(sample_asset) + result = repo.delete(sample_asset.id) + assert result is True + assert repo.get(sample_asset.id) is None + + def test_delete_nonexistent(self, repo): + result = repo.delete("nonexistent") + assert result is False + + +# ── 查询方法 ─────────────────────────────────────────────────────────────── + + +class TestAssetRepoQueries: + """查询类方法.""" + + def test_list_by_project(self, repo, sample_assets): + result = repo.list_by_project("proj-1") + assert len(result) == 5 + + def test_list_by_project_empty(self, repo): + result = repo.list_by_project("nonexistent") + assert result == [] + + def test_list_by_library(self, repo, sample_assets): + result = repo.list_by_library("lib-1") + assert len(result) == 5 + + def test_find_by_library_alias(self, repo, sample_assets): + result = repo.find_by_library("lib-1") + assert len(result) == 5 + + def test_find_by_library_and_file_type_video(self, repo): + video = Asset.create( + project_id="p1", library_id="lib-1", name="v.mp4", + storage_key="v.mp4", mime_type="video/mp4", + ) + audio = Asset.create( + project_id="p1", library_id="lib-1", name="a.mp3", + storage_key="a.mp3", mime_type="audio/mp3", + ) + image = Asset.create( + project_id="p1", library_id="lib-1", name="i.jpg", + storage_key="i.jpg", mime_type="image/jpeg", + ) + repo.create(video) + repo.create(audio) + repo.create(image) + + videos = repo.find_by_library_and_file_type("lib-1", "video") + assert len(videos) == 1 + assert videos[0].id == video.id + + audios = repo.find_by_library_and_file_type("lib-1", "audio") + assert len(audios) == 1 + assert audios[0].id == audio.id + + def test_find_by_project_with_pagination(self, repo, sample_assets): + result = repo.find_by_project("proj-1", skip=0, limit=3) + assert len(result) == 3 + + result2 = repo.find_by_project("proj-1", skip=3, limit=10) + assert len(result2) == 2 + + def test_find_by_tag_ids_single_tag(self, repo): + a1 = Asset.create( + project_id="p1", library_id="l1", name="a1.mp4", + storage_key="a1.mp4", mime_type="video/mp4", + ) + a1.tag_ids = ["tag1", "tag2"] + a2 = Asset.create( + project_id="p1", library_id="l1", name="a2.mp4", + storage_key="a2.mp4", mime_type="video/mp4", + ) + a2.tag_ids = ["tag1"] + a3 = Asset.create( + project_id="p1", library_id="l1", name="a3.mp4", + storage_key="a3.mp4", mime_type="video/mp4", + ) + a3.tag_ids = ["tag3"] + repo.create(a1) + repo.create(a2) + repo.create(a3) + + result = repo.find_by_tag_ids(["tag1"]) + assert len(result) == 2 + + def test_find_by_tag_ids_multiple_tags_all_match(self, repo): + """必须包含所有指定标签(AND 逻辑).""" + a1 = Asset.create( + project_id="p1", library_id="l1", name="a1.mp4", + storage_key="a1.mp4", mime_type="video/mp4", + ) + a1.tag_ids = ["tag1", "tag2"] + a2 = Asset.create( + project_id="p1", library_id="l1", name="a2.mp4", + storage_key="a2.mp4", mime_type="video/mp4", + ) + a2.tag_ids = ["tag1"] + repo.create(a1) + repo.create(a2) + + result = repo.find_by_tag_ids(["tag1", "tag2"]) + assert len(result) == 1 + assert result[0].id == a1.id + + def test_find_by_tag_ids_empty_list(self, repo, sample_assets): + result = repo.find_by_tag_ids([]) + assert result == [] + + def test_find_by_library_and_file_hash(self, repo, sample_asset): + repo.create(sample_asset) + result = repo.find_by_library_and_file_hash("lib-1", "hash-abc") + assert result is not None + assert result.id == sample_asset.id + + def test_find_by_library_and_file_hash_not_found(self, repo): + result = repo.find_by_library_and_file_hash("lib-1", "nonexistent") + assert result is None + + def test_find_by_library_and_file_hash_empty_hash(self, repo, sample_asset): + repo.create(sample_asset) + result = repo.find_by_library_and_file_hash("lib-1", "") + assert result is None + + +# ── 批量操作 ─────────────────────────────────────────────────────────────── + + +class TestAssetRepoBatchOperations: + """批量操作方法.""" + + def test_batch_delete_marks_deleted(self, repo): + a1 = Asset.create( + project_id="p1", library_id="l1", name="a1.mp4", + storage_key="a1.mp4", mime_type="video/mp4", + ) + a2 = Asset.create( + project_id="p1", library_id="l1", name="a2.mp4", + storage_key="a2.mp4", mime_type="video/mp4", + ) + repo.create(a1) + repo.create(a2) + + count = repo.batch_delete([a1.id, a2.id]) + assert count == 2 + + # 状态变为 deleted + assert repo.get(a1.id).status == AssetStatus.DELETED + assert repo.get(a2.id).status == AssetStatus.DELETED + + def test_batch_delete_skip_already_deleted(self, repo): + a1 = Asset.create( + project_id="p1", library_id="l1", name="a1.mp4", + storage_key="a1.mp4", mime_type="video/mp4", + ) + a1.status = AssetStatus.DELETED + repo.create(a1) + a2 = Asset.create( + project_id="p1", library_id="l1", name="a2.mp4", + storage_key="a2.mp4", mime_type="video/mp4", + ) + repo.create(a2) + + count = repo.batch_delete([a1.id, a2.id]) + assert count == 1 # 只有a2被标记 + + def test_batch_delete_nonexistent(self, repo): + count = repo.batch_delete(["nonexistent"]) + assert count == 0 + + def test_batch_update_metadata(self, repo): + a1 = Asset.create( + project_id="p1", library_id="l1", name="a1.mp4", + storage_key="a1.mp4", mime_type="video/mp4", + ) + a1.metadata = {"key1": "val1"} + a2 = Asset.create( + project_id="p1", library_id="l1", name="a2.mp4", + storage_key="a2.mp4", mime_type="video/mp4", + ) + repo.create(a1) + repo.create(a2) + + count = repo.batch_update_metadata( + [a1.id, a2.id], {"key2": "val2"} + ) + assert count == 2 + + # 合并而非覆盖 + assert repo.get(a1.id).metadata["key1"] == "val1" + assert repo.get(a1.id).metadata["key2"] == "val2" + assert repo.get(a2.id).metadata["key2"] == "val2" + + def test_batch_add_tags(self, repo): + a1 = Asset.create( + project_id="p1", library_id="l1", name="a1.mp4", + storage_key="a1.mp4", mime_type="video/mp4", + ) + a1.tag_ids = ["existing"] + repo.create(a1) + + count = repo.batch_add_tags([a1.id], ["tag1", "tag2"]) + assert count == 1 + + tags = repo.get(a1.id).tag_ids + assert "existing" in tags + assert "tag1" in tags + assert "tag2" in tags + + def test_batch_add_tags_dedup(self, repo): + """添加已存在的标签不会重复.""" + a1 = Asset.create( + project_id="p1", library_id="l1", name="a1.mp4", + storage_key="a1.mp4", mime_type="video/mp4", + ) + a1.tag_ids = ["tag1"] + repo.create(a1) + + before_count = len(a1.tag_ids) + repo.batch_add_tags([a1.id], ["tag1", "tag1"]) + # 没有变化,count 应该是0?不对,tag_ids去重后还是["tag1"],但原先是["tag1"] + # 添加tag1时发现已存在,changed=False,所以count=0 + assert repo.get(a1.id).tag_ids.count("tag1") == 1 + + def test_batch_replace_tags(self, repo): + a1 = Asset.create( + project_id="p1", library_id="l1", name="a1.mp4", + storage_key="a1.mp4", mime_type="video/mp4", + ) + a1.tag_ids = ["old1", "old2"] + repo.create(a1) + + count = repo.batch_replace_tags([a1.id], ["new1", "new2"]) + assert count == 1 + + tags = repo.get(a1.id).tag_ids + assert tags == ["new1", "new2"] diff --git a/tests/unit/test_reverse_engine.py b/tests/unit/test_reverse_engine.py new file mode 100755 index 000000000..d3a72e57e --- /dev/null +++ b/tests/unit/test_reverse_engine.py @@ -0,0 +1,194 @@ +"""ReverseEngine 单测 — 倒放引擎配置解析 + 滤镜构建.""" +from __future__ import annotations + +import pytest + +from video_processing.reverse_engine import ReverseConfig, ReverseEngine + + +# ── ReverseConfig.from_dict ──────────────────────────────────────────────── + + +class TestReverseConfigFromDict: + """ReverseConfig.from_dict 配置解析.""" + + def test_none_returns_disabled(self): + config = ReverseConfig.from_dict(None) + assert config.enabled is False + assert config.reverse_video is True + assert config.reverse_audio is True + + def test_empty_dict_returns_disabled(self): + config = ReverseConfig.from_dict({}) + assert config.enabled is False + + def test_enabled_false(self): + config = ReverseConfig.from_dict({"enabled": False}) + assert config.enabled is False + + def test_enabled_true_defaults(self): + """只传 enabled=True,video/audio 默认都开.""" + config = ReverseConfig.from_dict({"enabled": True}) + assert config.enabled is True + assert config.reverse_video is True + assert config.reverse_audio is True + + def test_disable_video_only(self): + config = ReverseConfig.from_dict({ + "enabled": True, + "reverse_video": False, + "reverse_audio": True, + }) + assert config.enabled is True + assert config.reverse_video is False + assert config.reverse_audio is True + + def test_disable_audio_only(self): + config = ReverseConfig.from_dict({ + "enabled": True, + "reverse_video": True, + "reverse_audio": False, + }) + assert config.enabled is True + assert config.reverse_video is True + assert config.reverse_audio is False + + def test_both_disabled(self): + config = ReverseConfig.from_dict({ + "enabled": True, + "reverse_video": False, + "reverse_audio": False, + }) + assert config.enabled is True + assert config.reverse_video is False + assert config.reverse_audio is False + + def test_invalid_type_falls_back_default(self): + """传入非字典类型(如列表),捕获 TypeError,返回默认配置.""" + config = ReverseConfig.from_dict([1, 2, 3]) # type: ignore[arg-type] + assert config.enabled is False + + def test_attribute_error_falls_back(self): + """没有 .get() 方法的对象,捕获 AttributeError,返回默认配置.""" + config = ReverseConfig.from_dict(123) # type: ignore[arg-type] + assert config.enabled is False + + def test_truthy_values(self): + """非布尔真值也能被 bool() 转换.""" + config = ReverseConfig.from_dict({ + "enabled": 1, + "reverse_video": 1, + "reverse_audio": "yes", + }) + assert config.enabled is True + assert config.reverse_video is True + assert config.reverse_audio is True + + def test_falsy_values(self): + """非布尔假值也能被 bool() 转换.""" + config = ReverseConfig.from_dict({ + "enabled": True, + "reverse_video": 0, + "reverse_audio": "", + }) + assert config.enabled is True + assert config.reverse_video is False + assert config.reverse_audio is False + + +# ── ReverseEngine.build_video_filter ─────────────────────────────────────── + + +class TestBuildVideoFilter: + """ReverseEngine.build_video_filter 视频倒放滤镜构建.""" + + def test_disabled_returns_empty(self): + config = ReverseConfig(enabled=False) + result = ReverseEngine.build_video_filter(config) + assert result == "" + + def test_enabled_but_video_off_returns_empty(self): + config = ReverseConfig(enabled=True, reverse_video=False, reverse_audio=True) + result = ReverseEngine.build_video_filter(config) + assert result == "" + + def test_enabled_normal_duration(self): + config = ReverseConfig(enabled=True) + result = ReverseEngine.build_video_filter(config, duration=10.0) + assert result == "reverse" + + def test_zero_duration(self): + config = ReverseConfig(enabled=True) + result = ReverseEngine.build_video_filter(config, duration=0.0) + assert result == "reverse" + + def test_exactly_max_safe_duration(self): + """刚好等于上限,允许.""" + config = ReverseConfig(enabled=True) + result = ReverseEngine.build_video_filter( + config, duration=ReverseEngine.MAX_SAFE_DURATION + ) + assert result == "reverse" + + def test_exceeds_max_safe_duration_skips(self): + """超过上限,跳过倒放,返回空字符串.""" + config = ReverseConfig(enabled=True) + result = ReverseEngine.build_video_filter( + config, duration=ReverseEngine.MAX_SAFE_DURATION + 1 + ) + assert result == "" + + def test_negative_duration(self): + """负时长应该不会触发上限,但仍会正常返回 reverse.""" + config = ReverseConfig(enabled=True) + result = ReverseEngine.build_video_filter(config, duration=-5.0) + assert result == "reverse" + + +# ── ReverseEngine.build_audio_filter ─────────────────────────────────────── + + +class TestBuildAudioFilter: + """ReverseEngine.build_audio_filter 音频倒放滤镜构建.""" + + def test_disabled_returns_empty(self): + config = ReverseConfig(enabled=False) + result = ReverseEngine.build_audio_filter(config) + assert result == "" + + def test_enabled_but_audio_off_returns_empty(self): + config = ReverseConfig(enabled=True, reverse_video=True, reverse_audio=False) + result = ReverseEngine.build_audio_filter(config) + assert result == "" + + def test_enabled_normal_duration(self): + config = ReverseConfig(enabled=True) + result = ReverseEngine.build_audio_filter(config, duration=10.0) + assert result == "areverse" + + def test_zero_duration(self): + config = ReverseConfig(enabled=True) + result = ReverseEngine.build_audio_filter(config, duration=0.0) + assert result == "areverse" + + def test_exactly_max_safe_duration(self): + config = ReverseConfig(enabled=True) + result = ReverseEngine.build_audio_filter( + config, duration=ReverseEngine.MAX_SAFE_DURATION + ) + assert result == "areverse" + + def test_exceeds_max_safe_duration_skips(self): + config = ReverseConfig(enabled=True) + result = ReverseEngine.build_audio_filter( + config, duration=ReverseEngine.MAX_SAFE_DURATION + 1 + ) + assert result == "" + + def test_both_video_and_audio_disabled(self): + """两个都关,两个滤镜都为空.""" + config = ReverseConfig( + enabled=True, reverse_video=False, reverse_audio=False + ) + assert ReverseEngine.build_video_filter(config, 10) == "" + assert ReverseEngine.build_audio_filter(config, 10) == "" diff --git a/tests/unit/test_task_enqueue.py b/tests/unit/test_task_enqueue.py new file mode 100755 index 000000000..9c50f4319 --- /dev/null +++ b/tests/unit/test_task_enqueue.py @@ -0,0 +1,304 @@ +"""task_enqueue 单测 — 队列限流 + 安全入队逻辑.""" +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from app.core.task_enqueue import ( + USER_PENDING_LIMIT, + GLOBAL_PENDING_LIMIT, + UserPendingLimitExceeded, + GlobalQueueFull, + check_queue_limits, + safe_enqueue_generation_task, +) + + +# ── Fixtures / Helpers ───────────────────────────────────────────────────── + + +class MockRepository: + """Mock 任务仓储,用计数器模拟 pending 数量.""" + + def __init__(self, global_count: int = 0, user_count: int = 0): + self._global = global_count + self._user = user_count + self.update_called = 0 + + def count_pending_total(self) -> int: + return self._global + + def count_pending_by_user(self, user_id: str) -> int: + return self._user + + def update(self, task): + self.update_called += 1 + + +def make_mock_task(task_id: str = "task-1"): + task = MagicMock() + task.id = task_id + task.status = "pending" + task.mark_failed = MagicMock() + return task + + +# ── check_queue_limits ──────────────────────────────────────────────────── + + +class TestCheckQueueLimits: + """check_queue_limits 预检查限流.""" + + def test_below_limits_passes(self): + repo = MockRepository(global_count=5, user_count=1) + # 不抛异常就是通过 + check_queue_limits("user-1", repo) + + def test_global_at_limit_raises(self): + """达到全局上限即拒绝.""" + repo = MockRepository( + global_count=GLOBAL_PENDING_LIMIT, user_count=1 + ) + with pytest.raises(GlobalQueueFull) as exc_info: + check_queue_limits("user-1", repo) + assert exc_info.value.pending_count == GLOBAL_PENDING_LIMIT + assert exc_info.value.limit == GLOBAL_PENDING_LIMIT + + def test_global_over_limit_raises(self): + repo = MockRepository( + global_count=GLOBAL_PENDING_LIMIT + 1, user_count=1 + ) + with pytest.raises(GlobalQueueFull): + check_queue_limits("user-1", repo) + + def test_user_at_limit_raises(self): + """达到用户上限即拒绝.""" + repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT) + with pytest.raises(UserPendingLimitExceeded) as exc_info: + check_queue_limits("user-1", repo) + assert exc_info.value.user_id == "user-1" + assert exc_info.value.pending_count == USER_PENDING_LIMIT + assert exc_info.value.limit == USER_PENDING_LIMIT + + def test_user_over_limit_raises(self): + repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT + 1) + with pytest.raises(UserPendingLimitExceeded): + check_queue_limits("user-1", repo) + + def test_global_priority_over_user(self): + """全局和用户都超限时,优先抛全局异常.""" + repo = MockRepository( + global_count=GLOBAL_PENDING_LIMIT + 1, + user_count=USER_PENDING_LIMIT + 1, + ) + with pytest.raises(GlobalQueueFull): + check_queue_limits("user-1", repo) + + def test_empty_user_id_skips_user_check(self): + """user_id 为空时跳过用户级检查.""" + repo = MockRepository(global_count=5, user_count=999) + # 不抛异常 = 通过(只检查全局) + check_queue_limits("", repo) + + def test_custom_limits(self): + """支持自定义限流阈值.""" + repo = MockRepository(global_count=5, user_count=5) + # 默认阈值下 user 5 > 3 会被拒 + with pytest.raises(UserPendingLimitExceeded): + check_queue_limits("u1", repo) + + # 自定义更高阈值就能通过 + check_queue_limits("u1", repo, user_pending_limit=10, global_pending_limit=10) + + +# ── safe_enqueue_generation_task ────────────────────────────────────────── + + +class TestSafeEnqueueGenerationTask: + """safe_enqueue_generation_task 安全入队.""" + + @patch("app.core.task_enqueue.celery_app") + def test_success_path(self, mock_celery): + """正常路径:入队前检查通过 → 发送Celery → 入队后检查通过.""" + repo = MockRepository(global_count=1, user_count=1) + task = make_mock_task() + + result = safe_enqueue_generation_task( + task, repo, user_id="user-1" + ) + + assert result is True + mock_celery.send_task.assert_called_once_with( + "worker.generate_video", args=[task.id] + ) + task.mark_failed.assert_not_called() + + @patch("app.core.task_enqueue.celery_app") + def test_no_user_id_skips_user_check(self, mock_celery): + """不传 user_id 跳过用户级限流.""" + repo = MockRepository(global_count=1, user_count=999) + task = make_mock_task() + + result = safe_enqueue_generation_task(task, repo, user_id="") + assert result is True + + @patch("app.core.task_enqueue.celery_app") + def test_precheck_global_over_marks_failed(self, mock_celery): + """入队前全局超限:标记 failed,抛异常.""" + repo = MockRepository( + global_count=GLOBAL_PENDING_LIMIT + 1, user_count=0 + ) + task = make_mock_task() + + with pytest.raises(GlobalQueueFull): + safe_enqueue_generation_task(task, repo, user_id="user-1") + + task.mark_failed.assert_called_once() + mock_celery.send_task.assert_not_called() + assert repo.update_called == 1 + + @patch("app.core.task_enqueue.celery_app") + def test_precheck_user_over_marks_failed(self, mock_celery): + """入队前用户超限:标记 failed,抛异常.""" + repo = MockRepository( + global_count=5, user_count=USER_PENDING_LIMIT + 1 + ) + task = make_mock_task() + + with pytest.raises(UserPendingLimitExceeded): + safe_enqueue_generation_task(task, repo, user_id="user-1") + + task.mark_failed.assert_called_once() + mock_celery.send_task.assert_not_called() + + @patch("app.core.task_enqueue.celery_app") + def test_celery_send_false_returns_false(self, mock_celery): + """Celery 发送失败:返回 False,任务标记 failed.""" + repo = MockRepository(global_count=1, user_count=1) + task = make_mock_task() + mock_celery.send_task.side_effect = Exception("celery down") + + result = safe_enqueue_generation_task(task, repo, user_id="user-1") + + assert result is False + task.mark_failed.assert_called_once() + assert "入队失败" in task.mark_failed.call_args[0][0] + + @patch("app.core.task_enqueue.celery_app") + def test_celery_send_failure_update_also_fails(self, mock_celery): + """Celery 发送失败 + mark_failed 更新也失败:不崩溃.""" + repo = MockRepository(global_count=1, user_count=1) + repo.update = MagicMock(side_effect=Exception("db down")) + task = make_mock_task() + mock_celery.send_task.side_effect = Exception("celery down") + + result = safe_enqueue_generation_task(task, repo, user_id="user-1") + + assert result is False + # 不抛异常就是胜利 + + @patch("app.core.task_enqueue.celery_app") + def test_postcheck_global_over_rollback(self, mock_celery): + """入队后全局超限(并发竞态):回滚标记 failed,抛异常.""" + # 入队前刚好通过,但入队后再查发现超限 + call_count = [0] + + def count_pending_total_side_effect(): + call_count[0] += 1 + if call_count[0] == 1: # 入队前检查 + return GLOBAL_PENDING_LIMIT # 等于上限,用 > 判断所以通过 + return GLOBAL_PENDING_LIMIT + 1 # 入队后再查,超限 + + repo = MockRepository(global_count=GLOBAL_PENDING_LIMIT, user_count=0) + repo.count_pending_total = MagicMock( + side_effect=count_pending_total_side_effect + ) + task = make_mock_task() + + with pytest.raises(GlobalQueueFull): + safe_enqueue_generation_task(task, repo, user_id="user-1") + + # 异常是 GlobalQueueFull 类型,且任务已被标记为 failed(含"入队后"原因) + task.mark_failed.assert_called_once() + assert "入队后" in task.mark_failed.call_args[0][0] + mock_celery.send_task.assert_called_once() + + @patch("app.core.task_enqueue.celery_app") + def test_postcheck_user_over_rollback(self, mock_celery): + """入队后用户超限:回滚标记 failed,抛异常.""" + repo = MockRepository( + global_count=5, user_count=USER_PENDING_LIMIT + ) + # 入队前用 > 判断,等于上限通过;入队后模拟并发超限 + original_user_count = repo.count_pending_by_user + + call_count = [0] + + def count_by_user_side_effect(user_id): + call_count[0] += 1 + if call_count[0] <= 1: # 入队前 + return USER_PENDING_LIMIT # 用 > 判断,等于时通过 + return USER_PENDING_LIMIT + 1 # 入队后,超限 + + repo.count_pending_by_user = MagicMock( + side_effect=count_by_user_side_effect + ) + task = make_mock_task() + + with pytest.raises(UserPendingLimitExceeded): + safe_enqueue_generation_task(task, repo, user_id="user-1") + + task.mark_failed.assert_called_once() + + @patch("app.core.task_enqueue.celery_app") + def test_log_task_status_enabled(self, mock_celery): + """log_task_status=True 时日志中包含状态.""" + repo = MockRepository(global_count=1, user_count=1) + task = make_mock_task() + + result = safe_enqueue_generation_task( + task, repo, user_id="user-1", log_task_status=True + ) + assert result is True + + @patch("app.core.task_enqueue.celery_app") + def test_custom_limits_in_enqueue(self, mock_celery): + """自定义限流阈值用于入队检查.""" + repo = MockRepository(global_count=5, user_count=5) + task = make_mock_task() + + # 默认阈值下用户 5 > 3 会被拒 + with pytest.raises(UserPendingLimitExceeded): + safe_enqueue_generation_task(task, repo, user_id="user-1") + + # 重置 mock 计数 + task.mark_failed.reset_mock() + + # 调大阈值后通过 + result = safe_enqueue_generation_task( + task, + repo, + user_id="user-1", + user_pending_limit=10, + global_pending_limit=10, + ) + assert result is True + + +# ── 异常类 ──────────────────────────────────────────────────────────────── + + +class TestExceptionClasses: + """异常类消息格式.""" + + def test_user_pending_limit_message(self): + exc = UserPendingLimitExceeded("u1", 5, 3) + assert "u1" in str(exc) + assert "5" in str(exc) + assert "3" in str(exc) + + def test_global_queue_full_message(self): + exc = GlobalQueueFull(25, 20) + assert "25" in str(exc) + assert "20" in str(exc) -- 2.54.0