test(P3-1): 第60波 reverse引擎+in-memory素材+task_enqueue单测(+85) #852
+185
@@ -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
|
||||
+318
@@ -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"]
|
||||
Executable
+194
@@ -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) == ""
|
||||
Executable
+304
@@ -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)
|
||||
Reference in New Issue
Block a user