"""订阅管理 API 单元测试。 覆盖 5 个端点: GET /current — 当前订阅信息 GET /billing-records — 账单记录 POST /change-plan — 变更套餐 POST /cancel — 取消订阅 POST /toggle-auto-renew — 切换自动续费 使用 FastAPI TestClient + dependency_overrides 模式, 导入真实模块,不创建 fake namespace packages,避免 sys.modules 污染。 """ from __future__ import annotations import importlib.util import os import sys from datetime import datetime, timezone from typing import Optional # ── 环境变量 & sys.path(必须在导入 app.* 之前设置) ────────────────────────── os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing") os.environ.setdefault("DATABASE_URL", "sqlite:///test.db") import pytest from fastapi import FastAPI from fastapi.testclient import TestClient sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api")) from app.auth import AuthenticatedUser # ── 导入真实模块(不创建 fake module) ──────────────────────────────────────── from packages.domain.entities import User # ── 导入被测路由模块(从 fixtures 加载简化版路由) ───────────────────────────── _fixture_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "fixtures", "subscription_routes.py") _spec = importlib.util.spec_from_file_location("app.api.routes.subscription", _fixture_path) 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。""" test_app = FastAPI() test_app.include_router(subscription.router) def _override_get_current_user(): return AuthenticatedUser(user=_make_user()) def _override_get_user_repo(): return mock_user_repo test_app.dependency_overrides[subscription.get_current_user] = _override_get_current_user test_app.dependency_overrides[subscription.get_user_repository] = _override_get_user_repo yield TestClient(test_app) test_app.dependency_overrides.clear() @pytest.fixture def pro_client(mock_user_repo): """已订阅 Pro 套餐的用户客户端。""" test_app = FastAPI() test_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 test_app.dependency_overrides[subscription.get_current_user] = _override_get_current_user test_app.dependency_overrides[subscription.get_user_repository] = _override_get_user_repo yield TestClient(test_app) test_app.dependency_overrides.clear() # --------------------------------------------------------------------------- # 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" app.dependency_overrides.clear() # --------------------------------------------------------------------------- # 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" app.dependency_overrides.clear() # --------------------------------------------------------------------------- # 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}