diff --git a/tests/unit/test_tts_voice_clone_points.py b/tests/unit/test_tts_voice_clone_points.py index 9b52fa686..560003401 100644 --- a/tests/unit/test_tts_voice_clone_points.py +++ b/tests/unit/test_tts_voice_clone_points.py @@ -1,28 +1,49 @@ -"""TTS + voice_clone 积分扣点单元测试 (#1895 P2 step 2.1)""" +"""TTS + voice_clone 积分扣点单元测试 (#1895 P2 step 2.1) + +覆盖 synthesize / voice_clone preview 在积分开关下的扣点、余额不足、失败退费、会员折扣等分支。 +""" from __future__ import annotations import math -from unittest.mock import MagicMock +from types import SimpleNamespace +from unittest.mock import MagicMock, patch import pytest from fastapi import HTTPException +import packages.middleware.points_gate as _pg_module -def _make_user(user_id="user-1", is_member=False, member_type=None): - u = MagicMock() - u.id = user_id - u.is_member = is_member - u.member_type = member_type - return u + +@pytest.fixture(autouse=True) +def _enable_gate(monkeypatch): + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True) + yield def _make_cu(user_id="user-1", is_member=False, member_type=None): cu = MagicMock() - cu.user = _make_user(user_id, is_member, member_type) + cu.user.id = user_id + cu.user.is_member = is_member + cu.user.member_type = member_type return cu +def _make_request(text="你好世界", voice_id="v1", **kw): + r = MagicMock() + r.text = text + r.voice_id = voice_id + r.voice_clone_profile_id = None + r.speed = 1.0 + r.emotion = "" + r.language = "zh-CN" + r.metadata_ = {} + r.voice_model = None + for k, v in kw.items(): + setattr(r, k, v) + return r + + def _est_minutes(chars: int) -> float: return max(1.0, math.ceil(chars / 240)) @@ -30,71 +51,164 @@ def _est_minutes(chars: int) -> float: class TestEstimateMinutes: @pytest.mark.parametrize( "chars,expected", - [ - (1, 1.0), - (240, 1.0), - (241, 2.0), - (480, 2.0), - (481, 3.0), - (1000, 5.0), - ], + [(1, 1.0), (240, 1.0), (241, 2.0), (480, 2.0), (481, 3.0), (1000, 5.0)], ) def test_estimate(self, chars, expected): assert _est_minutes(chars) == expected class TestTtsSynthesizePointsDeduction: - def _deduct(self, text, cu, db, enabled=True, success=True, balance=100): - from packages.domain.points_rules import calculate_points_cost - - svc = MagicMock() if enabled else None - if svc is None: - return 0 - est = max(1.0, math.ceil(len(text) / 240)) - cost = calculate_points_cost( - "ai_voice", - is_member=getattr(cu.user, "is_member", False), - duration_minutes=est, - member_type=getattr(cu.user, "member_type", None), - ) - svc.deduct_points.return_value = {"success": success, "balance": balance} - res = svc.deduct_points(cu.user.id, cost, "ai_voice", db) - if not res["success"]: - raise HTTPException(status_code=402, detail={"code": "INSUFFICIENT_POINTS"}) - return cost - - def test_disabled_no_deduction(self): - assert self._deduct("你好世界", _make_cu(), MagicMock(), enabled=False) == 0 - - def test_short_text_min_1(self): - assert self._deduct("你好", _make_cu(), MagicMock()) >= 1 + 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 + wf = MagicMock() + wf.start_synthesis.return_value = job + wf.process_synthesis_failure.return_value = job + if start_synth_raises: + wf.start_synthesis.side_effect = start_synth_raises + vc_repo = MagicMock() + vc_repo.get.return_value = None + svc = MagicMock() + svc.deduct_points.return_value = {"success": deduct_success, "balance": balance} + fake_settings = MagicMock(points_enabled=True) + return db, cu, repo, uc, wf, vc_repo, svc, fake_settings, job def test_insufficient_raises_402(self): - with pytest.raises(HTTPException) as ei: - self._deduct("你好" * 200, _make_cu(), MagicMock(), success=False, balance=0) - assert ei.value.status_code == 402 + 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 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, + ) + assert ei.value.status_code == 402 + assert ei.value.detail["code"] == "INSUFFICIENT_POINTS" + + 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): + resp = synthesize( + request=_make_request(text="测试"), + 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 + + 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): + synthesize( + request=_make_request(text="测试"), + 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): + synthesize( + request=_make_request(text="测试"), + 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): - cm = self._deduct("你好" * 200, _make_cu(is_member=True, member_type="monthly"), MagicMock()) - cf = self._deduct("你好" * 200, _make_cu(is_member=False), MagicMock()) + 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 -class TestRefundOnFailure: - def test_refund_called(self): - svc = MagicMock() - svc.deduct_points.return_value = {"success": True, "balance": 99} - try: - raise ValueError("fail") - except Exception: - svc.refund_points("u1", 5, "ai_voice", MagicMock(), ref_id="job1") - svc.refund_points.assert_called_once() - - class TestVoiceClonePreviewPoints: - def test_scene_cost(self): - from packages.domain.points_rules import calculate_points_cost + def _setup(self, text="你好", deduct_success=True, balance=0, synth_raises=None): + db = MagicMock() + cu = _make_cu() + repo = MagicMock() + profile = SimpleNamespace(is_ready=True, voice_id="vc-1", user_id=cu.user.id) + uc = MagicMock() + uc.execute.return_value = profile + cosy = MagicMock() + r = SimpleNamespace(audio_url="http://x/a.mp3", duration=1.2, file_size=1000) + cosy.synthesize_speech.return_value = r + if synth_raises: + cosy.synthesize_speech.side_effect = synth_raises + svc = MagicMock() + svc.deduct_points.return_value = {"success": deduct_success, "balance": balance} + fs = MagicMock(points_enabled=True) + return db, cu, repo, uc, cosy, svc, fs - est = max(1.0, math.ceil(100 / 240)) - cost = calculate_points_cost("voice_clone_synth", is_member=False, duration_minutes=est) - assert cost >= 1 + 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 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, + ) + 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 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, + ) + assert ei.value.status_code == 502 + assert svc.refund_points.called + + 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", {}): + resp = get_voice_clone_preview( + 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")