diff --git a/apps/api/app/api/routes/scripts_ai.py b/apps/api/app/api/routes/scripts_ai.py index 495551e7a..2f8d18f07 100644 --- a/apps/api/app/api/routes/scripts_ai.py +++ b/apps/api/app/api/routes/scripts_ai.py @@ -27,7 +27,10 @@ from app.services.script_asr_service import ( transcribe_to_text, ) from fastapi import APIRouter, Depends, HTTPException, status +from sqlalchemy.orm import Session +from app.dependencies import get_db_session +from packages.middleware.points_gate import points_gate from packages.shared.ai_client import get_doubao_client logger = logging.getLogger(__name__) @@ -62,9 +65,11 @@ def _validate_douyin_url(url: str) -> None: "/extract-from-douyin", response_model=ExtractFromDouyinResponse, ) +@points_gate("douyin_extract") def extract_from_douyin( request: ExtractFromDouyinRequest, - authenticated_user: AuthenticatedUser = Depends(get_current_user), + current_user: AuthenticatedUser = Depends(get_current_user), + db: Session = Depends(get_db_session), ) -> ExtractFromDouyinResponse: """从抖音视频下载无水印视频并通过 ASR 提取文案.""" source_url = request.url.strip() @@ -138,9 +143,11 @@ def extract_from_douyin( "/ai-rewrite", response_model=AiRewriteResponse, ) +@points_gate("ai_rewrite") def ai_rewrite( request: AiRewriteRequest, - authenticated_user: AuthenticatedUser = Depends(get_current_user), + current_user: AuthenticatedUser = Depends(get_current_user), + db: Session = Depends(get_db_session), ) -> AiRewriteResponse: """使用豆包大模型改写文案.""" content = (request.content or "").strip() @@ -206,9 +213,11 @@ def ai_rewrite( "/ai-generate-titles", response_model=AiGenerateTitlesResponse, ) +@points_gate("ai_title") def ai_generate_titles( request: AiGenerateTitlesRequest, - authenticated_user: AuthenticatedUser = Depends(get_current_user), + current_user: AuthenticatedUser = Depends(get_current_user), + db: Session = Depends(get_db_session), ) -> AiGenerateTitlesResponse: """使用现有 generate_smart_titles 生成标题.""" content = (request.content or "").strip() diff --git a/tests/unit/test_scripts_ai.py b/tests/unit/test_scripts_ai.py index 3cf946e2b..737b8c603 100644 --- a/tests/unit/test_scripts_ai.py +++ b/tests/unit/test_scripts_ai.py @@ -13,6 +13,8 @@ from __future__ import annotations import sys from unittest.mock import MagicMock, patch +import packages.middleware.points_gate as _pg_module + import pydantic import pytest @@ -28,6 +30,18 @@ def _make_auth_user(user_id: str = "u1"): return auth +@pytest.fixture(autouse=True) +def _disable_points_gate(monkeypatch): + """默认关闭积分闸门,避免影响既有用例。""" + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: False) + yield + + +@pytest.fixture +def mock_db(): + return MagicMock() + + def _mock_youtube_dl( extract_info_return=None, extract_info_side_effect=None, @@ -82,7 +96,7 @@ class TestExtractFromDouyin: req = ExtractFromDouyinRequest(url="https://v.douyin.com/xxxxx/") auth = _make_auth_user() - result = extract_from_douyin(request=req, authenticated_user=auth) + result = extract_from_douyin(request=req, current_user=auth) assert result.text == "这是一段测试文案内容" assert result.duration_seconds == 120.5 @@ -112,7 +126,7 @@ class TestExtractFromDouyin: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - extract_from_douyin(request=req, authenticated_user=auth) + extract_from_douyin(request=req, current_user=auth) assert exc_info.value.status_code == 400 @patch("tempfile.TemporaryDirectory") @@ -136,7 +150,7 @@ class TestExtractFromDouyin: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - extract_from_douyin(request=req, authenticated_user=auth) + extract_from_douyin(request=req, current_user=auth) assert exc_info.value.status_code == 502 @patch("app.api.routes.scripts_ai.transcribe_to_text") @@ -169,7 +183,7 @@ class TestExtractFromDouyin: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - extract_from_douyin(request=req, authenticated_user=auth) + extract_from_douyin(request=req, current_user=auth) assert exc_info.value.status_code == 503 @patch("app.api.routes.scripts_ai.transcribe_to_text") @@ -202,7 +216,7 @@ class TestExtractFromDouyin: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - extract_from_douyin(request=req, authenticated_user=auth) + extract_from_douyin(request=req, current_user=auth) assert exc_info.value.status_code == 502 @@ -225,7 +239,7 @@ class TestAiRewrite: req = AiRewriteRequest(content="原始文案内容", style="口语化") auth = _make_auth_user() - result = ai_rewrite(request=req, authenticated_user=auth) + result = ai_rewrite(request=req, current_user=auth) assert result.original == "原始文案内容" assert result.rewritten == "改写后的文案内容,口语化风格" @@ -242,7 +256,7 @@ class TestAiRewrite: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - ai_rewrite(request=req, authenticated_user=auth) + ai_rewrite(request=req, current_user=auth) assert exc_info.value.status_code == 400 @patch("app.api.routes.scripts_ai.get_doubao_client") @@ -261,7 +275,7 @@ class TestAiRewrite: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - ai_rewrite(request=req, authenticated_user=auth) + ai_rewrite(request=req, current_user=auth) assert exc_info.value.status_code == 502 @patch("app.api.routes.scripts_ai.get_doubao_client") @@ -279,7 +293,7 @@ class TestAiRewrite: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - ai_rewrite(request=req, authenticated_user=auth) + ai_rewrite(request=req, current_user=auth) assert exc_info.value.status_code == 502 @@ -301,7 +315,7 @@ class TestAiGenerateTitles: req = AiGenerateTitlesRequest(content="这是一段关于美食的文案", count=3) auth = _make_auth_user() - result = ai_generate_titles(request=req, authenticated_user=auth) + result = ai_generate_titles(request=req, current_user=auth) assert len(result.titles) == 3 assert all(isinstance(t, str) for t in result.titles) @@ -339,12 +353,12 @@ class TestAiGenerateTitles: # count=5 req = AiGenerateTitlesRequest(content="测试内容", count=5) - result = ai_generate_titles(request=req, authenticated_user=auth) + result = ai_generate_titles(request=req, current_user=auth) assert len(result.titles) <= 5 # count=1 req = AiGenerateTitlesRequest(content="测试内容", count=1) - result = ai_generate_titles(request=req, authenticated_user=auth) + result = ai_generate_titles(request=req, current_user=auth) assert len(result.titles) >= 1 def test_generate_titles_empty_content(self): @@ -357,7 +371,7 @@ class TestAiGenerateTitles: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - ai_generate_titles(request=req, authenticated_user=auth) + ai_generate_titles(request=req, current_user=auth) assert exc_info.value.status_code == 400 @patch("app.services.ai_service.get_doubao_client") @@ -372,7 +386,7 @@ class TestAiGenerateTitles: req = AiGenerateTitlesRequest(content="测试文案内容") auth = _make_auth_user() - result = ai_generate_titles(request=req, authenticated_user=auth) + result = ai_generate_titles(request=req, current_user=auth) assert len(result.titles) == 3 diff --git a/tests/unit/test_scripts_ai_points.py b/tests/unit/test_scripts_ai_points.py new file mode 100644 index 000000000..70967fafa --- /dev/null +++ b/tests/unit/test_scripts_ai_points.py @@ -0,0 +1,76 @@ +"""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)