diff --git a/tests/unit/test_tts_voice_clone_points.py b/tests/unit/test_tts_voice_clone_points.py index 8e5b65c59..560003401 100644 --- a/tests/unit/test_tts_voice_clone_points.py +++ b/tests/unit/test_tts_voice_clone_points.py @@ -58,15 +58,14 @@ class TestEstimateMinutes: class TestTtsSynthesizePointsDeduction: - def _setup(self, text="你好", deduct_success=True, balance=0, start_synth_raises=None, send_task_raises=None): + def _setup(self, text="你好", deduct_success=True, balance=0, + start_synth_raises=None, send_task_raises=None): db = MagicMock() cu = _make_cu() repo = MagicMock() import enum - class _S(enum.Enum): processing = "processing" - job = SimpleNamespace(id="job-1", status=_S.processing, metadata={}) uc = MagicMock() uc.execute.return_value = job @@ -83,23 +82,18 @@ class TestTtsSynthesizePointsDeduction: return db, cu, repo, uc, wf, vc_repo, svc, fake_settings, job def test_insufficient_raises_402(self): - db, cu, repo, uc, wf, vc_repo, svc, fs, _ = self._setup(text="你好" * 200, deduct_success=False, balance=0) + db, cu, repo, uc, wf, vc_repo, svc, fs, _ = self._setup( + text="你好" * 200, deduct_success=False, balance=0) from app.api.routes.tts import synthesize - - with ( - patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc), - patch("app.api.routes.tts.TTSWorkflowService", return_value=wf), - patch("app.api.routes.tts.PointsService", return_value=svc), - patch("app.api.routes.tts.settings", fs), - ): + with patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc), \ + patch("app.api.routes.tts.TTSWorkflowService", return_value=wf), \ + patch("app.api.routes.tts.PointsService", return_value=svc), \ + patch("app.api.routes.tts.settings", fs): with pytest.raises(HTTPException) as ei: synthesize( request=_make_request(text="你好" * 200), - authenticated_user=cu, - db=db, - repository=repo, - cosyvoice_service=MagicMock(), - voice_clone_repo=vc_repo, + authenticated_user=cu, db=db, repository=repo, + cosyvoice_service=MagicMock(), voice_clone_repo=vc_repo, ) assert ei.value.status_code == 402 assert ei.value.detail["code"] == "INSUFFICIENT_POINTS" @@ -107,21 +101,15 @@ class TestTtsSynthesizePointsDeduction: def test_success_deducts_points(self): db, cu, repo, uc, wf, vc_repo, svc, fs, job = self._setup() from app.api.routes.tts import synthesize - - with ( - patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc), - patch("app.api.routes.tts.TTSWorkflowService", return_value=wf), - patch("app.api.routes.tts.PointsService", return_value=svc), - patch("app.api.routes.tts.celery_app.send_task") as _st, - patch("app.api.routes.tts.settings", fs), - ): + with patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc), \ + patch("app.api.routes.tts.TTSWorkflowService", return_value=wf), \ + patch("app.api.routes.tts.PointsService", return_value=svc), \ + patch("app.api.routes.tts.celery_app.send_task") as _st, \ + patch("app.api.routes.tts.settings", fs): resp = synthesize( request=_make_request(text="测试"), - authenticated_user=cu, - db=db, - repository=repo, - cosyvoice_service=MagicMock(), - voice_clone_repo=vc_repo, + authenticated_user=cu, db=db, repository=repo, + cosyvoice_service=MagicMock(), voice_clone_repo=vc_repo, ) svc.deduct_points.assert_called_once() assert resp.job_id == job.id @@ -129,48 +117,35 @@ class TestTtsSynthesizePointsDeduction: def test_synthesis_failure_refunds(self): db, cu, repo, uc, wf, vc_repo, svc, fs, _ = self._setup(start_synth_raises=RuntimeError("boom")) from app.api.routes.tts import synthesize - - with ( - patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc), - patch("app.api.routes.tts.TTSWorkflowService", return_value=wf), - patch("app.api.routes.tts.PointsService", return_value=svc), - patch("app.api.routes.tts.celery_app.send_task"), - patch("app.api.routes.tts.settings", fs), - ): + with patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc), \ + patch("app.api.routes.tts.TTSWorkflowService", return_value=wf), \ + patch("app.api.routes.tts.PointsService", return_value=svc), \ + patch("app.api.routes.tts.celery_app.send_task"), \ + patch("app.api.routes.tts.settings", fs): synthesize( request=_make_request(text="测试"), - authenticated_user=cu, - db=db, - repository=repo, - cosyvoice_service=MagicMock(), - voice_clone_repo=vc_repo, + authenticated_user=cu, db=db, repository=repo, + cosyvoice_service=MagicMock(), voice_clone_repo=vc_repo, ) assert svc.refund_points.called def test_celery_send_failure_refunds(self): db, cu, repo, uc, wf, vc_repo, svc, fs, _ = self._setup(send_task_raises=RuntimeError("celery down")) from app.api.routes.tts import synthesize - - with ( - patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc), - patch("app.api.routes.tts.TTSWorkflowService", return_value=wf), - patch("app.api.routes.tts.PointsService", return_value=svc), - patch("app.api.routes.tts.celery_app.send_task", side_effect=RuntimeError("celery down")), - patch("app.api.routes.tts.settings", fs), - ): + with patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc), \ + patch("app.api.routes.tts.TTSWorkflowService", return_value=wf), \ + patch("app.api.routes.tts.PointsService", return_value=svc), \ + patch("app.api.routes.tts.celery_app.send_task", side_effect=RuntimeError("celery down")), \ + patch("app.api.routes.tts.settings", fs): synthesize( request=_make_request(text="测试"), - authenticated_user=cu, - db=db, - repository=repo, - cosyvoice_service=MagicMock(), - voice_clone_repo=vc_repo, + authenticated_user=cu, db=db, repository=repo, + cosyvoice_service=MagicMock(), voice_clone_repo=vc_repo, ) assert svc.refund_points.called def test_member_cheaper(self): from packages.domain.points_rules import calculate_points_cost - cf = calculate_points_cost("ai_voice", is_member=False, duration_minutes=2) cm = calculate_points_cost("ai_voice", is_member=True, member_type="monthly", duration_minutes=2) assert cm < cf @@ -197,48 +172,29 @@ class TestVoiceClonePreviewPoints: def test_insufficient_raises_402(self): db, cu, repo, uc, cosy, svc, fs = self._setup(deduct_success=False, balance=0) from app.api.routes.voice_clones import get_voice_clone_preview - - with ( - patch("app.api.routes.voice_clones.GetVoiceCloneUseCase", return_value=uc), - patch("app.api.routes.voice_clones.PointsService", return_value=svc), - patch("app.api.routes.voice_clones.settings", fs), - patch("app.api.routes.voice_clones._clone_preview_cache", {}), - ): + with patch("app.api.routes.voice_clones.GetVoiceCloneUseCase", return_value=uc), \ + patch("app.api.routes.voice_clones.PointsService", return_value=svc), \ + patch("app.api.routes.voice_clones.settings", fs), \ + patch("app.api.routes.voice_clones._clone_preview_cache", {}): with pytest.raises(HTTPException) as ei: get_voice_clone_preview( - clone_id="c1", - text="你好", - speed=1.0, - emotion="", - authenticated_user=cu, - db=db, - repository=repo, - cosyvoice=cosy, + clone_id="c1", text="你好", speed=1.0, emotion="", + authenticated_user=cu, db=db, repository=repo, cosyvoice=cosy, ) assert ei.value.status_code == 402 def test_synth_cosyvoice_error_refunds_and_raises_502(self): from packages.application.cosyvoice_service import CosyVoiceError - db, cu, repo, uc, cosy, svc, fs = self._setup(synth_raises=CosyVoiceError("fail")) from app.api.routes.voice_clones import get_voice_clone_preview - - with ( - patch("app.api.routes.voice_clones.GetVoiceCloneUseCase", return_value=uc), - patch("app.api.routes.voice_clones.PointsService", return_value=svc), - patch("app.api.routes.voice_clones.settings", fs), - patch("app.api.routes.voice_clones._clone_preview_cache", {}), - ): + with patch("app.api.routes.voice_clones.GetVoiceCloneUseCase", return_value=uc), \ + patch("app.api.routes.voice_clones.PointsService", return_value=svc), \ + patch("app.api.routes.voice_clones.settings", fs), \ + patch("app.api.routes.voice_clones._clone_preview_cache", {}): with pytest.raises(HTTPException) as ei: get_voice_clone_preview( - clone_id="c1", - text="你好", - speed=1.0, - emotion="", - authenticated_user=cu, - db=db, - repository=repo, - cosyvoice=cosy, + clone_id="c1", text="你好", speed=1.0, emotion="", + authenticated_user=cu, db=db, repository=repo, cosyvoice=cosy, ) assert ei.value.status_code == 502 assert svc.refund_points.called @@ -246,22 +202,13 @@ class TestVoiceClonePreviewPoints: def test_success_returns_audio(self): db, cu, repo, uc, cosy, svc, fs = self._setup() from app.api.routes.voice_clones import get_voice_clone_preview - - with ( - patch("app.api.routes.voice_clones.GetVoiceCloneUseCase", return_value=uc), - patch("app.api.routes.voice_clones.PointsService", return_value=svc), - patch("app.api.routes.voice_clones.settings", fs), - patch("app.api.routes.voice_clones._clone_preview_cache", {}), - ): + with patch("app.api.routes.voice_clones.GetVoiceCloneUseCase", return_value=uc), \ + patch("app.api.routes.voice_clones.PointsService", return_value=svc), \ + patch("app.api.routes.voice_clones.settings", fs), \ + patch("app.api.routes.voice_clones._clone_preview_cache", {}): resp = get_voice_clone_preview( - clone_id="c1", - text="你好", - speed=1.0, - emotion="", - authenticated_user=cu, - db=db, - repository=repo, - cosyvoice=cosy, + clone_id="c1", text="你好", speed=1.0, emotion="", + authenticated_user=cu, db=db, repository=repo, cosyvoice=cosy, ) svc.deduct_points.assert_called_once() assert resp.audio_url.startswith("http")