diff --git a/apps/api/app/api/routes/ai_avatar_render.py b/apps/api/app/api/routes/ai_avatar_render.py index eceb30d24..52fc98abb 100644 --- a/apps/api/app/api/routes/ai_avatar_render.py +++ b/apps/api/app/api/routes/ai_avatar_render.py @@ -14,7 +14,6 @@ import logging from datetime import UTC, datetime from app.auth import AuthenticatedUser, get_current_user -from packages.middleware.points_gate import points_gate from app.dependencies import get_db_session from app.schemas.ai_avatar_render import ( AiAvatarRenderJobResponse, @@ -30,6 +29,8 @@ from app.services.ai_avatar_render_service import ( from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy.orm import Session +from packages.middleware.points_gate import points_gate + logger = logging.getLogger(__name__) router = APIRouter() diff --git a/apps/api/app/api/routes/generation_preview.py b/apps/api/app/api/routes/generation_preview.py index 1154de555..b518b16d4 100755 --- a/apps/api/app/api/routes/generation_preview.py +++ b/apps/api/app/api/routes/generation_preview.py @@ -8,7 +8,6 @@ from __future__ import annotations import logging from app.auth import AuthenticatedUser, get_current_user -from packages.middleware.points_gate import points_gate from app.core.storage import get_storage_service from app.core.task_enqueue import ( GLOBAL_PENDING_LIMIT, @@ -44,6 +43,7 @@ from packages.application import ( GetGenerationTaskUseCase, ListGeneratedVideosByTaskUseCase, ) +from packages.middleware.points_gate import points_gate logger = logging.getLogger(__name__) diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index 8153bef90..ac4ae0d6b 100755 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -4,7 +4,6 @@ from typing import Any from app.api.routes._helpers import check_project_access from app.auth import AuthenticatedUser, get_current_user -from packages.middleware.points_gate import points_gate from app.core.storage import OSSStorageService, get_storage_service from app.core.task_enqueue import ( GLOBAL_PENDING_LIMIT, @@ -43,6 +42,7 @@ from packages.application import ( ListGeneratedVideosByTaskUseCase, ) from packages.domain.smart_match import smart_select_assets +from packages.middleware.points_gate import points_gate logger = logging.getLogger(__name__) diff --git a/packages/middleware/points_gate.py b/packages/middleware/points_gate.py index 54f8b7187..405fa35ee 100644 --- a/packages/middleware/points_gate.py +++ b/packages/middleware/points_gate.py @@ -81,10 +81,10 @@ def points_gate( return decorator - def _filter_kwargs(func: Callable, kwargs: dict) -> dict: """过滤掉目标函数签名不接受的 kwargs(避免 TypeError)。""" import inspect + try: sig = inspect.signature(func) params = sig.parameters @@ -95,6 +95,7 @@ def _filter_kwargs(func: Callable, kwargs: dict) -> dict: except (ValueError, TypeError): return kwargs + def _extract_kwargs(func: Callable, args: tuple, kwargs: dict) -> dict: """将位置参数映射到函数签名中的参数名,便于统一按 kwargs 提取。""" sig = inspect.signature(func) diff --git a/tests/unit/test_ai_avatar_render_points.py b/tests/unit/test_ai_avatar_render_points.py index 5299b1599..6ae5eea8d 100644 --- a/tests/unit/test_ai_avatar_render_points.py +++ b/tests/unit/test_ai_avatar_render_points.py @@ -1,4 +1,5 @@ """AI数字人渲染 积分扣点单元测试 (#1895 P2 step 2.6)""" + from __future__ import annotations from unittest.mock import MagicMock, patch @@ -17,17 +18,19 @@ def _enable(monkeypatch): class TestAiAvatarRenderPoints: def test_ai_digital_human_per_unit(self): from packages.domain.points_rules import calculate_points_cost + cost = calculate_points_cost("ai_digital_human", is_member=False, duration_minutes=1) assert cost >= 15 def test_decorator_attached(self): from app.api.routes.ai_avatar_render import create_render_job + assert hasattr(create_render_job, "__wrapped__"), "missing @points_gate" def test_insufficient_raises_402(self): - from fastapi import HTTPException from app.api.routes.ai_avatar_render import create_render_job from app.schemas.ai_avatar_render import CreateAiAvatarRenderRequest + from fastapi import HTTPException db = MagicMock() cu = MagicMock() diff --git a/tests/unit/test_generation_preview.py b/tests/unit/test_generation_preview.py index 9ac87ff79..69523f2bd 100644 --- a/tests/unit/test_generation_preview.py +++ b/tests/unit/test_generation_preview.py @@ -583,7 +583,6 @@ def _disable_points_gate(monkeypatch): yield - def _make_user(user_id="test_user_001"): """构造 mock AuthenticatedUser""" mock_user = MagicMock() diff --git a/tests/unit/test_generation_preview_points.py b/tests/unit/test_generation_preview_points.py index 2767e3f79..196800d51 100644 --- a/tests/unit/test_generation_preview_points.py +++ b/tests/unit/test_generation_preview_points.py @@ -1,4 +1,5 @@ """视频预览生成 积分扣点单元测试 (#1895 P2 step 2.5)""" + from __future__ import annotations from unittest.mock import MagicMock, patch @@ -17,13 +18,14 @@ def _enable(monkeypatch): class TestGenerationPreviewPoints: def test_ai_video_cost(self): from packages.domain.points_rules import calculate_points_cost + assert calculate_points_cost("ai_video", is_member=False) == 4 assert calculate_points_cost("ai_video", is_member=True, member_type="monthly") == 2 def test_insufficient_raises_402(self): - from fastapi import HTTPException from app.api.routes.generation_preview import create_preview_generation_task from app.schemas.generation_task import CreatePreviewGenerationTaskRequest + from fastapi import HTTPException db = MagicMock() cu = MagicMock() @@ -38,7 +40,9 @@ class TestGenerationPreviewPoints: MS.return_value = svc with pytest.raises(HTTPException) as ei: create_preview_generation_task( - request=req, authenticated_user=cu, db=db, + request=req, + authenticated_user=cu, + db=db, generation_task_repository=MagicMock(), asset_repo=MagicMock(), ) @@ -46,4 +50,5 @@ class TestGenerationPreviewPoints: def test_decorator_attached(self): from app.api.routes.generation_preview import create_preview_generation_task + assert hasattr(create_preview_generation_task, "__wrapped__"), "missing @points_gate" diff --git a/tests/unit/test_generation_tasks.py b/tests/unit/test_generation_tasks.py index 3475f9f60..239e1480d 100755 --- a/tests/unit/test_generation_tasks.py +++ b/tests/unit/test_generation_tasks.py @@ -6,6 +6,7 @@ from unittest.mock import MagicMock import pytest +import packages.middleware.points_gate as _pg_module from packages.application.generation_tasks import ( CreateGenerationTaskCommand, CreateGenerationTaskUseCase, @@ -17,8 +18,6 @@ from packages.application.generation_tasks import ( ) from packages.domain import GenerationTask -import packages.middleware.points_gate as _pg_module - @pytest.fixture(autouse=True) def _disable_points_gate(monkeypatch): @@ -27,7 +26,6 @@ def _disable_points_gate(monkeypatch): yield - @pytest.fixture def mock_repo(): return MagicMock() diff --git a/tests/unit/test_generation_tasks_points.py b/tests/unit/test_generation_tasks_points.py index 001095be4..2599398ad 100644 --- a/tests/unit/test_generation_tasks_points.py +++ b/tests/unit/test_generation_tasks_points.py @@ -1,4 +1,5 @@ """视频生成 积分扣点单元测试 (#1895 P2 step 2.4)""" + from __future__ import annotations from unittest.mock import MagicMock, patch @@ -17,19 +18,21 @@ def _enable(monkeypatch): class TestGenerationTasksPoints: def test_ai_video_base_cost(self): from packages.domain.points_rules import calculate_points_cost + assert calculate_points_cost("ai_video", is_member=False) == 4 assert calculate_points_cost("ai_video", is_member=True, member_type="monthly") == 2 def test_ai_video_quantity_scales(self): from packages.domain.points_rules import calculate_points_cost + c1 = calculate_points_cost("ai_video", is_member=False, quantity=1) c3 = calculate_points_cost("ai_video", is_member=False, quantity=3) assert c3 > c1 def test_insufficient_raises_402(self): - from fastapi import HTTPException from app.api.routes.generation_tasks import create_generation_task from app.schemas.generation_task import CreateGenerationTaskRequest + from fastapi import HTTPException db = MagicMock() cu = MagicMock() @@ -44,12 +47,17 @@ class TestGenerationTasksPoints: MS.return_value = svc with pytest.raises(HTTPException) as ei: create_generation_task( - request=req, authenticated_user=cu, db=db, - generation_task_repository=MagicMock(), project_repository=MagicMock(), - asset_library_repository=MagicMock(), asset_repository=MagicMock(), + request=req, + authenticated_user=cu, + db=db, + generation_task_repository=MagicMock(), + project_repository=MagicMock(), + asset_library_repository=MagicMock(), + asset_repository=MagicMock(), ) assert ei.value.status_code == 402 def test_decorator_attached(self): from app.api.routes.generation_tasks import create_generation_task + assert hasattr(create_generation_task, "__wrapped__"), "missing @points_gate"