feat(#1894): 清理文案库多余字段 + 废弃标题库 API + 清理旧定价档位
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 49s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m11s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m42s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 3m44s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 4m16s
CI/CD Pipeline / Validate - Style (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
AI Code Review / AI Code Review (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been cancelled
PR Automation / Auto Approve on CI Green (pull_request) Has been cancelled

P0: 删除 scripts 表 title_text/title_category/title_config 三列
- migration 078_drop_script_title_fields 批量删列
- schemas/script.py: ScriptResponse/Create/Update 精简到 title/content/segments/tags
- services/script_service.py: 移除 get_title_config_for_script 和三字段参数
- routes/scripts.py: _to_response 不再映射三字段,create/update 不传三字段
- models.py: ScriptModel 移除三列;保留 TemplateModel/GenerationTaskModel/AiAvatarRenderJob 的 title_config(渲染层配置)

P1: 废弃 /api/v1/titles/*
- 所有子路径统一返回 HTTP 410 Gone + Deprecation/Sunset header
- 前端应直接调 GET /api/v1/scripts 取文案列表,用 title 作为标题

P1: 清理旧定价档位
- 新会员体系统一为 free/monthly/quarterly/yearly(#1951/#1955)
- schemas/subscription.py: MembershipType 替代旧 PlanType,BillingCycle 新增 QUARTERLY
- routes/subscription.py: _PLAN_NAMES 新映射;payment-callback 接受 quarterly(90天);旧档位 plan 值拒绝
- routes/_helpers.py: get_user_plan() 归一化 standard/basic→monthly、pro/premium/enterprise→quarterly
- packages/domain/quota.py: QUOTA_TIERS 改为 free/monthly/quarterly/yearly;保留 basic/standard/premium/pro/enterprise 别名兼容历史数据
- 单元测试 15527 passed
This commit is contained in:
xiaoxia-agent
2026-09-16 17:52:38 +08:00
parent 2d7f1c3a71
commit 83616ae4a3
17 changed files with 530 additions and 1037 deletions
@@ -0,0 +1,33 @@
"""#1894: drop obsolete script title fields (title_text/title_category/title_config)
Revision ID: 078_drop_script_title_fields
Revises: 077_merge_title_libs
Create Date: 2026-09-16
口播文案(scripts)不再自带配套标题、标题分类和标题样式字段。
智能剪辑 / AI 数字人等生成场景各自通过入参配置标题,不再从文案读取。
保留字段:title(名称)、content(正文)、segments(分段)、tags(标签)。
"""
import sqlalchemy as sa
from alembic import op
revision = "078_drop_script_title_fields"
down_revision = "077_merge_title_libs"
branch_labels = None
depends_on = None
def upgrade() -> None:
with op.batch_alter_table("scripts") as batch:
batch.drop_column("title_config")
batch.drop_column("title_category")
batch.drop_column("title_text")
def downgrade() -> None:
with op.batch_alter_table("scripts") as batch:
batch.add_column(sa.Column("title_text", sa.String(500), nullable=False, server_default=""))
batch.add_column(sa.Column("title_category", sa.String(50), nullable=False, server_default=""))
batch.add_column(sa.Column("title_config", sa.JSON, nullable=False, server_default="{}"))
+17 -2
View File
@@ -25,12 +25,27 @@ def check_project_access(project_id: str, user_id: str, project_repository) -> N
raise HTTPException(status_code=403, detail="无权访问该项目")
_LEGACY_PLANS = {"standard", "pro", "enterprise", "basic", "premium"}
def get_user_plan(user_id: str, user_repository: UserRepository) -> str:
"""获取用户的订阅计划名称。"""
"""获取用户的会员类型,兼容旧档位值。
旧档位 standard/pro/enterprise/basic/premium 统一映射到当前体系:
- standard/basic → monthly
- pro/premium/enterprise → quarterly
"""
user = user_repository.find_by_id(user_id)
if user is None:
return "free"
return getattr(user, "subscription_plan", "free") or "free"
plan = getattr(user, "subscription_plan", "free") or "free"
if plan in {"standard", "basic"}:
return "monthly"
if plan in {"pro", "premium", "enterprise"}:
return "quarterly"
if plan not in {"free", "monthly", "quarterly", "yearly"}:
return "free"
return plan
def require_project_and_library(
-9
View File
@@ -36,9 +36,6 @@ def _to_response(script) -> ScriptResponse:
for s in segments
],
tags=script.tags or [],
title_text=getattr(script, "title_text", "") or "",
title_category=getattr(script, "title_category", "") or "",
title_config=getattr(script, "title_config", None) or {},
created_at=script.created_at,
updated_at=script.updated_at,
)
@@ -73,9 +70,6 @@ def create_script(
content=request.content,
segments=[s.model_dump() for s in request.segments],
tags=request.tags,
title_text=request.title_text or "",
title_category=request.title_category or "",
title_config=request.title_config or {},
)
return _to_response(script)
@@ -110,9 +104,6 @@ def update_script(
content=request.content,
segments=[s.model_dump() for s in request.segments] if request.segments is not None else None,
tags=request.tags,
title_text=request.title_text,
title_category=request.title_category,
title_config=request.title_config,
)
except ScriptNotFoundError as exc:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Script not found") from exc
+54 -70
View File
@@ -11,8 +11,10 @@ from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_user_repository
from app.schemas.subscription import (
BillingRecord,
BillingCycle,
ChangePlanRequest,
ChangePlanResponse,
MembershipType,
SimpleResponse,
SubscriptionInfo,
ToggleAutoRenewRequest,
@@ -26,43 +28,18 @@ logger = logging.getLogger(__name__)
router = APIRouter()
# ============ 配额定义(硬编码,后续可迁移到配置中心) ============
# ============ 会员展示名称(与 packages.domain.points_rules.MEMBERSHIP_PRICES 对应)============
PLAN_QUOTAS = {
"free": {"max_projects": 3, "max_storage_gb": 10},
"standard": {"max_projects": 10, "max_storage_gb": 50},
"pro": {"max_projects": -1, "max_storage_gb": 100},
"enterprise": {"max_projects": -1, "max_storage_gb": 1000},
_PLAN_NAMES: dict[str, str] = {
MembershipType.FREE: "免费用户",
MembershipType.MONTHLY: "月卡会员",
MembershipType.QUARTERLY: "季卡会员",
MembershipType.YEARLY: "年卡会员",
}
# ============ Helper Functions ============
def _get_plan_name(plan_id: str) -> str:
"""获取套餐显示名称"""
plan_names = {
"free": "体验版",
"standard": "标准版",
"pro": "专业版",
"enterprise": "企业版",
}
return plan_names.get(plan_id, "未知套餐")
def _get_plan_price(plan_id: str, billing_cycle: str) -> float:
"""获取套餐价格"""
prices = {
("free", "monthly"): 0,
("free", "yearly"): 0,
("standard", "monthly"): 99,
("standard", "yearly"): 999,
("pro", "monthly"): 299,
("pro", "yearly"): 2999,
("enterprise", "monthly"): 999,
("enterprise", "yearly"): 9999,
}
return prices.get((plan_id, billing_cycle), 0)
return _PLAN_NAMES.get(plan_id, "免费用户")
def _build_subscription_info(user: AuthenticatedUser) -> SubscriptionInfo:
@@ -75,15 +52,20 @@ def _build_subscription_info(user: AuthenticatedUser) -> SubscriptionInfo:
period_start = now.isoformat()
period_end = now.isoformat()
plan_id = user.user.subscription_plan or MembershipType.FREE
# 旧档位(standard/pro/enterprise)统一降级为 monthly,避免前端炸掉
if plan_id in {"standard", "pro", "enterprise"}:
plan_id = MembershipType.MONTHLY
return SubscriptionInfo(
id=f"sub-{user.user.id[:8]}",
plan_id=user.user.subscription_plan or "free",
plan_name=_get_plan_name(user.user.subscription_plan or "free"),
plan_id=plan_id,
plan_name=_get_plan_name(plan_id),
status=user.user.subscription_status or "active",
billing_cycle="monthly",
billing_cycle=plan_id if plan_id != MembershipType.FREE else BillingCycle.MONTHLY,
current_period_start=period_start,
current_period_end=period_end,
amount=_get_plan_price(user.user.subscription_plan or "free", "monthly"),
amount=0 if plan_id == MembershipType.FREE else 0, # 金额由前端 /plans 接口展示
auto_renew=True,
created_at=user.user.created_at.isoformat() if user.user.created_at else now.isoformat(),
)
@@ -115,11 +97,11 @@ def list_membership_plans(
days = info["duration_days"]
monthly_cents = round(info["price_cents"] * 30 / days)
features: dict[str, Any] = {"max_resolution": "1080p"}
if plan_id == "monthly":
if plan_id == MembershipType.MONTHLY:
features.update({"free_clips_daily": 2})
elif plan_id == "quarterly":
elif plan_id == MembershipType.QUARTERLY:
features.update({"free_clips_daily": 5})
elif plan_id == "yearly":
elif plan_id == MembershipType.YEARLY:
features.update({"free_clips_daily": "unlimited"})
plans.append({
"plan_id": plan_id,
@@ -151,7 +133,7 @@ async def get_billing_records(
return [
BillingRecord(
id=r.id,
plan_name=r.plan_name,
plan_name=_get_plan_name(r.plan_name),
amount=r.amount,
billing_cycle=r.billing_cycle,
status=r.status,
@@ -165,6 +147,10 @@ async def get_billing_records(
session.close()
_VALID_PLANS = {MembershipType.MONTHLY, MembershipType.QUARTERLY, MembershipType.YEARLY}
_VALID_CYCLES = {BillingCycle.MONTHLY, BillingCycle.QUARTERLY, BillingCycle.YEARLY}
@router.post("/change-plan", response_model=ChangePlanResponse)
async def change_plan(
request: ChangePlanRequest,
@@ -173,47 +159,45 @@ async def change_plan(
) -> ChangePlanResponse:
"""变更订阅套餐(升级/降级)"""
# TODO: 接入支付验证(支付宝/微信支付)
valid_plans = {"free", "standard", "pro", "enterprise"}
if request.target_plan_id not in valid_plans:
target_plan = request.target_plan_id
if target_plan not in _VALID_PLANS:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"无效的套餐ID。支持的套餐: {', '.join(valid_plans)}",
detail=f"无效的会员类型。支持: {', '.join(sorted(_VALID_PLANS))}",
)
valid_cycles = {"monthly", "yearly"}
if request.billing_cycle not in valid_cycles:
if request.billing_cycle not in _VALID_CYCLES:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="无效的计费周期。支持: monthly, yearly",
detail=f"无效的计费周期。支持: {', '.join(sorted(_VALID_CYCLES))}",
)
user = current_user.user
current_plan = user.subscription_plan or "free"
target_plan = request.target_plan_id
current_plan = user.subscription_plan or MembershipType.FREE
# 旧档位归一化,避免永远显示"您已经是xxx"
if current_plan in {"standard", "pro", "enterprise"}:
current_plan = MembershipType.MONTHLY
if current_plan == target_plan:
return ChangePlanResponse(
success=False,
message=f"您已经是 {_get_plan_name(target_plan)}",
message=f"您已经是{_get_plan_name(target_plan)}",
)
# 通过 dataclasses.replace 创建新实例(不直接修改 dataclass)
quotas = PLAN_QUOTAS.get(target_plan, PLAN_QUOTAS["free"])
updated_user = replace(
user,
subscription_plan=target_plan,
subscription_status="active",
max_projects=quotas["max_projects"],
max_storage_gb=quotas["max_storage_gb"],
max_projects=-1, # 付费会员不限项目数
max_storage_gb=100,
)
user_repository.save(updated_user)
# 用更新后的用户构造响应
refreshed_auth_user = AuthenticatedUser(user=updated_user)
return ChangePlanResponse(
success=True,
message=f"套餐已成功变更为 {_get_plan_name(target_plan)}",
message=f"套餐已成功变更为{_get_plan_name(target_plan)}",
new_subscription=_build_subscription_info(refreshed_auth_user),
)
@@ -225,10 +209,11 @@ async def cancel_subscription(
) -> SimpleResponse:
"""取消订阅"""
user = current_user.user
if user.subscription_plan == "free":
plan_id = user.subscription_plan or MembershipType.FREE
if plan_id == MembershipType.FREE:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="体验版无需取消",
detail="免费用户无需取消订阅",
)
updated_user = replace(user, subscription_status="cancelled")
@@ -236,7 +221,7 @@ async def cancel_subscription(
return SimpleResponse(
success=True,
message="订阅已取消,当前周期结束后停止服务",
message="订阅已取消,当前周期结束后将降级为免费用户",
)
@@ -262,11 +247,14 @@ async def payment_callback(
if SessionLocal is None:
raise HTTPException(status_code=500, detail="Database not available")
# 仅接受当前会员体系的 plan 值
if plan not in _VALID_PLANS:
raise HTTPException(status_code=400, detail=f"未知的会员类型: {plan}")
session = SessionLocal()
try:
repo = SQLAlchemyBillingRepository(session)
# 创建账单记录
record_id = uuid.uuid4().hex
repo.create(
{
@@ -279,19 +267,20 @@ async def payment_callback(
}
)
# 在事务中标记支付成功并更新订阅
repo.mark_paid(record_id, payment_method, payment_id)
# 计算到期时间
days = 365 if billing_cycle == "yearly" else 30
days_map = {BillingCycle.MONTHLY: 30, BillingCycle.QUARTERLY: 90, BillingCycle.YEARLY: 365}
days = days_map.get(billing_cycle, 30)
expires_at = datetime.now(UTC) + timedelta(days=days)
repo.update_subscription_on_payment(user_id, plan, expires_at)
return {"success": True, "message": "支付成功", "record_id": record_id}
except HTTPException:
session.rollback()
raise
except Exception as e:
session.rollback()
logger.error(f"支付回调处理失败: user_id={user_id}, plan={plan}, error={e}")
# 不返回原始异常信息,避免泄漏内部实现细节
logger.error("支付回调处理失败: user_id=%s, plan=%s, error=%s", user_id, plan, e)
raise HTTPException(status_code=500, detail="支付处理失败,请稍后重试") from e
finally:
session.close()
@@ -303,10 +292,5 @@ async def toggle_auto_renew(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> SimpleResponse:
"""切换自动续费"""
# TODO: 实际需要在数据库中存储 auto_renew 字段
status_text = "已开启自动续费" if request.enabled else "已关闭自动续费"
return SimpleResponse(
success=True,
message=status_text,
)
return SimpleResponse(success=True, message=status_text)
+20 -228
View File
@@ -1,243 +1,35 @@
"""Title library CRUD routes.
"""Title library routes — DEPRECATED (#1894).
.. deprecated::
标题库 API 已废弃(#1894),标题配置已整合到 scripts 模型。
所有接口保留向后兼容,但返回 Warning header 并记录日志。
独立标题库已废弃。前端应直接调用 GET /api/v1/scripts 获取文案列表,
取每条文案的 `title` 字段作为标题候选。
所有 /api/v1/titles 端点统一返回 HTTP 410 Gone。
"""
from __future__ import annotations
import logging
from typing import Optional
from app.api.routes._helpers import get_user_plan
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session, get_user_repository
from app.schemas.title_library import (
CreateTitleLibraryRequest,
ListTitleLibraryResponse,
TitleLibraryItemResponse,
UpdateTitleLibraryRequest,
)
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.title_library_repository import SQLAlchemyTitleLibraryRepository
from packages.application.title_library.commands import (
CreateTitleLibraryCommand,
PickTitleCommand,
UpdateTitleLibraryCommand,
)
from packages.application.title_library.use_cases import (
CreateTitleLibraryUseCase,
DeleteTitleLibraryUseCase,
GetTitleLibraryUseCase,
ListTitleLibraryUseCase,
NotFoundError,
PickTitleUseCase,
QuotaExceededError,
UpdateTitleLibraryUseCase,
)
from packages.ports.user_repository import UserRepository
from fastapi import APIRouter, Response, status
router = APIRouter()
logger = logging.getLogger(__name__)
_DEPRECATION_WARNING = (
'299 - "Title library API is deprecated; migrate to scripts.title_text/'
'title_category/title_config (issue #1894)"'
_GONE_MESSAGE = (
"标题库 API 已废弃(#1894):独立标题库已合并进文案库,"
"请使用 GET /api/v1/scripts 获取文案列表并取 title 字段作为标题。"
)
def _deprecation_headers() -> dict:
"""返回 deprecation Warning header (ASCII-only, RFC 7234 §5.5)."""
return {"Warning": _DEPRECATION_WARNING, "Deprecation": "true"}
def _gone(response: Response) -> dict:
response.status_code = status.HTTP_410_GONE
response.headers["Deprecation"] = "true"
response.headers["Sunset"] = "Tue, 16 Sep 2026 00:00:00 GMT"
return {"error": {"code": "GONE", "message": _GONE_MESSAGE}}
def _log_deprecation(endpoint: str) -> None:
logger.warning("[Deprecated] title_library API 调用: %s — %s", endpoint, _DEPRECATION_WARNING)
@router.api_route("", methods=["GET", "POST", "PUT", "DELETE", "PATCH"])
def titles_root_gone(response: Response) -> dict:
return _gone(response)
def _get_title_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyTitleLibraryRepository:
return SQLAlchemyTitleLibraryRepository(session)
def _to_response(item) -> TitleLibraryItemResponse:
return TitleLibraryItemResponse(
id=item.id,
user_id=item.user_id,
name=item.name,
text=item.text,
category=item.category,
description=item.description,
tags=item.tags,
usage_count=item.usage_count,
is_active=item.is_active,
created_at=item.created_at,
updated_at=item.updated_at,
)
@router.get("", response_model=ListTitleLibraryResponse)
def list_titles(
response: Response,
category: Optional[str] = Query(None),
skip: int = Query(0, ge=0),
limit: int = Query(50, ge=1, le=200),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
) -> ListTitleLibraryResponse:
"""[Deprecated] 请使用 scripts API 的 title_text/title_category 字段替代."""
_log_deprecation("list_titles")
for k, v in _deprecation_headers().items():
response.headers[k] = v
user_id = authenticated_user.user.id
use_case = ListTitleLibraryUseCase(title_repository)
items = use_case.execute(user_id, category=category, skip=skip, limit=limit)
total = title_repository.count_by_user(user_id)
return ListTitleLibraryResponse(
items=[_to_response(i) for i in items],
total=total,
)
@router.post("/pick", response_model=TitleLibraryItemResponse)
def pick_title(
response: Response,
category: Optional[str] = Query(None, description="按分类筛选,不填则从全部标题中选"),
exclude_ids: Optional[str] = Query(
None,
description="排除的标题ID(逗号分隔),用于批量生成时避免重复",
),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
) -> TitleLibraryItemResponse:
"""[Deprecated] 请使用 scripts API 的 title_text/title_category 字段替代.
智能选择一个标题。
策略:优先使用次数少的,从最少的前5个中随机选一个,兼顾公平和多样性。
"""
_log_deprecation("pick_title")
for k, v in _deprecation_headers().items():
response.headers[k] = v
user_id = authenticated_user.user.id
exclude_list: list[str] = []
if exclude_ids:
exclude_list = [t.strip() for t in exclude_ids.split(",") if t.strip()]
use_case = PickTitleUseCase(title_repository)
item = use_case.execute(
PickTitleCommand(
user_id=user_id,
category=category,
exclude_ids=exclude_list,
)
)
if item is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="标题库为空,请先添加标题",
)
return _to_response(item)
@router.get("/{title_id}", response_model=TitleLibraryItemResponse)
def get_title(
title_id: str,
response: Response,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
) -> TitleLibraryItemResponse:
"""[Deprecated] 请使用 scripts API 的 title_text/title_category 字段替代."""
_log_deprecation("get_title")
for k, v in _deprecation_headers().items():
response.headers[k] = v
user_id = authenticated_user.user.id
use_case = GetTitleLibraryUseCase(title_repository)
item = use_case.execute(title_id, user_id)
if item is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Title not found")
return _to_response(item)
@router.post("", response_model=TitleLibraryItemResponse, status_code=status.HTTP_201_CREATED)
def create_title(
response: Response,
request: CreateTitleLibraryRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
user_repository: UserRepository = Depends(get_user_repository),
) -> TitleLibraryItemResponse:
"""[Deprecated] 请使用 scripts API 的 title_text/title_category 字段替代."""
_log_deprecation("create_title")
for k, v in _deprecation_headers().items():
response.headers[k] = v
user_id = authenticated_user.user.id
plan_name = get_user_plan(user_id, user_repository)
command = CreateTitleLibraryCommand(
user_id=user_id,
name=request.name,
text=request.text,
category=request.category,
description=request.description,
tags=request.tags,
)
use_case = CreateTitleLibraryUseCase(title_repository)
try:
item = use_case.execute(command, plan_name=plan_name)
except QuotaExceededError as exc:
raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail=f"标题库配额已满({exc.used}/{exc.limit}),请升级套餐",
) from exc
return _to_response(item)
@router.put("/{title_id}", response_model=TitleLibraryItemResponse)
def update_title(
title_id: str,
response: Response,
request: UpdateTitleLibraryRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
) -> TitleLibraryItemResponse:
"""[Deprecated] 请使用 scripts API 的 title_text/title_category 字段替代."""
_log_deprecation("update_title")
for k, v in _deprecation_headers().items():
response.headers[k] = v
user_id = authenticated_user.user.id
command = UpdateTitleLibraryCommand(
title_id=title_id,
user_id=user_id,
name=request.name,
text=request.text,
category=request.category,
description=request.description,
tags=request.tags,
)
use_case = UpdateTitleLibraryUseCase(title_repository)
try:
item = use_case.execute(command)
except NotFoundError as _e:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Title not found") from _e
return _to_response(item)
@router.delete("/{title_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response)
def delete_title(
title_id: str,
response: Response,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
) -> Response:
"""[Deprecated] 请使用 scripts API 的 title_text/title_category 字段替代."""
_log_deprecation("delete_title")
for k, v in _deprecation_headers().items():
response.headers[k] = v
user_id = authenticated_user.user.id
use_case = DeleteTitleLibraryUseCase(title_repository)
deleted = use_case.execute(title_id, user_id)
if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Title not found")
return
@router.api_route("/{path:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"])
def titles_subpath_gone(response: Response, path: str) -> dict:
return _gone(response)
+1 -10
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Any, Dict, Optional
from typing import Optional
from pydantic import BaseModel, Field
@@ -22,9 +22,6 @@ class ScriptResponse(BaseModel):
content: str
segments: list[ScriptSegment] = Field(default_factory=list)
tags: list[str] = Field(default_factory=list)
title_text: str = ""
title_category: str = ""
title_config: Dict[str, Any] = Field(default_factory=dict)
created_at: datetime
updated_at: datetime
@@ -39,9 +36,6 @@ class CreateScriptRequest(BaseModel):
content: str = ""
segments: list[ScriptSegment] = Field(default_factory=list)
tags: list[str] = Field(default_factory=list)
title_text: str = ""
title_category: str = ""
title_config: Optional[Dict[str, Any]] = None
class UpdateScriptRequest(BaseModel):
@@ -49,6 +43,3 @@ class UpdateScriptRequest(BaseModel):
content: Optional[str] = None
segments: Optional[list[ScriptSegment]] = None
tags: Optional[list[str]] = None
title_text: Optional[str] = None
title_category: Optional[str] = None
title_config: Optional[Dict[str, Any]] = None
+15 -7
View File
@@ -6,16 +6,23 @@ from typing import Optional
from pydantic import BaseModel, Field
# ============ Enums / Types ============
# 会员体系(#1951/#1955 实装):
# free — 免费用户
# monthly — 月卡
# quarterly — 季卡
# yearly — 年卡
# 已废弃档位:standard / pro / enterprise(保留常量名便于识别旧字段,但不在 API 中暴露)
class PlanType(str):
"""套餐类型"""
class MembershipType(str):
"""会员类型(与 packages.domain.points_rules.MEMBERSHIP_PRICES 一致)"""
FREE = "free"
STANDARD = "standard"
PRO = "pro"
ENTERPRISE = "enterprise"
MONTHLY = "monthly"
QUARTERLY = "quarterly"
YEARLY = "yearly"
class SubscriptionStatus(str):
@@ -40,6 +47,7 @@ class BillingCycle(str):
"""计费周期"""
MONTHLY = "monthly"
QUARTERLY = "quarterly"
YEARLY = "yearly"
@@ -95,8 +103,8 @@ class SimpleResponse(BaseModel):
class ChangePlanRequest(BaseModel):
"""升级/降级请求"""
target_plan_id: str = Field(..., description="目标套餐ID")
billing_cycle: str = Field(..., description="计费周期: monthly/yearly")
target_plan_id: str = Field(..., description="目标会员类型: monthly/quarterly/yearly")
billing_cycle: str = Field(..., description="计费周期: monthly/quarterly/yearly")
class ToggleAutoRenewRequest(BaseModel):
-25
View File
@@ -51,9 +51,6 @@ class ScriptService:
content: str = "",
segments: list | None = None,
tags: list | None = None,
title_text: str = "",
title_category: str = "",
title_config: dict | None = None,
) -> ScriptModel:
script = ScriptModel(
id=str(uuid.uuid4()),
@@ -62,9 +59,6 @@ class ScriptService:
content=content,
segments=segments if segments is not None else [],
tags=tags if tags is not None else [],
title_text=title_text or "",
title_category=title_category or "",
title_config=title_config if title_config is not None else {},
)
self.db.add(script)
self.db.commit()
@@ -89,9 +83,6 @@ class ScriptService:
content: Optional[str] = None,
segments: Optional[list] = None,
tags: Optional[list] = None,
title_text: Optional[str] = None,
title_category: Optional[str] = None,
title_config: Optional[dict] = None,
) -> ScriptModel:
script = self.get_script(script_id, user_id)
if title is not None:
@@ -102,27 +93,11 @@ class ScriptService:
script.segments = segments
if tags is not None:
script.tags = tags
if title_text is not None:
script.title_text = title_text
if title_category is not None:
script.title_category = title_category
if title_config is not None:
script.title_config = title_config
script.updated_at = datetime.now(UTC)
self.db.commit()
self.db.refresh(script)
return script
# ── title config ─────────────────────────────────────────────────────
def get_title_config_for_script(self, script_id: str, user_id: str) -> dict:
"""从 script 读取标题配置,返回可直接用于渲染的 title_config dict."""
script = self.get_script(script_id, user_id)
config = dict(script.title_config or {})
if not config.get("text") and script.title_text:
config["text"] = script.title_text
return config
# ── delete ────────────────────────────────────────────────────────────
def delete_script(self, script_id: str, user_id: str) -> bool:
@@ -671,10 +671,6 @@ class ScriptModel(Base):
content = Column(Text, nullable=False, default="")
segments = Column(JSON, nullable=False, default=list)
tags = Column(JSON, nullable=False, default=list)
# #1894: 废弃标题库整合到文案库 — 标题配置字段
title_text = Column(String(500), nullable=False, default="")
title_category = Column(String(50), nullable=False, default="")
title_config = Column(JSON, nullable=False, default=dict)
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
+40 -17
View File
@@ -1,13 +1,14 @@
"""Quota system with registry pattern.
Four subscription tiers with different limits:
- free: 2GB storage, 5 videos/month, 3 concurrent, 3 templates, 50 titles, 10 voiceovers, no AI voice
- basic: 20GB storage, 30 videos/month, 10 concurrent, 15 templates, 500 titles, 100 voiceovers, AI voice
- premium: 100GB storage, 100 videos/month, 20 concurrent, unlimited templates, 500 titles, 100 voiceovers, AI voice
- pro: Same as premium (alias for premium tier)
Member tiers (see packages.domain.points_rules.MEMBERSHIP_PRICES):
- free: 2GB storage, 5 videos/month, 3 concurrent, 3 templates, 10 voiceovers, no AI voice
- monthly: 月卡会员(同 basic 级别)
- quarterly: 季卡会员(同 premium 级别)
- yearly: 年卡会员(同 premium 级别,更多每日免费额度)
旧档位(standard/pro/enterprise/basic/premium)已在 #1894 清理,统一为 free/monthly/quarterly/yearly。
Quota dimensions are registered by modules via the ModuleRegistry,
and checked against the user's subscription plan.
and checked against the user's membership type.
"""
from __future__ import annotations
@@ -59,7 +60,6 @@ QUOTA_TIERS: dict[str, QuotaTier] = {
QuotaDimension.VIDEOS_PER_MONTH: 5,
QuotaDimension.MAX_CONCURRENT: 3,
QuotaDimension.MAX_TEMPLATES: 3,
QuotaDimension.MAX_TITLES: 50,
QuotaDimension.MAX_VOICEOVERS: 10,
QuotaDimension.AI_VOICE_ENABLED: 0,
QuotaDimension.AI_VOICE_CREDITS: 0,
@@ -68,14 +68,14 @@ QUOTA_TIERS: dict[str, QuotaTier] = {
QuotaDimension.DEDUP_REPORT_ENABLED: 0,
},
),
"basic": QuotaTier(
name="basic",
# 月卡会员:基础付费档(原 basic)
"monthly": QuotaTier(
name="monthly",
limits={
QuotaDimension.STORAGE_GB: 20,
QuotaDimension.VIDEOS_PER_MONTH: 30,
QuotaDimension.MAX_CONCURRENT: 10,
QuotaDimension.MAX_TEMPLATES: 15,
QuotaDimension.MAX_TITLES: 500,
QuotaDimension.MAX_VOICEOVERS: 100,
QuotaDimension.AI_VOICE_ENABLED: 1,
QuotaDimension.AI_VOICE_CREDITS: 100,
@@ -84,14 +84,14 @@ QUOTA_TIERS: dict[str, QuotaTier] = {
QuotaDimension.DEDUP_REPORT_ENABLED: 0,
},
),
"premium": QuotaTier(
name="premium",
# 季卡会员:高级付费档(原 premium)
"quarterly": QuotaTier(
name="quarterly",
limits={
QuotaDimension.STORAGE_GB: 100,
QuotaDimension.VIDEOS_PER_MONTH: 100,
QuotaDimension.MAX_CONCURRENT: 20,
QuotaDimension.MAX_TEMPLATES: float("inf"), # 不限量
QuotaDimension.MAX_TITLES: 500,
QuotaDimension.MAX_TEMPLATES: float("inf"),
QuotaDimension.MAX_VOICEOVERS: 100,
QuotaDimension.AI_VOICE_ENABLED: 1,
QuotaDimension.AI_VOICE_CREDITS: 500,
@@ -100,9 +100,32 @@ QUOTA_TIERS: dict[str, QuotaTier] = {
QuotaDimension.DEDUP_REPORT_ENABLED: 1,
},
),
# 年卡会员:同季卡配额 + 每日不限免费条数(由前端/积分规则实现)
"yearly": QuotaTier(
name="yearly",
limits={
QuotaDimension.STORAGE_GB: 100,
QuotaDimension.VIDEOS_PER_MONTH: float("inf"),
QuotaDimension.MAX_CONCURRENT: 20,
QuotaDimension.MAX_TEMPLATES: float("inf"),
QuotaDimension.MAX_VOICEOVERS: 200,
QuotaDimension.AI_VOICE_ENABLED: 1,
QuotaDimension.AI_VOICE_CREDITS: 2000,
QuotaDimension.BATCH_EXPORT_ENABLED: 1,
QuotaDimension.MULTI_PLATFORM_ENABLED: 1,
QuotaDimension.DEDUP_REPORT_ENABLED: 1,
},
),
}
# pro 套餐与 premium 配额相同,使用别名引用避免重复维护
QUOTA_TIERS["pro"] = QUOTA_TIERS["premium"]
# #1894: 旧档位别名(basic/standard → monthly, premium/pro/enterprise → quarterly)
# 历史 DB 数据、单测和内部模块可能仍在传旧 plan_name;这里保留别名保证配额查询不炸。
# 新代码请统一使用 free/monthly/quarterly/yearly。
QUOTA_TIERS["basic"] = QUOTA_TIERS["monthly"]
QUOTA_TIERS["standard"] = QUOTA_TIERS["monthly"]
QUOTA_TIERS["premium"] = QUOTA_TIERS["quarterly"]
QUOTA_TIERS["pro"] = QUOTA_TIERS["quarterly"]
QUOTA_TIERS["enterprise"] = QUOTA_TIERS["quarterly"]
class QuotaWarningLevel:
@@ -216,7 +239,7 @@ class QuotaChecker:
"""检查指定维度的配额使用情况
Args:
plan_name: 用户套餐等级 (free/basic/premium)
plan_name: 会员类型 (free/monthly/quarterly/yearly)
dimension: 配额维度
used: 当前已使用量
+18 -11
View File
@@ -354,13 +354,18 @@ class TestQuotaRegistry:
assert len(reg.list_dimensions()) == len(QuotaDimension)
def test_list_tiers(self):
"""四个套餐等级."""
"""套餐等级包含核心四档 + 旧档位别名."""
reg = QuotaRegistry()
tiers = reg.list_tiers()
assert "pro" in tiers
assert "free" in tiers
assert "basic" in tiers
assert len(tiers) == 4
assert "monthly" in tiers
assert "quarterly" in tiers
assert "yearly" in tiers
assert "basic" in tiers # alias → monthly
assert "premium" in tiers # alias → quarterly
assert "pro" in tiers # alias → quarterly
assert "standard" in tiers # alias → monthly
assert "enterprise" in tiers # alias → quarterly
def test_get_tier_existing(self):
"""获取已有的套餐."""
@@ -370,9 +375,11 @@ class TestQuotaRegistry:
assert tier.name == "free"
def test_get_tier_nonexistent(self):
"""不存在的套餐返回 None"""
"""不存在的套餐返回 None(enterprise 现为 quarterly 别名)."""
from packages.domain.quota import QUOTA_TIERS
reg = QuotaRegistry()
assert reg.get_tier("enterprise") is None
assert reg.get_tier("totally_unknown_plan_xyz") is None
assert reg.get_tier("enterprise") is QUOTA_TIERS["quarterly"]
def test_get_limit_existing(self):
"""获取已有限制."""
@@ -380,9 +387,10 @@ class TestQuotaRegistry:
assert reg.get_limit("free", QuotaDimension.STORAGE_GB) == 2
def test_get_limit_nonexistent_plan(self):
"""不存在的套餐 fallback 到 free 配额"""
"""不存在的套餐返回 0;enterprise 现为 quarterly 别名,返回 100."""
reg = QuotaRegistry()
assert reg.get_limit("enterprise", QuotaDimension.STORAGE_GB) == 0
assert reg.get_limit("totally_unknown_plan_xyz", QuotaDimension.STORAGE_GB) == 0
assert reg.get_limit("enterprise", QuotaDimension.STORAGE_GB) == 100
def test_get_limit_unknown_dimension(self):
"""未知维度返回 0."""
@@ -506,11 +514,10 @@ class TestQuotaChecker:
assert result.usage_percent == 0.0
def test_check_unknown_plan(self):
"""未知套餐,限制为0."""
"""未知套餐,限制为0(enterprise现为quarterly别名,这里用一个真不存在的名)."""
checker = QuotaChecker()
result = checker.check("enterprise", QuotaDimension.STORAGE_GB, 0)
result = checker.check("totally_unknown_plan_xyz", QuotaDimension.STORAGE_GB, 0)
assert result.limit == 0
# used=0, limit=0 → 0 < 0 is False → allowed=False
assert result.allowed is False
assert result.warning_level == QuotaWarningLevel.NORMAL
+34 -46
View File
@@ -31,7 +31,7 @@ from fastapi.testclient import TestClient
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
from app.api.routes.subscription import _get_plan_name, _get_plan_price, router
from app.api.routes.subscription import _get_plan_name, router
# ---------------------------------------------------------------------------
# 1. Mock Billing Repository
@@ -140,8 +140,8 @@ class TestPaymentCallbackSuccess:
@patch("packages.adapters.sqlalchemy_impl.billing_repository.SQLAlchemyBillingRepository")
@patch("packages.adapters.sqlalchemy_impl.session.SessionLocal")
def test_monthly_pro_payment_success(self, MockSession, MockRepo):
"""Pro 套餐月付支付成功。"""
def test_monthly_payment_success(self, MockSession, MockRepo):
"""月卡支付成功。"""
mock_repo = MockBillingRepository()
MockRepo.return_value = mock_repo
MockSession.return_value = MagicMock()
@@ -154,7 +154,7 @@ class TestPaymentCallbackSuccess:
"/subscription/payment-callback",
params={
"user_id": "user-001",
"plan": "pro",
"plan": "monthly",
"billing_cycle": "monthly",
"amount": 299.0,
"payment_method": "alipay",
@@ -175,12 +175,12 @@ class TestPaymentCallbackSuccess:
# 验证订阅更新
assert mock_repo.update_subscription_count == 1
assert "user-001" in mock_repo.updated_subscriptions
assert mock_repo.updated_subscriptions["user-001"]["plan"] == "pro"
assert mock_repo.updated_subscriptions["user-001"]["plan"] == "monthly"
@patch("packages.adapters.sqlalchemy_impl.billing_repository.SQLAlchemyBillingRepository")
@patch("packages.adapters.sqlalchemy_impl.session.SessionLocal")
def test_yearly_standard_payment_success(self, MockSession, MockRepo):
"""标准版年付支付成功。"""
def test_yearly_payment_success(self, MockSession, MockRepo):
"""年卡支付成功。"""
mock_repo = MockBillingRepository()
MockRepo.return_value = mock_repo
MockSession.return_value = MagicMock()
@@ -193,7 +193,7 @@ class TestPaymentCallbackSuccess:
"/subscription/payment-callback",
params={
"user_id": "user-002",
"plan": "standard",
"plan": "monthly",
"billing_cycle": "yearly",
"amount": 999.0,
"payment_method": "wechat",
@@ -204,7 +204,7 @@ class TestPaymentCallbackSuccess:
assert resp.status_code == 200
data = resp.json()
assert data["success"] is True
assert mock_repo.updated_subscriptions["user-002"]["plan"] == "standard"
assert mock_repo.updated_subscriptions["user-002"]["plan"] == "monthly"
# 年付到期时间应为约 365 天后
expires_at = mock_repo.updated_subscriptions["user-002"]["expires_at"]
expected = datetime.now(UTC) + timedelta(days=365)
@@ -212,8 +212,8 @@ class TestPaymentCallbackSuccess:
@patch("packages.adapters.sqlalchemy_impl.billing_repository.SQLAlchemyBillingRepository")
@patch("packages.adapters.sqlalchemy_impl.session.SessionLocal")
def test_enterprise_payment_success(self, MockSession, MockRepo):
"""企业版支付成功。"""
def test_quarterly_payment_success(self, MockSession, MockRepo):
"""季卡支付成功。"""
mock_repo = MockBillingRepository()
MockRepo.return_value = mock_repo
MockSession.return_value = MagicMock()
@@ -226,8 +226,8 @@ class TestPaymentCallbackSuccess:
"/subscription/payment-callback",
params={
"user_id": "user-003",
"plan": "enterprise",
"billing_cycle": "monthly",
"plan": "quarterly",
"billing_cycle": "quarterly",
"amount": 999.0,
"payment_method": "bank_transfer",
"payment_id": "ent_20240101_003",
@@ -236,7 +236,7 @@ class TestPaymentCallbackSuccess:
assert resp.status_code == 200
assert resp.json()["success"] is True
assert mock_repo.updated_subscriptions["user-003"]["plan"] == "enterprise"
assert mock_repo.updated_subscriptions["user-003"]["plan"] == "quarterly"
@patch("packages.adapters.sqlalchemy_impl.billing_repository.SQLAlchemyBillingRepository")
@patch("packages.adapters.sqlalchemy_impl.session.SessionLocal")
@@ -254,7 +254,7 @@ class TestPaymentCallbackSuccess:
"/subscription/payment-callback",
params={
"user_id": "user-004",
"plan": "standard",
"plan": "monthly",
"billing_cycle": "monthly",
"amount": 99.0,
},
@@ -288,7 +288,7 @@ class TestPaymentCallbackIdempotency:
params = {
"user_id": "user-idem-1",
"plan": "pro",
"plan": "monthly",
"billing_cycle": "monthly",
"amount": 299.0,
"payment_id": "pay_dup_001",
@@ -351,7 +351,7 @@ class TestPaymentCallbackValidation:
resp = client.post(
"/subscription/payment-callback",
params={"plan": "pro", "billing_cycle": "monthly", "amount": 299.0},
params={"plan": "monthly", "billing_cycle": "monthly", "amount": 299.0},
)
assert resp.status_code == 422
@@ -385,7 +385,7 @@ class TestPaymentCallbackValidation:
resp = client.post(
"/subscription/payment-callback",
params={"user_id": "u1", "plan": "pro", "billing_cycle": "monthly"},
params={"user_id": "u1", "plan": "monthly", "billing_cycle": "monthly"},
)
assert resp.status_code == 422
@@ -405,7 +405,7 @@ class TestPaymentCallbackValidation:
"/subscription/payment-callback",
params={
"user_id": "u1",
"plan": "pro",
"plan": "monthly",
"billing_cycle": "monthly",
"amount": -100.0,
},
@@ -420,35 +420,23 @@ class TestPaymentCallbackValidation:
class TestHelperFunctions:
"""订阅辅助函数测试。"""
"""订阅辅助函数测试 — #1894 新档位 free/monthly/quarterly/yearly。"""
def test_get_plan_name_all_plans(self):
"""所有套餐名称映射正确。"""
assert _get_plan_name("free") == "体验版"
assert _get_plan_name("standard") == "标准版"
assert _get_plan_name("pro") == "专业版"
assert _get_plan_name("enterprise") == "企业版"
assert _get_plan_name("free") == "免费用户"
assert _get_plan_name("monthly") == "月卡会员"
assert _get_plan_name("quarterly") == "季卡会员"
assert _get_plan_name("yearly") == "年卡会员"
def test_get_plan_name_unknown(self):
"""未知套餐返回「未知套餐」。"""
assert _get_plan_name("unknown") == "未知套餐"
assert _get_plan_name("") == "未知套餐"
def test_get_plan_price_all_combinations(self):
"""所有套餐价格映射正确。"""
assert _get_plan_price("free", "monthly") == 0
assert _get_plan_price("free", "yearly") == 0
assert _get_plan_price("standard", "monthly") == 99
assert _get_plan_price("standard", "yearly") == 999
assert _get_plan_price("pro", "monthly") == 299
assert _get_plan_price("pro", "yearly") == 2999
assert _get_plan_price("enterprise", "monthly") == 999
assert _get_plan_price("enterprise", "yearly") == 9999
def test_get_plan_price_unknown(self):
"""未知组合返回 0。"""
assert _get_plan_price("unknown", "monthly") == 0
assert _get_plan_price("pro", "weekly") == 0
def test_get_plan_name_unknown_defaults_free(self):
"""未知套餐返回默认「免费用户」。"""
assert _get_plan_name("unknown") == "免费用户"
assert _get_plan_name("") == "免费用户"
# legacy 旧值不直接命中 → 也回落免费用户(实际会被 _helpers.get_user_plan 归一化到 monthly/quarterly)
assert _get_plan_name("standard") == "免费用户"
assert _get_plan_name("pro") == "免费用户"
assert _get_plan_name("enterprise") == "免费用户"
# ---------------------------------------------------------------------------
@@ -568,12 +556,12 @@ class TestMockBillingRepository:
repo = MockBillingRepository()
expires = datetime.now(UTC) + timedelta(days=30)
repo.update_subscription_on_payment("user-001", "pro", expires)
repo.update_subscription_on_payment("user-001", "monthly", expires)
assert repo.update_subscription_count == 1
assert "user-001" in repo.updated_subscriptions
sub = repo.updated_subscriptions["user-001"]
assert sub["plan"] == "pro"
assert sub["plan"] == "monthly"
assert sub["status"] == "active"
assert sub["expires_at"] == expires
+89 -139
View File
@@ -1,4 +1,4 @@
"""Quota 配额系统单测 — 全维度覆盖."""
"""Quota 配额系统单测 — #1894 档位清理后版本 (free/monthly/quarterly/yearly)."""
from __future__ import annotations
@@ -19,6 +19,7 @@ from packages.domain.quota import (
quota_registry,
)
# ── 枚举与常量 ──────────────────────────────────────────────────────────────
@@ -39,7 +40,6 @@ class TestQuotaDimension:
assert QuotaDimension.AI_VOICE_ENABLED == "ai_voice_enabled"
def test_all_dimensions_count(self):
# 至少包含内置的几个核心维度
dims = list(QuotaDimension)
assert len(dims) >= 7
@@ -82,11 +82,13 @@ class TestQuotaTier:
assert tier.is_unlimited("storage_gb") is False
def test_is_unlimited_undefined_defaults_true(self):
# 未定义的维度,get 默认为 inf → is_unlimited 返回 True
tier = QuotaTier(name="test", limits={})
assert tier.is_unlimited("unknown") is True
# ── QuotaTiers — #1894 新档位: free / monthly / quarterly / yearly ─────────
class TestQuotaTiers:
def test_free_tier_exists(self):
assert "free" in QUOTA_TIERS
@@ -95,36 +97,56 @@ class TestQuotaTiers:
assert free.get_limit("storage_gb") == 2
assert free.get_limit("videos_per_month") == 5
def test_basic_tier_exists(self):
assert "basic" in QUOTA_TIERS
basic = QUOTA_TIERS["basic"]
assert basic.get_limit("storage_gb") == 20
assert basic.get_limit("videos_per_month") == 30
assert basic.get_limit("ai_voice_enabled") == 1
def test_monthly_tier_exists(self):
assert "monthly" in QUOTA_TIERS
monthly = QUOTA_TIERS["monthly"]
assert monthly.get_limit("storage_gb") == 20
assert monthly.get_limit("videos_per_month") == 30
assert monthly.get_limit("ai_voice_enabled") == 1
def test_premium_tier_exists(self):
assert "premium" in QUOTA_TIERS
premium = QUOTA_TIERS["premium"]
assert premium.get_limit("storage_gb") == 100
assert premium.get_limit("videos_per_month") == 100
def test_quarterly_tier_exists(self):
assert "quarterly" in QUOTA_TIERS
quarterly = QUOTA_TIERS["quarterly"]
assert quarterly.get_limit("storage_gb") == 100
assert quarterly.get_limit("videos_per_month") == 100
def test_premium_templates_unlimited(self):
premium = QUOTA_TIERS["premium"]
assert premium.is_unlimited("max_templates") is True
def test_yearly_tier_exists(self):
assert "yearly" in QUOTA_TIERS
yearly = QUOTA_TIERS["yearly"]
assert yearly.get_limit("storage_gb") == 100
assert yearly.is_unlimited("videos_per_month") is True
assert yearly.get_limit("ai_voice_credits") == 2000
def test_quarterly_templates_unlimited(self):
quarterly = QUOTA_TIERS["quarterly"]
assert quarterly.is_unlimited("max_templates") is True
def test_yearly_templates_unlimited(self):
yearly = QUOTA_TIERS["yearly"]
assert yearly.is_unlimited("max_templates") is True
def test_free_ai_voice_disabled(self):
free = QUOTA_TIERS["free"]
assert free.get_limit("ai_voice_enabled") == 0
def test_basic_ai_voice_enabled(self):
basic = QUOTA_TIERS["basic"]
assert basic.get_limit("ai_voice_enabled") == 1
def test_monthly_ai_voice_enabled(self):
monthly = QUOTA_TIERS["monthly"]
assert monthly.get_limit("ai_voice_enabled") == 1
def test_storage_increases_with_tier(self):
free = QUOTA_TIERS["free"].get_limit("storage_gb")
basic = QUOTA_TIERS["basic"].get_limit("storage_gb")
premium = QUOTA_TIERS["premium"].get_limit("storage_gb")
assert free < basic < premium
monthly = QUOTA_TIERS["monthly"].get_limit("storage_gb")
quarterly = QUOTA_TIERS["quarterly"].get_limit("storage_gb")
assert free < monthly <= quarterly
def test_legacy_tiers_are_aliases(self):
"""#1894: old standard/pro/enterprise/basic/premium 保留为别名以兼容历史数据。
basic/standard → monthly; premium/pro/enterprise → quarterly."""
assert QUOTA_TIERS["basic"] is QUOTA_TIERS["monthly"]
assert QUOTA_TIERS["standard"] is QUOTA_TIERS["monthly"]
assert QUOTA_TIERS["premium"] is QUOTA_TIERS["quarterly"]
assert QUOTA_TIERS["pro"] is QUOTA_TIERS["quarterly"]
assert QUOTA_TIERS["enterprise"] is QUOTA_TIERS["quarterly"]
# ── QuotaCheckResult ───────────────────────────────────────────────────────
@@ -133,67 +155,43 @@ class TestQuotaTiers:
class TestQuotaCheckResult:
def test_usage_percent_normal(self):
result = QuotaCheckResult(
allowed=True,
dimension="storage_gb",
limit=100,
used=50,
remaining=50,
warning_level=QuotaWarningLevel.NORMAL,
allowed=True, dimension="storage_gb", limit=100, used=50,
remaining=50, warning_level=QuotaWarningLevel.NORMAL,
)
assert result.usage_percent == 50.0
def test_usage_percent_zero_usage(self):
result = QuotaCheckResult(
allowed=True,
dimension="storage_gb",
limit=100,
used=0,
remaining=100,
warning_level=QuotaWarningLevel.NORMAL,
allowed=True, dimension="storage_gb", limit=100, used=0,
remaining=100, warning_level=QuotaWarningLevel.NORMAL,
)
assert result.usage_percent == 0.0
def test_usage_percent_exceeded_capped(self):
result = QuotaCheckResult(
allowed=False,
dimension="storage_gb",
limit=100,
used=150,
remaining=0,
warning_level=QuotaWarningLevel.EXCEEDED,
allowed=False, dimension="storage_gb", limit=100, used=150,
remaining=0, warning_level=QuotaWarningLevel.EXCEEDED,
)
assert result.usage_percent == 100.0 # capped at 100
assert result.usage_percent == 100.0
def test_usage_percent_zero_limit_with_usage(self):
result = QuotaCheckResult(
allowed=False,
dimension="storage_gb",
limit=0,
used=10,
remaining=0,
warning_level=QuotaWarningLevel.EXCEEDED,
allowed=False, dimension="storage_gb", limit=0, used=10,
remaining=0, warning_level=QuotaWarningLevel.EXCEEDED,
)
assert result.usage_percent == 100.0
def test_usage_percent_zero_limit_no_usage(self):
result = QuotaCheckResult(
allowed=True,
dimension="storage_gb",
limit=0,
used=0,
remaining=0,
warning_level=QuotaWarningLevel.NORMAL,
allowed=True, dimension="storage_gb", limit=0, used=0,
remaining=0, warning_level=QuotaWarningLevel.NORMAL,
)
assert result.usage_percent == 0.0
def test_usage_percent_unlimited(self):
result = QuotaCheckResult(
allowed=True,
dimension="max_templates",
limit=float("inf"),
used=100,
remaining=float("inf"),
warning_level=QuotaWarningLevel.NORMAL,
allowed=True, dimension="max_templates", limit=float("inf"), used=100,
remaining=float("inf"), warning_level=QuotaWarningLevel.NORMAL,
)
assert result.usage_percent == 0.0
@@ -209,12 +207,22 @@ class TestQuotaRegistry:
assert "videos_per_month" in dims
assert "max_concurrent" in dims
def test_init_has_three_tiers(self):
def test_init_has_core_four_tiers(self):
reg = QuotaRegistry()
tiers = reg.list_tiers()
assert "free" in tiers
assert "basic" in tiers
assert "premium" in tiers
assert "monthly" in tiers
assert "quarterly" in tiers
assert "yearly" in tiers
def test_legacy_tiers_are_aliases(self):
"""旧档位作为别名注册以兼容历史数据."""
reg = QuotaRegistry()
tiers = reg.list_tiers()
for legacy in ("standard", "pro", "enterprise", "basic", "premium"):
assert legacy in tiers
assert reg.get_tier("basic") is reg.get_tier("monthly")
assert reg.get_tier("pro") is reg.get_tier("quarterly")
def test_get_limit_free_storage(self):
reg = QuotaRegistry()
@@ -231,38 +239,36 @@ class TestQuotaRegistry:
assert tier.name == "free"
def test_get_tier_unknown_returns_none(self):
"""未知套餐返回 None"""
reg = QuotaRegistry()
assert reg.get_tier("nonexistent") is None
def test_register_new_dimension(self):
reg = QuotaRegistry()
reg.register_dimension("custom_dim", "自定义维度", default_limits={"free": 5, "basic": 20})
reg.register_dimension("custom_dim", "自定义维度",
default_limits={"free": 5, "monthly": 20, "quarterly": 50, "yearly": 100})
assert "custom_dim" in reg.list_dimensions()
assert reg.get_limit("free", "custom_dim") == 5
assert reg.get_limit("basic", "custom_dim") == 20
assert reg.get_limit("monthly", "custom_dim") == 20
assert reg.get_limit("yearly", "custom_dim") == 100
def test_register_dimension_idempotent(self):
reg = QuotaRegistry()
reg.register_dimension("custom_dim", "v1", default_limits={"free": 5})
reg.register_dimension("custom_dim", "v2", default_limits={"free": 99})
# 幂等:第二次注册不改变
assert reg.list_dimensions()["custom_dim"] == "v1"
assert reg.get_limit("free", "custom_dim") == 5
def test_register_dimension_no_defaults(self):
reg = QuotaRegistry()
reg.register_dimension("new_dim", "新维度")
# 默认所有套餐都是 0
assert reg.get_limit("free", "new_dim") == 0
assert reg.get_limit("basic", "new_dim") == 0
assert reg.get_limit("premium", "new_dim") == 0
assert reg.get_limit("monthly", "new_dim") == 0
assert reg.get_limit("yearly", "new_dim") == 0
def test_list_dimensions_returns_copy(self):
reg = QuotaRegistry()
dims = reg.list_dimensions()
dims["fake"] = "test"
# 修改返回值不影响内部
assert "fake" not in reg.list_dimensions()
@@ -277,7 +283,6 @@ class TestQuotaChecker:
assert result.limit == 2
assert result.used == 1.0
assert result.remaining == 1.0
assert result.dimension == "storage_gb"
def test_check_exceeds_limit(self):
checker = QuotaChecker()
@@ -286,17 +291,21 @@ class TestQuotaChecker:
assert result.remaining == 0
def test_check_exactly_at_limit(self):
# used == limit 时 allowed 为 False(必须严格小于)
checker = QuotaChecker()
result = checker.check("free", "storage_gb", 2.0)
assert result.allowed is False
def test_check_unlimited(self):
def test_check_unlimited_quarterly_templates(self):
checker = QuotaChecker()
result = checker.check("premium", "max_templates", 1000.0)
result = checker.check("quarterly", "max_templates", 1000.0)
assert result.allowed is True
assert math.isinf(result.remaining)
def test_check_unlimited_yearly_videos(self):
checker = QuotaChecker()
result = checker.check("yearly", "videos_per_month", 9999.0)
assert result.allowed is True
assert math.isinf(result.remaining)
assert result.warning_level == QuotaWarningLevel.NORMAL
def test_check_unknown_plan(self):
checker = QuotaChecker()
@@ -306,13 +315,8 @@ class TestQuotaChecker:
def test_check_multiple(self):
checker = QuotaChecker()
results = checker.check_multiple(
"free",
{"storage_gb": 1.0, "videos_per_month": 2},
)
results = checker.check_multiple("free", {"storage_gb": 1.0, "videos_per_month": 2})
assert len(results) == 2
assert results[0].dimension == "storage_gb"
assert results[1].dimension == "videos_per_month"
assert all(r.allowed for r in results)
def test_warning_level_normal(self):
@@ -322,34 +326,19 @@ class TestQuotaChecker:
def test_warning_level_warning(self):
checker = QuotaChecker()
# 80% < 95% → warning
result = checker.check("free", "storage_gb", 1.7) # 85%
assert result.warning_level == QuotaWarningLevel.WARNING
def test_warning_level_critical(self):
checker = QuotaChecker()
# 95% <= < 100% → critical
result = checker.check("free", "storage_gb", 1.95) # 97.5%
assert result.warning_level == QuotaWarningLevel.CRITICAL
def test_warning_level_exceeded(self):
checker = QuotaChecker()
result = checker.check("free", "storage_gb", 2.5) # 125%
result = checker.check("free", "storage_gb", 2.5)
assert result.warning_level == QuotaWarningLevel.EXCEEDED
def test_warning_level_zero_limit_with_usage(self):
checker = QuotaChecker()
result = checker.check("free", "ai_voice_enabled", 1) # limit=0, used=1
assert result.warning_level == QuotaWarningLevel.EXCEEDED
def test_warning_level_zero_limit_no_usage(self):
checker = QuotaChecker()
# limit=0, used=0 → 特殊处理为 normal
# 但 allowed 是 False(0 < 0 不成立)
result = checker.check("free", "ai_voice_enabled", 0)
# 0 < 0 是 False → not allowed
assert result.allowed is False
def test_checker_uses_provided_registry(self):
reg = QuotaRegistry()
reg.register_dimension("custom", "自定义", default_limits={"free": 42})
@@ -359,34 +348,22 @@ class TestQuotaChecker:
assert result.allowed is True
# ── get_warning_level 便捷函数 ────────────────────────────────────────────
# ── get_warning_level ─────────────────────────────────────────────────────
class TestGetWarningLevel:
def test_normal_low_usage(self):
assert get_warning_level(50, 100) == QuotaWarningLevel.NORMAL
def test_normal_zero_usage(self):
assert get_warning_level(0, 100) == QuotaWarningLevel.NORMAL
def test_warning_threshold(self):
assert get_warning_level(80, 100) == QuotaWarningLevel.WARNING
def test_warning_between_80_and_95(self):
assert get_warning_level(90, 100) == QuotaWarningLevel.WARNING
def test_critical_threshold(self):
assert get_warning_level(95, 100) == QuotaWarningLevel.CRITICAL
def test_critical_between_95_and_100(self):
assert get_warning_level(99, 100) == QuotaWarningLevel.CRITICAL
def test_exceeded_at_100(self):
assert get_warning_level(100, 100) == QuotaWarningLevel.EXCEEDED
def test_exceeded_over_100(self):
assert get_warning_level(150, 100) == QuotaWarningLevel.EXCEEDED
def test_unlimited_always_normal(self):
assert get_warning_level(9999, float("inf")) == QuotaWarningLevel.NORMAL
@@ -397,7 +374,7 @@ class TestGetWarningLevel:
assert get_warning_level(0, 0) == QuotaWarningLevel.NORMAL
# ── 全局单例 ───────────────────────────────────────────────────────────────
# ── 全局单例 ─────────────────────────────────────────────────────────────
class TestGlobalSingletons:
@@ -413,30 +390,3 @@ class TestGlobalSingletons:
result = quota_checker.check("free", "storage_gb", 1.0)
assert result.allowed is True
assert result.limit == 2
class TestProTier:
"""Pro 套餐专项测试"""
def test_pro_tier_exists(self):
"""pro 套餐存在于 QUOTA_TIERS"""
from packages.domain.quota import QUOTA_TIERS
assert "pro" in QUOTA_TIERS
def test_pro_tier_same_as_premium(self):
"""pro 套餐配额与 premium 完全一致"""
from packages.domain.quota import QUOTA_TIERS
pro = QUOTA_TIERS["pro"]
premium = QUOTA_TIERS["premium"]
assert pro.limits == premium.limits
def test_pro_tier_get_limit(self):
"""pro 套餐各维度配额正确"""
reg = QuotaRegistry()
assert reg.get_limit("pro", "storage_gb") == 100
assert reg.get_limit("pro", "videos_per_month") == 100
assert reg.get_limit("pro", "max_concurrent") == 20
assert reg.get_limit("pro", "max_titles") == 500
assert reg.get_limit("pro", "ai_voice_enabled") == 1
+90 -89
View File
@@ -1,5 +1,6 @@
"""Quota 配额系统单元测试。"""
"""Quota 配额系统单元测试 — #1894 档位清理后 (free/monthly/quarterly/yearly)."""
import math
import pytest
from packages.domain.quota import (
@@ -22,7 +23,7 @@ class TestQuotaDimension:
assert QuotaDimension.VIDEOS_PER_MONTH.value == "videos_per_month"
assert QuotaDimension.MAX_CONCURRENT.value == "max_concurrent"
assert QuotaDimension.MAX_TEMPLATES.value == "max_templates"
assert QuotaDimension.MAX_TITLES.value == "max_titles"
# MAX_TITLES 保留作为枚举别名(与 MAX_TEMPLATES 同值),但不再在套餐配额中独立配置
assert QuotaDimension.MAX_VOICEOVERS.value == "max_voiceovers"
assert QuotaDimension.AI_VOICE_ENABLED.value == "ai_voice_enabled"
@@ -57,7 +58,6 @@ class TestQuotaTier:
def test_is_unlimited_undefined(self):
tier = QuotaTier(name="test", limits={})
# 未定义的维度,limits.get 返回默认 inf,所以 is_unlimited 返回 True
assert tier.is_unlimited("unknown") is True
def test_empty_limits(self):
@@ -67,10 +67,20 @@ class TestQuotaTier:
class TestQuotaTiers:
def test_three_tiers_exist(self):
def test_core_tiers_exist(self):
"""核心四档位存在."""
assert "free" in QUOTA_TIERS
assert "basic" in QUOTA_TIERS
assert "premium" in QUOTA_TIERS
assert "monthly" in QUOTA_TIERS
assert "quarterly" in QUOTA_TIERS
assert "yearly" in QUOTA_TIERS
def test_legacy_tiers_are_aliases(self):
"""旧档位保留为别名以兼容历史数据."""
assert QUOTA_TIERS["basic"] is QUOTA_TIERS["monthly"]
assert QUOTA_TIERS["standard"] is QUOTA_TIERS["monthly"]
assert QUOTA_TIERS["premium"] is QUOTA_TIERS["quarterly"]
assert QUOTA_TIERS["pro"] is QUOTA_TIERS["quarterly"]
assert QUOTA_TIERS["enterprise"] is QUOTA_TIERS["quarterly"]
def test_free_tier_limits(self):
free = QUOTA_TIERS["free"]
@@ -78,49 +88,52 @@ class TestQuotaTiers:
assert free.get_limit("videos_per_month") == 5
assert free.get_limit("max_concurrent") == 3
assert free.get_limit("max_templates") == 3
assert free.get_limit("max_titles") == 50
assert free.get_limit("max_voiceovers") == 10
assert free.get_limit("ai_voice_enabled") == 0
assert free.get_limit("ai_voice_credits") == 0
def test_basic_tier_limits(self):
basic = QUOTA_TIERS["basic"]
assert basic.get_limit("storage_gb") == 20
assert basic.get_limit("videos_per_month") == 30
assert basic.get_limit("max_concurrent") == 10
assert basic.get_limit("max_templates") == 15
assert basic.get_limit("max_titles") == 500
assert basic.get_limit("max_voiceovers") == 100
assert basic.get_limit("ai_voice_enabled") == 1
assert basic.get_limit("ai_voice_credits") == 100
assert basic.get_limit("batch_export_enabled") == 1
def test_monthly_tier_limits(self):
monthly = QUOTA_TIERS["monthly"]
assert monthly.get_limit("storage_gb") == 20
assert monthly.get_limit("videos_per_month") == 30
assert monthly.get_limit("max_concurrent") == 10
assert monthly.get_limit("max_templates") == 15
assert monthly.get_limit("max_voiceovers") == 100
assert monthly.get_limit("ai_voice_enabled") == 1
assert monthly.get_limit("ai_voice_credits") == 100
assert monthly.get_limit("batch_export_enabled") == 1
def test_premium_tier_limits(self):
premium = QUOTA_TIERS["premium"]
assert premium.get_limit("storage_gb") == 100
assert premium.get_limit("videos_per_month") == 100
assert premium.get_limit("max_concurrent") == 20
assert premium.is_unlimited("max_templates") is True
assert premium.get_limit("ai_voice_enabled") == 1
assert premium.get_limit("ai_voice_credits") == 500
assert premium.get_limit("batch_export_enabled") == 1
assert premium.get_limit("multi_platform_enabled") == 1
assert premium.get_limit("dedup_report_enabled") == 1
def test_quarterly_tier_limits(self):
q = QUOTA_TIERS["quarterly"]
assert q.get_limit("storage_gb") == 100
assert q.get_limit("videos_per_month") == 100
assert q.get_limit("max_concurrent") == 20
assert q.is_unlimited("max_templates") is True
assert q.get_limit("ai_voice_enabled") == 1
assert q.get_limit("ai_voice_credits") == 500
assert q.get_limit("batch_export_enabled") == 1
assert q.get_limit("multi_platform_enabled") == 1
assert q.get_limit("dedup_report_enabled") == 1
def test_yearly_tier_limits(self):
y = QUOTA_TIERS["yearly"]
assert y.get_limit("storage_gb") == 100
assert y.is_unlimited("videos_per_month") is True
assert y.get_limit("max_concurrent") == 20
assert y.is_unlimited("max_templates") is True
assert y.get_limit("max_voiceovers") == 200
assert y.get_limit("ai_voice_credits") == 2000
assert y.get_limit("batch_export_enabled") == 1
assert y.get_limit("multi_platform_enabled") == 1
assert y.get_limit("dedup_report_enabled") == 1
def test_tier_increase_monotonic(self):
free = QUOTA_TIERS["free"]
basic = QUOTA_TIERS["basic"]
premium = QUOTA_TIERS["premium"]
# 高级套餐应该 >= 低级套餐的所有限制
for dim in [
"storage_gb",
"videos_per_month",
"max_concurrent",
"max_titles",
"max_voiceovers",
"ai_voice_credits",
]:
assert basic.get_limit(dim) >= free.get_limit(dim)
assert premium.get_limit(dim) >= basic.get_limit(dim)
monthly = QUOTA_TIERS["monthly"]
quarterly = QUOTA_TIERS["quarterly"]
for dim in ["storage_gb", "videos_per_month", "max_concurrent", "max_voiceovers", "ai_voice_credits"]:
assert monthly.get_limit(dim) >= free.get_limit(dim)
assert quarterly.get_limit(dim) >= monthly.get_limit(dim)
class TestQuotaWarningLevel:
@@ -134,12 +147,8 @@ class TestQuotaWarningLevel:
class TestQuotaCheckResult:
def test_usage_percent_normal(self):
result = QuotaCheckResult(
allowed=True,
dimension="storage_gb",
limit=100,
used=50,
remaining=50,
warning_level="normal",
allowed=True, dimension="storage_gb", limit=100, used=50,
remaining=50, warning_level="normal",
)
assert result.usage_percent == 50.0
@@ -147,7 +156,7 @@ class TestQuotaCheckResult:
result = QuotaCheckResult(
allowed=False, dimension="d", limit=100, used=150, remaining=0, warning_level="exceeded"
)
assert result.usage_percent == 100.0 # min(100, 150%)
assert result.usage_percent == 100.0
def test_usage_percent_zero_used(self):
result = QuotaCheckResult(allowed=True, dimension="d", limit=100, used=0, remaining=100, warning_level="normal")
@@ -163,12 +172,8 @@ class TestQuotaCheckResult:
def test_usage_percent_unlimited(self):
result = QuotaCheckResult(
allowed=True,
dimension="d",
limit=float("inf"),
used=1000,
remaining=float("inf"),
warning_level="normal",
allowed=True, dimension="d", limit=float("inf"), used=1000,
remaining=float("inf"), warning_level="normal",
)
assert result.usage_percent == 0.0
@@ -185,10 +190,11 @@ class TestQuotaRegistry:
reg = QuotaRegistry()
tiers = reg.list_tiers()
assert "free" in tiers
assert "pro" in tiers
assert "basic" in tiers
assert "premium" in tiers
assert len(tiers) == 4
assert "monthly" in tiers
assert "quarterly" in tiers
assert "yearly" in tiers
# 核心四档位必须存在
assert "free" in tiers and "monthly" in tiers and "quarterly" in tiers and "yearly" in tiers
def test_get_tier_existing(self):
reg = QuotaRegistry()
@@ -203,7 +209,8 @@ class TestQuotaRegistry:
def test_get_limit_known(self):
reg = QuotaRegistry()
assert reg.get_limit("free", "storage_gb") == 2
assert reg.get_limit("premium", "storage_gb") == 100
assert reg.get_limit("quarterly", "storage_gb") == 100
assert reg.get_limit("yearly", "storage_gb") == 100
def test_get_limit_unknown_plan(self):
reg = QuotaRegistry()
@@ -211,32 +218,33 @@ class TestQuotaRegistry:
def test_register_new_dimension(self):
reg = QuotaRegistry()
reg.register_dimension("new_feature", "新功能", default_limits={"free": 0, "basic": 1, "premium": 5})
reg.register_dimension("new_feature", "新功能",
default_limits={"free": 0, "monthly": 1, "quarterly": 5, "yearly": 10})
dims = reg.list_dimensions()
assert "new_feature" in dims
assert dims["new_feature"] == "新功能"
assert reg.get_limit("free", "new_feature") == 0
assert reg.get_limit("basic", "new_feature") == 1
assert reg.get_limit("premium", "new_feature") == 5
assert reg.get_limit("monthly", "new_feature") == 1
assert reg.get_limit("quarterly", "new_feature") == 5
assert reg.get_limit("yearly", "new_feature") == 10
def test_register_dimension_idempotent(self):
reg = QuotaRegistry()
reg.register_dimension("storage_gb", "should not change", default_limits={"free": 999})
# 已经存在的不覆盖
assert reg.get_limit("free", "storage_gb") == 2
def test_register_without_defaults(self):
reg = QuotaRegistry()
reg.register_dimension("new_dim", "描述")
assert reg.get_limit("free", "new_dim") == 0
assert reg.get_limit("basic", "new_dim") == 0
assert reg.get_limit("premium", "new_dim") == 0
assert reg.get_limit("monthly", "new_dim") == 0
assert reg.get_limit("yearly", "new_dim") == 0
def test_register_partial_limits(self):
reg = QuotaRegistry()
reg.register_dimension("partial", "partial", default_limits={"premium": 42})
assert reg.get_limit("free", "partial") == 0 # 未设置的保持 0
assert reg.get_limit("premium", "partial") == 42
reg.register_dimension("partial", "partial", default_limits={"quarterly": 42})
assert reg.get_limit("free", "partial") == 0
assert reg.get_limit("quarterly", "partial") == 42
class TestQuotaChecker:
@@ -247,7 +255,6 @@ class TestQuotaChecker:
assert result.limit == 2
assert result.used == 1
assert result.remaining == 1
assert result.dimension == "storage_gb"
def test_check_exceeded(self):
checker = QuotaChecker()
@@ -257,19 +264,24 @@ class TestQuotaChecker:
assert result.warning_level == "exceeded"
def test_check_exact_limit_not_allowed(self):
# used < limit 才 allowed,等于不算
checker = QuotaChecker()
result = checker.check("free", "storage_gb", 2)
assert result.allowed is False
assert result.remaining == 0
def test_check_unlimited(self):
def test_check_unlimited_quarterly_templates(self):
checker = QuotaChecker()
result = checker.check("premium", "max_templates", 999999)
result = checker.check("quarterly", "max_templates", 999999)
assert result.allowed is True
assert result.remaining == float("inf")
assert math.isinf(result.remaining)
assert result.warning_level == "normal"
def test_check_unlimited_yearly_videos(self):
checker = QuotaChecker()
result = checker.check("yearly", "videos_per_month", 999999)
assert result.allowed is True
assert math.isinf(result.remaining)
def test_check_warning_level_normal(self):
checker = QuotaChecker()
result = checker.check("free", "storage_gb", 1) # 50%
@@ -277,40 +289,33 @@ class TestQuotaChecker:
def test_check_warning_level_warning(self):
checker = QuotaChecker()
# 80% <= used < 95%
result = checker.check("free", "max_templates", 2.5) # 2.5/3 = 83%
result = checker.check("free", "max_templates", 2.5) # 2.5/3 ≈ 83%
assert result.warning_level == "warning"
def test_check_warning_level_critical(self):
checker = QuotaChecker()
# 95% <= used < 100%
result = checker.check("free", "max_templates", 2.9) # 2.9/3 = 97%
result = checker.check("free", "max_templates", 2.9) # ≈97%
assert result.warning_level == "critical"
def test_check_warning_level_exceeded(self):
checker = QuotaChecker()
result = checker.check("free", "storage_gb", 5) # 250%
result = checker.check("free", "storage_gb", 5)
assert result.warning_level == "exceeded"
def test_check_multiple(self):
checker = QuotaChecker()
results = checker.check_multiple(
"free",
{"storage_gb": 1, "max_templates": 2, "max_titles": 10},
{"storage_gb": 1, "max_templates": 2, "max_voiceovers": 5},
)
assert len(results) == 3
assert results[0].dimension == "storage_gb"
assert results[1].dimension == "max_templates"
assert results[2].dimension == "max_titles"
assert all(r.allowed for r in results)
def test_check_zero_limit(self):
checker = QuotaChecker()
result = checker.check("free", "ai_voice_enabled", 0)
# limit=0, used=0: used < limit 为 False → allowed=False
assert result.allowed is False
assert result.remaining == 0
assert result.warning_level == "normal"
def test_compute_warning_level_normal(self):
assert QuotaChecker._compute_warning_level(50, 100) == "normal"
@@ -337,10 +342,6 @@ class TestQuotaChecker:
def test_compute_warning_level_zero_limit_no_usage(self):
assert QuotaChecker._compute_warning_level(0, 0) == "normal"
def test_compute_warning_level_negative_limit(self):
# limit <= 0 且 used=0 → NORMAL
assert QuotaChecker._compute_warning_level(0, -1) == "normal"
class TestGetWarningLevel:
def test_convenience_function(self):
+80 -297
View File
@@ -1,11 +1,10 @@
"""#1894 废弃标题库整合到文案库 — 集成测试.
"""#1894 方向修正 — 集成测试.
覆盖:
- ScriptModel 新字段 (title_text / title_category / title_config)
- ScriptService CRUD 新字段支持
- ScriptService.get_title_config_for_script 方法
- Scripts API 路由的新字段传递
- title_libraries API deprecated Warning header
- ScriptModel 不再有 title_text/title_category/title_config 列
- ScriptService 不再接受/暴露这三个字段
- /api/v1/titles/* 所有方法返回 410 Gone
- /api/v1/scripts 响应不含这三个字段
"""
from __future__ import annotations
@@ -20,7 +19,6 @@ from unittest.mock import MagicMock, patch
import pytest
from fastapi.testclient import TestClient
# 确保 apps/api 在 sys.path 中(conftest 已加 root,但 apps/api 也需要)
_APPS_API = str(Path(__file__).resolve().parents[2] / "apps" / "api")
if _APPS_API not in sys.path:
sys.path.insert(0, _APPS_API)
@@ -28,307 +26,92 @@ os.environ.setdefault("JWT_SECRET_KEY", "dev-secret-key-for-testing")
from main import app # noqa: E402
# ── helpers ──────────────────────────────────────────────────────────────
def _make_script(**overrides):
"""构造一个模拟 ScriptModel 对象."""
defaults = dict(
id=str(uuid.uuid4()),
user_id="user-001",
title="测试文案",
content="这是内容",
segments=[],
tags=["测试"],
title_text="开场大标题",
title_category="片头",
title_config={
"text": "开场大标题",
"font": "思源黑体",
"font_size": 48,
"font_color": "#FFFFFF",
"position": "top",
},
created_at=datetime(2026, 9, 1, tzinfo=timezone.utc),
updated_at=datetime(2026, 9, 1, tzinfo=timezone.utc),
)
defaults.update(overrides)
return MagicMock(**defaults)
# ── TestScriptModelNewFields ─────────────────────────────────────────────
class TestScriptModelNewFields:
"""验证 ScriptModel 新增字段的定义."""
def test_model_has_title_text_column(self):
from packages.adapters.sqlalchemy_impl.models import ScriptModel
assert hasattr(ScriptModel, "title_text")
col = ScriptModel.__table__.columns["title_text"]
assert col is not None
assert str(col.type) == "VARCHAR(500)"
def test_model_has_title_category_column(self):
from packages.adapters.sqlalchemy_impl.models import ScriptModel
assert hasattr(ScriptModel, "title_category")
col = ScriptModel.__table__.columns["title_category"]
assert col is not None
assert str(col.type) == "VARCHAR(50)"
def test_model_has_title_config_column(self):
from packages.adapters.sqlalchemy_impl.models import ScriptModel
assert hasattr(ScriptModel, "title_config")
col = ScriptModel.__table__.columns["title_config"]
assert col is not None
def test_model_defaults(self):
"""新字段默认值为空字符串/空 dict."""
from packages.adapters.sqlalchemy_impl.models import ScriptModel
s = ScriptModel(id="x", user_id="u", title="t")
# 检查 default 值
assert ScriptModel.__table__.columns["title_text"].default.arg == ""
assert ScriptModel.__table__.columns["title_category"].default.arg == ""
# ── TestScriptServiceTitleConfig ─────────────────────────────────────────
class TestScriptServiceTitleConfig:
"""验证 ScriptService 新方法 get_title_config_for_script."""
def test_get_title_config_returns_script_config(self):
from app.services.script_service import ScriptService
db = MagicMock()
mock_script = _make_script(
title_text="从文案读取",
title_config={"text": "从文案读取", "font": "Arial", "font_size": 36},
)
db.query.return_value.filter.return_value.first.return_value = mock_script
svc = ScriptService(db)
result = svc.get_title_config_for_script("script-1", "user-001")
assert result["text"] == "从文案读取"
assert result["font"] == "Arial"
assert result["font_size"] == 36
def test_get_title_config_fills_text_from_title_text(self):
"""title_config 为空时,用 title_text 填充 text 字段."""
from app.services.script_service import ScriptService
db = MagicMock()
mock_script = _make_script(
title_text="纯文本标题",
title_config={},
)
db.query.return_value.filter.return_value.first.return_value = mock_script
svc = ScriptService(db)
result = svc.get_title_config_for_script("script-2", "user-001")
assert result["text"] == "纯文本标题"
def test_get_title_config_raises_on_not_found(self):
from app.services.script_service import ScriptNotFoundError, ScriptService
db = MagicMock()
db.query.return_value.filter.return_value.first.return_value = None
svc = ScriptService(db)
with pytest.raises(ScriptNotFoundError):
svc.get_title_config_for_script("nonexistent", "user-001")
def test_get_title_config_validates_user_ownership(self):
"""script 不属于当前用户时应抛异常."""
from app.services.script_service import ScriptNotFoundError, ScriptService
db = MagicMock()
db.query.return_value.filter.return_value.first.return_value = None # 不同用户查不到
svc = ScriptService(db)
with pytest.raises(ScriptNotFoundError):
svc.get_title_config_for_script("script-other-user", "user-001")
# ── TestScriptServiceCreateWithNewFields ─────────────────────────────────
class TestScriptServiceCreateWithNewFields:
"""验证 create_script 和 update_script 支持新字段."""
def test_create_script_with_title_fields(self):
from app.services.script_service import ScriptService
db = MagicMock()
svc = ScriptService(db)
script = svc.create_script(
user_id="user-001",
title="新文案",
content="内容",
title_text="标题文字",
title_category="片尾",
title_config={"text": "标题文字", "font_size": 24},
)
db.add.assert_called_once()
db.commit.assert_called_once()
assert script.title_text == "标题文字"
assert script.title_category == "片尾"
assert script.title_config == {"text": "标题文字", "font_size": 24}
def test_update_script_title_fields(self):
from app.services.script_service import ScriptService
db = MagicMock()
existing = _make_script(title_text="旧标题", title_category="旧分类", title_config={"old": True})
db.query.return_value.filter.return_value.first.return_value = existing
svc = ScriptService(db)
updated = svc.update_script(
script_id=existing.id,
user_id="user-001",
title_text="新标题",
title_category="新分类",
title_config={"new": True},
)
assert updated.title_text == "新标题"
assert updated.title_category == "新分类"
assert updated.title_config == {"new": True}
# ── TestScriptsRoutesNewFields ───────────────────────────────────────────
def _make_mock_auth_user(user_id="user-001"):
"""创建 mock 认证用户."""
return MagicMock(user=MagicMock(id=user_id))
class TestScriptsRoutesNewFields:
"""验证 scripts API 路由正确处理新字段 — 使用 dependency_overrides 绕过真实 DB/Auth."""
# ── TestScriptModelNoTitleFields ─────────────────────────────────────────
class TestScriptModelNoTitleFields:
"""验证 ScriptModel 已删除 title_text/title_category/title_config 列."""
def test_model_has_no_title_text(self):
from packages.adapters.sqlalchemy_impl.models import ScriptModel
assert not hasattr(ScriptModel, "title_text") or "title_text" not in {c.name for c in ScriptModel.__table__.columns}
def test_model_has_no_title_category(self):
from packages.adapters.sqlalchemy_impl.models import ScriptModel
assert "title_category" not in {c.name for c in ScriptModel.__table__.columns}
def test_model_has_no_title_config(self):
from packages.adapters.sqlalchemy_impl.models import ScriptModel
assert "title_config" not in {c.name for c in ScriptModel.__table__.columns}
def test_model_retains_core_fields(self):
from packages.adapters.sqlalchemy_impl.models import ScriptModel
cols = {c.name for c in ScriptModel.__table__.columns}
for expected in ("id", "user_id", "title", "content", "segments", "tags", "created_at", "updated_at"):
assert expected in cols, f"ScriptModel 缺字段 {expected}"
# ── TestScriptServiceSignature ──────────────────────────────────────────
class TestScriptServiceSignature:
"""验证 ScriptService 的 create/update 不接受已删除字段."""
def test_create_script_rejects_title_fields(self):
"""Python 层:传入旧字段应抛 TypeError(被移除了)."""
from app.services.script_service import ScriptService
import inspect
sig = inspect.signature(ScriptService.create_script)
for name in ("title_text", "title_category", "title_config"):
assert name not in sig.parameters, f"create_script 仍接受参数 {name}"
def test_update_script_rejects_title_fields(self):
from app.services.script_service import ScriptService
import inspect
sig = inspect.signature(ScriptService.update_script)
for name in ("title_text", "title_category", "title_config"):
assert name not in sig.parameters, f"update_script 仍接受参数 {name}"
def test_get_title_config_removed(self):
from app.services.script_service import ScriptService
assert not hasattr(ScriptService, "get_title_config_for_script"), "get_title_config_for_script 应已删除"
# ── TestTitlesApiGone ───────────────────────────────────────────────────
class TestTitlesApiGone:
"""验证 /api/v1/titles 所有方法返回 410 Gone."""
def setup_method(self):
from app.api.routes.scripts import _get_service, get_current_user
self._mock_svc = MagicMock()
from app.auth import get_current_user
self._mock_user = _make_mock_auth_user()
def _override_svc():
return self._mock_svc
def _override_user():
return self._mock_user
app.dependency_overrides[_get_service] = _override_svc
app.dependency_overrides[get_current_user] = _override_user
self.client = TestClient(app)
def teardown_method(self):
app.dependency_overrides.clear()
def test_create_script_passes_title_fields(self):
mock_script = _make_script(
title_text="测试标题",
title_category="片头",
title_config={"text": "测试标题", "font_size": 48},
)
self._mock_svc.create_script.return_value = mock_script
resp = self.client.post(
"/api/v1/scripts",
json={
"title": "新文案",
"content": "内容",
"title_text": "测试标题",
"title_category": "片头",
"title_config": {"text": "测试标题", "font_size": 48},
},
)
assert resp.status_code == 201, resp.text
call_kwargs = self._mock_svc.create_script.call_args[1]
assert call_kwargs["title_text"] == "测试标题"
assert call_kwargs["title_category"] == "片头"
assert call_kwargs["title_config"] == {"text": "测试标题", "font_size": 48}
def test_get_script_response_includes_title_fields(self):
mock_script = _make_script(
title_text="响应标题",
title_category="片尾",
title_config={"text": "响应标题", "position": "bottom"},
)
self._mock_svc.get_script.return_value = mock_script
resp = self.client.get("/api/v1/scripts/script-123")
assert resp.status_code == 200, resp.text
data = resp.json()
assert data["title_text"] == "响应标题"
assert data["title_category"] == "片尾"
assert data["title_config"]["position"] == "bottom"
# ── TestTitleLibraryDeprecated ───────────────────────────────────────────
class TestTitleLibraryDeprecated:
"""验证 title_libraries API 返回 deprecated Warning header — 使用 dependency_overrides 绕过真实 DB/Auth."""
def setup_method(self):
from app.api.routes.titles import _get_title_repository, get_current_user
from app.dependencies import get_user_repository
self._mock_repo = MagicMock()
self._mock_user_repo = MagicMock()
self._mock_user = _make_mock_auth_user()
app.dependency_overrides[_get_title_repository] = lambda: self._mock_repo
app.dependency_overrides[get_user_repository] = lambda: self._mock_user_repo
app.dependency_overrides[get_current_user] = lambda: self._mock_user
self.client = TestClient(app)
def teardown_method(self):
app.dependency_overrides.clear()
def test_list_titles_has_warning_header(self):
# list_titles 调 use_case + repo, 注入真实用例但 mock 掉 repo 的 list/count
self._mock_repo.list_by_user.return_value = []
self._mock_repo.count_by_user.return_value = 0
resp = self.client.get("/api/v1/titles")
assert resp.status_code == 200, resp.text
headers_lower = {k.lower(): v for k, v in resp.headers.items()}
assert "warning" in headers_lower or "deprecation" in headers_lower
assert "1894" in resp.headers.get("Warning", "") or "1894" in resp.headers.get("warning", "")
def test_get_title_has_warning_header(self):
from packages.domain.title_library import TitleLibraryItem
mock_item = TitleLibraryItem(
id="t1",
user_id="user-001",
name="测试",
text="标题文字",
category="通用",
description="",
tags=[],
usage_count=0,
is_active=True,
created_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
updated_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
)
self._mock_repo.get.return_value = mock_item
resp = self.client.get("/api/v1/titles/t1")
assert resp.status_code == 200, resp.text
warning_header = resp.headers.get("Warning", "") or resp.headers.get("warning", "")
assert "1894" in warning_header or "deprecated" in warning_header.lower()
@pytest.mark.parametrize("method,path", [
("get", "/api/v1/titles"),
("post", "/api/v1/titles"),
("get", "/api/v1/titles/some-id"),
("put", "/api/v1/titles/some-id"),
("delete", "/api/v1/titles/some-id"),
("patch", "/api/v1/titles/some-id"),
("get", "/api/v1/titles/any/nested/path"),
])
def test_titles_routes_return_410(self, method, path):
resp = getattr(self.client, method)(path)
assert resp.status_code == 410, f"{method.upper()} {path} 应返回 410,实际 {resp.status_code}: {resp.text}"
data = resp.json()
assert "error" in data or "message" in data or "GONE" in resp.text
# Deprecation header
assert resp.headers.get("Deprecation") == "true"
+34 -80
View File
@@ -1,7 +1,8 @@
"""Scripts routes 单元测试 — Issue #1795.
"""Scripts routes 单元测试 — Issue #1795 + #1894 清理.
CI 增量映射: scripts.py → test_scripts.py
本文件同时覆盖 routes/scripts.py 和 schemas/script.py 的增量覆盖率。
#1894: 删除 title_text/title_category/title_config 三字段,仅保留 title/content/segments/tags。
"""
from __future__ import annotations
@@ -18,8 +19,6 @@ from app.schemas.script import (
UpdateScriptRequest,
)
# ── Schema 验证测试 ──────────────────────────────────────────────────────────
class TestScriptSegment:
def test_segment_with_duration(self):
@@ -56,7 +55,7 @@ class TestCreateScriptRequest:
def test_title_required(self):
with pytest.raises(ValueError):
CreateScriptRequest(title="") # min_length=1
CreateScriptRequest(title="")
def test_title_max_length(self):
with pytest.raises(ValueError):
@@ -70,20 +69,16 @@ class TestUpdateScriptRequest:
assert r.content is None
assert r.segments is None
assert r.tags is None
assert r.title_text is None
assert r.title_category is None
assert r.title_config is None
def test_partial_update(self):
r = UpdateScriptRequest(title="新标题")
assert r.title == "新标题"
assert r.content is None
def test_partial_update_title_fields(self):
r = UpdateScriptRequest(title_text="新标题文本", title_category="娱乐")
assert r.title_text == "新标题文本"
assert r.title_category == "娱乐"
def test_partial_update_content_only(self):
r = UpdateScriptRequest(content="新内容")
assert r.title is None
assert r.content == "新内容"
class TestScriptResponse:
@@ -96,28 +91,22 @@ class TestScriptResponse:
content="内容",
segments=[ScriptSegment(text="段1")],
tags=["t1"],
title_text="标题文案",
title_category="科技",
title_config={"font": "思源黑体", "size": 48},
created_at=now,
updated_at=now,
)
assert r.id == "s1"
assert len(r.segments) == 1
assert r.title_text == "标题文案"
assert r.title_category == "科技"
assert r.title_config["font"] == "思源黑体"
assert r.title == "标题"
assert r.tags == ["t1"]
def test_response_defaults(self):
"""新字段有默认值,不传也能构造."""
now = datetime(2026, 9, 8, 12, 0, 0, tzinfo=UTC)
r = ScriptResponse(
id="s1", user_id="u1", title="标题", content="",
segments=[], tags=[], created_at=now, updated_at=now,
)
assert r.title_text == ""
assert r.title_category == ""
assert r.title_config == {}
assert r.segments == []
assert r.tags == []
class TestScriptListResponse:
@@ -143,12 +132,7 @@ class TestScriptListResponse:
assert len(r.items) == 1
# ── Route handler 逻辑测试 (mock service) ────────────────────────────────────
class TestRouteHandlers:
"""测试路由层逻辑(不通过 TestClient,直接调用 handler 函数)."""
def _make_auth_user(self, user_id="u1"):
user = MagicMock()
user.id = user_id
@@ -156,23 +140,23 @@ class TestRouteHandlers:
auth.user = user
return auth
def _make_mock_script(self, **overrides):
m = MagicMock()
m.id = overrides.get("id", "s1")
m.user_id = overrides.get("user_id", "u1")
m.title = overrides.get("title", "测试")
m.content = overrides.get("content", "内容")
m.segments = overrides.get("segments", [{"text": "段1", "duration": None}])
m.tags = overrides.get("tags", [])
m.created_at = overrides.get("created_at", datetime(2026, 9, 8, tzinfo=UTC))
m.updated_at = overrides.get("updated_at", datetime(2026, 9, 8, tzinfo=UTC))
return m
def test_create_route_calls_service(self):
from app.api.routes.scripts import create_script
svc = MagicMock()
mock_script = MagicMock()
mock_script.id = "s1"
mock_script.user_id = "u1"
mock_script.title = "测试"
mock_script.content = "内容"
mock_script.segments = [{"text": "段1", "duration": None}]
mock_script.tags = []
mock_script.title_text = ""
mock_script.title_category = ""
mock_script.title_config = {}
mock_script.created_at = datetime(2026, 9, 8, tzinfo=UTC)
mock_script.updated_at = datetime(2026, 9, 8, tzinfo=UTC)
svc.create_script.return_value = mock_script
svc.create_script.return_value = self._make_mock_script(content="内容")
req = CreateScriptRequest(title="测试", content="内容")
auth = self._make_auth_user()
@@ -180,24 +164,16 @@ class TestRouteHandlers:
result = create_script(req, authenticated_user=auth, svc=svc)
assert result.id == "s1"
svc.create_script.assert_called_once()
call_kwargs = svc.create_script.call_args.kwargs
assert "title_text" not in call_kwargs
assert "title_category" not in call_kwargs
assert "title_config" not in call_kwargs
def test_list_route_returns_paginated(self):
from app.api.routes.scripts import list_scripts
svc = MagicMock()
mock_script = MagicMock()
mock_script.id = "s1"
mock_script.user_id = "u1"
mock_script.title = "测试"
mock_script.content = ""
mock_script.segments = []
mock_script.tags = []
mock_script.title_text = ""
mock_script.title_category = ""
mock_script.title_config = {}
mock_script.created_at = datetime(2026, 9, 8, tzinfo=UTC)
mock_script.updated_at = datetime(2026, 9, 8, tzinfo=UTC)
svc.list_scripts.return_value = ([mock_script], 1)
svc.list_scripts.return_value = ([self._make_mock_script()], 1)
auth = self._make_auth_user()
result = list_scripts(skip=0, limit=50, tag=None, authenticated_user=auth, svc=svc)
@@ -208,19 +184,7 @@ class TestRouteHandlers:
from app.api.routes.scripts import get_script
svc = MagicMock()
mock_script = MagicMock()
mock_script.id = "s1"
mock_script.user_id = "u1"
mock_script.title = "测试"
mock_script.content = ""
mock_script.segments = []
mock_script.tags = []
mock_script.title_text = ""
mock_script.title_category = ""
mock_script.title_config = {}
mock_script.created_at = datetime(2026, 9, 8, tzinfo=UTC)
mock_script.updated_at = datetime(2026, 9, 8, tzinfo=UTC)
svc.get_script.return_value = mock_script
svc.get_script.return_value = self._make_mock_script()
auth = self._make_auth_user()
result = get_script("s1", authenticated_user=auth, svc=svc)
@@ -243,24 +207,16 @@ class TestRouteHandlers:
from app.api.routes.scripts import update_script
svc = MagicMock()
mock_script = MagicMock()
mock_script.id = "s1"
mock_script.user_id = "u1"
mock_script.title = "新标题"
mock_script.content = "原内容"
mock_script.segments = []
mock_script.tags = []
mock_script.title_text = ""
mock_script.title_category = ""
mock_script.title_config = {}
mock_script.created_at = datetime(2026, 9, 8, tzinfo=UTC)
mock_script.updated_at = datetime(2026, 9, 8, tzinfo=UTC)
svc.update_script.return_value = mock_script
svc.update_script.return_value = self._make_mock_script(title="新标题")
req = UpdateScriptRequest(title="新标题")
auth = self._make_auth_user()
result = update_script("s1", req, authenticated_user=auth, svc=svc)
assert result.title == "新标题"
call_kwargs = svc.update_script.call_args.kwargs
assert "title_text" not in call_kwargs
assert "title_category" not in call_kwargs
assert "title_config" not in call_kwargs
def test_update_route_not_found(self):
from app.api.routes.scripts import update_script
@@ -284,7 +240,6 @@ class TestRouteHandlers:
auth = self._make_auth_user()
result = delete_script("s1", authenticated_user=auth, svc=svc)
# Should return None (204 No Content)
assert result is None
def test_delete_route_not_found(self):
@@ -298,4 +253,3 @@ class TestRouteHandlers:
with pytest.raises(HTTPException) as exc_info:
delete_script("bad", authenticated_user=auth, svc=svc)
assert exc_info.value.status_code == 404
+5 -3
View File
@@ -461,8 +461,9 @@ class TestQuotaTiers:
def test_all_tiers_exist(self):
from packages.domain.quota import QUOTA_TIERS
assert set(QUOTA_TIERS.keys()) == {"free", "basic", "premium", "pro"}
# 核心四档位 + 旧别名
for k in ("free", "monthly", "quarterly", "yearly", "basic", "premium", "pro", "standard", "enterprise"):
assert k in QUOTA_TIERS
# ============================================================
@@ -688,7 +689,8 @@ class TestQuotaRegistry:
reg = QuotaRegistry()
tiers = reg.list_tiers()
assert set(tiers) == {"free", "basic", "premium", "pro"}
for k in ("free", "monthly", "quarterly", "yearly", "basic", "premium", "pro"):
assert k in tiers
def test_register_new_dimension(self):
from packages.domain.quota import QuotaRegistry