diff --git a/tests/unit/test_transition_presets_domain.py b/tests/unit/test_transition_presets_domain.py new file mode 100755 index 000000000..546ff43b9 --- /dev/null +++ b/tests/unit/test_transition_presets_domain.py @@ -0,0 +1,180 @@ +"""transition_presets 模块单元测试.""" + +import pytest + +from domain.transition_presets import ( + TRANSITION_PRESET_LIBRARY, + TransitionPreset, + get_default_transition, + get_transition_preset, + list_transition_presets, +) + + +class TestTransitionPreset: + """TransitionPreset 数据类测试.""" + + def test_create_required_fields(self): + t = TransitionPreset(id="test_001", name="测试转场", category="basic") + assert t.id == "test_001" + assert t.name == "测试转场" + assert t.category == "basic" + # 默认值 + assert t.description == "" + assert t.tags == [] + assert t.transition == "fade" + assert t.default_duration == 0.5 + assert t.min_duration == 0.1 + assert t.max_duration == 3.0 + assert t.has_custom_params is False + + def test_create_all_fields(self): + t = TransitionPreset( + id="test_002", + name="完整转场", + category="slide", + description="测试描述", + tags=["标签1", "标签2"], + transition="slideleft", + default_duration=1.0, + min_duration=0.3, + max_duration=2.5, + has_custom_params=True, + ) + assert t.category == "slide" + assert t.description == "测试描述" + assert t.tags == ["标签1", "标签2"] + assert t.transition == "slideleft" + assert t.default_duration == 1.0 + assert t.min_duration == 0.3 + assert t.max_duration == 2.5 + assert t.has_custom_params is True + + def test_frozen_immutable(self): + t = TransitionPreset(id="test", name="测试", category="basic") + with pytest.raises(Exception): + t.name = "修改" # type: ignore[misc] + + def test_tags_default_new_list(self): + t1 = TransitionPreset(id="1", name="a", category="basic") + t2 = TransitionPreset(id="2", name="b", category="basic") + assert t1.tags is not t2.tags + assert t1.tags == [] + + +class TestTransitionPresetLibrary: + """TRANSITION_PRESET_LIBRARY 预设库测试.""" + + def test_not_empty(self): + assert len(TRANSITION_PRESET_LIBRARY) > 0 + + def test_all_unique_ids(self): + ids = [t.id for t in TRANSITION_PRESET_LIBRARY] + assert len(ids) == len(set(ids)), "转场 ID 不能重复" + + def test_all_are_transition_preset_instances(self): + for t in TRANSITION_PRESET_LIBRARY: + assert isinstance(t, TransitionPreset) + + def test_contains_basic_categories(self): + cats = {t.category for t in TRANSITION_PRESET_LIBRARY} + assert "basic" in cats + assert "fade" in cats + + def test_duration_constraints_valid(self): + """每个预设的 min <= default <= max.""" + for t in TRANSITION_PRESET_LIBRARY: + assert t.min_duration <= t.default_duration, f"{t.id}: min > default" + assert t.default_duration <= t.max_duration, f"{t.id}: default > max" + + def test_none_transition_zero_duration(self): + t = get_transition_preset("transition_none") + assert t is not None + assert t.default_duration == 0.0 + assert t.min_duration == 0.0 + assert t.max_duration == 0.0 + + +class TestGetTransitionPreset: + """get_transition_preset 函数测试.""" + + def test_existing_id(self): + t = get_transition_preset("transition_fade") + assert t is not None + assert t.id == "transition_fade" + assert t.name == "淡入淡出" + assert t.category == "fade" + + def test_nonexistent_id(self): + assert get_transition_preset("nonexistent") is None + + def test_empty_string(self): + assert get_transition_preset("") is None + + +class TestListTransitionPresets: + """list_transition_presets 函数测试.""" + + def test_no_filters_returns_all(self): + result = list_transition_presets() + assert len(result) == len(TRANSITION_PRESET_LIBRARY) + + def test_filter_by_category_basic(self): + result = list_transition_presets(category="basic") + assert len(result) >= 2 + for t in result: + assert t.category == "basic" + + def test_filter_by_category_fade(self): + result = list_transition_presets(category="fade") + assert len(result) >= 3 + for t in result: + assert t.category == "fade" + + def test_filter_by_unknown_category_returns_empty(self): + result = list_transition_presets(category="nonexistent") + assert result == [] + + def test_filter_by_keyword_name(self): + result = list_transition_presets(keyword="淡入") + assert len(result) >= 1 + assert any(t.name == "淡入淡出" for t in result) + + def test_filter_by_keyword_description(self): + result = list_transition_presets(keyword="经典") + assert len(result) >= 1 + + def test_filter_by_keyword_tag(self): + result = list_transition_presets(keyword="电影感") + assert len(result) >= 1 + + def test_filter_keyword_case_insensitive(self): + r1 = list_transition_presets(keyword="FADE") + r2 = list_transition_presets(keyword="fade") + assert len(r1) == len(r2) + + def test_filter_keyword_no_match(self): + result = list_transition_presets(keyword="xyz_nonexistent_12345") + assert result == [] + + def test_combined_category_and_keyword(self): + result = list_transition_presets(category="fade", keyword="黑场") + assert len(result) >= 1 + for t in result: + assert t.category == "fade" + + def test_combined_no_match(self): + result = list_transition_presets(category="basic", keyword="黑场") + assert result == [] + + +class TestGetDefaultTransition: + """get_default_transition 函数测试.""" + + def test_returns_none_transition(self): + t = get_default_transition() + assert t.id == "transition_none" + assert t.name == "无转场" + + def test_returns_transition_preset_instance(self): + assert isinstance(get_default_transition(), TransitionPreset) diff --git a/tests/unit/test_tts_job_domain.py b/tests/unit/test_tts_job_domain.py new file mode 100755 index 000000000..2d087cc49 --- /dev/null +++ b/tests/unit/test_tts_job_domain.py @@ -0,0 +1,394 @@ +"""tts_job 领域模型单元测试.""" + +import pytest + +from domain.tts_job import TERMINAL_STATUSES, TTSJob, TTSJobStatus + + +class TestTTSJobStatus: + """TTSJobStatus 枚举测试.""" + + def test_values(self): + assert TTSJobStatus.PENDING == "pending" + assert TTSJobStatus.PROCESSING == "processing" + assert TTSJobStatus.COMPLETED == "completed" + assert TTSJobStatus.FAILED == "failed" + assert TTSJobStatus.CANCELLED == "cancelled" + + def test_terminal_statuses(self): + assert TTSJobStatus.COMPLETED in TERMINAL_STATUSES + assert TTSJobStatus.FAILED in TERMINAL_STATUSES + assert TTSJobStatus.CANCELLED in TERMINAL_STATUSES + assert TTSJobStatus.PENDING not in TERMINAL_STATUSES + assert TTSJobStatus.PROCESSING not in TERMINAL_STATUSES + + +class TestTTSJobCreate: + """TTSJob.create 工厂方法测试.""" + + def test_create_with_required_fields(self): + job = TTSJob.create(user_id="user_001", input_text="你好世界") + assert job.id + assert len(job.id) == 32 + assert job.user_id == "user_001" + assert job.input_text == "你好世界" + assert job.status == TTSJobStatus.PENDING + assert job.voice_id == "" + assert job.sample_rate == 22050 + assert job.format == "mp3" + assert job.retry_count == 0 + assert job.max_retries == 3 + assert job.metadata == {} + assert job.created_at is not None + assert job.updated_at is not None + + def test_create_with_all_fields(self): + job = TTSJob.create( + user_id="user_002", + input_text="测试文本", + voice_id="voice_001", + voice_model="cosyvoice", + project_id="proj_001", + voice_clone_profile_id="clone_001", + sample_rate=16000, + format="wav", + max_retries=5, + metadata={"key": "value"}, + ) + assert job.voice_id == "voice_001" + assert job.voice_model == "cosyvoice" + assert job.project_id == "proj_001" + assert job.voice_clone_profile_id == "clone_001" + assert job.sample_rate == 16000 + assert job.format == "wav" + assert job.max_retries == 5 + assert job.metadata == {"key": "value"} + + def test_create_strips_strings(self): + job = TTSJob.create( + user_id=" user_003 ", + input_text=" 测试文本 ", + voice_id=" voice_001 ", + voice_model=" cosyvoice ", + project_id=" proj_001 ", + voice_clone_profile_id=" clone_001 ", + format="wav", + ) + assert job.user_id == "user_003" + assert job.input_text == "测试文本" + assert job.voice_id == "voice_001" + assert job.voice_model == "cosyvoice" + assert job.project_id == "proj_001" + assert job.voice_clone_profile_id == "clone_001" + assert job.format == "wav" + + def test_create_empty_user_id_raises(self): + with pytest.raises(ValueError, match="user_id"): + TTSJob.create(user_id="", input_text="test") + + def test_create_whitespace_user_id_raises(self): + with pytest.raises(ValueError, match="user_id"): + TTSJob.create(user_id=" ", input_text="test") + + def test_create_empty_input_text_raises(self): + with pytest.raises(ValueError, match="input_text"): + TTSJob.create(user_id="u", input_text="") + + def test_create_input_text_too_long_raises(self): + long_text = "a" * 10001 + with pytest.raises(ValueError, match="10000"): + TTSJob.create(user_id="u", input_text=long_text) + + def test_create_input_text_at_limit_ok(self): + text = "a" * 10000 + job = TTSJob.create(user_id="u", input_text=text) + assert job.input_text == text + + def test_create_invalid_format_raises(self): + with pytest.raises(ValueError, match="不支持的输出格式"): + TTSJob.create(user_id="u", input_text="t", format="flac") + + def test_create_supported_formats(self): + for fmt in ["mp3", "wav", "pcm"]: + job = TTSJob.create(user_id="u", input_text="t", format=fmt) + assert job.format == fmt + + def test_create_none_metadata_defaults_to_empty_dict(self): + job = TTSJob.create(user_id="u", input_text="t", metadata=None) + assert job.metadata == {} + + def test_create_ids_are_unique(self): + j1 = TTSJob.create(user_id="u", input_text="t") + j2 = TTSJob.create(user_id="u", input_text="t") + assert j1.id != j2.id + + +class TestTTSJobStateMachine: + """TTSJob 状态机测试.""" + + @pytest.fixture + def pending_job(self): + return TTSJob.create(user_id="user_001", input_text="测试") + + def test_initial_status_is_pending(self, pending_job): + assert pending_job.status == TTSJobStatus.PENDING + assert not pending_job.is_terminal + + def test_pending_to_processing(self, pending_job): + pending_job.mark_processing() + assert pending_job.status == TTSJobStatus.PROCESSING + assert pending_job.started_at is not None + assert pending_job.error_message == "" + + def test_pending_can_fail_directly(self, pending_job): + """pending 可以直接到 failed(比如入参校验失败)""" + pending_job.mark_failed("校验失败") + assert pending_job.status == TTSJobStatus.FAILED + assert pending_job.error_message == "校验失败" + + def test_pending_can_be_cancelled(self, pending_job): + pending_job.mark_cancelled() + assert pending_job.status == TTSJobStatus.CANCELLED + + def test_processing_to_completed(self, pending_job): + pending_job.mark_processing() + pending_job.mark_completed(output_audio_url="https://example.com/out.mp3") + assert pending_job.status == TTSJobStatus.COMPLETED + assert pending_job.output_audio_url == "https://example.com/out.mp3" + assert pending_job.completed_at is not None + assert pending_job.error_message == "" + + def test_processing_to_failed(self, pending_job): + pending_job.mark_processing() + pending_job.mark_failed("API 超时") + assert pending_job.status == TTSJobStatus.FAILED + assert pending_job.error_message == "API 超时" + + def test_processing_can_be_cancelled(self, pending_job): + pending_job.mark_processing() + pending_job.mark_cancelled() + assert pending_job.status == TTSJobStatus.CANCELLED + + def test_completed_is_terminal(self, pending_job): + pending_job.mark_processing() + pending_job.mark_completed(output_audio_url="https://example.com/out.mp3") + assert pending_job.is_terminal + assert pending_job.is_completed + + def test_failed_is_terminal_but_retryable(self, pending_job): + pending_job.mark_processing() + pending_job.mark_failed("error") + assert pending_job.is_terminal + assert pending_job.is_retryable + + def test_cancelled_is_terminal_and_not_retryable(self, pending_job): + pending_job.mark_cancelled() + assert pending_job.is_terminal + assert not pending_job.is_retryable + + def test_invalid_transition_completed_to_processing_raises(self, pending_job): + pending_job.mark_processing() + pending_job.mark_completed(output_audio_url="https://example.com/out.mp3") + with pytest.raises(ValueError, match="非法状态转换"): + pending_job.mark_processing() + + def test_invalid_transition_completed_to_failed_raises(self, pending_job): + pending_job.mark_processing() + pending_job.mark_completed(output_audio_url="https://example.com/out.mp3") + with pytest.raises(ValueError, match="非法状态转换"): + pending_job.mark_failed("test") + + def test_cancelled_cannot_transition(self, pending_job): + pending_job.mark_cancelled() + with pytest.raises(ValueError): + pending_job.mark_processing() + with pytest.raises(ValueError): + pending_job.mark_failed("test") + + def test_transition_to_with_string(self, pending_job): + """transition_to 支持字符串参数""" + pending_job.transition_to("processing") + assert pending_job.status == TTSJobStatus.PROCESSING + + def test_transition_to_invalid_string_raises(self, pending_job): + with pytest.raises(ValueError, match="无效状态"): + pending_job.transition_to("invalid_status") + + def test_state_transition_updates_updated_at(self, pending_job): + old_updated = pending_job.updated_at + import time + + time.sleep(0.001) + pending_job.mark_processing() + assert pending_job.updated_at > old_updated + + +class TestTTSJobRetry: + """TTSJob 重试逻辑测试.""" + + def test_failed_can_retry(self): + job = TTSJob.create(user_id="u", input_text="t", max_retries=3) + job.mark_processing() + job.mark_failed("error") + assert job.is_retryable + assert job.retry_count == 0 + + def test_prepare_retry_resets_to_pending(self): + job = TTSJob.create(user_id="u", input_text="t") + job.mark_processing() + job.mark_failed("error") + + job.prepare_retry() + assert job.status == TTSJobStatus.PENDING + assert job.retry_count == 1 + assert job.error_message == "" + assert job.started_at is None + assert job.completed_at is None + + def test_retry_up_to_max_retries(self): + job = TTSJob.create(user_id="u", input_text="t", max_retries=2) + # 第 1 次失败 + 重试 → retry_count=1,还可以重试 + job.mark_processing() + job.mark_failed("e1") + assert job.is_retryable + job.prepare_retry() + assert job.retry_count == 1 + + # 第 2 次失败 → retry_count=1,还是 failed 状态,还可以重试(max_retries=2) + job.mark_processing() + job.mark_failed("e2") + assert job.is_retryable # retry_count=1 < max_retries=2 + job.prepare_retry() + assert job.retry_count == 2 + + # 第 3 次失败 → retry_count=2,达到上限,不可重试 + job.mark_processing() + job.mark_failed("e3") + assert not job.is_retryable # retry_count=2 == max_retries=2 + + def test_retry_exceed_max_raises(self): + job = TTSJob.create(user_id="u", input_text="t", max_retries=1) + job.mark_processing() + job.mark_failed("e") + job.prepare_retry() # 第 1 次重试,用完了 + + job.mark_processing() + job.mark_failed("e2") + with pytest.raises(ValueError, match="不可重试"): + job.prepare_retry() + + def test_pending_not_retryable(self): + job = TTSJob.create(user_id="u", input_text="t") + assert not job.is_retryable + with pytest.raises(ValueError, match="不可重试"): + job.prepare_retry() + + def test_completed_not_retryable(self): + job = TTSJob.create(user_id="u", input_text="t") + job.mark_processing() + job.mark_completed(output_audio_url="https://example.com/out.mp3") + assert not job.is_retryable + with pytest.raises(ValueError, match="不可重试"): + job.prepare_retry() + + def test_cancelled_not_retryable(self): + job = TTSJob.create(user_id="u", input_text="t") + job.mark_cancelled() + assert not job.is_retryable + + +class TestTTSJobMarkCompleted: + """mark_completed 方法测试.""" + + def test_requires_output_url(self): + job = TTSJob.create(user_id="u", input_text="t") + job.mark_processing() + with pytest.raises(ValueError, match="output_audio_url"): + job.mark_completed(output_audio_url="") + + def test_sets_all_fields(self): + job = TTSJob.create(user_id="u", input_text="t") + job.mark_processing() + job.mark_completed( + output_audio_url="https://example.com/out.mp3", + output_audio_key="audio/001.mp3", + duration=30.5, + file_size=102400, + ) + assert job.output_audio_url == "https://example.com/out.mp3" + assert job.output_audio_key == "audio/001.mp3" + assert job.duration == 30.5 + assert job.file_size == 102400 + assert job.completed_at is not None + + def test_strips_whitespace(self): + job = TTSJob.create(user_id="u", input_text="t") + job.mark_processing() + job.mark_completed( + output_audio_url=" https://example.com/out.mp3 ", + output_audio_key=" audio/001.mp3 ", + ) + assert job.output_audio_url == "https://example.com/out.mp3" + assert job.output_audio_key == "audio/001.mp3" + + +class TestTTSJobIsCompleted: + """is_completed 属性测试.""" + + def test_completed_with_url_is_completed(self): + job = TTSJob.create(user_id="u", input_text="t") + job.mark_processing() + job.mark_completed(output_audio_url="https://example.com/out.mp3") + assert job.is_completed + + def test_completed_without_url_not_completed(self): + """极端情况:completed 状态但没有 URL(理论不会发生)""" + job = TTSJob.create(user_id="u", input_text="t") + job.mark_processing() + job.transition_to(TTSJobStatus.COMPLETED) # 直接转,不设 URL + assert not job.is_completed + + def test_pending_not_completed(self): + job = TTSJob.create(user_id="u", input_text="t") + assert not job.is_completed + + +class TestTTSJobToDict: + """to_dict 序列化测试.""" + + def test_pending_job_to_dict(self): + job = TTSJob.create(user_id="user_001", input_text="测试文本", voice_id="v001") + d = job.to_dict() + assert d["id"] == job.id + assert d["user_id"] == "user_001" + assert d["status"] == "pending" + assert d["input_text"] == "测试文本" + assert d["voice_id"] == "v001" + assert d["retry_count"] == 0 + assert d["is_retryable"] is False + assert d["is_completed"] is False + assert d["metadata"] == {} + assert d["started_at"] is None + assert d["completed_at"] is None + assert d["created_at"] is not None + assert d["updated_at"] is not None + + def test_completed_job_to_dict(self): + job = TTSJob.create(user_id="u", input_text="t") + job.mark_processing() + job.mark_completed(output_audio_url="https://example.com/out.mp3", duration=10.0) + d = job.to_dict() + assert d["status"] == "completed" + assert d["output_audio_url"] == "https://example.com/out.mp3" + assert d["duration"] == 10.0 + assert d["is_completed"] is True + assert d["started_at"] is not None + assert d["completed_at"] is not None + + def test_failed_job_to_dict(self): + job = TTSJob.create(user_id="u", input_text="t") + job.mark_failed("出错了") + d = job.to_dict() + assert d["status"] == "failed" + assert d["error_message"] == "出错了" + assert d["is_retryable"] is True