"""scripts_ai 积分扣点单元测试 (#1895 P2 step 2.3)""" from __future__ import annotations from unittest.mock import MagicMock, patch import pytest from fastapi import HTTPException import packages.middleware.points_gate as _pg_module def _make_cu(user_id="u1", is_member=False, member_type=None): cu = MagicMock() cu.user.id = user_id cu.user.is_member = is_member cu.user.member_type = member_type return cu @pytest.fixture(autouse=True) def _enable_gate(monkeypatch): monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True) yield class TestScriptsAiPointsGate: """测试 scripts_ai 三个端点都挂了 @points_gate 并正确扣费。""" @pytest.mark.parametrize( "scene,endpoint_fn_name", [ ("douyin_extract", "extract_from_douyin"), ("ai_rewrite", "ai_rewrite"), ("ai_title", "ai_generate_titles"), ], ) def test_insufficient_points_raises_402(self, scene, endpoint_fn_name): """积分不足时抛 402。""" from app.api.routes import scripts_ai from app.schemas.scripts_ai import ( AiGenerateTitlesRequest, AiRewriteRequest, ExtractFromDouyinRequest, ) fn = getattr(scripts_ai, endpoint_fn_name) db = MagicMock() cu = _make_cu() if scene == "douyin_extract": req = ExtractFromDouyinRequest(url="https://v.douyin.com/abc/") elif scene == "ai_rewrite": req = AiRewriteRequest(content="测试文案") else: req = AiGenerateTitlesRequest(content="测试文案", count=3) with patch("packages.domain.points_service.PointsService") as MockSvc: svc = MagicMock() svc.deduct_points.return_value = {"success": False, "balance": 0} MockSvc.return_value = svc with pytest.raises(HTTPException) as ei: fn(request=req, current_user=cu, db=db) assert ei.value.status_code == 402 def test_disabled_passthrough_no_user_error(self, monkeypatch): """关闭时不需要 user/db 也能被装饰器透传(验证 gate 关闭零副作用)。""" from app.api.routes import scripts_ai from app.schemas.scripts_ai import AiRewriteRequest monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: False) fn = scripts_ai.ai_rewrite # 不带 db/current_user 也应透传(后续业务逻辑可能报错但不是 401/500 gate 错误) with pytest.raises(Exception) as ei: fn(request=AiRewriteRequest(content="x"), current_user=None, db=None) # 不应是 gate 抛的 401/500 assert isinstance(ei.value, AttributeError) or ei.value.status_code not in (401, 500)