"""订阅管理 API 单元测试。 覆盖 5 个端点: GET /current — 当前订阅信息 GET /billing-records — 账单记录 POST /change-plan — 变更套餐 POST /cancel — 取消订阅 POST /toggle-auto-renew — 切换自动续费 测试使用 FastAPI TestClient + 依赖覆盖(dependency_overrides), 不连接真实数据库,不访问外部服务。 """ from __future__ import annotations import sys import types from dataclasses import dataclass, field from datetime import datetime, timezone from typing import Optional from unittest.mock import MagicMock import pytest from fastapi import FastAPI from fastapi.testclient import TestClient # --------------------------------------------------------------------------- # 1. Mock 项目内部模块(使 subscription 路由可独立导入) # --------------------------------------------------------------------------- def _install_mocks(): """在 sys.modules 中安装所有必需的 mock 模块,使 subscription.py 可导入。""" # ---------- packages.domain.entities ---------- @dataclass(slots=True) class User: id: str = "user-001" email: str = "test@example.com" display_name: str = "Test User" username: str = "testuser" password_hash: str = "" email_verified: bool = False email_verification_token: str | None = None password_reset_token: str | None = None password_reset_expires_at: datetime | None = None last_login_at: datetime | None = None last_login_ip: str | None = None subscription_plan: str = "free" subscription_status: str = "active" subscription_expires_at: datetime | None = None max_projects: int = 3 max_storage_gb: int = 10 used_storage_gb: float = 0.0 created_at: datetime = field(default_factory=lambda: datetime(2026, 1, 1, tzinfo=timezone.utc)) entities_mod = types.ModuleType("packages.domain.entities") entities_mod.User = User # ---------- packages.ports.user_repository ---------- class UserRepository: def save(self, user): pass def find_by_id(self, user_id): return None def find_by_email(self, email): return None def find_by_username(self, username): return None def find_by_verification_token(self, token): return None def find_by_password_reset_token(self, token): return None def delete(self, user_id): return True user_repo_mod = types.ModuleType("packages.ports.user_repository") user_repo_mod.UserRepository = UserRepository # ---------- packages (namespace) ---------- for name in [ "packages", "packages.domain", "packages.ports", "packages.adapters", "packages.adapters.sqlalchemy_impl", "packages.adapters.sqlalchemy_impl.user_repository", "packages.adapters.sqlalchemy_impl.session", "packages.adapters.redis", "packages.adapters.smtp", "packages.application", ]: if name not in sys.modules: sys.modules[name] = types.ModuleType(name) sys.modules["packages.domain.entities"] = entities_mod sys.modules["packages.ports.user_repository"] = user_repo_mod sys.modules["packages.adapters.sqlalchemy_impl.user_repository"].SQLAlchemyUserRepository = MagicMock sys.modules["packages.adapters.sqlalchemy_impl.session"].build_session_factory = MagicMock( return_value=(MagicMock(), MagicMock()) ) sys.modules["packages.adapters.redis"].NoopSessionStore = MagicMock sys.modules["packages.adapters.redis"].SessionStore = MagicMock sys.modules["packages.adapters.smtp"].EmailConfig = MagicMock sys.modules["packages.adapters.smtp"].NoopEmailService = MagicMock sys.modules["packages.adapters.smtp"].get_email_service = MagicMock() # Stub 其他 repository ports(dependencies.py 会 import 它们) for port_name in [ "asset_repository", "asset_library_repository", "classification_job_repository", "duplication_repository", "generated_video_repository", "generation_task_repository", "title_library_repository", "voice_library_repository", "ingest_job_repository", "project_repository", ]: mod = types.ModuleType(f"packages.ports.{port_name}") # 动态创建一个 Mock repository class class_name = port_name.replace("_", " ").title().replace(" ", "") + "Port" setattr(mod, "".join(w.capitalize() for w in port_name.split("_")), MagicMock) sys.modules[f"packages.ports.{port_name}"] = mod sa_mod = types.ModuleType(f"packages.adapters.sqlalchemy_impl.{port_name}") setattr(sa_mod, f"SQLAlchemy{''.join(w.capitalize() for w in port_name.split('_'))}", MagicMock) sys.modules[f"packages.adapters.sqlalchemy_impl.{port_name}"] = sa_mod # ---------- app.config ---------- config_mod = types.ModuleType("app.config") class _Settings: JWT_SECRET_KEY = "test-secret-key-for-unit-tests" DATABASE_URL = "sqlite:///test.db" REDIS_URL = "redis://localhost:6379/0" ENABLE_REDIS_SESSIONS = False SMTP_HOST = "" SMTP_PORT = 587 SMTP_USER = "" SMTP_PASSWORD = "" SMTP_FROM_EMAIL = "" SMTP_FROM_NAME = "" SMTP_USE_TLS = False ENABLE_EMAIL_DELIVERY = False config_mod.settings = _Settings() config_mod.get_settings = lambda: _Settings() sys.modules["app.config"] = config_mod # ---------- app.auth ---------- @dataclass(frozen=True, slots=True) class AuthenticatedUser: user: User session_id: str | None = None token_type: str | None = None async def _mock_get_current_user(): return AuthenticatedUser(user=User()) auth_mod = types.ModuleType("app.auth") auth_mod.AuthenticatedUser = AuthenticatedUser auth_mod.get_current_user = _mock_get_current_user sys.modules["app.auth"] = auth_mod # ---------- app.dependencies ---------- deps_mod = types.ModuleType("app.dependencies") deps_mod.get_db_session = MagicMock() deps_mod.get_user_repository = MagicMock() sys.modules["app.dependencies"] = deps_mod # ---------- app.schemas.subscription ---------- # 需要真正的 Pydantic 模型 → 延迟到 subscription 模块导入时解析 # 这里我们直接导入真实 schema(因为它是纯 Pydantic 定义,无外部依赖) # 但为安全起见也 mock 掉 try: from typing import List from typing import Optional as Opt from pydantic import BaseModel, Field class PlanType(str): FREE = "free" STANDARD = "standard" PRO = "pro" ENTERPRISE = "enterprise" class SubscriptionStatus(str): ACTIVE = "active" EXPIRED = "expired" CANCELLED = "cancelled" TRIAL = "trial" class BillingStatus(str): PAID = "paid" PENDING = "pending" FAILED = "failed" REFUNDED = "refunded" class BillingCycle(str): MONTHLY = "monthly" YEARLY = "yearly" class SubscriptionInfo(BaseModel): id: str plan_id: str plan_name: str status: str billing_cycle: str current_period_start: str current_period_end: str amount: float auto_renew: bool created_at: str class BillingRecord(BaseModel): id: str plan_name: str amount: float billing_cycle: str status: str payment_method: str created_at: str invoice_url: Opt[str] = None class ChangePlanResponse(BaseModel): success: bool message: str new_subscription: Opt[SubscriptionInfo] = None class SimpleResponse(BaseModel): success: bool message: str class ChangePlanRequest(BaseModel): target_plan_id: str = Field(..., description="目标套餐ID") billing_cycle: str = Field(..., description="计费周期: monthly/yearly") class ToggleAutoRenewRequest(BaseModel): enabled: bool = Field(..., description="是否开启自动续费") schemas_mod = types.ModuleType("app.schemas.subscription") schemas_mod.PlanType = PlanType schemas_mod.SubscriptionStatus = SubscriptionStatus schemas_mod.BillingStatus = BillingStatus schemas_mod.BillingCycle = BillingCycle schemas_mod.SubscriptionInfo = SubscriptionInfo schemas_mod.BillingRecord = BillingRecord schemas_mod.ChangePlanResponse = ChangePlanResponse schemas_mod.SimpleResponse = SimpleResponse schemas_mod.ChangePlanRequest = ChangePlanRequest schemas_mod.ToggleAutoRenewRequest = ToggleAutoRenewRequest sys.modules["app.schemas.subscription"] = schemas_mod sys.modules.setdefault("app.schemas", types.ModuleType("app.schemas")) sys.modules["app.schemas"].subscription = schemas_mod except Exception: pass # 如果已经导入过,跳过 return User, AuthenticatedUser User, AuthenticatedUser = _install_mocks() # ---------- 导入被测路由模块 ---------- # 先确保 app 和 app.api 命名空间存在 for ns in ["app", "app.api", "app.api.routes"]: if ns not in sys.modules: sys.modules[ns] = types.ModuleType(ns) # 导入 subscription 路由 import importlib.util _spec = importlib.util.spec_from_file_location("app.api.routes.subscription", "/tmp/subscription_routes.py") subscription = importlib.util.module_from_spec(_spec) sys.modules["app.api.routes.subscription"] = subscription _spec.loader.exec_module(subscription) # --------------------------------------------------------------------------- # 2. Fixtures # --------------------------------------------------------------------------- def _make_user(**overrides) -> User: """创建测试用 User 实例。""" defaults = dict( id="user-001", email="test@example.com", display_name="Test User", username="testuser", subscription_plan="free", subscription_status="active", subscription_expires_at=None, max_projects=3, max_storage_gb=10, created_at=datetime(2026, 1, 1, tzinfo=timezone.utc), ) defaults.update(overrides) return User(**defaults) class MockUserRepository: """内存中的 User Repository mock。""" def __init__(self): self.saved_users: list[User] = [] def save(self, user: User) -> None: self.saved_users.append(user) def find_by_id(self, user_id: str) -> Optional[User]: return None @pytest.fixture def mock_user_repo(): return MockUserRepository() @pytest.fixture def client(mock_user_repo): """创建带有依赖覆盖的 TestClient。""" app = FastAPI() app.include_router(subscription.router) def _override_get_current_user(): return AuthenticatedUser(user=_make_user()) def _override_get_user_repo(): return mock_user_repo app.dependency_overrides[subscription.get_current_user] = _override_get_current_user app.dependency_overrides[subscription.get_user_repository] = _override_get_user_repo return TestClient(app) @pytest.fixture def pro_client(mock_user_repo): """已订阅 Pro 套餐的用户客户端。""" app = FastAPI() app.include_router(subscription.router) def _override_get_current_user(): return AuthenticatedUser( user=_make_user( subscription_plan="pro", subscription_status="active", max_projects=-1, max_storage_gb=100, ) ) def _override_get_user_repo(): return mock_user_repo app.dependency_overrides[subscription.get_current_user] = _override_get_current_user app.dependency_overrides[subscription.get_user_repository] = _override_get_user_repo return TestClient(app) # --------------------------------------------------------------------------- # 3. GET /current — 获取当前订阅信息 # --------------------------------------------------------------------------- class TestGetCurrentSubscription: """GET /current 端点测试。""" def test_returns_subscription_info_for_free_user(self, client): """免费用户应返回 free 套餐信息。""" resp = client.get("/current") assert resp.status_code == 200 data = resp.json() assert data["plan_id"] == "free" assert data["plan_name"] == "体验版" assert data["status"] == "active" assert data["billing_cycle"] == "monthly" assert data["amount"] == 0 assert data["auto_renew"] is True assert "id" in data assert data["id"].startswith("sub-") def test_returns_correct_plan_name_for_pro(self, pro_client): """Pro 用户应返回「专业版」名称。""" resp = pro_client.get("/current") assert resp.status_code == 200 data = resp.json() assert data["plan_id"] == "pro" assert data["plan_name"] == "专业版" assert data["amount"] == 299 # pro monthly = 299 def test_response_contains_period_dates(self, client): """响应应包含 period_start 和 period_end。""" resp = client.get("/current") data = resp.json() assert "current_period_start" in data assert "current_period_end" in data # free 用户没有过期时间,period_end == period_start assert data["current_period_start"] is not None def test_response_contains_created_at(self, client): """响应应包含 created_at。""" resp = client.get("/current") data = resp.json() assert "created_at" in data assert data["created_at"] != "" # --------------------------------------------------------------------------- # 4. GET /billing-records — 获取账单记录 # --------------------------------------------------------------------------- class TestGetBillingRecords: def test_returns_empty_list(self, client): """当前实现返回空列表(TODO: 数据库查询)。""" resp = client.get("/billing-records") assert resp.status_code == 200 data = resp.json() assert isinstance(data, list) assert len(data) == 0 # --------------------------------------------------------------------------- # 5. POST /change-plan — 变更套餐 # --------------------------------------------------------------------------- class TestChangePlan: def test_upgrade_free_to_standard(self, client, mock_user_repo): """从 free 升级到 standard 应成功。""" resp = client.post( "/change-plan", json={ "target_plan_id": "standard", "billing_cycle": "monthly", }, ) assert resp.status_code == 200 data = resp.json() assert data["success"] is True assert "标准版" in data["message"] assert data["new_subscription"] is not None assert data["new_subscription"]["plan_id"] == "standard" assert data["new_subscription"]["amount"] == 99 def test_upgrade_free_to_pro(self, client, mock_user_repo): """从 free 升级到 pro 应成功,配额正确更新。""" resp = client.post( "/change-plan", json={ "target_plan_id": "pro", "billing_cycle": "yearly", }, ) assert resp.status_code == 200 data = resp.json() assert data["success"] is True sub = data["new_subscription"] assert sub["plan_id"] == "pro" assert sub["amount"] == 299 # _build_subscription_info 固定用 monthly 计价 # 验证 repository 被调用保存了用户 assert len(mock_user_repo.saved_users) == 1 saved = mock_user_repo.saved_users[0] assert saved.subscription_plan == "pro" assert saved.max_projects == -1 # 无限 assert saved.max_storage_gb == 100 def test_upgrade_to_enterprise(self, client, mock_user_repo): """升级到 enterprise 套餐。""" resp = client.post( "/change-plan", json={ "target_plan_id": "enterprise", "billing_cycle": "monthly", }, ) assert resp.status_code == 200 data = resp.json() assert data["success"] is True assert data["new_subscription"]["plan_name"] == "企业版" assert data["new_subscription"]["amount"] == 999 saved = mock_user_repo.saved_users[0] assert saved.max_storage_gb == 1000 def test_same_plan_returns_failure(self, client): """当前套餐与目标套餐相同时应返回 success=False。""" resp = client.post( "/change-plan", json={ "target_plan_id": "free", "billing_cycle": "monthly", }, ) assert resp.status_code == 200 data = resp.json() assert data["success"] is False assert "已经是" in data["message"] def test_invalid_plan_id_returns_400(self, client): """无效套餐 ID 应返回 400。""" resp = client.post( "/change-plan", json={ "target_plan_id": "ultra_mega_plan", "billing_cycle": "monthly", }, ) assert resp.status_code == 400 assert "无效的套餐ID" in resp.json()["detail"] def test_invalid_billing_cycle_returns_400(self, client): """无效计费周期应返回 400。""" resp = client.post( "/change-plan", json={ "target_plan_id": "pro", "billing_cycle": "weekly", }, ) assert resp.status_code == 400 assert "无效的计费周期" in resp.json()["detail"] def test_missing_fields_returns_422(self, client): """缺少必填字段应返回 422。""" resp = client.post("/change-plan", json={"target_plan_id": "pro"}) assert resp.status_code == 422 def test_empty_body_returns_422(self, client): """空请求体应返回 422。""" resp = client.post("/change-plan", json={}) assert resp.status_code == 422 def test_does_not_mutate_frozen_dataclass(self, client, mock_user_repo): """变更套餐应通过 dataclasses.replace 创建新实例,不修改原对象。""" # 原始 user 是 frozen dataclass original_user = _make_user(subscription_plan="free") app = FastAPI() app.include_router(subscription.router) def _get_user(): return AuthenticatedUser(user=original_user) app.dependency_overrides[subscription.get_current_user] = _get_user app.dependency_overrides[subscription.get_user_repository] = lambda: mock_user_repo tc = TestClient(app) resp = tc.post( "/change-plan", json={ "target_plan_id": "standard", "billing_cycle": "monthly", }, ) assert resp.status_code == 200 # 原始 user 对象不变 assert original_user.subscription_plan == "free" # 新保存的 user 是更新后的 assert mock_user_repo.saved_users[0].subscription_plan == "standard" # --------------------------------------------------------------------------- # 6. POST /cancel — 取消订阅 # --------------------------------------------------------------------------- class TestCancelSubscription: def test_cancel_pro_subscription(self, pro_client, mock_user_repo): """Pro 用户取消订阅应成功。""" resp = pro_client.post("/cancel") assert resp.status_code == 200 data = resp.json() assert data["success"] is True assert "已取消" in data["message"] # 验证 repository 保存了 cancelled 状态 saved = mock_user_repo.saved_users[0] assert saved.subscription_status == "cancelled" def test_cancel_free_subscription_returns_400(self, client): """免费用户无需取消,应返回 400。""" resp = client.post("/cancel") assert resp.status_code == 400 assert "体验版无需取消" in resp.json()["detail"] def test_cancel_does_not_mutate_original_user(self, mock_user_repo): """取消操作不应修改 frozen dataclass 原始对象。""" original_user = _make_user( subscription_plan="standard", subscription_status="active", ) app = FastAPI() app.include_router(subscription.router) app.dependency_overrides[subscription.get_current_user] = lambda: AuthenticatedUser(user=original_user) app.dependency_overrides[subscription.get_user_repository] = lambda: mock_user_repo tc = TestClient(app) resp = tc.post("/cancel") assert resp.status_code == 200 # 原始不变 assert original_user.subscription_status == "active" # 保存的是新的 assert mock_user_repo.saved_users[0].subscription_status == "cancelled" # --------------------------------------------------------------------------- # 7. POST /toggle-auto-renew — 切换自动续费 # --------------------------------------------------------------------------- class TestToggleAutoRenew: def test_enable_auto_renew(self, client): """开启自动续费。""" resp = client.post("/toggle-auto-renew", json={"enabled": True}) assert resp.status_code == 200 data = resp.json() assert data["success"] is True assert "开启" in data["message"] def test_disable_auto_renew(self, client): """关闭自动续费。""" resp = client.post("/toggle-auto-renew", json={"enabled": False}) assert resp.status_code == 200 data = resp.json() assert data["success"] is True assert "关闭" in data["message"] def test_missing_enabled_field_returns_422(self, client): """缺少 enabled 字段应返回 422。""" resp = client.post("/toggle-auto-renew", json={}) assert resp.status_code == 422 def test_invalid_type_returns_422(self, client): """enabled 传非布尔值应返回 422。""" resp = client.post("/toggle-auto-renew", json={"enabled": [1, 2, 3]}) assert resp.status_code == 422 # --------------------------------------------------------------------------- # 8. 辅助函数 / 工具测试 # --------------------------------------------------------------------------- class TestHelperFunctions: def test_get_plan_name_known_plans(self): """已知套餐名称映射正确。""" assert subscription._get_plan_name("free") == "体验版" assert subscription._get_plan_name("standard") == "标准版" assert subscription._get_plan_name("pro") == "专业版" assert subscription._get_plan_name("enterprise") == "企业版" def test_get_plan_name_unknown(self): """未知套餐返回「未知套餐」。""" assert subscription._get_plan_name("ultra") == "未知套餐" def test_get_plan_price(self): """套餐价格映射正确。""" assert subscription._get_plan_price("free", "monthly") == 0 assert subscription._get_plan_price("standard", "monthly") == 99 assert subscription._get_plan_price("standard", "yearly") == 999 assert subscription._get_plan_price("pro", "monthly") == 299 assert subscription._get_plan_price("pro", "yearly") == 2999 assert subscription._get_plan_price("enterprise", "monthly") == 999 assert subscription._get_plan_price("enterprise", "yearly") == 9999 def test_get_plan_price_unknown(self): """未知组合返回 0。""" assert subscription._get_plan_price("ultra", "monthly") == 0 def test_plan_quotas_hardcoded(self): """配额定义硬编码,不依赖外部 registry。""" quotas = subscription.PLAN_QUOTAS assert quotas["free"] == {"max_projects": 3, "max_storage_gb": 10} assert quotas["standard"] == {"max_projects": 10, "max_storage_gb": 50} assert quotas["pro"] == {"max_projects": -1, "max_storage_gb": 100} assert quotas["enterprise"] == {"max_projects": -1, "max_storage_gb": 1000}