test(P3-1): 第60波 reverse引擎+in-memory素材+task_enqueue单测(+85) #852

Merged
xiaoxia merged 1 commits from test/wave60-more-infra-and-adapters into develop 2026-07-24 22:06:07 +08:00
4 changed files with 1001 additions and 0 deletions
+185
View File
@@ -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
View File
@@ -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"]
+194
View File
@@ -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=Truevideo/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) == ""
+304
View File
@@ -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)