Files
xiaoxia-saas/tests/integration/test_subscription_api.py
T
xiaoxia 1217d8cef0
Deploy / Build Production Runtime Images (push) Has been skipped
Deploy / Deploy Production (push) Has been skipped
Deploy / Production Browser E2E (push) Has been skipped
Deploy / Deploy Staging (push) Failing after 210h35m44s
CI/CD Pipeline / Frontend Lint (push) Failing after 210h36m11s
CI/CD Pipeline / Validate Code Quality And Tests (push) Failing after 210h36m17s
style: apply isort formatting to pass CI validation (#129)
2026-06-30 18:17:44 +08:00

693 lines
24 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""订阅管理 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 portsdependencies.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}