"""points_gate 中间件单元测试 (#1895)""" from __future__ import annotations import asyncio from unittest.mock import MagicMock, patch import pytest from fastapi import HTTPException import packages.middleware.points_gate as _pg_module from packages.middleware.points_gate import _execute_with_gate, _extract_kwargs, points_gate @pytest.fixture(autouse=True) def _enable_points_gate(monkeypatch): """测试用:强制开启 points_gate,绕过 POINTS_ENABLED 默认关闭。""" monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True) yield def _make_user(user_id="user-1", is_member=False, member_type=None): user = MagicMock() user.id = user_id user.is_member = is_member user.member_type = member_type return user def _make_current_user(user_id="user-1", is_member=False, member_type=None): cu = MagicMock() cu.user = _make_user(user_id, is_member, member_type) return cu class TestExtractKwargs: def test_basic_extraction(self): def fn(a, b, c=None): pass result = _extract_kwargs(fn, (1, 2), {"c": 3}) assert result == {"a": 1, "b": 2, "c": 3} class TestPointsGateSync: def test_no_user_raises_401(self): @points_gate("ai_rewrite") def my_func(db=None): return "ok" with pytest.raises(HTTPException) as exc_info: my_func(db=MagicMock()) assert exc_info.value.status_code == 401 def test_no_db_raises_500(self): @points_gate("ai_rewrite") def my_func(current_user=None, db=None): return "ok" with pytest.raises(HTTPException) as exc_info: my_func(current_user=_make_current_user(), db=None) assert exc_info.value.status_code == 500 def test_zero_cost_scene_passes_through(self): @points_gate("voice_clone_train") def my_func(current_user=None, db=None, **kwargs): return kwargs.get("_points_deducted", -1) mock_db = MagicMock() cu = _make_current_user() result = my_func(current_user=cu, db=mock_db) assert result == 0 class TestPointsGateExecuteLogic: def test_insufficient_points_raises_402(self): cu = _make_current_user() db = MagicMock() mock_svc = MagicMock() mock_svc.deduct_points.return_value = {"success": False, "balance": 2, "transaction_id": None} def my_func(current_user=cu, db=db, **kwargs): return "ok" with patch("packages.domain.points_service.PointsService", return_value=mock_svc): with pytest.raises(HTTPException) as exc_info: _execute_with_gate( my_func, (), {"current_user": cu, "db": db}, "ai_rewrite", None, None, None, is_async=False ) assert exc_info.value.status_code == 402 def test_free_scene_passes_through(self): cu = _make_current_user() db = MagicMock() def my_func(current_user=cu, db=db, **kwargs): return "result" result = _execute_with_gate( my_func, (), {"current_user": cu, "db": db}, "voice_clone_train", None, None, None, is_async=False ) assert result == "result" def test_per_unit_fixed_cost(self): cu = _make_current_user() db = MagicMock() mock_svc = MagicMock() mock_svc.deduct_points.return_value = {"success": True, "balance": 90, "transaction_id": "t1"} def my_func(current_user=cu, db=db, **kwargs): return kwargs.get("_points_deducted", 0) with patch("packages.domain.points_service.PointsService", return_value=mock_svc): result = _execute_with_gate( my_func, (), {"current_user": cu, "db": db}, "ai_rewrite", per_unit=10, unit_field=None, quantity_field=None, is_async=False, ) assert result == 10 mock_svc.deduct_points.assert_called_once() def test_refund_on_failure(self): cu = _make_current_user() db = MagicMock() mock_svc = MagicMock() mock_svc.deduct_points.return_value = {"success": True, "balance": 90, "transaction_id": "t1"} def failing_func(current_user=cu, db=db, **kwargs): raise RuntimeError("business error") with patch("packages.domain.points_service.PointsService", return_value=mock_svc): with pytest.raises(RuntimeError, match="business error"): _execute_with_gate( failing_func, (), {"current_user": cu, "db": db}, "ai_rewrite", per_unit=10, unit_field=None, quantity_field=None, is_async=False, ) mock_svc.refund_points.assert_called_once() def test_ai_video_free_quota_for_free_user(self): cu = _make_current_user(is_member=False) db = MagicMock() mock_svc = MagicMock() mock_svc.check_daily_free_clip.return_value = True mock_svc.record_daily_free_clip.return_value = True def my_func(current_user=cu, db=db, **kwargs): return kwargs.get("_is_free_quota", False) with patch("packages.domain.points_service.PointsService", return_value=mock_svc): result = _execute_with_gate( my_func, (), {"current_user": cu, "db": db}, "ai_video", None, None, None, is_async=False ) assert result is True class TestPointsGateAsync: @pytest.mark.asyncio async def test_async_func_supported(self): cu = _make_current_user() db = MagicMock() mock_svc = MagicMock() mock_svc.deduct_points.return_value = {"success": True, "balance": 90, "transaction_id": "t1"} @points_gate("ai_rewrite", per_unit=5) async def my_async_func(current_user=None, db=None, **kwargs): return kwargs.get("_points_deducted", 0) with patch("packages.domain.points_service.PointsService", return_value=mock_svc): result = await my_async_func(current_user=cu, db=db) assert result == 5