diff --git a/packages/middleware/points_gate.py b/packages/middleware/points_gate.py index bb9282e28..ea8ddb6e2 100644 --- a/packages/middleware/points_gate.py +++ b/packages/middleware/points_gate.py @@ -91,13 +91,16 @@ def points_gate( is_async = asyncio.iscoroutinefunction(func) if is_async: + async def wrapper(*args: Any, **kwargs: Any) -> Any: if not _pg_enabled(): # noqa: F821 return await func(*args, **kwargs) return await _pg_execute( # noqa: F821 func, args, kwargs, scene_key, per_unit, unit_field, quantity_field, True ) + else: + def wrapper(*args: Any, **kwargs: Any) -> Any: if not _pg_enabled(): # noqa: F821 return func(*args, **_pg_filter(func, kwargs)) # noqa: F821 @@ -126,8 +129,7 @@ def points_gate( # 为稳妥起见,直接把 wrapper code 的 co_names 映射到新名——复杂度过高, # 这里采用「确保短名没冲突」策略:如果冲突就抛异常让开发者改名。 raise RuntimeError( - f"points_gate: name collision in {func.__module__}.{func.__name__}: " - f"'{k}' already defined" + f"points_gate: name collision in {func.__module__}.{func.__name__}: " f"'{k}' already defined" ) merged_globals[k] = v diff --git a/tests/unit/test_lipsync_points.py b/tests/unit/test_lipsync_points.py index 97d0c4338..5acb81789 100644 --- a/tests/unit/test_lipsync_points.py +++ b/tests/unit/test_lipsync_points.py @@ -99,7 +99,7 @@ class TestLipsyncPointsDeduction: # ── 直接调用 create_lipsync_job 覆盖扣点/402/退费分支 ── import importlib from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import patch import packages.middleware.points_gate as _pg_module @@ -113,9 +113,16 @@ def _do_enable(monkeypatch): def _body(**kw): b = MagicMock() defaults = dict( - video_url="http://x/v.mp4", audio_url=None, audio_duration=None, - sentence_timings=None, voice_id=None, script_text="你好世界", - speed=1.0, emotion="", enable_video_loop=False, project_id=None, + video_url="http://x/v.mp4", + audio_url=None, + audio_duration=None, + sentence_timings=None, + voice_id=None, + script_text="你好世界", + speed=1.0, + emotion="", + enable_video_loop=False, + project_id=None, ) defaults.update(kw) for k, v in defaults.items(): @@ -135,13 +142,16 @@ class TestLipsyncEndpointPoints: def test_insufficient_raises_402(self, monkeypatch): _do_enable(monkeypatch) from app.api.routes.lipsync import create_lipsync_job + db = MagicMock() svc = MagicMock() ps = MagicMock() ps.deduct_points.return_value = {"success": False, "balance": 0} fs = MagicMock(points_enabled=True) - with patch("app.api.routes.lipsync.PointsService", return_value=ps), \ - patch("app.api.routes.lipsync.settings", fs): + with ( + patch("app.api.routes.lipsync.PointsService", return_value=ps), + patch("app.api.routes.lipsync.settings", fs), + ): with pytest.raises(HTTPException) as ei: create_lipsync_job(body=_body(script_text="你" * 500), current_user=_cu(), db=db, svc=svc) assert ei.value.status_code == 402 @@ -149,14 +159,17 @@ class TestLipsyncEndpointPoints: def test_value_error_refunds(self, monkeypatch): _do_enable(monkeypatch) from app.api.routes.lipsync import create_lipsync_job + db = MagicMock() svc = MagicMock() svc.create_job.side_effect = ValueError("bad input") ps = MagicMock() ps.deduct_points.return_value = {"success": True, "balance": 99} fs = MagicMock(points_enabled=True) - with patch("app.api.routes.lipsync.PointsService", return_value=ps), \ - patch("app.api.routes.lipsync.settings", fs): + with ( + patch("app.api.routes.lipsync.PointsService", return_value=ps), + patch("app.api.routes.lipsync.settings", fs), + ): with pytest.raises(HTTPException) as ei: create_lipsync_job(body=_body(), current_user=_cu(), db=db, svc=svc) assert ei.value.status_code == 400 @@ -166,14 +179,17 @@ class TestLipsyncEndpointPoints: _do_enable(monkeypatch) from app.api.routes.lipsync import create_lipsync_job from app.services.mediakit_client import MediaKitError + db = MagicMock() svc = MagicMock() svc.create_job.side_effect = MediaKitError("fail", code="InvalidInput") ps = MagicMock() ps.deduct_points.return_value = {"success": True, "balance": 99} fs = MagicMock(points_enabled=True) - with patch("app.api.routes.lipsync.PointsService", return_value=ps), \ - patch("app.api.routes.lipsync.settings", fs): + with ( + patch("app.api.routes.lipsync.PointsService", return_value=ps), + patch("app.api.routes.lipsync.settings", fs), + ): with pytest.raises(HTTPException) as ei: create_lipsync_job(body=_body(), current_user=_cu(), db=db, svc=svc) assert ei.value.status_code == 400 @@ -182,14 +198,17 @@ class TestLipsyncEndpointPoints: def test_generic_exception_refunds(self, monkeypatch): _do_enable(monkeypatch) from app.api.routes.lipsync import create_lipsync_job + db = MagicMock() svc = MagicMock() svc.create_job.side_effect = RuntimeError("boom") ps = MagicMock() ps.deduct_points.return_value = {"success": True, "balance": 99} fs = MagicMock(points_enabled=True) - with patch("app.api.routes.lipsync.PointsService", return_value=ps), \ - patch("app.api.routes.lipsync.settings", fs): + with ( + patch("app.api.routes.lipsync.PointsService", return_value=ps), + patch("app.api.routes.lipsync.settings", fs), + ): with pytest.raises(HTTPException) as ei: create_lipsync_job(body=_body(), current_user=_cu(), db=db, svc=svc) assert ei.value.status_code == 400 @@ -198,7 +217,9 @@ class TestLipsyncEndpointPoints: def test_audio_duration_estimation(self, monkeypatch): _do_enable(monkeypatch) from app.api.routes.lipsync import create_lipsync_job + from packages.domain.points_rules import calculate_points_cost + db = MagicMock() svc = MagicMock() job = SimpleNamespace(id="job-1", status="queued") @@ -206,11 +227,15 @@ class TestLipsyncEndpointPoints: ps = MagicMock() ps.deduct_points.return_value = {"success": True, "balance": 99} fs = MagicMock(points_enabled=True) - with patch("app.api.routes.lipsync.PointsService", return_value=ps), \ - patch("app.api.routes.lipsync.settings", fs): + with ( + patch("app.api.routes.lipsync.PointsService", return_value=ps), + patch("app.api.routes.lipsync.settings", fs), + ): create_lipsync_job( body=_body(audio_url="http://x/a.mp3", audio_duration=180, script_text=None), - current_user=_cu(), db=db, svc=svc, + current_user=_cu(), + db=db, + svc=svc, ) # 180 seconds -> 3 minutes; assert deduct called with cost >= 15*3 args = ps.deduct_points.call_args[0] diff --git a/tests/unit/test_tts_voice_clone_points.py b/tests/unit/test_tts_voice_clone_points.py index 560003401..8e5b65c59 100644 --- a/tests/unit/test_tts_voice_clone_points.py +++ b/tests/unit/test_tts_voice_clone_points.py @@ -58,14 +58,15 @@ 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 @@ -82,18 +83,23 @@ 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" @@ -101,15 +107,21 @@ 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 @@ -117,35 +129,48 @@ 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 @@ -172,29 +197,48 @@ 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 @@ -202,13 +246,22 @@ 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")