Compare commits

..

1 Commits

Author SHA1 Message Date
xiaoxia 28a2322f73 test(wave210): 验证码服务单测补全 +70测
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 38s
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 / Validate - Type Check (mypy) (pull_request) Successful in 1m21s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m43s
Preview Deploy / Deploy Preview Environment (pull_request) Failing after 1m3s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m34s
CI/CD Pipeline / Validate - Code Quality (pull_request) Successful in 3m50s
AI Code Review / AI Code Review (pull_request) Successful in 3m46s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m49s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
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 / Unit Tests (pull_request) Successful in 4m1s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 3m7s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 5m48s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 5m31s
CI/CD Pipeline / CI Gate (pull_request) Successful in 17s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 40s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 39s
覆盖范围:
- generate: 成功/自定义code/自定义TTL/参数校验/频控(冷却+每日上限)
- verify: 成功/错误/不存在/过期/已使用/尝试次数/consume开关
- validate_phone: 各种合法/非法号码格式
- normalize_phone: +86前缀/空格处理
- validate_email: 各种合法/非法邮箱格式
- VerificationCode实体: is_expired/is_used/is_valid/生命周期

70 test cases, 8 test classes
2026-07-30 07:20:16 +08:00
36 changed files with 2994 additions and 4595 deletions
+188
View File
@@ -0,0 +1,188 @@
/**
* 认证相关 API
*/
import axios from "axios"
import apiClient from "./client"
// 类型定义
export interface LoginRequest {
email: string
password: string
}
export interface LoginResponse {
access_token: string
refresh_token?: string | null
token_type: string
expires_in: number
user_id: string
email: string
username: string
display_name: string
}
export interface RegisterRequest {
email: string
password: string
username: string
display_name?: string
}
export interface User {
id: string
user_id: string
email: string
username: string
display_name: string
is_email_verified: boolean
email_verified: boolean
created_at?: string
}
export interface UserResponse {
id?: string
user_id?: string
email: string
username: string
display_name: string
is_email_verified?: boolean
email_verified?: boolean
created_at?: string
}
export const normalizeUser = (data: UserResponse): User => {
const userId = data.id ?? data.user_id ?? ""
const emailVerified = data.is_email_verified ?? data.email_verified ?? false
return {
id: userId,
user_id: userId,
email: data.email,
username: data.username,
display_name: data.display_name,
is_email_verified: emailVerified,
email_verified: emailVerified,
created_at: data.created_at,
}
}
// 登录
export const login = async (data: LoginRequest): Promise<LoginResponse> => {
const response = await apiClient.post("/auth/login", data)
return response.data
}
// 刷新 access_token(使用裸 axios 避免拦截器递归)
export const refreshAccessToken = async (refreshToken: string): Promise<LoginResponse> => {
const baseURL = apiClient.defaults.baseURL ?? ""
const response = await axios.post(`${baseURL}/auth/refresh`, {
refresh_token: refreshToken,
})
return response.data
}
// 注册
export const register = async (data: RegisterRequest): Promise<{ message: string }> => {
const response = await apiClient.post("/auth/register", data)
return response.data
}
// 登出
export const logout = async (): Promise<void> => {
await apiClient.post("/auth/logout")
}
// 获取当前用户
export const getCurrentUser = async (): Promise<User> => {
const response = await apiClient.get<UserResponse>("/auth/me")
return normalizeUser(response.data)
}
// 请求密码重置
export const requestPasswordReset = async (email: string): Promise<{ message: string }> => {
const response = await apiClient.post("/auth/forgot-password", { email })
return response.data
}
// 重置密码
export const resetPassword = async (
token: string,
newPassword: string,
): Promise<{ message: string }> => {
const response = await apiClient.post("/auth/reset-password", {
token,
new_password: newPassword,
})
return response.data
}
// 验证邮箱
export const verifyEmail = async (token: string): Promise<{ message: string }> => {
const response = await apiClient.post("/auth/verify-email", { token })
return response.data
}
/* ========== 微信登录 ========== */
export interface WechatAuthUrlResponse {
auth_url: string
state: string
}
export interface WechatCallbackResponse {
access_token: string
refresh_token?: string | null
user_id: string
display_name: string
avatar_url: string
is_new_user: boolean
binding_complete: boolean
expires_in: number
}
export interface SendVerificationCodeRequest {
target: "email" | "phone"
value: string
purpose: "bind" | "login" | "reset_password"
}
export interface BindContactRequest {
email?: string
email_code?: string
phone?: string
phone_code?: string
}
export interface BindContactResponse {
success: boolean
user: User
}
// 获取微信授权链接
export const getWechatAuthUrl = async (): Promise<WechatAuthUrlResponse> => {
const response = await apiClient.get("/auth/wechat/url")
return response.data
}
// 微信回调登录
export const wechatCallback = async (
code: string,
state: string,
): Promise<WechatCallbackResponse> => {
const response = await apiClient.post("/auth/wechat/callback", { code, state })
return response.data
}
// 发送验证码
export const sendVerificationCode = async (
data: SendVerificationCodeRequest,
): Promise<{ message: string }> => {
const response = await apiClient.post("/auth/send-verification-code", data)
return response.data
}
// 绑定联系方式
export const bindContact = async (data: BindContactRequest): Promise<BindContactResponse> => {
const response = await apiClient.post("/auth/bind-contact", data)
return response.data
}
-20
View File
@@ -1,20 +0,0 @@
import apiClient from "../client"
import type { SendVerificationCodeRequest, BindContactRequest, BindContactResponse } from "./types"
/**
* 发送验证码
*/
export const sendVerificationCode = async (
data: SendVerificationCodeRequest,
): Promise<{ message: string }> => {
const response = await apiClient.post("/auth/send-verification-code", data)
return response.data
}
/**
* 绑定联系方式
*/
export const bindContact = async (data: BindContactRequest): Promise<BindContactResponse> => {
const response = await apiClient.post("/auth/bind-contact", data)
return response.data
}
-11
View File
@@ -1,11 +0,0 @@
import apiClient from "../client"
import type { User, UserResponse } from "./types"
import { normalizeUser } from "./user"
/**
* 获取当前用户
*/
export const getCurrentUser = async (): Promise<User> => {
const response = await apiClient.get<UserResponse>("/auth/me")
return normalizeUser(response.data)
}
-9
View File
@@ -1,9 +0,0 @@
import apiClient from "../client"
/**
* 验证邮箱
*/
export const verifyEmail = async (token: string): Promise<{ message: string }> => {
const response = await apiClient.post("/auth/verify-email", { token })
return response.data
}
-39
View File
@@ -1,39 +0,0 @@
/**
* 认证相关 API
* 保持向后兼容,从子模块 re-export
*/
// 类型
export type {
LoginRequest,
LoginResponse,
RegisterRequest,
User,
UserResponse,
WechatAuthUrlResponse,
WechatCallbackResponse,
SendVerificationCodeRequest,
BindContactRequest,
BindContactResponse,
} from "./types"
// 用户工具函数
export { normalizeUser } from "./user"
// 登录/注册/登出/刷新
export { login, refreshAccessToken, register, logout } from "./login"
// 当前用户
export { getCurrentUser } from "./currentUser"
// 密码重置
export { requestPasswordReset, resetPassword } from "./password"
// 邮箱验证
export { verifyEmail } from "./email"
// 微信登录
export { getWechatAuthUrl, wechatCallback } from "./wechat"
// 联系方式
export { sendVerificationCode, bindContact } from "./contact"
-37
View File
@@ -1,37 +0,0 @@
import axios from "axios"
import apiClient from "../client"
import type { LoginRequest, LoginResponse, RegisterRequest } from "./types"
/**
* 登录
*/
export const login = async (data: LoginRequest): Promise<LoginResponse> => {
const response = await apiClient.post("/auth/login", data)
return response.data
}
/**
* 刷新 access_token(使用裸 axios 避免拦截器递归)
*/
export const refreshAccessToken = async (refreshToken: string): Promise<LoginResponse> => {
const baseURL = apiClient.defaults.baseURL ?? ""
const response = await axios.post(`${baseURL}/auth/refresh`, {
refresh_token: refreshToken,
})
return response.data
}
/**
* 注册
*/
export const register = async (data: RegisterRequest): Promise<{ message: string }> => {
const response = await apiClient.post("/auth/register", data)
return response.data
}
/**
* 登出
*/
export const logout = async (): Promise<void> => {
await apiClient.post("/auth/logout")
}
-23
View File
@@ -1,23 +0,0 @@
import apiClient from "../client"
/**
* 请求密码重置
*/
export const requestPasswordReset = async (email: string): Promise<{ message: string }> => {
const response = await apiClient.post("/auth/forgot-password", { email })
return response.data
}
/**
* 重置密码
*/
export const resetPassword = async (
token: string,
newPassword: string,
): Promise<{ message: string }> => {
const response = await apiClient.post("/auth/reset-password", {
token,
new_password: newPassword,
})
return response.data
}
-82
View File
@@ -1,82 +0,0 @@
/**
* 认证相关类型定义
*/
export interface LoginRequest {
email: string
password: string
}
export interface LoginResponse {
access_token: string
refresh_token?: string | null
token_type: string
expires_in: number
user_id: string
email: string
username: string
display_name: string
}
export interface RegisterRequest {
email: string
password: string
username: string
display_name?: string
}
export interface User {
id: string
user_id: string
email: string
username: string
display_name: string
is_email_verified: boolean
email_verified: boolean
created_at?: string
}
export interface UserResponse {
id?: string
user_id?: string
email: string
username: string
display_name: string
is_email_verified?: boolean
email_verified?: boolean
created_at?: string
}
export interface WechatAuthUrlResponse {
auth_url: string
state: string
}
export interface WechatCallbackResponse {
access_token: string
refresh_token?: string | null
user_id: string
display_name: string
avatar_url: string
is_new_user: boolean
binding_complete: boolean
expires_in: number
}
export interface SendVerificationCodeRequest {
target: "email" | "phone"
value: string
purpose: "bind" | "login" | "reset_password"
}
export interface BindContactRequest {
email?: string
email_code?: string
phone?: string
phone_code?: string
}
export interface BindContactResponse {
success: boolean
user: User
}
-20
View File
@@ -1,20 +0,0 @@
import type { User, UserResponse } from "./types"
/**
* 规范化用户数据,兼容不同后端返回格式
*/
export const normalizeUser = (data: UserResponse): User => {
const userId = data.id ?? data.user_id ?? ""
const emailVerified = data.is_email_verified ?? data.email_verified ?? false
return {
id: userId,
user_id: userId,
email: data.email,
username: data.username,
display_name: data.display_name,
is_email_verified: emailVerified,
email_verified: emailVerified,
created_at: data.created_at,
}
}
-21
View File
@@ -1,21 +0,0 @@
import apiClient from "../client"
import type { WechatAuthUrlResponse, WechatCallbackResponse } from "./types"
/**
* 获取微信授权链接
*/
export const getWechatAuthUrl = async (): Promise<WechatAuthUrlResponse> => {
const response = await apiClient.get("/auth/wechat/url")
return response.data
}
/**
* 微信回调登录
*/
export const wechatCallback = async (
code: string,
state: string,
): Promise<WechatCallbackResponse> => {
const response = await apiClient.post("/auth/wechat/callback", { code, state })
return response.data
}
@@ -1,76 +0,0 @@
import React from "react"
import type { MediaAsset } from "@/api/template-editor"
import { MATERIAL_TYPE_LABELS, MATERIAL_TYPE_ICONS } from "@/api/template-editor"
import { formatSize, formatDuration, getQualityLevel } from "../utils"
export interface AssetCardProps {
asset: MediaAsset
isSelected: boolean
isDragging: boolean
isDragOver: boolean
showBatchSelect: boolean
onToggleSelect: (e: React.MouseEvent) => void
onMouseEnter: (e: React.MouseEvent) => void
onMouseLeave: () => void
}
export const AssetCard: React.FC<AssetCardProps> = ({
asset,
isSelected,
isDragging,
isDragOver,
showBatchSelect,
onToggleSelect,
onMouseEnter,
onMouseLeave,
}) => {
const qualityLevel = getQualityLevel(asset.quality_score)
return (
<>
{/* 缩略图 */}
<div className="as-card-thumb">
{asset.thumbnail_url ? (
<img src={asset.thumbnail_url} alt={asset.name} loading="lazy" />
) : (
<span className="as-card-thumb-icon">{MATERIAL_TYPE_ICONS[asset.type]}</span>
)}
{/* Checkbox */}
{showBatchSelect && (
<span
data-checkbox
className={`as-card-checkbox${isSelected ? " checked" : ""}`}
onClick={onToggleSelect}
/>
)}
{/* 类型角标 */}
<span className="as-card-type-badge">{MATERIAL_TYPE_LABELS[asset.type]}</span>
{/* 时长角标 */}
{asset.duration != null && (
<span className="as-card-duration">{formatDuration(asset.duration)}</span>
)}
{/* 质量分角标 */}
{asset.quality_score != null && (
<span
className={`as-card-quality ${qualityLevel}`}
title={`质量分: ${asset.quality_score}`}
>
{asset.quality_score}
</span>
)}
</div>
{/* 信息 */}
<div className="as-card-info">
<p className="as-card-name" title={asset.name}>
{asset.name}
</p>
<div className="as-card-meta">{formatSize(asset.size)}</div>
</div>
</>
)
}
@@ -1,67 +0,0 @@
import React from "react"
import type { MediaAsset } from "@/api/template-editor"
import { MATERIAL_TYPE_LABELS, MATERIAL_TYPE_ICONS } from "@/api/template-editor"
import { formatSize, formatDuration, getQualityColor } from "../utils"
export interface AssetListItemProps {
asset: MediaAsset
isSelected: boolean
isDragging: boolean
isDragOver: boolean
showBatchSelect: boolean
onToggleSelect: (e: React.MouseEvent) => void
onMouseEnter: (e: React.MouseEvent) => void
onMouseLeave: () => void
}
export const AssetListItem: React.FC<AssetListItemProps> = ({
asset,
isSelected,
isDragging,
isDragOver,
showBatchSelect,
onToggleSelect,
onMouseEnter,
onMouseLeave,
}) => {
return (
<>
{/* 拖拽手柄 */}
<span className="as-list-item-drag" title="拖拽排序">
</span>
{/* Checkbox */}
{showBatchSelect && (
<span
data-checkbox
className={`as-list-item-checkbox${isSelected ? " checked" : ""}`}
onClick={onToggleSelect}
/>
)}
{/* 图标 */}
<span className="as-list-item-icon">{MATERIAL_TYPE_ICONS[asset.type]}</span>
{/* 信息 */}
<div className="as-list-item-info">
<div className="as-list-item-name">{asset.name}</div>
<div className="as-list-item-meta">
{MATERIAL_TYPE_LABELS[asset.type]}
{asset.duration != null && ` · ${formatDuration(asset.duration)}`}
{asset.size != null && ` · ${formatSize(asset.size)}`}
</div>
</div>
{/* 质量分 */}
{asset.quality_score != null && (
<span
className="as-list-item-quality"
style={{ color: getQualityColor(asset.quality_score) }}
>
{asset.quality_score}
</span>
)}
</>
)
}
@@ -1,14 +0,0 @@
import type { MediaAsset } from "@/api/template-editor"
export interface AssetSelectorProps {
assets: MediaAsset[]
selectedIds?: string[]
onSelectionChange?: (ids: string[]) => void
onAssetDragStart?: (asset: MediaAsset) => void
onReorder?: (fromIdx: number, toIdx: number) => void
showQualityFilter?: boolean
showBatchSelect?: boolean
compact?: boolean
}
export type ViewMode = "grid" | "list"
@@ -1,41 +0,0 @@
/** 格式化文件大小 */
export const formatSize = (bytes?: number): string => {
if (!bytes) return ""
if (bytes < 1024) return `${bytes}B`
if (bytes < 1024 * 1024) return `${(bytes / 1024).toFixed(1)}KB`
return `${(bytes / (1024 * 1024)).toFixed(1)}MB`
}
/** 格式化时长 */
export const formatDuration = (seconds?: number): string => {
if (!seconds) return ""
const m = Math.floor(seconds / 60)
const s = Math.floor(seconds % 60)
return m > 0 ? `${m}:${s.toString().padStart(2, "0")}` : `${s}s`
}
/** 获取质量分等级 */
export const getQualityLevel = (score?: number): string => {
if (score == null) return "none"
if (score >= 90) return "excellent"
if (score >= 70) return "good"
if (score >= 50) return "fair"
return "poor"
}
/** 质量分颜色 */
export const getQualityColor = (score?: number): string => {
if (score == null) return "var(--text-secondary)"
if (score >= 90) return "var(--success-color, #10b981)"
if (score >= 70) return "var(--primary-color, #6366f1)"
if (score >= 50) return "var(--warning-color, #f59e0b)"
return "var(--error-color, #ef4444)"
}
/** 类型筛选选项 */
export const TYPE_OPTIONS = [
{ value: "", label: "全部类型" },
{ value: "video", label: "🎬 视频" },
{ value: "image", label: "🖼️ 图片" },
{ value: "audio", label: "🎵 音频" },
]
@@ -159,4 +159,3 @@
flex-wrap: wrap;
}
}
+182
View File
@@ -0,0 +1,182 @@
/**
* PageHead - 页面头部组件(Task 1.4
*
* 功能:
* - 页面标题展示
* - 面包屑导航(自动根据路由生成,也支持手动传入)
* - 右侧操作按钮区(slot,由页面自行填充)
* - 响应式:移动端简化布局(隐藏面包屑,缩小标题)
*
* 复用 global.css 中已有的 .xx-page-head 基础样式,
* 补充面包屑、操作区等扩展样式。
*/
import React from "react"
import { useLocation, useNavigate, Link } from "react-router-dom"
import { RightOutlined, HomeOutlined } from "@ant-design/icons"
import "./PageHead.css"
/* ── 类型定义 ─────────────────────────────────────────────── */
/** 面包屑项 */
export interface BreadcrumbItem {
/** 显示文字 */
label: string
/** 路由路径,不传则为当前页(不可点击) */
path?: string
}
/** PageHead 组件属性 */
export interface PageHeadProps {
/** 页面标题 */
title: string
/** 页面描述(可选,显示在标题下方) */
description?: React.ReactNode
/** 面包屑项(可选,不传则自动根据路由生成) */
breadcrumb?: BreadcrumbItem[]
/** 右侧操作区内容(按钮等) */
actions?: React.ReactNode
/** 是否隐藏面包屑 */
hideBreadcrumb?: boolean
}
/* ── 路由 → 标题映射(用于自动生成面包屑) ────────────────── */
const ROUTE_TITLE_MAP: Record<string, string> = {
"/app/dashboard": "首页",
"/app/generate": "智能剪辑",
"/app/assets": "视频库",
"/app/voices": "配音库",
"/app/titles": "标题库",
"/app/products": "成片库",
"/app/templates": "模板库",
"/app/history": "任务历史",
"/app/admin": "控制台",
"/app/admin/users": "用户管理",
"/app/admin/analytics": "数据分析",
"/app/admin/monitor": "系统监控",
"/app/admin/logs": "系统日志",
"/app/subscription": "订阅管理",
"/app/subscription/upgrade": "升级订阅",
"/app/subscription/billing": "账单管理",
"/app/profile": "个人设置",
"/app/editing-planner": "模板制作",
"/app/my-templates": "我的模板",
"/app/voice-clone": "我的音色",
"/app/voice-materials": "配音库",
"/app/accounts": "账号管理",
"/app/duplication": "查重",
"/app/duplication/results": "查重结果",
}
/* ── 自动生成面包屑 ─────────────────────────────────────── */
/** 根据当前路径生成面包屑 */
const generateBreadcrumb = (pathname: string): BreadcrumbItem[] => {
const items: BreadcrumbItem[] = [{ label: "首页", path: "/app/dashboard" }]
// 首页本身不需要面包屑
if (pathname === "/app" || pathname === "/app/dashboard") {
return items
}
// 逐级拆分路径,生成中间层级
const segments = pathname.split("/").filter(Boolean)
let currentPath = ""
for (let i = 0; i < segments.length; i++) {
currentPath += `/${segments[i]}`
const title = ROUTE_TITLE_MAP[currentPath]
if (title) {
// 最后一级不带 path(当前页面,不可点击)
const isLast = i === segments.length - 1
items.push({
label: title,
path: isLast ? undefined : currentPath,
})
} else {
// 动态路由段(如 :id),用路径片段做 label
const isLast = i === segments.length - 1
items.push({
label: segments[i],
path: isLast ? undefined : currentPath,
})
}
}
return items
}
/* ── 组件 ───────────────────────────────────────────────── */
const PageHead: React.FC<PageHeadProps> = ({
title,
description,
breadcrumb,
actions,
hideBreadcrumb = false,
}) => {
const location = useLocation()
const navigate = useNavigate()
// 使用传入的面包屑或自动生成
const breadcrumbItems = breadcrumb ?? generateBreadcrumb(location.pathname)
// 首页不显示面包屑
const showBreadcrumb =
!hideBreadcrumb &&
breadcrumbItems.length > 1 &&
location.pathname !== "/app" &&
location.pathname !== "/app/dashboard"
return (
<header className="xx-page-head">
<div className="xx-page-head-left">
{/* 面包屑导航 */}
{showBreadcrumb && (
<nav className="xx-page-breadcrumb" aria-label="面包屑导航">
<ol>
{breadcrumbItems.map((item, index) => {
const isLast = index === breadcrumbItems.length - 1
return (
<li key={`${item.label}-${index}`} className="xx-page-breadcrumb-item">
{index > 0 && <RightOutlined className="xx-page-breadcrumb-separator" />}
{item.path && !isLast ? (
<Link
to={item.path}
className="xx-page-breadcrumb-link"
onClick={(e) => {
e.preventDefault()
navigate(item.path!)
}}
>
{index === 0 ? <HomeOutlined className="xx-page-breadcrumb-home" /> : null}
<span>{item.label}</span>
</Link>
) : (
<span className="xx-page-breadcrumb-current" aria-current="page">
{index === 0 ? <HomeOutlined className="xx-page-breadcrumb-home" /> : null}
<span>{item.label}</span>
</span>
)}
</li>
)
})}
</ol>
</nav>
)}
{/* 标题 + 描述 */}
<div className="xx-page-head-title">
<h2>{title}</h2>
{description && <p>{description}</p>}
</div>
</div>
{/* 右侧操作区 */}
{actions && <div className="xx-page-head-actions">{actions}</div>}
</header>
)
}
export default PageHead
@@ -1,27 +0,0 @@
/** 路由 → 标题映射(用于自动生成面包屑) */
export const ROUTE_TITLE_MAP: Record<string, string> = {
"/app/dashboard": "首页",
"/app/generate": "智能剪辑",
"/app/assets": "视频库",
"/app/voices": "配音库",
"/app/titles": "标题库",
"/app/products": "成片库",
"/app/templates": "模板库",
"/app/history": "任务历史",
"/app/admin": "控制台",
"/app/admin/users": "用户管理",
"/app/admin/analytics": "数据分析",
"/app/admin/monitor": "系统监控",
"/app/admin/logs": "系统日志",
"/app/subscription": "订阅管理",
"/app/subscription/upgrade": "升级订阅",
"/app/subscription/billing": "账单管理",
"/app/profile": "个人设置",
"/app/editing-planner": "模板制作",
"/app/my-templates": "我的模板",
"/app/voice-clone": "我的音色",
"/app/voice-materials": "配音库",
"/app/accounts": "账号管理",
"/app/duplication": "查重",
"/app/duplication/results": "查重结果",
}
@@ -1,91 +0,0 @@
/**
* PageHead - 页面头部组件(Task 1.4
*
* 功能:
* - 页面标题展示
* - 面包屑导航(自动根据路由生成,也支持手动传入)
* - 右侧操作按钮区(slot,由页面自行填充)
* - 响应式:移动端简化布局(隐藏面包屑,缩小标题)
*
* 复用 global.css 中已有的 .xx-page-head 基础样式,
* 补充面包屑、操作区等扩展样式。
*/
import React from "react"
import { useLocation, useNavigate, Link } from "react-router-dom"
import { RightOutlined, HomeOutlined } from "@ant-design/icons"
import type { PageHeadProps } from "./types"
import { generateBreadcrumb } from "./utils"
import "./PageHead.css"
const PageHead: React.FC<PageHeadProps> = ({
title,
description,
breadcrumb,
actions,
hideBreadcrumb = false,
}) => {
const location = useLocation()
const navigate = useNavigate()
// 使用传入的面包屑或自动生成
const breadcrumbItems = breadcrumb ?? generateBreadcrumb(location.pathname)
// 首页不显示面包屑
const showBreadcrumb =
!hideBreadcrumb &&
breadcrumbItems.length > 1 &&
location.pathname !== "/app" &&
location.pathname !== "/app/dashboard"
return (
<header className="xx-page-head">
<div className="xx-page-head-left">
{/* 面包屑导航 */}
{showBreadcrumb && (
<nav className="xx-page-breadcrumb" aria-label="面包屑导航">
<ol>
{breadcrumbItems.map((item, index) => {
const isLast = index === breadcrumbItems.length - 1
return (
<li key={`${item.label}-${index}`} className="xx-page-breadcrumb-item">
{index > 0 && <RightOutlined className="xx-page-breadcrumb-separator" />}
{item.path && !isLast ? (
<Link
to={item.path}
className="xx-page-breadcrumb-link"
onClick={(e) => {
e.preventDefault()
navigate(item.path!)
}}
>
{index === 0 ? <HomeOutlined className="xx-page-breadcrumb-home" /> : null}
<span>{item.label}</span>
</Link>
) : (
<span className="xx-page-breadcrumb-current" aria-current="page">
{index === 0 ? <HomeOutlined className="xx-page-breadcrumb-home" /> : null}
<span>{item.label}</span>
</span>
)}
</li>
)
})}
</ol>
</nav>
)}
{/* 标题 + 描述 */}
<div className="xx-page-head-title">
<h2>{title}</h2>
{description && <p>{description}</p>}
</div>
</div>
{/* 右侧操作区 */}
{actions && <div className="xx-page-head-actions">{actions}</div>}
</header>
)
}
export default PageHead
export type { BreadcrumbItem, PageHeadProps } from "./types"
@@ -1,23 +0,0 @@
import type React from "react"
/** 面包屑项 */
export interface BreadcrumbItem {
/** 显示文字 */
label: string
/** 路由路径,不传则为当前页(不可点击) */
path?: string
}
/** PageHead 组件属性 */
export interface PageHeadProps {
/** 页面标题 */
title: string
/** 页面描述(可选,显示在标题下方) */
description?: React.ReactNode
/** 面包屑项(可选,不传则自动根据路由生成) */
breadcrumb?: BreadcrumbItem[]
/** 右侧操作区内容(按钮等) */
actions?: React.ReactNode
/** 是否隐藏面包屑 */
hideBreadcrumb?: boolean
}
@@ -1,39 +0,0 @@
import { ROUTE_TITLE_MAP } from "./constants"
import type { BreadcrumbItem } from "./types"
/** 根据当前路径生成面包屑 */
export const generateBreadcrumb = (pathname: string): BreadcrumbItem[] => {
const items: BreadcrumbItem[] = [{ label: "首页", path: "/app/dashboard" }]
// 首页本身不需要面包屑
if (pathname === "/app" || pathname === "/app/dashboard") {
return items
}
// 逐级拆分路径,生成中间层级
const segments = pathname.split("/").filter(Boolean)
let currentPath = ""
for (let i = 0; i < segments.length; i++) {
currentPath += `/${segments[i]}`
const title = ROUTE_TITLE_MAP[currentPath]
if (title) {
// 最后一级不带 path(当前页面,不可点击)
const isLast = i === segments.length - 1
items.push({
label: title,
path: isLast ? undefined : currentPath,
})
} else {
// 动态路由段(如 :id),用路径片段做 label
const isLast = i === segments.length - 1
items.push({
label: segments[i],
path: isLast ? undefined : currentPath,
})
}
}
return items
}
@@ -1,39 +0,0 @@
import React from "react"
import type { StickerItem } from "@/pages/editing-planner/types"
import { TEXT_PRESET_STYLES } from "@/pages/editing-planner/constants/sticker"
interface StickerPreviewProps {
sticker: StickerItem
}
export const StickerPreview: React.FC<StickerPreviewProps> = ({ sticker }) => (
<div className="sticker-preview-box">
<div
className="sticker-preview-item"
style={{
left: `${sticker.x}%`,
top: `${sticker.y}%`,
width: `${sticker.width}%`,
height: `${sticker.height}%`,
transform: `translate(-50%, -50%) rotate(${sticker.rotation}deg)`,
opacity: sticker.opacity / 100,
fontSize: sticker.type === "text" ? `${sticker.font_size}px` : undefined,
...TEXT_PRESET_STYLES[sticker.text_preset],
}}
>
{sticker.type === "emoji" && sticker.content}
{sticker.type === "text" && sticker.content}
{sticker.type === "image" && (
<img
src={sticker.content}
alt="sticker"
style={{
width: "100%",
height: "100%",
objectFit: "contain",
}}
/>
)}
</div>
</div>
)
+75 -5
View File
@@ -2,9 +2,11 @@
* 选中贴纸的属性编辑器
*/
import React from "react"
import type { StickerItem } from "@/pages/editing-planner/types"
import { StickerPreview } from "./StickerPreview"
import { TextStickerPropsEditor } from "./TextStickerPropsEditor"
import type { StickerItem, TextStickerPreset } from "@/pages/editing-planner/types"
import {
TEXT_PRESET_STYLES,
TEXT_STICKER_PRESET_LABELS,
} from "@/pages/editing-planner/constants/sticker"
interface StickerPropsEditorProps {
sticker: StickerItem
@@ -119,10 +121,78 @@ const StickerPropsEditor: React.FC<StickerPropsEditorProps> = ({
</div>
{/* 文字贴纸特有属性 */}
<TextStickerPropsEditor sticker={sticker} onUpdate={onUpdate} />
{sticker.type === "text" && (
<>
<div className="sticker-prop-row">
<span className="sticker-prop-label"></span>
<select
className="sticker-prop-select"
value={sticker.text_preset}
onChange={(e) =>
onUpdate(sticker.id, { text_preset: e.target.value as TextStickerPreset })
}
>
{(Object.keys(TEXT_STICKER_PRESET_LABELS) as TextStickerPreset[]).map((p) => (
<option key={p} value={p}>
{TEXT_STICKER_PRESET_LABELS[p]}
</option>
))}
</select>
</div>
<div className="sticker-prop-row">
<span className="sticker-prop-label"></span>
<input
type="range"
className="sticker-prop-slider"
min={12}
max={72}
value={sticker.font_size}
onChange={(e) => onUpdate(sticker.id, { font_size: Number(e.target.value) })}
/>
<span className="sticker-prop-value">{sticker.font_size}px</span>
</div>
<div className="sticker-prop-row">
<span className="sticker-prop-label"></span>
<input
type="color"
className="sticker-prop-color"
value={sticker.text_color}
onChange={(e) => onUpdate(sticker.id, { text_color: e.target.value })}
/>
</div>
</>
)}
{/* 预览 */}
<StickerPreview sticker={sticker} />
<div className="sticker-preview-box">
<div
className="sticker-preview-item"
style={{
left: `${sticker.x}%`,
top: `${sticker.y}%`,
width: `${sticker.width}%`,
height: `${sticker.width}%`,
transform: `translate(-50%, -50%) rotate(${sticker.rotation}deg)`,
opacity: sticker.opacity / 100,
fontSize: sticker.type === "text" ? `${sticker.font_size}px` : undefined,
...TEXT_PRESET_STYLES[sticker.text_preset],
}}
>
{sticker.type === "emoji" && sticker.content}
{sticker.type === "text" && sticker.content}
{sticker.type === "image" && (
<img
src={sticker.content}
alt="sticker"
style={{
width: "100%",
height: "100%",
objectFit: "contain",
}}
/>
)}
</div>
</div>
</div>
)
}
@@ -1,57 +0,0 @@
import React from "react"
import type { StickerItem, TextStickerPreset } from "@/pages/editing-planner/types"
import { TEXT_STICKER_PRESET_LABELS } from "@/pages/editing-planner/constants/sticker"
interface TextStickerPropsEditorProps {
sticker: StickerItem
onUpdate: (id: string, partial: Partial<StickerItem>) => void
}
export const TextStickerPropsEditor: React.FC<TextStickerPropsEditorProps> = ({
sticker,
onUpdate,
}) => {
if (sticker.type !== "text") return null
return (
<>
<div className="sticker-prop-row">
<span className="sticker-prop-label"></span>
<select
className="sticker-prop-select"
value={sticker.text_preset}
onChange={(e) =>
onUpdate(sticker.id, { text_preset: e.target.value as TextStickerPreset })
}
>
{(Object.keys(TEXT_STICKER_PRESET_LABELS) as TextStickerPreset[]).map((p) => (
<option key={p} value={p}>
{TEXT_STICKER_PRESET_LABELS[p]}
</option>
))}
</select>
</div>
<div className="sticker-prop-row">
<span className="sticker-prop-label"></span>
<input
type="range"
className="sticker-prop-slider"
min={12}
max={72}
value={sticker.font_size}
onChange={(e) => onUpdate(sticker.id, { font_size: Number(e.target.value) })}
/>
<span className="sticker-prop-value">{sticker.font_size}px</span>
</div>
<div className="sticker-prop-row">
<span className="sticker-prop-label"></span>
<input
type="color"
className="sticker-prop-color"
value={sticker.text_color}
onChange={(e) => onUpdate(sticker.id, { text_color: e.target.value })}
/>
</div>
</>
)
}
+1 -4
View File
@@ -1,7 +1,3 @@
/**
* Auth API 测试
* 对应 api/auth/ 目录化后的模块
*/
import { describe, expect, it, vi, beforeEach } from "vitest"
import {
normalizeUser,
@@ -14,6 +10,7 @@ import {
resetPassword,
verifyEmail,
} from "@/api/auth"
const mockPost = vi.fn()
const mockGet = vi.fn()
const mockAxiosPost = vi.fn()
-2
View File
@@ -61,8 +61,6 @@ import "@/pages/editing-planner/components/pip-config/LayerConfig"
import "@/pages/editing-planner/components/sticker/StickerLibrary"
import "@/pages/editing-planner/components/sticker/StickerList"
import "@/pages/editing-planner/components/sticker/StickerPropsEditor"
import "@/pages/editing-planner/components/sticker/StickerPreview"
import "@/pages/editing-planner/components/sticker/TextStickerPropsEditor"
import "@/pages/editing-planner/components/filter/FilterPresetGrid"
import "@/pages/editing-planner/components/filter/FilterManualAdjust"
import "@/pages/editing-planner/components/intro-outro/IntroOutroBlock"
-281
View File
@@ -1,281 +0,0 @@
"""edit_template 剪辑模板实体单测."""
from datetime import datetime, timezone
import pytest
from domain.edit_template import EditTemplate, EditTemplateStatus
from domain.editing_mode import EditingMode
# ── EditTemplateStatus 枚举 ──────────────────────────────────────────────────
class TestEditTemplateStatus:
"""EditTemplateStatus 枚举"""
def test_enum_values(self):
assert EditTemplateStatus.ACTIVE.value == "active"
assert EditTemplateStatus.INACTIVE.value == "inactive"
def test_is_str_enum(self):
assert isinstance(EditTemplateStatus.ACTIVE, str)
assert EditTemplateStatus.ACTIVE == "active"
def test_from_string(self):
assert EditTemplateStatus("active") == EditTemplateStatus.ACTIVE
assert EditTemplateStatus("inactive") == EditTemplateStatus.INACTIVE
def test_invalid_raises(self):
with pytest.raises(ValueError):
EditTemplateStatus("deleted")
# ── EditTemplate.create 工厂方法 ────────────────────────────────────────────
class TestEditTemplateCreate:
"""EditTemplate.create 工厂方法"""
def test_minimal_create(self):
t = EditTemplate.create("测试模板")
assert t.id is not None
assert len(t.id) == 32 # uuid4 hex
assert t.name == "测试模板"
assert t.description == ""
assert t.template_type == "default"
assert t.editing_mode == "one_take"
assert t.config == {}
assert t.preview_url == ""
assert t.sort_weight == 0
assert t.status == EditTemplateStatus.ACTIVE
assert t.version == 1
def test_unique_ids(self):
t1 = EditTemplate.create("模板A")
t2 = EditTemplate.create("模板B")
assert t1.id != t2.id
def test_custom_fields(self):
t = EditTemplate.create(
"自定义模板",
description="这是一个自定义模板",
template_type="story",
editing_mode="one_take",
config={"key": "value"},
preview_url="https://example.com/preview.mp4",
sort_weight=100,
status=EditTemplateStatus.INACTIVE,
version=2,
)
assert t.name == "自定义模板"
assert t.description == "这是一个自定义模板"
assert t.template_type == "story"
assert t.editing_mode == "one_take"
assert t.config == {"key": "value"}
assert t.preview_url == "https://example.com/preview.mp4"
assert t.sort_weight == 100
assert t.status == EditTemplateStatus.INACTIVE
assert t.version == 2
def test_name_stripped(self):
t = EditTemplate.create(" 带空格的模板 ")
assert t.name == "带空格的模板"
def test_empty_name_raises(self):
with pytest.raises(ValueError, match="名称"):
EditTemplate.create("")
def test_whitespace_only_name_raises(self):
with pytest.raises(ValueError):
EditTemplate.create(" ")
def test_invalid_editing_mode_raises(self):
with pytest.raises(ValueError, match="editing_mode"):
EditTemplate.create("测试", editing_mode="invalid_mode")
def test_empty_editing_mode_falls_back_to_default(self):
t = EditTemplate.create("测试", editing_mode="")
assert t.editing_mode == "one_take"
def test_whitespace_editing_mode_falls_back(self):
t = EditTemplate.create("测试", editing_mode=" ")
assert t.editing_mode == "one_take"
def test_editing_mode_stripped(self):
t = EditTemplate.create("测试", editing_mode=" one_take ")
assert t.editing_mode == "one_take"
def test_description_stripped(self):
t = EditTemplate.create("测试", description=" 描述 ")
assert t.description == "描述"
def test_template_type_stripped(self):
t = EditTemplate.create("测试", template_type=" vlog ")
assert t.template_type == "vlog"
def test_empty_template_type_falls_back(self):
t = EditTemplate.create("测试", template_type="")
assert t.template_type == "default"
def test_none_config_becomes_empty_dict(self):
t = EditTemplate.create("测试", config=None)
assert t.config == {}
assert isinstance(t.config, dict)
def test_preview_url_stripped(self):
t = EditTemplate.create("测试", preview_url=" https://x.com/a.mp4 ")
assert t.preview_url == "https://x.com/a.mp4"
def test_timestamps_are_utc(self):
t = EditTemplate.create("测试")
assert t.created_at.tzinfo is not None
assert t.updated_at.tzinfo is not None
def test_created_at_equals_updated_at_on_create(self):
t = EditTemplate.create("测试")
# 创建时两个时间应该非常接近
diff = abs((t.updated_at - t.created_at).total_seconds())
assert diff < 1.0
# ── 状态操作 ─────────────────────────────────────────────────────────────────
class TestEditTemplateStatusOperations:
"""EditTemplate 状态操作"""
def test_activate_sets_active(self):
t = EditTemplate.create("测试", status=EditTemplateStatus.INACTIVE)
t.activate()
assert t.status == EditTemplateStatus.ACTIVE
assert t.is_active is True
def test_deactivate_sets_inactive(self):
t = EditTemplate.create("测试")
t.deactivate()
assert t.status == EditTemplateStatus.INACTIVE
assert t.is_active is False
def test_is_active_true(self):
t = EditTemplate.create("测试")
assert t.is_active is True
def test_is_active_false(self):
t = EditTemplate.create("测试", status=EditTemplateStatus.INACTIVE)
assert t.is_active is False
def test_activate_updates_updated_at(self):
t = EditTemplate.create("测试", status=EditTemplateStatus.INACTIVE)
old_updated = t.updated_at
t.activate()
assert t.updated_at >= old_updated
def test_deactivate_updates_updated_at(self):
t = EditTemplate.create("测试")
old_updated = t.updated_at
t.deactivate()
assert t.updated_at >= old_updated
# ── 版本操作 ─────────────────────────────────────────────────────────────────
class TestEditTemplateVersion:
"""EditTemplate 版本操作"""
def test_bump_version_increments(self):
t = EditTemplate.create("测试")
assert t.version == 1
t.bump_version()
assert t.version == 2
def test_bump_version_multiple(self):
t = EditTemplate.create("测试", version=5)
t.bump_version()
t.bump_version()
t.bump_version()
assert t.version == 8
def test_bump_version_updates_updated_at(self):
t = EditTemplate.create("测试")
old_updated = t.updated_at
t.bump_version()
assert t.updated_at >= old_updated
# ── dataclass 基础特性 ───────────────────────────────────────────────────────
class TestEditTemplateBasics:
"""EditTemplate 基础特性"""
def test_slots_no_extra_attrs(self):
t = EditTemplate.create("测试")
with pytest.raises(AttributeError):
t.nonexistent_field = "value"
def test_direct_construction_minimal(self):
# 最小构造:仅必填字段 + 状态,其余走默认值
t = EditTemplate(
id="custom_id",
name="直接构造",
status=EditTemplateStatus.ACTIVE,
)
assert t.id == "custom_id"
assert t.name == "直接构造"
assert t.status == EditTemplateStatus.ACTIVE
# 默认值检查
assert t.description == ""
assert t.config == {}
assert t.version == 1
assert t.editing_mode == EditingMode.ONE_TAKE.value
assert isinstance(t.created_at, datetime)
assert isinstance(t.updated_at, datetime)
def test_direct_construction_full(self):
# 完整构造:所有字段都传
now = datetime(2025, 1, 1, tzinfo=timezone.utc)
t = EditTemplate(
id="full_id",
name="完整构造",
description="测试描述",
template_type="custom",
editing_mode=EditingMode.PIP.value,
config={"key": "value"},
preview_url="https://example.com/preview.jpg",
sort_weight=100,
status=EditTemplateStatus.INACTIVE,
version=3,
created_at=now,
updated_at=now,
)
assert t.id == "full_id"
assert t.name == "完整构造"
assert t.description == "测试描述"
assert t.template_type == "custom"
assert t.editing_mode == EditingMode.PIP.value
assert t.config == {"key": "value"}
assert t.preview_url == "https://example.com/preview.jpg"
assert t.sort_weight == 100
assert t.status == EditTemplateStatus.INACTIVE
assert t.version == 3
assert t.created_at == now
assert t.updated_at == now
def test_config_is_independent(self):
# 不同实例的 config 应该是独立的 dict
t1 = EditTemplate.create("模板1")
t2 = EditTemplate.create("模板2")
t1.config["key"] = "value"
assert "key" not in t2.config
def test_equality(self):
# 两个不同实例即使内容相同也不等(id不同)
t1 = EditTemplate.create("同名模板")
t2 = EditTemplate.create("同名模板")
assert t1 != t2
def test_same_id_equal(self):
now = datetime.now(timezone.utc)
t1 = EditTemplate(id="same", name="同名", created_at=now, updated_at=now)
t2 = EditTemplate(id="same", name="同名", created_at=now, updated_at=now)
assert t1 == t2
+204 -198
View File
@@ -1,55 +1,58 @@
"""视频分享领域模型单元测试 — wave215"""
"""video_share 视频分享领域实体单测."""
from __future__ import annotations
import re
from datetime import datetime, timedelta, timezone
import pytest
from packages.domain.video_share import (
from domain.video_share import (
VideoShare,
_hash_password,
generate_share_token,
)
# ── 密码哈希 ─────────────────────────────────────────────────────────────────
# ── _hash_password ───────────────────────────────────────────────────────────
class TestHashPassword:
"""_hash_password 函数"""
def test_empty_password_returns_empty(self):
assert _hash_password("") == ""
def test_same_password_same_hash(self):
h1 = _hash_password("secret123")
h2 = _hash_password("secret123")
assert h1 == h2
assert h1 != ""
def test_none_password_returns_empty(self):
assert _hash_password(None) == ""
def test_different_password_different_hash(self):
h1 = _hash_password("pass1")
h2 = _hash_password("pass2")
def test_same_password_same_hash(self):
h1 = _hash_password("mypassword")
h2 = _hash_password("mypassword")
assert h1 == h2
def test_different_passwords_different_hashes(self):
h1 = _hash_password("password1")
h2 = _hash_password("password2")
assert h1 != h2
def test_hash_is_sha256_hex(self):
def test_hash_is_hex_string(self):
h = _hash_password("test")
assert len(h) == 64
assert re.match(r"^[0-9a-f]{64}$", h)
assert isinstance(h, str)
assert len(h) == 64 # SHA-256 hex
int(h, 16) # 应该能被解析为16进制
def test_hash_contains_salt(self):
# 直接SHA-256("test") vs 加盐后的结果应该不同
import hashlib
# 直接SHA-256(password) 应该不等于加盐后的
from hashlib import sha256
direct = hashlib.sha256(b"test").hexdigest()
salted = _hash_password("test")
assert direct != salted
raw = sha256("mypass".encode()).hexdigest()
salted = _hash_password("mypass")
assert raw != salted
# ── Token 生成 ──────────────────────────────────────────────────────────────
# ── generate_share_token ─────────────────────────────────────────────────────
class TestGenerateShareToken:
def test_default_length_12(self):
"""generate_share_token 函数"""
def test_default_length(self):
token = generate_share_token()
assert len(token) == 12
@@ -57,228 +60,231 @@ class TestGenerateShareToken:
token = generate_share_token(20)
assert len(token) == 20
def test_url_friendly_no_ambiguous_chars(self):
# 不应包含容易混淆的字符:i, l, o, I, L, O, 0, 1
token = generate_share_token(100)
for ch in "ilO01":
assert ch not in token
def test_short_token(self):
token = generate_share_token(6)
assert len(token) == 6
def test_alphanumeric_only(self):
def test_url_friendly_chars(self):
token = generate_share_token(100)
# 不应该有容易混淆的字符 i,l,o,0,1
assert "i" not in token
assert "l" not in token
assert "o" not in token
assert "0" not in token
assert "1" not in token
def test_unique_tokens(self):
tokens = {generate_share_token() for _ in range(100)}
assert len(tokens) == 100 # 应该都是唯一的
def test_alphanumeric(self):
token = generate_share_token(50)
assert token.isalnum()
def test_two_tokens_different(self):
# 随机生成的两个token应该不同
t1 = generate_share_token()
t2 = generate_share_token()
assert t1 != t2
# ── VideoShare.create ───────────────────────────────────────────────────────
class TestVideoShareCreate:
def test_basic_create(self):
share = VideoShare.create(video_id="v1", user_id="u1")
assert share.id is not None
assert share.video_id == "v1"
assert share.user_id == "u1"
assert share.share_token is not None
assert len(share.share_token) == 12
assert share.password_hash is None
assert share.expires_at is None
assert share.view_count == 0
assert share.download_count == 0
assert share.is_active is True
assert share.created_at is not None
assert share.updated_at is not None
"""VideoShare.create 工厂方法"""
def test_create_with_password(self):
share = VideoShare.create(video_id="v1", user_id="u1", password="secret")
assert share.password_hash is not None
assert share.password_hash != "secret"
assert len(share.password_hash) == 64
def test_minimal_create(self):
s = VideoShare.create(video_id="vid_001", user_id="user_001")
assert s.id is not None
assert len(s.id) == 32 # uuid4 hex
assert s.video_id == "vid_001"
assert s.user_id == "user_001"
assert s.share_token is not None
assert len(s.share_token) == 12
assert s.password_hash is None
assert s.expires_at is None
assert s.view_count == 0
assert s.download_count == 0
assert s.is_active is True
def test_create_with_empty_password_no_hash(self):
share = VideoShare.create(video_id="v1", user_id="u1", password="")
assert share.password_hash is None
def test_with_password(self):
s = VideoShare.create(video_id="v1", user_id="u1", password="secret123")
assert s.password_hash is not None
assert s.password_hash != "secret123" # 不是明文
assert len(s.password_hash) == 64 # SHA-256
def test_create_with_expires_at(self):
def test_with_expiry(self):
future = datetime.now(timezone.utc) + timedelta(days=7)
share = VideoShare.create(video_id="v1", user_id="u1", expires_at=future)
assert share.expires_at == future
s = VideoShare.create(video_id="v1", user_id="u1", expires_at=future)
assert s.expires_at == future
def test_create_past_expires_at_raises(self):
past = datetime.now(timezone.utc) - timedelta(days=1)
with pytest.raises(ValueError, match="expires_at cannot be in the past"):
VideoShare.create(video_id="v1", user_id="u1", expires_at=past)
def test_create_empty_video_id_raises(self):
with pytest.raises(ValueError, match="video_id cannot be empty"):
def test_empty_video_id_raises(self):
with pytest.raises(ValueError, match="video_id"):
VideoShare.create(video_id="", user_id="u1")
def test_create_whitespace_video_id_raises(self):
with pytest.raises(ValueError, match="video_id cannot be empty"):
def test_whitespace_video_id_raises(self):
with pytest.raises(ValueError):
VideoShare.create(video_id=" ", user_id="u1")
def test_create_empty_user_id_raises(self):
with pytest.raises(ValueError, match="user_id cannot be empty"):
def test_empty_user_id_raises(self):
with pytest.raises(ValueError, match="user_id"):
VideoShare.create(video_id="v1", user_id="")
def test_create_strips_whitespace(self):
share = VideoShare.create(video_id=" v1 ", user_id=" u1 ")
assert share.video_id == "v1"
assert share.user_id == "u1"
def test_past_expiry_raises(self):
past = datetime.now(timezone.utc) - timedelta(hours=1)
with pytest.raises(ValueError, match="past"):
VideoShare.create(video_id="v1", user_id="u1", expires_at=past)
def test_create_unique_id_each_time(self):
def test_video_id_stripped(self):
s = VideoShare.create(video_id=" vid_123 ", user_id="u1")
assert s.video_id == "vid_123"
def test_user_id_stripped(self):
s = VideoShare.create(video_id="v1", user_id=" user_456 ")
assert s.user_id == "user_456"
def test_unique_ids(self):
s1 = VideoShare.create(video_id="v1", user_id="u1")
s2 = VideoShare.create(video_id="v1", user_id="u1")
assert s1.id != s2.id
def test_create_unique_token_each_time(self):
def test_unique_tokens(self):
s1 = VideoShare.create(video_id="v1", user_id="u1")
s2 = VideoShare.create(video_id="v1", user_id="u1")
assert s1.share_token != s2.share_token
# ── has_password ────────────────────────────────────────────────────────────
def test_timestamps_set(self):
s = VideoShare.create(video_id="v1", user_id="u1")
assert s.created_at.tzinfo is not None
assert s.updated_at.tzinfo is not None
class TestVideoShareHasPassword:
def test_no_password(self):
share = VideoShare.create(video_id="v1", user_id="u1")
assert share.has_password is False
def test_with_password(self):
share = VideoShare.create(video_id="v1", user_id="u1", password="pass")
assert share.has_password is True
def test_empty_password_none(self):
share = VideoShare.create(video_id="v1", user_id="u1", password="")
assert share.has_password is False
# ── VideoShare 属性方法 ─────────────────────────────────────────────────────
# ── is_expired ──────────────────────────────────────────────────────────────
class TestVideoShareProperties:
"""VideoShare 属性方法"""
def test_has_password_true(self):
s = VideoShare.create(video_id="v1", user_id="u1", password="pass")
assert s.has_password is True
class TestVideoShareIsExpired:
def test_no_expiry_never_expired(self):
share = VideoShare.create(video_id="v1", user_id="u1")
assert share.is_expired is False
def test_has_password_false(self):
s = VideoShare.create(video_id="v1", user_id="u1")
assert s.has_password is False
def test_future_expiry_not_expired(self):
def test_is_expired_false_no_expiry(self):
s = VideoShare.create(video_id="v1", user_id="u1")
assert s.is_expired is False
def test_is_expired_false_future_expiry(self):
future = datetime.now(timezone.utc) + timedelta(hours=1)
share = VideoShare.create(video_id="v1", user_id="u1", expires_at=future)
assert share.is_expired is False
s = VideoShare.create(video_id="v1", user_id="u1", expires_at=future)
assert s.is_expired is False
def test_past_expiry_is_expired(self):
share = VideoShare.create(video_id="v1", user_id="u1")
share.expires_at = datetime.now(timezone.utc) - timedelta(seconds=1)
assert share.is_expired is True
def test_is_expired_true_past_expiry(self):
# 直接构造一个已过期的
past = datetime.now(timezone.utc) - timedelta(hours=1)
s = VideoShare(
id="test",
video_id="v1",
user_id="u1",
share_token="abc",
expires_at=past,
)
assert s.is_expired is True
def test_is_accessible_true(self):
s = VideoShare.create(video_id="v1", user_id="u1")
assert s.is_accessible is True
def test_is_accessible_false_inactive(self):
s = VideoShare.create(video_id="v1", user_id="u1")
s.is_active = False
assert s.is_accessible is False
def test_is_accessible_false_expired(self):
past = datetime.now(timezone.utc) - timedelta(hours=1)
s = VideoShare(
id="test",
video_id="v1",
user_id="u1",
share_token="abc",
expires_at=past,
)
assert s.is_accessible is False
# ── is_accessible ───────────────────────────────────────────────────────────
# ── VideoShare 方法 ─────────────────────────────────────────────────────────
class TestVideoShareIsAccessible:
def test_active_no_expiry_accessible(self):
share = VideoShare.create(video_id="v1", user_id="u1")
assert share.is_accessible is True
class TestVideoShareMethods:
"""VideoShare 方法"""
def test_revoked_not_accessible(self):
share = VideoShare.create(video_id="v1", user_id="u1")
share.is_active = False
assert share.is_accessible is False
def test_verify_password_no_password_true(self):
s = VideoShare.create(video_id="v1", user_id="u1")
assert s.verify_password("anything") is True
assert s.verify_password("") is True
def test_expired_not_accessible(self):
share = VideoShare.create(video_id="v1", user_id="u1")
share.expires_at = datetime.now(timezone.utc) - timedelta(days=1)
assert share.is_accessible is False
def test_verify_password_correct(self):
s = VideoShare.create(video_id="v1", user_id="u1", password="mypass")
assert s.verify_password("mypass") is True
def test_revoked_and_expired_not_accessible(self):
share = VideoShare.create(video_id="v1", user_id="u1")
share.is_active = False
share.expires_at = datetime.now(timezone.utc) - timedelta(days=1)
assert share.is_accessible is False
def test_verify_password_wrong(self):
s = VideoShare.create(video_id="v1", user_id="u1", password="mypass")
assert s.verify_password("wrongpass") is False
def test_verify_password_empty_false(self):
s = VideoShare.create(video_id="v1", user_id="u1", password="mypass")
assert s.verify_password("") is False
# ── verify_password ─────────────────────────────────────────────────────────
def test_increment_view_count(self):
s = VideoShare.create(video_id="v1", user_id="u1")
assert s.view_count == 0
s.increment_view_count()
assert s.view_count == 1
s.increment_view_count()
assert s.view_count == 2
def test_increment_download_count(self):
s = VideoShare.create(video_id="v1", user_id="u1")
assert s.download_count == 0
s.increment_download_count()
assert s.download_count == 1
s.increment_download_count()
assert s.download_count == 2
class TestVideoShareVerifyPassword:
def test_no_password_any_pass_ok(self):
share = VideoShare.create(video_id="v1", user_id="u1")
assert share.verify_password("anything") is True
assert share.verify_password("") is True
def test_no_password_none_ok(self):
share = VideoShare.create(video_id="v1", user_id="u1")
assert share.verify_password("") is True
def test_correct_password(self):
share = VideoShare.create(video_id="v1", user_id="u1", password="mysecret")
assert share.verify_password("mysecret") is True
def test_wrong_password(self):
share = VideoShare.create(video_id="v1", user_id="u1", password="mysecret")
assert share.verify_password("wrong") is False
def test_empty_password_with_protection(self):
share = VideoShare.create(video_id="v1", user_id="u1", password="mysecret")
assert share.verify_password("") is False
def test_password_case_sensitive(self):
share = VideoShare.create(video_id="v1", user_id="u1", password="Secret")
assert share.verify_password("secret") is False
assert share.verify_password("Secret") is True
# ── 计数方法 ────────────────────────────────────────────────────────────────
class TestVideoShareCounters:
def test_increment_view(self):
share = VideoShare.create(video_id="v1", user_id="u1")
assert share.view_count == 0
share.increment_view_count()
assert share.view_count == 1
share.increment_view_count()
assert share.view_count == 2
def test_increment_download(self):
share = VideoShare.create(video_id="v1", user_id="u1")
assert share.download_count == 0
share.increment_download_count()
assert share.download_count == 1
share.increment_download_count()
assert share.download_count == 2
def test_counters_independent(self):
share = VideoShare.create(video_id="v1", user_id="u1")
share.increment_view_count()
share.increment_view_count()
share.increment_download_count()
assert share.view_count == 2
assert share.download_count == 1
# ── revoke ──────────────────────────────────────────────────────────────────
class TestVideoShareRevoke:
def test_revoke_sets_inactive(self):
share = VideoShare.create(video_id="v1", user_id="u1")
assert share.is_active is True
share.revoke()
assert share.is_active is False
def test_revoke(self):
s = VideoShare.create(video_id="v1", user_id="u1")
assert s.is_active is True
s.revoke()
assert s.is_active is False
def test_revoke_makes_inaccessible(self):
share = VideoShare.create(video_id="v1", user_id="u1")
share.revoke()
assert share.is_accessible is False
s = VideoShare.create(video_id="v1", user_id="u1")
assert s.is_accessible is True
s.revoke()
assert s.is_accessible is False
def test_revoke_idempotent(self):
share = VideoShare.create(video_id="v1", user_id="u1")
share.revoke()
share.revoke() # 第二次也不报错
assert share.is_active is False
# ── dataclass 基础特性 ───────────────────────────────────────────────────────
class TestVideoShareBasics:
"""VideoShare 基础特性"""
def test_slots_no_extra_attrs(self):
s = VideoShare.create(video_id="v1", user_id="u1")
with pytest.raises(AttributeError):
s.nonexistent = "value"
def test_direct_construction(self):
s = VideoShare(
id="custom_id",
video_id="v1",
user_id="u1",
share_token="abc123",
)
assert s.id == "custom_id"
assert s.share_token == "abc123"
def test_equality_same_id(self):
now = datetime.now(timezone.utc)
s1 = VideoShare(id="same", video_id="v1", user_id="u1", share_token="t", created_at=now, updated_at=now)
s2 = VideoShare(id="same", video_id="v1", user_id="u1", share_token="t", created_at=now, updated_at=now)
assert s1 == s2
File diff suppressed because it is too large Load Diff
+298 -387
View File
@@ -1,68 +1,63 @@
"""JWT 服务与处理器单元测试."""
"""JWT 服务单元测试 — wave130."""
from __future__ import annotations
import time
from datetime import datetime, timedelta, timezone
import jwt
import jwt as pyjwt
import pytest
from jwt.exceptions import ExpiredSignatureError, InvalidTokenError
from packages.application.auth.jwt_handler import (
JWTHandler,
configure_jwt_handler,
get_jwt_handler,
)
from packages.application.auth.jwt_service import (
JWTConfig,
JWTService,
TokenType,
)
# ── 测试常量 ──────────────────────────────────────────────────────────────────
TEST_SECRET = "test-secret-key-for-unit-testing-only-not-for-production"
STRONG_SECRET = "x" * 32 # 满足长度要求的测试密钥
# ── 测试常量 ────────────────────────────────────────────────────────────────
# ── JWTConfig 测试 ───────────────────────────────────────────────────────────
TEST_SECRET = "test-secret-key-for-unit-testing-only-1234567890"
TEST_ALGORITHM = "HS256"
# ── JWTConfig 配置 ──────────────────────────────────────────────────────────
class TestJWTConfig:
"""JWTConfig 配置类测试"""
def test_init_with_valid_secret(self):
config = JWTConfig(secret_key=STRONG_SECRET)
assert config.SECRET_KEY == STRONG_SECRET
def test_normal_config(self):
config = JWTConfig(secret_key=TEST_SECRET)
assert config.SECRET_KEY == TEST_SECRET
assert config.ALGORITHM == "HS256"
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 15
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 7
def test_init_custom_values(self):
def test_custom_config(self):
config = JWTConfig(
secret_key=STRONG_SECRET,
secret_key=TEST_SECRET,
algorithm="HS384",
access_token_expire_minutes=60,
refresh_token_expire_days=14,
refresh_token_expire_days=30,
)
assert config.SECRET_KEY == STRONG_SECRET
assert config.ALGORITHM == "HS384"
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 60
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 14
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 30
def test_empty_secret_raises(self):
with pytest.raises(ValueError, match="secret_key must be provided"):
JWTConfig(secret_key="")
def test_whitespace_only_secret_raises(self):
with pytest.raises(ValueError, match="secret_key must be provided"):
def test_whitespace_secret_raises(self):
with pytest.raises(ValueError):
JWTConfig(secret_key=" ")
def test_none_secret_raises(self):
with pytest.raises(ValueError, match="secret_key must be provided"):
JWTConfig(secret_key=None)
with pytest.raises(ValueError):
JWTConfig(secret_key=None) # type: ignore
@pytest.mark.parametrize(
"insecure_secret",
"bad_secret",
[
"your-secret-key-change-in-production",
"your-secret-key",
@@ -73,407 +68,323 @@ class TestJWTConfig:
"Your-Secret-Key",
],
)
def test_insecure_default_secret_raises(self, insecure_secret):
def test_insecure_defaults_rejected(self, bad_secret):
with pytest.raises(ValueError, match="insecure"):
JWTConfig(secret_key=insecure_secret)
def test_zero_expire_minutes_allowed(self):
config = JWTConfig(secret_key=STRONG_SECRET, access_token_expire_minutes=0)
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 0
def test_negative_expire_days_allowed(self):
# 配置类不校验合理性,由业务层判断
config = JWTConfig(secret_key=STRONG_SECRET, refresh_token_expire_days=-1)
assert config.REFRESH_TOKEN_EXPIRE_DAYS == -1
JWTConfig(secret_key=bad_secret)
# ── JWTService 初始化测试 ────────────────────────────────────────────────────
# ── JWTService 初始化 ──────────────────────────────────────────────────────
class TestJWTServiceInit:
"""JWTService 初始化测试"""
def test_init_with_config(self):
config = JWTConfig(secret_key=STRONG_SECRET)
def test_with_config_works(self):
config = JWTConfig(secret_key=TEST_SECRET)
service = JWTService(config)
assert service.config is config
def test_init_none_config_raises(self):
with pytest.raises(ValueError, match="JWTService requires a JWTConfig"):
def test_none_config_raises(self):
with pytest.raises(ValueError, match="JWTService requires"):
JWTService(None)
# ── TokenType 测试 ───────────────────────────────────────────────────────────
# ── create_access_token ────────────────────────────────────────────────────
class TestCreateAccessToken:
def setup_method(self):
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
def test_creates_valid_jwt(self):
token = self.service.create_access_token(user_id="user123")
assert isinstance(token, str)
assert len(token) > 0
# JWT 格式:xxx.yyy.zzz
assert token.count(".") == 2
def test_payload_contains_user_id(self):
token = self.service.create_access_token(user_id="user_001")
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
assert payload["sub"] == "user_001"
def test_payload_contains_role(self):
token = self.service.create_access_token(user_id="u1", role="admin")
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
assert payload["role"] == "admin"
def test_default_role_empty(self):
token = self.service.create_access_token(user_id="u1")
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
assert payload["role"] == ""
def test_token_type_is_access(self):
token = self.service.create_access_token(user_id="u1")
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
assert payload["type"] == TokenType.ACCESS
def test_has_iat_and_exp(self):
token = self.service.create_access_token(user_id="u1")
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
assert "iat" in payload
assert "exp" in payload
assert payload["exp"] > payload["iat"]
def test_expiration_correct(self):
"""过期时间大约等于当前时间 + 配置的分钟数."""
config = JWTConfig(secret_key=TEST_SECRET, access_token_expire_minutes=30)
service = JWTService(config)
before = datetime.now(timezone.utc)
token = service.create_access_token(user_id="u1")
after = datetime.now(timezone.utc)
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
exp = datetime.fromtimestamp(payload["exp"], tz=timezone.utc)
min_expected = before + timedelta(minutes=30) - timedelta(seconds=1)
max_expected = after + timedelta(minutes=30) + timedelta(seconds=1)
assert min_expected <= exp <= max_expected
def test_additional_claims_included(self):
extra = {"email": "test@example.com", "org_id": "org_001", "level": 5}
token = self.service.create_access_token(user_id="u1", additional_claims=extra)
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
assert payload["email"] == "test@example.com"
assert payload["org_id"] == "org_001"
assert payload["level"] == 5
def test_additional_claims_none(self):
token = self.service.create_access_token(user_id="u1", additional_claims=None)
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
assert "email" not in payload
def test_signed_with_correct_key(self):
token = self.service.create_access_token(user_id="u1")
# 用正确的密钥可以解码
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
assert payload["sub"] == "u1"
# 用错误的密钥无法解码
with pytest.raises(InvalidTokenError):
pyjwt.decode(token, "wrong-secret", algorithms=["HS256"])
# ── create_refresh_token ───────────────────────────────────────────────────
class TestCreateRefreshToken:
def setup_method(self):
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
def test_creates_valid_token(self):
token = self.service.create_refresh_token(user_id="u1", session_id="sess_001")
assert isinstance(token, str)
assert token.count(".") == 2
def test_payload_contains_session_id(self):
token = self.service.create_refresh_token(user_id="u1", session_id="sess_abc")
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
assert payload["session_id"] == "sess_abc"
assert payload["sub"] == "u1"
def test_token_type_is_refresh(self):
token = self.service.create_refresh_token(user_id="u1", session_id="s1")
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
assert payload["type"] == TokenType.REFRESH
def test_refresh_expiration_days(self):
config = JWTConfig(secret_key=TEST_SECRET, refresh_token_expire_days=7)
service = JWTService(config)
before = datetime.now(timezone.utc)
token = service.create_refresh_token(user_id="u1", session_id="s1")
after = datetime.now(timezone.utc)
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
exp = datetime.fromtimestamp(payload["exp"], tz=timezone.utc)
min_exp = before + timedelta(days=7) - timedelta(seconds=1)
max_exp = after + timedelta(days=7, seconds=1)
assert min_exp <= exp <= max_exp
# ── verify_token ───────────────────────────────────────────────────────────
class TestVerifyToken:
def setup_method(self):
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
def test_valid_token_returns_payload(self):
token = self.service.create_access_token(user_id="u1")
payload = self.service.verify_token(token)
assert payload["sub"] == "u1"
def test_expired_token_raises(self):
# 创建一个 1 秒过期的 token
config = JWTConfig(secret_key=TEST_SECRET, access_token_expire_minutes=1)
service = JWTService(config)
token = service.create_access_token(user_id="u1")
# 等待过期(用 pyjwt 直接构造过期 token 更可靠)
expired_payload = {
"sub": "u1",
"type": "access",
"exp": datetime.now(timezone.utc) - timedelta(seconds=10),
}
expired_token = pyjwt.encode(expired_payload, TEST_SECRET, algorithm="HS256")
with pytest.raises(ExpiredSignatureError, match="expired"):
self.service.verify_token(expired_token)
def test_invalid_token_raises(self):
with pytest.raises(InvalidTokenError, match="Invalid token"):
self.service.verify_token("not-a-valid-jwt-token")
def test_wrong_signature_raises(self):
token = pyjwt.encode({"sub": "u1"}, "different-secret", algorithm="HS256")
with pytest.raises(InvalidTokenError):
self.service.verify_token(token)
def test_tampered_payload_raises(self):
token = self.service.create_access_token(user_id="u1")
# 尝试篡改:JWT 有签名保护,篡改会导致验证失败
parts = token.split(".")
assert len(parts) == 3
# 把 payload 部分替换(不会成功,因为签名不对)
import base64
fake_payload = base64.urlsafe_b64encode(b'{"sub":"admin","role":"admin"}').rstrip(b"=").decode()
tampered = f"{parts[0]}.{fake_payload}.{parts[2]}"
with pytest.raises(InvalidTokenError):
self.service.verify_token(tampered)
# ── verify_access_token ─────────────────────────────────────────────────────
class TestVerifyAccessToken:
def setup_method(self):
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
def test_access_token_passes(self):
token = self.service.create_access_token(user_id="u1", role="user")
payload = self.service.verify_access_token(token)
assert payload["sub"] == "u1"
assert payload["type"] == "access"
def test_refresh_token_rejected(self):
token = self.service.create_refresh_token(user_id="u1", session_id="s1")
with pytest.raises(ValueError, match="Token type must be 'access'"):
self.service.verify_access_token(token)
def test_expired_token_raises(self):
expired_payload = {
"sub": "u1",
"type": "access",
"exp": datetime.now(timezone.utc) - timedelta(seconds=10),
}
token = pyjwt.encode(expired_payload, TEST_SECRET, algorithm="HS256")
with pytest.raises(ExpiredSignatureError):
self.service.verify_access_token(token)
# ── verify_refresh_token ────────────────────────────────────────────────────
class TestVerifyRefreshToken:
def setup_method(self):
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
def test_refresh_token_passes(self):
token = self.service.create_refresh_token(user_id="u1", session_id="sess_001")
payload = self.service.verify_refresh_token(token)
assert payload["sub"] == "u1"
assert payload["session_id"] == "sess_001"
def test_access_token_rejected(self):
token = self.service.create_access_token(user_id="u1")
with pytest.raises(ValueError, match="Token type must be 'refresh'"):
self.service.verify_refresh_token(token)
def test_has_session_id(self):
token = self.service.create_refresh_token(user_id="u1", session_id="custom_sess")
payload = self.service.verify_refresh_token(token)
assert payload["session_id"] == "custom_sess"
# ── TokenType 常量 ──────────────────────────────────────────────────────────
class TestTokenType:
"""TokenType 常量测试"""
def test_access_value(self):
assert TokenType.ACCESS == "access"
def test_refresh_value(self):
assert TokenType.REFRESH == "refresh"
def test_access_and_refresh_different(self):
def test_different_types(self):
assert TokenType.ACCESS != TokenType.REFRESH
# ── JWTService create_access_token 测试 ─────────────────────────────────────
# ── 多算法支持 ──────────────────────────────────────────────────────────────
class TestCreateAccessToken:
"""创建 access_token 测试"""
@pytest.fixture
def service(self):
return JWTService(JWTConfig(secret_key=STRONG_SECRET))
def test_creates_valid_jwt_string(self, service):
token = service.create_access_token(user_id="user-123")
assert isinstance(token, str)
assert len(token) > 0
def test_token_contains_user_id_as_sub(self, service):
token = service.create_access_token(user_id="user-123")
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
assert payload["sub"] == "user-123"
def test_token_type_is_access(self, service):
token = service.create_access_token(user_id="user-123")
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
assert payload["type"] == TokenType.ACCESS
def test_default_role_is_empty_string(self, service):
token = service.create_access_token(user_id="user-123")
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
assert payload["role"] == ""
def test_custom_role(self, service):
token = service.create_access_token(user_id="user-123", role="admin")
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
assert payload["role"] == "admin"
def test_has_iat_and_exp(self, service):
token = service.create_access_token(user_id="user-123")
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
assert "iat" in payload
assert "exp" in payload
assert payload["exp"] > payload["iat"]
def test_expire_matches_config(self, service):
token = service.create_access_token(user_id="user-123")
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
iat = datetime.fromtimestamp(payload["iat"], tz=timezone.utc)
exp = datetime.fromtimestamp(payload["exp"], tz=timezone.utc)
delta = exp - iat
assert delta.total_seconds() == 15 * 60 # 15分钟
def test_custom_expire_time(self):
config = JWTConfig(secret_key=STRONG_SECRET, access_token_expire_minutes=30)
class TestDifferentAlgorithms:
def test_hs384_works(self):
config = JWTConfig(secret_key=TEST_SECRET * 2, algorithm="HS384")
service = JWTService(config)
token = service.create_access_token(user_id="user-123")
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
delta = payload["exp"] - payload["iat"]
assert delta == 30 * 60
def test_additional_claims(self, service):
extra = {"custom_field": "value", "another": 42}
token = service.create_access_token(user_id="user-123", additional_claims=extra)
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
assert payload["custom_field"] == "value"
assert payload["another"] == 42
def test_additional_claims_can_override_standard(self, service):
# additional_claims 可以覆盖标准字段(由调用者负责)
token = service.create_access_token(
user_id="user-123",
additional_claims={"sub": "overridden"},
)
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
assert payload["sub"] == "overridden"
def test_additional_claims_none_is_same_as_empty(self, service):
token = service.create_access_token(user_id="user-123", additional_claims=None)
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
assert payload["sub"] == "user-123"
def test_uses_correct_algorithm(self):
config = JWTConfig(secret_key=STRONG_SECRET, algorithm="HS384")
service = JWTService(config)
token = service.create_access_token(user_id="u1")
# 用 HS256 解码应该失败
with pytest.raises(InvalidTokenError):
jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
# 用 HS384 解码应该成功
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS384"])
assert payload["sub"] == "u1"
# ── JWTService create_refresh_token 测试 ────────────────────────────────────
class TestCreateRefreshToken:
"""创建 refresh_token 测试"""
@pytest.fixture
def service(self):
return JWTService(JWTConfig(secret_key=STRONG_SECRET))
def test_creates_valid_string(self, service):
token = service.create_refresh_token(user_id="u1", session_id="s1")
assert isinstance(token, str)
assert len(token) > 0
def test_contains_user_id_and_session_id(self, service):
token = service.create_refresh_token(user_id="u1", session_id="sess-abc")
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
assert payload["sub"] == "u1"
assert payload["session_id"] == "sess-abc"
def test_token_type_is_refresh(self, service):
token = service.create_refresh_token(user_id="u1", session_id="s1")
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
assert payload["type"] == TokenType.REFRESH
def test_has_iat_and_exp(self, service):
token = service.create_refresh_token(user_id="u1", session_id="s1")
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
assert "iat" in payload
assert "exp" in payload
assert payload["exp"] > payload["iat"]
def test_expire_matches_config_days(self, service):
token = service.create_refresh_token(user_id="u1", session_id="s1")
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
delta = payload["exp"] - payload["iat"]
assert delta == 7 * 24 * 60 * 60 # 7天
def test_custom_refresh_expire_days(self):
config = JWTConfig(secret_key=STRONG_SECRET, refresh_token_expire_days=30)
service = JWTService(config)
token = service.create_refresh_token(user_id="u1", session_id="s1")
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
delta = payload["exp"] - payload["iat"]
assert delta == 30 * 24 * 60 * 60
# ── JWTService verify_token 测试 ────────────────────────────────────────────
class TestVerifyToken:
"""通用 Token 验证测试"""
@pytest.fixture
def service(self):
return JWTService(JWTConfig(secret_key=STRONG_SECRET))
def test_verify_valid_access_token(self, service):
token = service.create_access_token(user_id="u1")
payload = service.verify_token(token)
assert payload["sub"] == "u1"
assert payload["type"] == TokenType.ACCESS
def test_verify_valid_refresh_token(self, service):
token = service.create_refresh_token(user_id="u1", session_id="s1")
payload = service.verify_token(token)
assert payload["sub"] == "u1"
assert payload["session_id"] == "s1"
def test_verify_expired_token_raises(self, service):
config = JWTConfig(secret_key=STRONG_SECRET, access_token_expire_minutes=0)
svc = JWTService(config)
token = svc.create_access_token(user_id="u1")
# 0 分钟过期,立即过期
time.sleep(0.1) # 稍微等一下确保过期
with pytest.raises(ExpiredSignatureError):
svc.verify_token(token)
def test_verify_wrong_secret_raises(self, service):
token = service.create_access_token(user_id="u1")
other_service = JWTService(JWTConfig(secret_key="different-secret-1234567890"))
with pytest.raises(InvalidTokenError):
other_service.verify_token(token)
def test_verify_tampered_token_raises(self, service):
token = service.create_access_token(user_id="u1")
# 篡改 token 中间部分
parts = token.split(".")
assert len(parts) == 3
tampered = parts[0] + "." + parts[1][:-1] + "A." + parts[2]
with pytest.raises(InvalidTokenError):
service.verify_token(tampered)
def test_verify_empty_string_raises(self, service):
with pytest.raises(InvalidTokenError):
service.verify_token("")
def test_verify_garbage_string_raises(self, service):
with pytest.raises(InvalidTokenError):
service.verify_token("not.a.valid.jwt.token")
def test_verify_returns_dict(self, service):
token = service.create_access_token(user_id="u1", role="admin")
payload = service.verify_token(token)
assert isinstance(payload, dict)
assert "sub" in payload
assert "role" in payload
# ── JWTService verify_access_token 测试 ─────────────────────────────────────
class TestVerifyAccessToken:
"""Access Token 专属验证测试"""
@pytest.fixture
def service(self):
return JWTService(JWTConfig(secret_key=STRONG_SECRET))
def test_valid_access_token_passes(self, service):
token = service.create_access_token(user_id="u1", role="admin")
payload = service.verify_access_token(token)
assert payload["sub"] == "u1"
assert payload["role"] == "admin"
def test_refresh_token_fails_type_check(self, service):
token = service.create_refresh_token(user_id="u1", session_id="s1")
with pytest.raises(ValueError, match="Token type must be 'access'"):
service.verify_access_token(token)
def test_token_without_type_field_raises(self, service):
# 手动构造一个没有 type 字段的 token
payload_data = {"sub": "u1", "iat": 1000, "exp": 9999999999}
token = jwt.encode(payload_data, STRONG_SECRET, algorithm="HS256")
with pytest.raises(ValueError, match="Token type must be 'access'"):
service.verify_access_token(token)
def test_expired_access_token_raises_expired_error(self, service):
config = JWTConfig(secret_key=STRONG_SECRET, access_token_expire_minutes=0)
svc = JWTService(config)
token = svc.create_access_token(user_id="u1")
time.sleep(0.1)
with pytest.raises(ExpiredSignatureError):
svc.verify_access_token(token)
# ── JWTService verify_refresh_token 测试 ────────────────────────────────────
class TestVerifyRefreshToken:
"""Refresh Token 专属验证测试"""
@pytest.fixture
def service(self):
return JWTService(JWTConfig(secret_key=STRONG_SECRET))
def test_valid_refresh_token_passes(self, service):
token = service.create_refresh_token(user_id="u1", session_id="s1")
payload = service.verify_refresh_token(token)
assert payload["sub"] == "u1"
assert payload["session_id"] == "s1"
def test_access_token_fails_type_check(self, service):
token = service.create_access_token(user_id="u1")
with pytest.raises(ValueError, match="Token type must be 'refresh'"):
service.verify_refresh_token(token)
def test_token_without_type_field_raises(self, service):
payload_data = {"sub": "u1", "session_id": "s1", "iat": 1000, "exp": 9999999999}
token = jwt.encode(payload_data, STRONG_SECRET, algorithm="HS256")
with pytest.raises(ValueError, match="Token type must be 'refresh'"):
service.verify_refresh_token(token)
def test_expired_refresh_token_raises(self):
config = JWTConfig(secret_key=STRONG_SECRET, refresh_token_expire_days=0)
def test_hs512_works(self):
config = JWTConfig(secret_key=TEST_SECRET * 3, algorithm="HS512")
service = JWTService(config)
token = service.create_refresh_token(user_id="u1", session_id="s1")
# 0天过期,应该立即使exp <= iat
with pytest.raises(ExpiredSignatureError):
service.verify_refresh_token(token)
# ── JWTHandler 委托层测试 ────────────────────────────────────────────────────
class TestJWTHandler:
"""JWTHandler 委托层测试"""
def test_init_creates_handler(self):
handler = JWTHandler(secret_key=STRONG_SECRET)
assert handler is not None
def test_create_and_verify_access_token(self):
handler = JWTHandler(secret_key=STRONG_SECRET)
token = handler.create_access_token(user_id="u1", role="user")
payload = handler.verify_access_token(token)
assert payload["sub"] == "u1"
assert payload["role"] == "user"
def test_verify_token_generic(self):
handler = JWTHandler(secret_key=STRONG_SECRET)
token = handler.create_access_token(user_id="u1")
payload = handler.verify_token(token)
token = service.create_access_token(user_id="u1")
payload = service.verify_token(token)
assert payload["sub"] == "u1"
def test_custom_algorithm(self):
handler = JWTHandler(secret_key=STRONG_SECRET, algorithm="HS384")
token = handler.create_access_token(user_id="u1")
payload = handler.verify_access_token(token)
assert payload["sub"] == "u1"
def test_algorithm_mismatch_fails(self):
config_hs256 = JWTConfig(secret_key=TEST_SECRET, algorithm="HS256")
config_hs384 = JWTConfig(secret_key=TEST_SECRET, algorithm="HS384")
service_256 = JWTService(config_hs256)
service_384 = JWTService(config_hs384)
def test_custom_expire_minutes(self):
handler = JWTHandler(secret_key=STRONG_SECRET, access_token_expire_minutes=45)
token = handler.create_access_token(user_id="u1")
payload = handler.verify_access_token(token)
delta = payload["exp"] - payload["iat"]
assert delta == 45 * 60
def test_additional_claims_passthrough(self):
handler = JWTHandler(secret_key=STRONG_SECRET)
extra = {"org_id": "org-1", "plan": "pro"}
token = handler.create_access_token("u1", additional_claims={"org_id": "org-1"})
payload = handler.verify_access_token(
token := handler.create_access_token("u1", additional_claims={"org_id": "org-1"})
)
# 这里直接测试更简洁
payload = handler.verify_access_token(handler.create_access_token("u1", additional_claims={"x": 1}))
assert payload["x"] == 1
token = service_256.create_access_token(user_id="u1")
with pytest.raises(InvalidTokenError):
service_384.verify_token(token)
# ── 全局 JWT handler 测试 ───────────────────────────────────────────────────
# ── 边界:空用户ID等 ────────────────────────────────────────────────────────
class TestGlobalJWTHandler:
"""全局 JWT Handler 配置与获取测试"""
class TestEdgeCases:
def setup_method(self):
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
def test_configure_creates_handler(self):
handler = configure_jwt_handler(secret_key=STRONG_SECRET)
assert isinstance(handler, JWTHandler)
def test_empty_user_id(self):
token = self.service.create_access_token(user_id="")
payload = self.service.verify_access_token(token)
assert payload["sub"] == ""
def test_get_after_configure_works(self):
configure_jwt_handler(secret_key=STRONG_SECRET)
handler = get_jwt_handler()
assert isinstance(handler, JWTHandler)
token = handler.create_access_token(user_id="u1")
payload = handler.verify_access_token(token)
assert payload["sub"] == "u1"
def test_long_user_id(self):
long_id = "x" * 1000
token = self.service.create_access_token(user_id=long_id)
payload = self.service.verify_access_token(token)
assert payload["sub"] == long_id
def test_get_before_configure_raises(self):
# 重置全局状态(通过设置 None 模拟未配置)
import packages.application.auth.jwt_handler as mod
def test_special_chars_in_user_id(self):
uid = "user@#$%^&*()_+-=[]{}|;:',.<>?/`~"
token = self.service.create_access_token(user_id=uid)
payload = self.service.verify_access_token(token)
assert payload["sub"] == uid
mod._default_handler = None
with pytest.raises(RuntimeError, match="JWT handler not configured"):
get_jwt_handler()
def test_unicode_user_id(self):
uid = "用户_测试_123_🎉"
token = self.service.create_access_token(user_id=uid)
payload = self.service.verify_access_token(token)
assert payload["sub"] == uid
def test_configure_returns_same_as_get(self):
h1 = configure_jwt_handler(secret_key=STRONG_SECRET)
h2 = get_jwt_handler()
assert h1 is h2
def test_reconfigure_replaces_handler(self):
h1 = configure_jwt_handler(secret_key=STRONG_SECRET)
h2 = configure_jwt_handler(secret_key=STRONG_SECRET + "_new")
assert h1 is not h2
assert get_jwt_handler() is h2
def test_many_additional_claims(self):
claims = {f"key_{i}": f"value_{i}" for i in range(50)}
token = self.service.create_access_token(user_id="u1", additional_claims=claims)
payload = self.service.verify_access_token(token)
for i in range(50):
assert payload[f"key_{i}"] == f"value_{i}"
File diff suppressed because it is too large Load Diff
+260 -450
View File
@@ -1,4 +1,6 @@
"""密码重置 Use Case 单元测试."""
"""密码重置 UseCase 单元测试."""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from unittest.mock import MagicMock, patch
@@ -13,477 +15,285 @@ from packages.application.auth.password_reset_use_case import (
)
from packages.domain.entities import User
# ── Test Fixtures ────────────────────────────────────────────────────────────
def _make_user(
user_id="user-1",
email="user@example.com",
username="testuser",
display_name="Test User",
password_hash="hashed_password_123",
):
"""创建一个测试用户."""
return User(
id=user_id,
email=email,
display_name=display_name,
username=username,
password_hash=password_hash,
)
@pytest.fixture
def mock_user_repo():
"""mock 用户仓储."""
repo = MagicMock()
repo.find_by_email.return_value = None
repo.find_by_password_reset_token.return_value = None
repo.save.return_value = None
return repo
return MagicMock()
@pytest.fixture
def mock_email_service():
"""mock 邮件服务."""
svc = MagicMock()
svc.send_password_reset_email.return_value = (True, None)
return svc
# ── RequestPasswordResetUseCase 测试 ────────────────────────────────────────
class TestRequestPasswordReset:
"""请求密码重置用例测试"""
def test_request_success_sends_email(self, mock_user_repo, mock_email_service):
"""成功请求时发送重置邮件."""
user = _make_user()
mock_user_repo.find_by_email.return_value = user
uc = RequestPasswordResetUseCase(
user_repository=mock_user_repo,
base_url="https://app.example.com",
email_service=mock_email_service,
)
req = RequestPasswordResetRequest(email="user@example.com")
ok, error = uc.execute(req)
assert ok is True
assert error is None
# 用户被更新了 reset_token
mock_user_repo.save.assert_called_once()
saved_user = mock_user_repo.save.call_args[0][0]
assert saved_user.password_reset_token is not None
assert len(saved_user.password_reset_token) > 0
assert saved_user.password_reset_expires_at is not None
# 邮件发送了
mock_email_service.send_password_reset_email.assert_called_once()
call_kwargs = mock_email_service.send_password_reset_email.call_args.kwargs
assert call_kwargs["to_email"] == "user@example.com"
assert "reset-password?token=" in call_kwargs["reset_url"]
assert "https://app.example.com" in call_kwargs["reset_url"]
def test_request_nonexistent_user_returns_success(self, mock_user_repo, mock_email_service):
"""用户不存在时也返回成功(不暴露用户存在性)."""
mock_user_repo.find_by_email.return_value = None
uc = RequestPasswordResetUseCase(
user_repository=mock_user_repo,
base_url="https://app.example.com",
email_service=mock_email_service,
)
req = RequestPasswordResetRequest(email="nonexistent@example.com")
ok, error = uc.execute(req)
assert ok is True
assert error is None
# 不保存任何东西
mock_user_repo.save.assert_not_called()
# 不发邮件
mock_email_service.send_password_reset_email.assert_not_called()
def test_request_empty_email(self, mock_user_repo, mock_email_service):
"""空邮箱返回错误."""
uc = RequestPasswordResetUseCase(
user_repository=mock_user_repo,
base_url="https://app.example.com",
email_service=mock_email_service,
)
req = RequestPasswordResetRequest(email="")
ok, error = uc.execute(req)
assert ok is False
assert "Email is required" in error
def test_request_email_normalized(self, mock_user_repo, mock_email_service):
"""邮箱会被规范化(小写+去空格)."""
user = _make_user()
mock_user_repo.find_by_email.return_value = user
uc = RequestPasswordResetUseCase(
user_repository=mock_user_repo,
base_url="https://app.example.com",
email_service=mock_email_service,
)
req = RequestPasswordResetRequest(email=" USER@Example.COM ")
ok, _ = uc.execute(req)
assert ok is True
# find_by_email 收到的是小写的
mock_user_repo.find_by_email.assert_called_with("user@example.com")
def test_request_token_expiry_custom_hours(self, mock_user_repo, mock_email_service):
"""自定义令牌过期时间."""
user = _make_user()
mock_user_repo.find_by_email.return_value = user
uc = RequestPasswordResetUseCase(
user_repository=mock_user_repo,
base_url="https://app.example.com",
token_expire_hours=6,
email_service=mock_email_service,
)
req = RequestPasswordResetRequest(email="user@example.com")
before = datetime.now(timezone.utc)
ok, _ = uc.execute(req)
after = datetime.now(timezone.utc)
assert ok is True
saved_user = mock_user_repo.save.call_args[0][0]
expires_at = saved_user.password_reset_expires_at
# 过期时间应该在 ~6 小时后
expected_min = before + timedelta(hours=6)
expected_max = after + timedelta(hours=6)
assert expected_min <= expires_at <= expected_max
def test_request_email_failure_returns_success(self, mock_user_repo, mock_email_service):
"""邮件发送失败不影响返回结果(安全考虑)."""
user = _make_user()
mock_user_repo.find_by_email.return_value = user
mock_email_service.send_password_reset_email.return_value = (False, "SMTP error")
uc = RequestPasswordResetUseCase(
user_repository=mock_user_repo,
base_url="https://app.example.com",
email_service=mock_email_service,
)
req = RequestPasswordResetRequest(email="user@example.com")
ok, error = uc.execute(req)
assert ok is True
assert error is None
# token 仍然保存了
mock_user_repo.save.assert_called_once()
def test_request_email_exception_does_not_propagate(self, mock_user_repo, mock_email_service):
"""邮件服务异常不向外传播."""
user = _make_user()
mock_user_repo.find_by_email.return_value = user
mock_email_service.send_password_reset_email.side_effect = Exception("SMTP down")
uc = RequestPasswordResetUseCase(
user_repository=mock_user_repo,
base_url="https://app.example.com",
email_service=mock_email_service,
)
req = RequestPasswordResetRequest(email="user@example.com")
ok, error = uc.execute(req)
assert ok is True
assert error is None
def test_request_username_uses_display_name_fallback(self, mock_user_repo, mock_email_service):
"""用户名为空时用 display_name."""
user = _make_user(username="", display_name="Display Name")
mock_user_repo.find_by_email.return_value = user
uc = RequestPasswordResetUseCase(
user_repository=mock_user_repo,
base_url="https://app.example.com",
email_service=mock_email_service,
)
req = RequestPasswordResetRequest(email="user@example.com")
uc.execute(req)
call_kwargs = mock_email_service.send_password_reset_email.call_args.kwargs
assert call_kwargs["username"] == "Display Name"
def test_request_uses_username_when_available(self, mock_user_repo, mock_email_service):
"""有用户名时用用户名."""
user = _make_user(username="myusername", display_name="Display Name")
mock_user_repo.find_by_email.return_value = user
uc = RequestPasswordResetUseCase(
user_repository=mock_user_repo,
base_url="https://app.example.com",
email_service=mock_email_service,
)
req = RequestPasswordResetRequest(email="user@example.com")
uc.execute(req)
call_kwargs = mock_email_service.send_password_reset_email.call_args.kwargs
assert call_kwargs["username"] == "myusername"
def test_request_generates_unique_tokens(self, mock_user_repo, mock_email_service):
"""每次请求生成不同的令牌."""
user = _make_user()
mock_user_repo.find_by_email.return_value = user
uc = RequestPasswordResetUseCase(
user_repository=mock_user_repo,
base_url="https://app.example.com",
email_service=mock_email_service,
)
tokens = []
for _ in range(3):
req = RequestPasswordResetRequest(email="user@example.com")
uc.execute(req)
saved_user = mock_user_repo.save.call_args[0][0]
tokens.append(saved_user.password_reset_token)
assert len(set(tokens)) == 3 # 三个不同的令牌
def test_request_general_exception_returns_error(self, mock_user_repo, mock_email_service):
"""其他异常返回错误信息."""
mock_user_repo.find_by_email.side_effect = Exception("DB connection error")
uc = RequestPasswordResetUseCase(
user_repository=mock_user_repo,
base_url="https://app.example.com",
email_service=mock_email_service,
)
req = RequestPasswordResetRequest(email="user@example.com")
ok, error = uc.execute(req)
assert ok is False
assert "failed" in error.lower()
# ── ResetPasswordUseCase 测试 ───────────────────────────────────────────────
class TestResetPassword:
"""重置密码用例测试"""
def test_reset_success(self, mock_user_repo):
"""成功重置密码."""
user = _make_user()
user.password_reset_token = "valid-token-123"
user.password_reset_expires_at = datetime.now(timezone.utc) + timedelta(hours=1)
mock_user_repo.find_by_password_reset_token.return_value = user
uc = ResetPasswordUseCase(user_repository=mock_user_repo)
# mock password_hasher 和 password_validator
with (
patch("packages.application.auth.password_reset_use_case.password_hasher") as mock_hasher,
patch("packages.application.auth.password_reset_use_case.password_validator") as mock_validator,
):
mock_validator.validate.return_value = (True, None)
mock_hasher.hash_password.return_value = "new_hashed_password"
req = ResetPasswordRequest(token="valid-token-123", new_password="NewPass123!")
ok, error = uc.execute(req)
assert ok is True
assert error is None
# 密码被更新
mock_user_repo.save.assert_called_once()
saved_user = mock_user_repo.save.call_args[0][0]
assert saved_user.password_hash == "new_hashed_password"
# 令牌被清除
assert saved_user.password_reset_token is None
assert saved_user.password_reset_expires_at is None
def test_reset_empty_token(self, mock_user_repo):
"""空令牌返回错误."""
uc = ResetPasswordUseCase(user_repository=mock_user_repo)
req = ResetPasswordRequest(token="", new_password="NewPass123!")
ok, error = uc.execute(req)
assert ok is False
assert "token is required" in error.lower()
def test_reset_empty_password(self, mock_user_repo):
"""空密码返回错误."""
uc = ResetPasswordUseCase(user_repository=mock_user_repo)
req = ResetPasswordRequest(token="valid-token", new_password="")
ok, error = uc.execute(req)
assert ok is False
assert "password is required" in error.lower()
def test_reset_invalid_token(self, mock_user_repo):
"""无效令牌返回错误."""
mock_user_repo.find_by_password_reset_token.return_value = None
uc = ResetPasswordUseCase(user_repository=mock_user_repo)
with patch("packages.application.auth.password_reset_use_case.password_validator") as mock_validator:
mock_validator.validate.return_value = (True, None)
req = ResetPasswordRequest(token="invalid-token", new_password="NewPass123!")
ok, error = uc.execute(req)
assert ok is False
assert "Invalid or expired" in error
def test_reset_expired_token(self, mock_user_repo):
"""过期令牌返回错误."""
user = _make_user()
user.password_reset_token = "expired-token"
user.password_reset_expires_at = datetime.now(timezone.utc) - timedelta(hours=1)
mock_user_repo.find_by_password_reset_token.return_value = user
uc = ResetPasswordUseCase(user_repository=mock_user_repo)
with patch("packages.application.auth.password_reset_use_case.password_validator") as mock_validator:
mock_validator.validate.return_value = (True, None)
req = ResetPasswordRequest(token="expired-token", new_password="NewPass123!")
ok, error = uc.execute(req)
assert ok is False
assert "expired" in error.lower()
def test_reset_naive_datetime_treated_as_utc(self, mock_user_repo):
"""不带时区的过期时间被当作 UTC 处理."""
user = _make_user()
user.password_reset_token = "token-123"
# 用 naive datetime(无时区),应该被当作 UTC
user.password_reset_expires_at = datetime.utcnow() - timedelta(hours=1) # type: ignore
mock_user_repo.find_by_password_reset_token.return_value = user
uc = ResetPasswordUseCase(user_repository=mock_user_repo)
with patch("packages.application.auth.password_reset_use_case.password_validator") as mock_validator:
mock_validator.validate.return_value = (True, None)
req = ResetPasswordRequest(token="token-123", new_password="NewPass123!")
ok, error = uc.execute(req)
assert ok is False
assert "expired" in error.lower()
def test_reset_weak_password_fails(self, mock_user_repo):
"""弱密码被拒绝."""
user = _make_user()
user.password_reset_token = "valid-token"
user.password_reset_expires_at = datetime.now(timezone.utc) + timedelta(hours=1)
mock_user_repo.find_by_password_reset_token.return_value = user
uc = ResetPasswordUseCase(user_repository=mock_user_repo)
with patch("packages.application.auth.password_reset_use_case.password_validator") as mock_validator:
mock_validator.validate.return_value = (False, "Password too short")
req = ResetPasswordRequest(token="valid-token", new_password="123")
ok, error = uc.execute(req)
assert ok is False
assert "too short" in error.lower()
# 密码没被更新
mock_user_repo.save.assert_not_called()
def test_reset_no_expires_at_still_works(self, mock_user_repo):
"""没有过期时间时视为不过期."""
user = _make_user()
user.password_reset_token = "valid-token"
user.password_reset_expires_at = None
mock_user_repo.find_by_password_reset_token.return_value = user
uc = ResetPasswordUseCase(user_repository=mock_user_repo)
with (
patch("packages.application.auth.password_reset_use_case.password_hasher") as mock_hasher,
patch("packages.application.auth.password_reset_use_case.password_validator") as mock_validator,
):
mock_validator.validate.return_value = (True, None)
mock_hasher.hash_password.return_value = "newhash"
req = ResetPasswordRequest(token="valid-token", new_password="NewPass123!")
ok, error = uc.execute(req)
assert ok is True
assert error is None
def test_reset_exception_returns_error(self, mock_user_repo):
"""异常情况返回错误信息."""
mock_user_repo.find_by_password_reset_token.side_effect = Exception("DB error")
uc = ResetPasswordUseCase(user_repository=mock_user_repo)
req = ResetPasswordRequest(token="token", new_password="NewPass123!")
ok, error = uc.execute(req)
assert ok is False
assert "failed" in error.lower()
def test_reset_clears_token_on_success(self, mock_user_repo):
"""成功重置后令牌被清除,防止重复使用."""
user = _make_user()
user.password_reset_token = "valid-token"
user.password_reset_expires_at = datetime.now(timezone.utc) + timedelta(hours=1)
mock_user_repo.find_by_password_reset_token.return_value = user
uc = ResetPasswordUseCase(user_repository=mock_user_repo)
with (
patch("packages.application.auth.password_reset_use_case.password_hasher") as mock_hasher,
patch("packages.application.auth.password_reset_use_case.password_validator") as mock_validator,
):
mock_validator.validate.return_value = (True, None)
mock_hasher.hash_password.return_value = "newhash"
req = ResetPasswordRequest(token="valid-token", new_password="NewPass123!")
ok, _ = uc.execute(req)
assert ok is True
saved_user = mock_user_repo.save.call_args[0][0]
assert saved_user.password_reset_token is None
assert saved_user.password_reset_expires_at is None
def test_reset_hashes_new_password(self, mock_user_repo):
"""密码被哈希后保存."""
user = _make_user()
user.password_reset_token = "valid-token"
user.password_reset_expires_at = datetime.now(timezone.utc) + timedelta(hours=1)
mock_user_repo.find_by_password_reset_token.return_value = user
uc = ResetPasswordUseCase(user_repository=mock_user_repo)
with (
patch("packages.application.auth.password_reset_use_case.password_hasher") as mock_hasher,
patch("packages.application.auth.password_reset_use_case.password_validator") as mock_validator,
):
mock_validator.validate.return_value = (True, None)
mock_hasher.hash_password.return_value = "hashed_abcdef"
req = ResetPasswordRequest(token="valid-token", new_password="MyNewPass123!")
uc.execute(req)
mock_hasher.hash_password.assert_called_once_with("MyNewPass123!")
saved_user = mock_user_repo.save.call_args[0][0]
assert saved_user.password_hash == "hashed_abcdef"
# ── RequestPasswordResetRequest 测试 ────────────────────────────────────────
@pytest.fixture
def sample_user():
user = User(
id="user_001",
email="user@example.com",
display_name="测试用户",
username="testuser",
password_hash="old_hash",
)
user.password_reset_token = None
user.password_reset_expires_at = None
return user
class TestRequestPasswordResetRequest:
"""请求数据类测试"""
"""RequestPasswordResetRequest 测试"""
def test_email_stripped_and_lowercased(self):
req = RequestPasswordResetRequest(email=" USER@Example.COM ")
def test_email_lowercased_and_stripped(self):
"""邮箱转小写并去空格"""
req = RequestPasswordResetRequest(" User@Example.COM ")
assert req.email == "user@example.com"
def test_email_already_lowercase(self):
req = RequestPasswordResetRequest(email="user@example.com")
assert req.email == "user@example.com"
def test_empty_email(self):
"""空邮箱"""
req = RequestPasswordResetRequest("")
assert req.email == ""
# ── ResetPasswordRequest 测试 ───────────────────────────────────────────────
class TestRequestPasswordResetUseCase:
"""RequestPasswordResetUseCase 测试"""
def test_request_success(self, mock_user_repo, mock_email_service, sample_user):
"""请求重置成功"""
mock_user_repo.find_by_email.return_value = sample_user
mock_user_repo.save.return_value = sample_user
use_case = RequestPasswordResetUseCase(
mock_user_repo,
base_url="https://example.com",
email_service=mock_email_service,
)
request = RequestPasswordResetRequest("user@example.com")
success, error = use_case.execute(request)
assert success is True
assert error is None
assert sample_user.password_reset_token is not None
assert len(sample_user.password_reset_token) > 0
assert sample_user.password_reset_expires_at is not None
mock_user_repo.save.assert_called_once()
mock_email_service.send_password_reset_email.assert_called_once()
def test_request_user_not_found_returns_success(self, mock_user_repo, mock_email_service):
"""用户不存在也返回成功(安全考虑,不暴露用户存在性)"""
mock_user_repo.find_by_email.return_value = None
use_case = RequestPasswordResetUseCase(
mock_user_repo,
base_url="https://example.com",
email_service=mock_email_service,
)
request = RequestPasswordResetRequest("nonexistent@example.com")
success, error = use_case.execute(request)
assert success is True
assert error is None
mock_user_repo.save.assert_not_called()
mock_email_service.send_password_reset_email.assert_not_called()
def test_request_empty_email_returns_error(self, mock_user_repo, mock_email_service):
"""空邮箱返回错误"""
use_case = RequestPasswordResetUseCase(
mock_user_repo,
base_url="https://example.com",
email_service=mock_email_service,
)
request = RequestPasswordResetRequest("")
success, error = use_case.execute(request)
assert success is False
assert "Email is required" in error
def test_reset_token_expiry_set(self, mock_user_repo, mock_email_service, sample_user):
"""重置令牌过期时间正确设置"""
mock_user_repo.find_by_email.return_value = sample_user
mock_user_repo.save.return_value = sample_user
use_case = RequestPasswordResetUseCase(
mock_user_repo,
base_url="https://example.com",
token_expire_hours=2,
email_service=mock_email_service,
)
request = RequestPasswordResetRequest("user@example.com")
use_case.execute(request)
assert sample_user.password_reset_expires_at is not None
# 过期时间应该在约2小时后
expected = datetime.now(timezone.utc) + timedelta(hours=2)
diff = abs((sample_user.password_reset_expires_at - expected).total_seconds())
assert diff < 10 # 允许10秒误差
def test_email_contains_reset_url(self, mock_user_repo, mock_email_service, sample_user):
"""重置邮件包含正确的重置链接"""
mock_user_repo.find_by_email.return_value = sample_user
mock_user_repo.save.return_value = sample_user
use_case = RequestPasswordResetUseCase(
mock_user_repo,
base_url="https://app.example.com",
email_service=mock_email_service,
)
request = RequestPasswordResetRequest("user@example.com")
use_case.execute(request)
call_args = mock_email_service.send_password_reset_email.call_args
reset_url = call_args[1]["reset_url"] if "reset_url" in call_args[1] else call_args[0][2]
assert "https://app.example.com/reset-password?token=" in reset_url
def test_email_failure_does_not_affect_result(self, mock_user_repo, mock_email_service, sample_user):
"""邮件发送失败不影响返回结果(安全考虑)"""
mock_user_repo.find_by_email.return_value = sample_user
mock_user_repo.save.return_value = sample_user
mock_email_service.send_password_reset_email.return_value = (False, "SMTP error")
use_case = RequestPasswordResetUseCase(
mock_user_repo,
base_url="https://example.com",
email_service=mock_email_service,
)
request = RequestPasswordResetRequest("user@example.com")
success, error = use_case.execute(request)
assert success is True
assert error is None
def test_different_tokens_each_time(self, mock_user_repo, mock_email_service, sample_user):
"""每次请求生成不同的 token"""
mock_user_repo.find_by_email.return_value = sample_user
mock_user_repo.save.return_value = sample_user
use_case = RequestPasswordResetUseCase(
mock_user_repo,
base_url="https://example.com",
email_service=mock_email_service,
)
request = RequestPasswordResetRequest("user@example.com")
use_case.execute(request)
token1 = sample_user.password_reset_token
use_case.execute(request)
token2 = sample_user.password_reset_token
assert token1 != token2
class TestResetPasswordRequest:
"""重置密码请求数据类测试"""
"""ResetPasswordRequest 测试"""
def test_stores_token_and_password(self):
req = ResetPasswordRequest(token="token123", new_password="password123")
assert req.token == "token123"
assert req.new_password == "password123"
"""正确存储 token 和新密码"""
req = ResetPasswordRequest(token="abc123", new_password="NewPass1!")
assert req.token == "abc123"
assert req.new_password == "NewPass1!"
class TestResetPasswordUseCase:
"""ResetPasswordUseCase 测试"""
def test_reset_success(self, mock_user_repo, sample_user):
"""重置密码成功"""
sample_user.password_reset_token = "valid_token"
sample_user.password_reset_expires_at = datetime.now(timezone.utc) + timedelta(hours=1)
mock_user_repo.find_by_password_reset_token.return_value = sample_user
mock_user_repo.save.return_value = sample_user
use_case = ResetPasswordUseCase(mock_user_repo)
request = ResetPasswordRequest(token="valid_token", new_password="NewSecurePass1!")
success, error = use_case.execute(request)
assert success is True
assert error is None
assert sample_user.password_reset_token is None
assert sample_user.password_reset_expires_at is None
assert sample_user.password_hash != "old_hash"
mock_user_repo.save.assert_called_once()
def test_reset_empty_token(self, mock_user_repo):
"""空 token 返回错误"""
use_case = ResetPasswordUseCase(mock_user_repo)
request = ResetPasswordRequest(token="", new_password="NewPass1!")
success, error = use_case.execute(request)
assert success is False
assert "Reset token is required" in error
mock_user_repo.save.assert_not_called()
def test_reset_empty_password(self, mock_user_repo):
"""空密码返回错误"""
use_case = ResetPasswordUseCase(mock_user_repo)
request = ResetPasswordRequest(token="sometoken", new_password="")
success, error = use_case.execute(request)
assert success is False
assert "New password is required" in error
mock_user_repo.save.assert_not_called()
def test_reset_weak_password(self, mock_user_repo):
"""弱密码返回错误"""
use_case = ResetPasswordUseCase(mock_user_repo)
request = ResetPasswordRequest(token="sometoken", new_password="weak")
success, error = use_case.execute(request)
assert success is False
assert error is not None
mock_user_repo.save.assert_not_called()
def test_reset_invalid_token(self, mock_user_repo):
"""无效 token 返回错误"""
mock_user_repo.find_by_password_reset_token.return_value = None
use_case = ResetPasswordUseCase(mock_user_repo)
request = ResetPasswordRequest(token="invalid_token", new_password="NewPass1!")
success, error = use_case.execute(request)
assert success is False
assert "Invalid or expired" in error
mock_user_repo.save.assert_not_called()
def test_reset_expired_token(self, mock_user_repo, sample_user):
"""过期 token 返回错误"""
sample_user.password_reset_token = "expired_token"
sample_user.password_reset_expires_at = datetime.now(timezone.utc) - timedelta(hours=1)
mock_user_repo.find_by_password_reset_token.return_value = sample_user
use_case = ResetPasswordUseCase(mock_user_repo)
request = ResetPasswordRequest(token="expired_token", new_password="NewPass1!")
success, error = use_case.execute(request)
assert success is False
assert "expired" in error.lower()
mock_user_repo.save.assert_not_called()
def test_reset_naive_datetime_treated_as_utc(self, mock_user_repo, sample_user):
"""无时区的过期时间按 UTC 处理"""
sample_user.password_reset_token = "naive_token"
# 用无时区的时间,设置为过去
sample_user.password_reset_expires_at = datetime.utcnow() - timedelta(hours=1)
mock_user_repo.find_by_password_reset_token.return_value = sample_user
use_case = ResetPasswordUseCase(mock_user_repo)
request = ResetPasswordRequest(token="naive_token", new_password="NewPass1!")
success, error = use_case.execute(request)
assert success is False
assert "expired" in error.lower()
def test_reset_no_expiry_set(self, mock_user_repo, sample_user):
"""没有设置过期时间的 token 可以使用"""
sample_user.password_reset_token = "no_expiry_token"
sample_user.password_reset_expires_at = None
mock_user_repo.find_by_password_reset_token.return_value = sample_user
use_case = ResetPasswordUseCase(mock_user_repo)
request = ResetPasswordRequest(token="no_expiry_token", new_password="NewPass1!")
success, error = use_case.execute(request)
assert success is True
File diff suppressed because it is too large Load Diff
+68 -263
View File
@@ -1,295 +1,100 @@
"""TTS 文本分段工具单元测试."""
import pytest
"""text_splitter 单元测试."""
from packages.application.tts_job.text_splitter import split_text
class TestSplitTextEmpty:
"""空文本测试"""
def test_empty_string(self):
"""空字符串返回空列表."""
class TestSplitText:
def test_empty_text_returns_empty(self):
assert split_text("") == []
def test_only_whitespace(self):
"""纯空白文本返回空列表."""
assert split_text(" \n \t ") == []
def test_whitespace_only(self):
assert split_text(" \n\t ") == []
def test_none_not_allowed(self):
"""None 会抛出异常(不是我们的职责)."""
with pytest.raises(AttributeError):
split_text(None) # type: ignore
class TestSplitTextShort:
"""短文本测试"""
def test_short_text_one_segment(self):
"""短文本返回一个段落."""
text = "你好,世界。"
def test_short_text_single_segment(self):
text = "你好世界。"
result = split_text(text, max_chars=500)
assert len(result) == 1
assert result[0] == text
def test_exactly_max_chars(self):
"""刚好等于 max_chars 的文本返回一个段落."""
def test_exact_max_chars(self):
text = "a" * 500
result = split_text(text, max_chars=500)
assert len(result) == 1
assert len(result[0]) == 500
def test_one_under_max(self):
"""max_chars-1 的文本返回一个段落."""
text = "a" * 499
def test_splits_on_sentence_boundary(self):
# 两个长句子,各300字左右,超过50字阈值
sent1 = "" * 300 + ""
sent2 = "" * 300 + ""
text = sent1 + sent2
result = split_text(text, max_chars=500)
assert len(result) == 1
assert len(result[0]) == 499
assert len(result) == 2
assert result[0] == sent1
assert result[1] == sent2
class TestSplitTextSentenceBoundary:
"""句子边界分段测试"""
def test_split_at_period(self):
"""在句号处拆分."""
text = "第一句。第二句。第三句。"
# 每句5字符,max_chars=10,每次两句就接近10
result = split_text(text, max_chars=10)
def test_long_sentence_hard_cut(self):
# 一个超长句子,没有句末标点,会被硬切
text = "" * 800
result = split_text(text, max_chars=500)
assert len(result) >= 2
# 所有段落都不超过 max_chars
for seg in result:
assert len(seg) <= 10
def test_split_at_exclamation(self):
"""在感叹号处拆分."""
text = "好棒!真的好棒!太厉害了!"
result = split_text(text, max_chars=10)
assert len(result) >= 2
for seg in result:
assert len(seg) <= 10
def test_split_at_question(self):
"""在问号处拆分."""
text = "你好吗?你是谁?你在哪?"
result = split_text(text, max_chars=10)
assert len(result) >= 2
for seg in result:
assert len(seg) <= 10
def test_split_at_newline(self):
"""在换行处拆分."""
text = "第一段\n第二段\n第三段"
result = split_text(text, max_chars=10)
assert len(result) >= 2
for seg in result:
assert len(seg) <= 10
def test_split_at_semicolon(self):
"""在分号处拆分."""
text = "第一项;第二项;第三项;"
result = split_text(text, max_chars=10)
assert len(result) >= 2
def test_english_punctuation(self):
"""英文标点也能拆分."""
text = "Hello world. How are you? I am fine!"
result = split_text(text, max_chars=20)
assert len(result) >= 2
for seg in result:
assert len(seg) <= 20
def test_mixed_punctuation(self):
"""中英文标点混合."""
text = "你好!Hello. 你好吗?How are you?"
result = split_text(text, max_chars=15)
assert len(result) >= 2
class TestSplitTextForceSplit:
"""强制分段测试"""
def test_very_long_sentence_forced_split(self):
"""超长单句强制分段."""
text = "a" * 1000 # 没有标点的长文本
result = split_text(text, max_chars=100)
assert len(result) == 10
for seg in result:
assert len(seg) == 100
def test_mixed_long_and_short_sentences(self):
"""长短句混合."""
long = "" * 200
text = f"短句。{long}。短句。"
result = split_text(text, max_chars=100)
# 所有段都不超过100
for seg in result:
assert len(seg) <= 100
# 至少有3段(长句被强制拆分)
assert len(result) >= 3
class TestSplitTextMerging:
"""短段落合并测试"""
assert all(len(seg) <= 500 for seg in result)
# 合起来应该等于原文本
assert "".join(result) == text
def test_short_segments_merged(self):
"""多个短段合并为一个."""
# 生成5个短句,每句5字符,max_chars=100,应该合并成一段
text = "一。二。三。四。五。"
result = split_text(text, max_chars=100)
assert len(result) == 1
assert len(result[0]) <= 100
def test_merge_within_limit(self):
"""合并后不超过 max_chars."""
# 10个短句,每句4字符 = 40字符
text = "句子。" * 10
result = split_text(text, max_chars=100)
assert len(result) == 1
assert len(result[0]) <= 100
def test_merge_across_multiple(self):
"""多个短段依次合并."""
text = "短。" * 30 # 30个短句,每句2字符=60字符
result = split_text(text, max_chars=100)
assert len(result) == 1
assert len(result[0]) == 60 # 全部合并
class TestSplitTextChinese:
"""中文文本测试"""
def test_chinese_paragraph(self):
"""典型中文段落."""
text = (
"在一个阳光明媚的早晨,小明来到了公园。"
"他看到了很多人在锻炼身体。"
"有的人在跑步,有的人在打太极,还有的人在跳舞。"
"小明也加入了他们,开始了愉快的一天。"
)
result = split_text(text, max_chars=50)
assert len(result) >= 2
for seg in result:
assert len(seg) <= 50
# 重新拼回应该等于原文本(除了可能的空格处理)
combined = "".join(result)
assert combined == text.replace(" ", "") # strip 不影响中文字符
def test_chinese_long_paragraph(self):
"""长中文段落."""
text = "这是一个测试句子。" * 100 # 100个句子
# 多个短句应该被合并
sentences = [f"{i}句。" for i in range(10)]
text = "".join(sentences)
result = split_text(text, max_chars=200)
assert len(result) > 1
for seg in result:
assert len(seg) <= 200
# 总字符数不变
assert sum(len(s) for s in result) == len(text)
# 每句5字左右,10句才50字,应该合并成1段
assert len(result) < 10
assert len(result[0]) <= 200
def test_preserves_content(self):
text = "今天天气真好。我们去公园玩吧!你觉得怎么样?好的,走吧。"
result = split_text(text, max_chars=20)
# 合并后内容应一致
assert "".join(result) == text
class TestSplitTextCustomMaxChars:
"""自定义 max_chars 测试"""
def test_multiple_punctuation_types(self):
# 构造足够长的文本触发分段
text = "第一" * 30 + "" + "第二" * 30 + "" + "第三" * 30 + "" + "第四" * 30 + ""
result = split_text(text, max_chars=100)
assert len(result) >= 2
assert "".join(result) == text
def test_small_max_chars(self):
"""很小的 max_chars."""
text = "一二三四五六七八九十。"
result = split_text(text, max_chars=5)
for seg in result:
assert len(seg) <= 5
def test_custom_max_chars(self):
text = "a" * 100 + "" + "b" * 100 + ""
result = split_text(text, max_chars=150)
assert len(result) == 2
assert "a" in result[0]
assert "b" in result[1]
def test_large_max_chars(self):
"""很大的 max_chars(不拆分)."""
text = "这是一段测试文本。" * 10
result = split_text(text, max_chars=10000)
def test_newline_as_sentence_end(self):
text = "第一段\n第二段\n第三段"
result = split_text(text, max_chars=50)
assert len(result) >= 1
assert "".join(result) == text.strip()
def test_minimum_segment_length(self):
# 句子太短(<50字)不会立即分段
text = "短句一。短句二。短句三。"
result = split_text(text, max_chars=200)
assert len(result) == 1
def test_max_chars_zero(self):
"""max_chars=0 时的行为."""
text = "测试文本。"
# 0 会导致每加一个字符就触发强制分段
result = split_text(text, max_chars=0)
# 每个字符一段?或者至少有结果
assert isinstance(result, list)
assert len(result) > 0
def test_max_chars_one(self):
"""max_chars=1."""
text = "abc"
result = split_text(text, max_chars=1)
assert len(result) == 3
assert result == ["a", "b", "c"]
class TestSplitTextPreservesContent:
"""内容完整性测试"""
def test_preserves_all_chars(self):
"""分段后拼接等于原文(忽略空白调整)."""
text = "第一句。第二句!第三句?第四句。第五句。"
result = split_text(text, max_chars=10)
combined = "".join(result)
assert combined == text
def test_trailing_content_added(self):
# 最后一段不完整的句子也要加上
text = "完整的句子。剩余内容"
result = split_text(text, max_chars=50)
assert "".join(result) == text
def test_no_empty_segments(self):
"""没有空字符串段落."""
text = "句子。。。双标点。"
result = split_text(text, max_chars=10)
assert all(seg for seg in result) # 所有段非空
text = "。。。。。" # 全是标点
result = split_text(text, max_chars=2)
assert all(len(seg) > 0 for seg in result)
def test_stripped_segments(self):
"""段落首尾没有多余空白."""
text = " 第一句。 第二句。 "
result = split_text(text, max_chars=10)
for seg in result:
assert seg == seg.strip()
class TestSplitTextEdgeCases:
"""边界情况测试"""
def test_single_char(self):
"""单字符."""
result = split_text("", max_chars=10)
assert len(result) == 1
assert result[0] == ""
def test_only_punctuation(self):
"""纯标点."""
text = "。。。!!??"
result = split_text(text, max_chars=5)
assert len(result) >= 1
assert sum(len(s) for s in result) == len(text)
def test_numbers_and_symbols(self):
"""数字和符号."""
text = "第1章。第2节。第3段。"
result = split_text(text, max_chars=10)
assert len(result) >= 1
assert all(len(s) <= 10 for s in result)
def test_mixed_chinese_english(self):
"""中英文混合."""
text = "Hello你好World世界。Test测试。"
result = split_text(text, max_chars=10)
assert len(result) >= 2
assert all(len(s) <= 10 for s in result)
def test_consecutive_punctuation(self):
"""连续标点."""
text = "真的吗!?不对。。。好吧。"
def test_chinese_and_english_mixed(self):
text = "Hello世界。这是测试Test文本。Mixed混合。"
result = split_text(text, max_chars=20)
assert len(result) >= 1
combined = "".join(result)
assert combined == text
class TestSplitTextDefaultParams:
"""默认参数测试"""
def test_default_max_chars_is_500(self):
"""默认 max_chars=500."""
text = "a" * 500
result = split_text(text)
assert len(result) == 1
text2 = "a" * 501
result2 = split_text(text2)
assert len(result2) >= 2
assert len(result) >= 2
assert "".join(result) == text
+275 -338
View File
@@ -1,4 +1,4 @@
"""标题库 Use Cases 单元测试 — wave217"""
"""标题库 UseCase 单元测试."""
from __future__ import annotations
@@ -24,447 +24,384 @@ from packages.application.title_library.use_cases import (
from packages.domain.exceptions import NotFoundError, QuotaExceededError
from packages.domain.title_library import TitleLibraryItem
# ── helpers ──────────────────────────────────────────────────────────────────
def _make_item(
item_id="t1",
user_id="u1",
name="标题A",
text="这是一个标题",
category="default",
usage_count=0,
is_active=True,
description="",
tags=None,
metadata_=None,
):
def _make_item(id: str, name: str, text: str, usage_count: int = 0, category: str = "default") -> TitleLibraryItem:
return TitleLibraryItem(
id=item_id,
user_id=user_id,
id=id,
user_id="user_1",
name=name,
text=text,
category=category,
description=description,
tags=tags or [],
description="",
tags=[],
usage_count=usage_count,
is_active=is_active,
metadata_=metadata_ or {},
is_active=True,
metadata_={},
)
# ── ListTitleLibraryUseCase ─────────────────────────────────────────────────
@pytest.fixture
def mock_repo():
return MagicMock()
@pytest.fixture
def sample_item():
return _make_item("title_1", "爆款标题", "这是一个爆款标题文案", usage_count=5)
class TestListTitleLibraryUseCase:
def test_list_default_params(self):
items = [_make_item("t1"), _make_item("t2")]
repo = MagicMock()
repo.list_by_user.return_value = items
"""ListTitleLibraryUseCase 测试"""
uc = ListTitleLibraryUseCase(repo)
result = uc.execute("u1")
def test_list_returns_results(self, mock_repo, sample_item):
"""正常返回标题列表"""
mock_repo.list_by_user.return_value = [sample_item]
use_case = ListTitleLibraryUseCase(mock_repo)
assert len(result) == 2
repo.list_by_user.assert_called_once_with("u1", category=None, skip=0, limit=50)
result = use_case.execute("user_1")
def test_list_with_category(self):
repo = MagicMock()
repo.list_by_user.return_value = []
assert len(result) == 1
assert result[0].id == "title_1"
mock_repo.list_by_user.assert_called_once_with("user_1", category=None, skip=0, limit=50)
uc = ListTitleLibraryUseCase(repo)
result = uc.execute("u1", category="marketing")
def test_list_with_category(self, mock_repo, sample_item):
"""按分类过滤"""
mock_repo.list_by_user.return_value = [sample_item]
use_case = ListTitleLibraryUseCase(mock_repo)
use_case.execute("user_1", category="电商")
mock_repo.list_by_user.assert_called_once_with("user_1", category="电商", skip=0, limit=50)
def test_list_with_pagination(self, mock_repo, sample_item):
"""带分页参数"""
mock_repo.list_by_user.return_value = [sample_item]
use_case = ListTitleLibraryUseCase(mock_repo)
use_case.execute("user_1", skip=10, limit=20)
mock_repo.list_by_user.assert_called_once_with("user_1", category=None, skip=10, limit=20)
def test_empty_list(self, mock_repo):
"""空列表"""
mock_repo.list_by_user.return_value = []
use_case = ListTitleLibraryUseCase(mock_repo)
result = use_case.execute("user_1")
repo.list_by_user.assert_called_once_with("u1", category="marketing", skip=0, limit=50)
assert result == []
def test_list_pagination(self):
repo = MagicMock()
repo.list_by_user.return_value = []
uc = ListTitleLibraryUseCase(repo)
uc.execute("u1", skip=10, limit=20)
repo.list_by_user.assert_called_once_with("u1", category=None, skip=10, limit=20)
# ── GetTitleLibraryUseCase ──────────────────────────────────────────────────
class TestGetTitleLibraryUseCase:
def test_get_found(self):
item = _make_item()
repo = MagicMock()
repo.get.return_value = item
"""GetTitleLibraryUseCase 测试"""
uc = GetTitleLibraryUseCase(repo)
result = uc.execute("t1", "u1")
assert result.id == "t1"
repo.get.assert_called_once_with("t1", "u1")
def test_get_existing(self, mock_repo, sample_item):
"""获取存在的标题"""
mock_repo.get.return_value = sample_item
use_case = GetTitleLibraryUseCase(mock_repo)
def test_get_not_found(self):
repo = MagicMock()
repo.get.return_value = None
result = use_case.execute("title_1", "user_1")
assert result is not None
assert result.id == "title_1"
mock_repo.get.assert_called_once_with("title_1", "user_1")
def test_get_nonexistent_returns_none(self, mock_repo):
"""获取不存在的标题返回 None"""
mock_repo.get.return_value = None
use_case = GetTitleLibraryUseCase(mock_repo)
result = use_case.execute("nonexistent", "user_1")
uc = GetTitleLibraryUseCase(repo)
result = uc.execute("t999", "u1")
assert result is None
# ── CreateTitleLibraryUseCase ───────────────────────────────────────────────
class TestCreateTitleLibraryUseCase:
def test_create_success_free_plan_within_quota(self):
repo = MagicMock()
repo.count_by_user.return_value = 0 # 已用数量
repo.create.side_effect = lambda x: x
"""CreateTitleLibraryUseCase 测试"""
uc = CreateTitleLibraryUseCase(repo)
cmd = CreateTitleLibraryCommand(user_id="u1", name="好标题", text="这是一个好标题的内容")
def test_create_success(self, mock_repo, sample_item):
"""创建成功"""
mock_repo.count_by_user.return_value = 0
mock_repo.create.return_value = sample_item
use_case = CreateTitleLibraryUseCase(mock_repo)
with patch("packages.application.title_library.use_cases.quota_checker") as mock_qc:
mock_result = MagicMock()
mock_result.allowed = True
mock_result.limit = 10
mock_result.used = 0
mock_qc.check.return_value = mock_result
result = uc.execute(cmd, plan_name="free")
assert result.name == "好标题"
assert result.user_id == "u1"
repo.create.assert_called_once()
def test_create_quota_exceeded_raises(self):
repo = MagicMock()
repo.count_by_user.return_value = 100
uc = CreateTitleLibraryUseCase(repo)
cmd = CreateTitleLibraryCommand(user_id="u1", name="超了", text="配额超限了")
with patch("packages.application.title_library.use_cases.quota_checker") as mock_qc:
mock_result = MagicMock()
mock_result.allowed = False
mock_result.limit = 5
mock_result.used = 5
mock_qc.check.return_value = mock_result
with pytest.raises(QuotaExceededError):
uc.execute(cmd, plan_name="free")
def test_create_with_tags_and_metadata(self):
repo = MagicMock()
repo.count_by_user.return_value = 0
repo.create.side_effect = lambda x: x
uc = CreateTitleLibraryUseCase(repo)
cmd = CreateTitleLibraryCommand(
user_id="u1",
name="标题",
text="内容",
category="marketing",
description="描述",
tags=["tag1", "tag2"],
metadata_={"key": "value"},
command = CreateTitleLibraryCommand(
user_id="user_1",
name="新标题",
text="新标题文案",
category="default",
description="",
tags=[],
metadata_={},
)
result = use_case.execute(command, plan_name="free")
with patch("packages.application.title_library.use_cases.quota_checker") as mock_qc:
mock_result = MagicMock()
mock_result.allowed = True
mock_result.limit = 100
mock_result.used = 0
mock_qc.check.return_value = mock_result
assert result.id == "title_1"
mock_repo.count_by_user.assert_called_once_with("user_1")
mock_repo.create.assert_called_once()
result = uc.execute(cmd)
def test_create_quota_exceeded(self, mock_repo):
"""超过配额时抛出 QuotaExceededError"""
mock_repo.count_by_user.return_value = 9999
use_case = CreateTitleLibraryUseCase(mock_repo)
assert result.category == "marketing"
assert result.description == "描述"
assert result.tags == ["tag1", "tag2"]
assert result.metadata_ == {"key": "value"}
command = CreateTitleLibraryCommand(
user_id="user_1",
name="新标题",
text="文案",
category="default",
description="",
tags=[],
metadata_={},
)
with pytest.raises(QuotaExceededError):
use_case.execute(command, plan_name="free")
mock_repo.create.assert_not_called()
# ── UpdateTitleLibraryUseCase ───────────────────────────────────────────────
def test_create_with_tags_and_metadata(self, mock_repo, sample_item):
"""创建时带 tags 和 metadata_"""
mock_repo.count_by_user.return_value = 0
mock_repo.create.return_value = sample_item
use_case = CreateTitleLibraryUseCase(mock_repo)
command = CreateTitleLibraryCommand(
user_id="user_1",
name="带标签标题",
text="文案",
category="电商",
description="测试描述",
tags=["爆款", "促销"],
metadata_={"source": "manual"},
)
use_case.execute(command, plan_name="premium")
created = mock_repo.create.call_args[0][0]
assert isinstance(created, TitleLibraryItem)
assert created.name == "带标签标题"
assert created.category == "电商"
assert created.tags == ["爆款", "促销"]
assert created.metadata_ == {"source": "manual"}
class TestUpdateTitleLibraryUseCase:
def test_update_name(self):
existing = _make_item(name="old")
repo = MagicMock()
repo.get.return_value = existing
repo.update.return_value = existing
"""UpdateTitleLibraryUseCase 测试"""
uc = UpdateTitleLibraryUseCase(repo)
cmd = UpdateTitleLibraryCommand(title_id="t1", user_id="u1", name="new")
result = uc.execute(cmd)
def test_update_name(self, mock_repo, sample_item):
"""更新标题名称"""
mock_repo.get.return_value = sample_item
mock_repo.update.side_effect = lambda x: x
use_case = UpdateTitleLibraryUseCase(mock_repo)
assert result.name == "new"
repo.update.assert_called_once()
command = UpdateTitleLibraryCommand(title_id="title_1", user_id="user_1", name="新名称")
result = use_case.execute(command)
def test_update_text(self):
existing = _make_item(text="old")
repo = MagicMock()
repo.get.return_value = existing
repo.update.return_value = existing
assert result.name == "新名称"
# 其他字段不变
assert result.text == "这是一个爆款标题文案"
mock_repo.get.assert_called_once_with("title_1", "user_1")
mock_repo.update.assert_called_once()
uc = UpdateTitleLibraryUseCase(repo)
cmd = UpdateTitleLibraryCommand(title_id="t1", user_id="u1", text="new text")
result = uc.execute(cmd)
assert result.text == "new text"
def test_update_multiple_fields(self, mock_repo, sample_item):
"""同时更新多个字段"""
mock_repo.get.return_value = sample_item
mock_repo.update.side_effect = lambda x: x
use_case = UpdateTitleLibraryUseCase(mock_repo)
def test_update_category(self):
existing = _make_item(category="old")
repo = MagicMock()
repo.get.return_value = existing
repo.update.return_value = existing
command = UpdateTitleLibraryCommand(
title_id="title_1",
user_id="user_1",
text="新文案内容",
category="美食",
is_active=False,
)
result = use_case.execute(command)
uc = UpdateTitleLibraryUseCase(repo)
cmd = UpdateTitleLibraryCommand(title_id="t1", user_id="u1", category="new_cat")
result = uc.execute(cmd)
assert result.category == "new_cat"
def test_update_tags(self):
existing = _make_item(tags=["old"])
repo = MagicMock()
repo.get.return_value = existing
repo.update.return_value = existing
uc = UpdateTitleLibraryUseCase(repo)
cmd = UpdateTitleLibraryCommand(title_id="t1", user_id="u1", tags=["a", "b"])
result = uc.execute(cmd)
assert result.tags == ["a", "b"]
def test_update_is_active(self):
existing = _make_item(is_active=True)
repo = MagicMock()
repo.get.return_value = existing
repo.update.return_value = existing
uc = UpdateTitleLibraryUseCase(repo)
cmd = UpdateTitleLibraryCommand(title_id="t1", user_id="u1", is_active=False)
result = uc.execute(cmd)
assert result.text == "新文案内容"
assert result.category == "美食"
assert result.is_active is False
def test_update_not_found_raises(self):
repo = MagicMock()
repo.get.return_value = None
def test_update_nonexistent_raises(self, mock_repo):
"""更新不存在的标题抛出 NotFoundError"""
mock_repo.get.return_value = None
use_case = UpdateTitleLibraryUseCase(mock_repo)
uc = UpdateTitleLibraryUseCase(repo)
cmd = UpdateTitleLibraryCommand(title_id="t999", user_id="u1", name="x")
with pytest.raises(NotFoundError):
uc.execute(cmd)
command = UpdateTitleLibraryCommand(title_id="noexist", user_id="user_1", name="新名称")
with pytest.raises(NotFoundError, match="not found"):
use_case.execute(command)
def test_update_none_fields_not_modified(self):
existing = _make_item(name="keep", category="keep_cat", description="keep_desc")
repo = MagicMock()
repo.get.return_value = existing
repo.update.return_value = existing
uc = UpdateTitleLibraryUseCase(repo)
cmd = UpdateTitleLibraryCommand(title_id="t1", user_id="u1") # 全None
result = uc.execute(cmd)
assert result.name == "keep"
assert result.category == "keep_cat"
assert result.description == "keep_desc"
# ── DeleteTitleLibraryUseCase ───────────────────────────────────────────────
mock_repo.update.assert_not_called()
class TestDeleteTitleLibraryUseCase:
def test_delete_success(self):
repo = MagicMock()
repo.delete.return_value = True
"""DeleteTitleLibraryUseCase 测试"""
def test_delete_success(self, mock_repo):
"""删除成功"""
mock_repo.delete.return_value = True
use_case = DeleteTitleLibraryUseCase(mock_repo)
result = use_case.execute("title_1", "user_1")
uc = DeleteTitleLibraryUseCase(repo)
result = uc.execute("t1", "u1")
assert result is True
repo.delete.assert_called_once_with("t1", "u1")
mock_repo.delete.assert_called_once_with("title_1", "user_1")
def test_delete_not_found(self):
repo = MagicMock()
repo.delete.return_value = False
def test_delete_nonexistent_returns_false(self, mock_repo):
"""删除不存在的返回 False"""
mock_repo.delete.return_value = False
use_case = DeleteTitleLibraryUseCase(mock_repo)
result = use_case.execute("noexist", "user_1")
uc = DeleteTitleLibraryUseCase(repo)
result = uc.execute("t999", "u1")
assert result is False
# ── IncrementTitleUsageUseCase ──────────────────────────────────────────────
class TestIncrementTitleUsageUseCase:
def test_increment_default_1(self):
repo = MagicMock()
repo.increment_usage_count.return_value = True
"""IncrementTitleUsageUseCase 测试"""
uc = IncrementTitleUsageUseCase(repo)
cmd = IncrementTitleUsageCommand(title_id="t1", user_id="u1")
result = uc.execute(cmd)
def test_increment_positive(self, mock_repo):
"""正增量时调用 repository"""
mock_repo.increment_usage_count.return_value = True
use_case = IncrementTitleUsageUseCase(mock_repo)
command = IncrementTitleUsageCommand(title_id="title_1", user_id="user_1", increment=1)
result = use_case.execute(command)
assert result is True
repo.increment_usage_count.assert_called_once_with("t1", "u1", increment=1)
mock_repo.increment_usage_count.assert_called_once_with("title_1", "user_1", increment=1)
def test_increment_custom_amount(self):
repo = MagicMock()
repo.increment_usage_count.return_value = True
def test_increment_zero_returns_false(self, mock_repo):
"""增量为0返回False,不调用repository"""
use_case = IncrementTitleUsageUseCase(mock_repo)
uc = IncrementTitleUsageUseCase(repo)
cmd = IncrementTitleUsageCommand(title_id="t1", user_id="u1", increment=5)
result = uc.execute(cmd)
repo.increment_usage_count.assert_called_once_with("t1", "u1", increment=5)
def test_increment_zero_returns_false(self):
repo = MagicMock()
uc = IncrementTitleUsageUseCase(repo)
cmd = IncrementTitleUsageCommand(title_id="t1", user_id="u1", increment=0)
result = uc.execute(cmd)
command = IncrementTitleUsageCommand(title_id="title_1", user_id="user_1", increment=0)
result = use_case.execute(command)
assert result is False
repo.increment_usage_count.assert_not_called()
mock_repo.increment_usage_count.assert_not_called()
def test_increment_negative_returns_false(self):
repo = MagicMock()
def test_increment_negative_returns_false(self, mock_repo):
"""负增量返回False"""
use_case = IncrementTitleUsageUseCase(mock_repo)
uc = IncrementTitleUsageUseCase(repo)
cmd = IncrementTitleUsageCommand(title_id="t1", user_id="u1", increment=-1)
result = uc.execute(cmd)
command = IncrementTitleUsageCommand(title_id="title_1", user_id="user_1", increment=-1)
result = use_case.execute(command)
assert result is False
repo.increment_usage_count.assert_not_called()
mock_repo.increment_usage_count.assert_not_called()
def test_increment_large_number(self, mock_repo):
"""大增量值"""
mock_repo.increment_usage_count.return_value = True
use_case = IncrementTitleUsageUseCase(mock_repo)
# ── PickTitleUseCase ────────────────────────────────────────────────────────
command = IncrementTitleUsageCommand(title_id="title_1", user_id="user_1", increment=10)
use_case.execute(command)
mock_repo.increment_usage_count.assert_called_once_with("title_1", "user_1", increment=10)
class TestPickTitleUseCase:
def test_pick_from_multiple_returns_least_used_in_pool(self):
items = [
_make_item("t1", usage_count=10),
_make_item("t2", usage_count=1), # 最少
_make_item("t3", usage_count=5),
_make_item("t4", usage_count=3),
_make_item("t5", usage_count=8),
_make_item("t6", usage_count=2),
]
repo = MagicMock()
repo.list_by_user.return_value = items
"""PickTitleUseCase 智能选标题测试"""
uc = PickTitleUseCase(repo)
cmd = PickTitleCommand(user_id="u1")
def test_pick_from_multiple(self, mock_repo):
"""从多个标题中选一个(最少使用的前5个中随机)"""
items = [_make_item(f"t{i}", f"标题{i}", f"文案{i}", usage_count=i) for i in range(10)]
mock_repo.list_by_user.return_value = items
use_case = PickTitleUseCase(mock_repo)
# 由于有随机性,多次验证都在候选池(最少使用的5个)中
for _ in range(10):
result = uc.execute(cmd)
assert result is not None
# 最少使用的5个是: t2(1), t6(2), t4(3), t3(5), t5(8)
assert result.id in {"t1", "t2", "t3", "t4", "t5", "t6"}
# 选中的一定是使用次数最少的5个之一 (usage_count <= 8)
assert result.usage_count <= 8
command = PickTitleCommand(user_id="user_1")
result = use_case.execute(command)
# 验证查询参数
repo.list_by_user.assert_called()
call_args = repo.list_by_user.call_args
assert call_args[0][0] == "u1"
assert call_args[1]["is_active"] is True
assert result is not None
assert isinstance(result, TitleLibraryItem)
# 选出的应该是使用次数最少的前5个之一(0-4)
assert result.usage_count <= 4
mock_repo.list_by_user.assert_called_once()
def test_pick_empty_returns_none(self):
repo = MagicMock()
repo.list_by_user.return_value = []
def test_pick_empty_returns_none(self, mock_repo):
"""空标题库返回 None"""
mock_repo.list_by_user.return_value = []
use_case = PickTitleUseCase(mock_repo)
command = PickTitleCommand(user_id="user_1")
result = use_case.execute(command)
uc = PickTitleUseCase(repo)
cmd = PickTitleCommand(user_id="u1")
result = uc.execute(cmd)
assert result is None
def test_pick_single_item(self):
item = _make_item("t1")
repo = MagicMock()
repo.list_by_user.return_value = [item]
def test_pick_with_category(self, mock_repo):
"""按分类选标题"""
items = [_make_item("t1", "标题1", "文案1", category="美食")]
mock_repo.list_by_user.return_value = items
use_case = PickTitleUseCase(mock_repo)
uc = PickTitleUseCase(repo)
cmd = PickTitleCommand(user_id="u1")
result = uc.execute(cmd)
assert result.id == "t1"
def test_pick_with_category_filter(self):
repo = MagicMock()
repo.list_by_user.return_value = [_make_item("t1", category="marketing")]
uc = PickTitleUseCase(repo)
cmd = PickTitleCommand(user_id="u1", category="marketing")
result = uc.execute(cmd)
command = PickTitleCommand(user_id="user_1", category="美食")
result = use_case.execute(command)
assert result is not None
repo.list_by_user.assert_called_once()
assert repo.list_by_user.call_args[1]["category"] == "marketing"
call_kwargs = mock_repo.list_by_user.call_args[1]
assert call_kwargs["category"] == "美食"
assert call_kwargs["is_active"] is True
def test_pick_exclude_ids(self):
def test_pick_exclude_ids(self, mock_repo):
"""排除指定ID"""
items = [
_make_item("t1", usage_count=1),
_make_item("t2", usage_count=2),
_make_item("t3", usage_count=3),
_make_item("t1", "标题1", "文案1", usage_count=1),
_make_item("t2", "标题2", "文案2", usage_count=2),
_make_item("t3", "标题3", "文案3", usage_count=3),
]
repo = MagicMock()
repo.list_by_user.return_value = items
mock_repo.list_by_user.return_value = items
use_case = PickTitleUseCase(mock_repo)
uc = PickTitleUseCase(repo)
cmd = PickTitleCommand(user_id="u1", exclude_ids=["t1", "t2"])
command = PickTitleCommand(user_id="user_1", exclude_ids=["t1", "t2"])
result = use_case.execute(command)
# 排除 t1, t2 后只剩 t3
result = uc.execute(cmd)
# 排除两个后只剩t3
assert result.id == "t3"
def test_pick_exclude_all_fallback_to_all(self):
def test_pick_exclude_all_falls_back(self, mock_repo):
"""排除全部时从所有标题中选"""
items = [
_make_item("t1", usage_count=1),
_make_item("t2", usage_count=2),
_make_item("t1", "标题1", "文案1", usage_count=1),
_make_item("t2", "标题2", "文案2", usage_count=2),
]
repo = MagicMock()
repo.list_by_user.return_value = items
mock_repo.list_by_user.return_value = items
use_case = PickTitleUseCase(mock_repo)
uc = PickTitleUseCase(repo)
cmd = PickTitleCommand(user_id="u1", exclude_ids=["t1", "t2"])
command = PickTitleCommand(user_id="user_1", exclude_ids=["t1", "t2"])
result = use_case.execute(command)
# 排除后没了,回退到从全部选
result = uc.execute(cmd)
# 排除全部后fallback到全部,所以还是能选出一个
assert result is not None
assert result.id in {"t1", "t2"}
assert result.id in ("t1", "t2")
def test_pick_pool_size_is_5(self):
# 10个标题,使用次数从 1~10
items = [_make_item(f"t{i}", usage_count=i) for i in range(1, 11)]
repo = MagicMock()
repo.list_by_user.return_value = items
def test_pick_single_item(self, mock_repo):
"""只有一个标题时选它"""
item = _make_item("only", "唯一标题", "唯一文案", usage_count=10)
mock_repo.list_by_user.return_value = [item]
use_case = PickTitleUseCase(mock_repo)
uc = PickTitleUseCase(repo)
cmd = PickTitleCommand(user_id="u1")
command = PickTitleCommand(user_id="user_1")
result = use_case.execute(command)
# 运行多次,确保选中的都在前5个使用最少的里(t1~t5, usage 1~5
for _ in range(20):
result = uc.execute(cmd)
assert int(result.id[1:]) <= 5 # 只从前5个里选
assert result.id == "only"
def test_pick_fewer_than_pool_size(self):
# 只有3个标题,不足5个池大小
def test_pick_prefers_less_used(self, mock_repo):
"""倾向于选择使用次数少的"""
items = [
_make_item("t1", usage_count=3),
_make_item("t2", usage_count=1),
_make_item("t3", usage_count=2),
_make_item("t_used", "常用", "常用", usage_count=100),
_make_item("t_fresh", "新的", "新的", usage_count=0),
]
repo = MagicMock()
repo.list_by_user.return_value = items
uc = PickTitleUseCase(repo)
cmd = PickTitleCommand(user_id="u1")
mock_repo.list_by_user.return_value = items
use_case = PickTitleUseCase(mock_repo)
# 跑多次,验证使用少的出现在候选池里
results = set()
for _ in range(30):
result = uc.execute(cmd)
results.add(result.id)
for _ in range(20):
command = PickTitleCommand(user_id="user_1")
r = use_case.execute(command)
if r:
results.add(r.id)
# 3个都有可能被选中(随机性+少量样本,大概率至少出现2个)
assert len(results) >= 1
assert results.issubset({"t1", "t2", "t3"})
# 两个都在候选池(少于5个),所以都可能被选中
assert "t_used" in results or "t_fresh" in results
+357 -331
View File
@@ -1,4 +1,4 @@
"""视频分享 Use Cases 单元测试 — wave215"""
"""视频分享 UseCase 单元测试."""
from __future__ import annotations
@@ -22,7 +22,6 @@ from packages.application.video_share.use_cases import (
PasswordRequiredError,
RecordShareDownloadUseCase,
RevokeShareUseCase,
ShareAccessResult,
ShareExpiredError,
UpdateShareUseCase,
VideoNotFoundError,
@@ -30,459 +29,486 @@ from packages.application.video_share.use_cases import (
from packages.domain.generated_video import GeneratedVideo
from packages.domain.video_share import VideoShare
# ── helpers ──────────────────────────────────────────────────────────────────
@pytest.fixture
def mock_share_repo():
return MagicMock()
def _make_share(
video_id="v1",
user_id="u1",
password=None,
expires_at=None,
is_active=True,
view_count=0,
download_count=0,
):
@pytest.fixture
def mock_video_repo():
return MagicMock()
@pytest.fixture
def sample_video():
video = MagicMock(spec=GeneratedVideo)
video.id = "video_001"
video.user_id = "user_001"
return video
@pytest.fixture
def sample_share():
share = VideoShare.create(
video_id=video_id,
user_id=user_id,
password=password,
expires_at=expires_at,
video_id="video_001",
user_id="user_001",
)
share.is_active = is_active
share.view_count = view_count
share.download_count = download_count
return share
def _make_video(video_id="v1", user_id="u1", name="test.mp4", file_url="http://x/v.mp4"):
return GeneratedVideo(
id=video_id,
project_id="p1",
generation_task_id="t1",
name=name,
file_url=file_url,
file_size=1024,
duration=10.0,
width=1920,
height=1080,
fps=30.0,
user_id=user_id,
@pytest.fixture
def sample_share_with_password():
share = VideoShare.create(
video_id="video_001",
user_id="user_001",
password="secret123",
)
return share
# ── CreateShareUseCase ──────────────────────────────────────────────────────
@pytest.fixture
def sample_share_expired():
# 直接构造已过期的分享(不经过create方法的校验)
share = VideoShare(
id="share_expired_001",
video_id="video_001",
user_id="user_001",
share_token="expiredtoken123",
expires_at=datetime.now(timezone.utc) - timedelta(hours=1),
)
return share
class TestCreateShareUseCase:
def test_create_success(self):
video = _make_video()
share_repo = MagicMock()
video_repo = MagicMock()
video_repo.get.return_value = video
share_repo.create.side_effect = lambda s: s
"""CreateShareUseCase 测试"""
uc = CreateShareUseCase(share_repo, video_repo)
cmd = CreateShareCommand(video_id="v1", user_id="u1")
result = uc.execute(cmd)
def test_create_share_success(self, mock_share_repo, mock_video_repo, sample_video):
"""正常创建分享链接"""
mock_video_repo.get.return_value = sample_video
mock_share_repo.create.side_effect = lambda s: s
assert result.video_id == "v1"
assert result.user_id == "u1"
video_repo.get.assert_called_once_with("v1")
share_repo.create.assert_called_once()
use_case = CreateShareUseCase(mock_share_repo, mock_video_repo)
command = CreateShareCommand(video_id="video_001", user_id="user_001")
result = use_case.execute(command)
def test_create_with_password(self):
video = _make_video()
share_repo = MagicMock()
video_repo = MagicMock()
video_repo.get.return_value = video
share_repo.create.side_effect = lambda s: s
assert result.video_id == "video_001"
assert result.user_id == "user_001"
assert result.share_token is not None
assert result.has_password is False
mock_share_repo.create.assert_called_once()
uc = CreateShareUseCase(share_repo, video_repo)
cmd = CreateShareCommand(video_id="v1", user_id="u1", password="secret")
result = uc.execute(cmd)
def test_create_share_with_password(self, mock_share_repo, mock_video_repo, sample_video):
"""创建带密码的分享"""
mock_video_repo.get.return_value = sample_video
mock_share_repo.create.side_effect = lambda s: s
use_case = CreateShareUseCase(mock_share_repo, mock_video_repo)
command = CreateShareCommand(
video_id="video_001",
user_id="user_001",
password="mypassword",
)
result = use_case.execute(command)
assert result.has_password is True
assert result.password_hash is not None
def test_video_not_found_raises(self):
share_repo = MagicMock()
video_repo = MagicMock()
video_repo.get.return_value = None
def test_create_share_with_expiry(self, mock_share_repo, mock_video_repo, sample_video):
"""创建带有效期的分享"""
mock_video_repo.get.return_value = sample_video
mock_share_repo.create.side_effect = lambda s: s
uc = CreateShareUseCase(share_repo, video_repo)
cmd = CreateShareCommand(video_id="v999", user_id="u1")
future = datetime.now(timezone.utc) + timedelta(days=7)
use_case = CreateShareUseCase(mock_share_repo, mock_video_repo)
command = CreateShareCommand(
video_id="video_001",
user_id="user_001",
expires_at=future,
)
result = use_case.execute(command)
assert result.expires_at == future
def test_create_share_video_not_found(self, mock_share_repo, mock_video_repo):
"""视频不存在时抛出 VideoNotFoundError"""
mock_video_repo.get.return_value = None
use_case = CreateShareUseCase(mock_share_repo, mock_video_repo)
command = CreateShareCommand(video_id="nonexistent", user_id="user_001")
with pytest.raises(VideoNotFoundError):
uc.execute(cmd)
use_case.execute(command)
def test_wrong_user_video_not_found(self):
video = _make_video(user_id="u2")
share_repo = MagicMock()
video_repo = MagicMock()
video_repo.get.return_value = video
mock_share_repo.create.assert_not_called()
uc = CreateShareUseCase(share_repo, video_repo)
cmd = CreateShareCommand(video_id="v1", user_id="u1")
def test_create_share_wrong_user(self, mock_share_repo, mock_video_repo, sample_video):
"""非视频所有者创建分享失败"""
sample_video.user_id = "user_other"
mock_video_repo.get.return_value = sample_video
use_case = CreateShareUseCase(mock_share_repo, mock_video_repo)
command = CreateShareCommand(video_id="video_001", user_id="user_001")
with pytest.raises(VideoNotFoundError):
uc.execute(cmd)
use_case.execute(command)
def test_video_without_user_id_attribute(self):
# 视频没有user_id字段的情况
class SimpleVideo:
pass
video = SimpleVideo()
video.id = "v1"
share_repo = MagicMock()
video_repo = MagicMock()
video_repo.get.return_value = video
share_repo.create.side_effect = lambda s: s
uc = CreateShareUseCase(share_repo, video_repo)
cmd = CreateShareCommand(video_id="v1", user_id="u1")
result = uc.execute(cmd)
assert result is not None
# ── GetShareByTokenUseCase ──────────────────────────────────────────────────
mock_share_repo.create.assert_not_called()
class TestGetShareByTokenUseCase:
def test_get_success(self):
share = _make_share()
repo = MagicMock()
repo.get_by_token.return_value = share
"""GetShareByTokenUseCase 测试"""
uc = GetShareByTokenUseCase(repo)
result = uc.execute(share.share_token)
assert result.id == share.id
def test_get_share_success(self, mock_share_repo, sample_share):
"""通过 token 正常获取分享信息"""
mock_share_repo.get_by_token.return_value = sample_share
def test_not_found_raises(self):
repo = MagicMock()
repo.get_by_token.return_value = None
use_case = GetShareByTokenUseCase(mock_share_repo)
result = use_case.execute(sample_share.share_token)
assert result.id == sample_share.id
mock_share_repo.get_by_token.assert_called_once_with(sample_share.share_token)
def test_get_share_not_found(self, mock_share_repo):
"""token 不存在时抛出 NotFoundError"""
mock_share_repo.get_by_token.return_value = None
use_case = GetShareByTokenUseCase(mock_share_repo)
uc = GetShareByTokenUseCase(repo)
with pytest.raises(NotFoundError):
uc.execute("nonexistent")
use_case.execute("invalid_token")
def test_expired_share_raises(self):
share = _make_share()
share.expires_at = datetime.now(timezone.utc) - timedelta(days=1)
repo = MagicMock()
repo.get_by_token.return_value = share
def test_get_share_expired_raises(self, mock_share_repo, sample_share_expired):
"""已过期的分享不可访问"""
mock_share_repo.get_by_token.return_value = sample_share_expired
use_case = GetShareByTokenUseCase(mock_share_repo)
uc = GetShareByTokenUseCase(repo)
with pytest.raises(ShareExpiredError):
uc.execute(share.share_token)
def test_revoked_share_raises(self):
share = _make_share(is_active=False)
repo = MagicMock()
repo.get_by_token.return_value = share
uc = GetShareByTokenUseCase(repo)
with pytest.raises(ShareExpiredError):
uc.execute(share.share_token)
# ── AccessShareUseCase ──────────────────────────────────────────────────────
use_case.execute(sample_share_expired.share_token)
class TestAccessShareUseCase:
def test_access_no_password(self):
share = _make_share()
video = _make_video()
share_repo = MagicMock()
video_repo = MagicMock()
share_repo.get_by_token.return_value = share
video_repo.get.return_value = video
"""AccessShareUseCase 测试"""
uc = AccessShareUseCase(share_repo, video_repo)
result = uc.execute(share.share_token)
def test_access_without_password(self, mock_share_repo, mock_video_repo, sample_share, sample_video):
"""无密码分享直接访问成功"""
mock_share_repo.get_by_token.return_value = sample_share
mock_video_repo.get.return_value = sample_video
assert isinstance(result, ShareAccessResult)
assert result.share.id == share.id
assert result.video.id == video.id
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
result = use_case.execute(sample_share.share_token)
assert result.share.id == sample_share.id
assert result.video.id == "video_001"
assert result.password_verified is True
assert share.view_count == 1
share_repo.increment_view.assert_called_once_with(share.id)
mock_share_repo.increment_view.assert_called_once_with(sample_share.id)
assert sample_share.view_count == 1
def test_access_with_correct_password(self):
share = _make_share(password="secret")
video = _make_video()
share_repo = MagicMock()
video_repo = MagicMock()
share_repo.get_by_token.return_value = share
video_repo.get.return_value = video
def test_access_with_correct_password(
self, mock_share_repo, mock_video_repo, sample_share_with_password, sample_video
):
"""带密码分享输入正确密码访问成功"""
mock_share_repo.get_by_token.return_value = sample_share_with_password
mock_video_repo.get.return_value = sample_video
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
result = use_case.execute(sample_share_with_password.share_token, password="secret123")
uc = AccessShareUseCase(share_repo, video_repo)
result = uc.execute(share.share_token, password="secret")
assert result.password_verified is True
mock_share_repo.increment_view.assert_called_once()
def test_access_password_required_but_not_provided(self):
share = _make_share(password="secret")
share_repo = MagicMock()
video_repo = MagicMock()
share_repo.get_by_token.return_value = share
def test_access_password_required_but_not_provided(
self, mock_share_repo, mock_video_repo, sample_share_with_password
):
"""带密码分享不输入密码抛出 PasswordRequiredError"""
mock_share_repo.get_by_token.return_value = sample_share_with_password
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
uc = AccessShareUseCase(share_repo, video_repo)
with pytest.raises(PasswordRequiredError):
uc.execute(share.share_token)
use_case.execute(sample_share_with_password.share_token)
def test_access_wrong_password(self):
share = _make_share(password="secret")
share_repo = MagicMock()
video_repo = MagicMock()
share_repo.get_by_token.return_value = share
mock_share_repo.increment_view.assert_not_called()
def test_access_wrong_password(self, mock_share_repo, mock_video_repo, sample_share_with_password):
"""密码错误抛出 InvalidPasswordError"""
mock_share_repo.get_by_token.return_value = sample_share_with_password
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
uc = AccessShareUseCase(share_repo, video_repo)
with pytest.raises(InvalidPasswordError):
uc.execute(share.share_token, password="wrong")
use_case.execute(sample_share_with_password.share_token, password="wrongpass")
def test_access_expired_share(self):
share = _make_share()
share.expires_at = datetime.now(timezone.utc) - timedelta(days=1)
share_repo = MagicMock()
share_repo.get_by_token.return_value = share
mock_share_repo.increment_view.assert_not_called()
uc = AccessShareUseCase(share_repo, MagicMock())
with pytest.raises(ShareExpiredError):
uc.execute(share.share_token)
def test_access_share_not_found(self, mock_share_repo, mock_video_repo):
"""分享不存在抛出 NotFoundError"""
mock_share_repo.get_by_token.return_value = None
def test_access_share_not_found(self):
share_repo = MagicMock()
share_repo.get_by_token.return_value = None
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
uc = AccessShareUseCase(share_repo, MagicMock())
with pytest.raises(NotFoundError):
uc.execute("nonexistent")
use_case.execute("invalid_token")
def test_access_video_not_found(self):
share = _make_share()
share_repo = MagicMock()
video_repo = MagicMock()
share_repo.get_by_token.return_value = share
video_repo.get.return_value = None
def test_access_expired_share(self, mock_share_repo, mock_video_repo, sample_share_expired):
"""已过期分享不可访问"""
mock_share_repo.get_by_token.return_value = sample_share_expired
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
with pytest.raises(ShareExpiredError):
use_case.execute(sample_share_expired.share_token)
mock_share_repo.increment_view.assert_not_called()
def test_access_video_not_found(self, mock_share_repo, mock_video_repo, sample_share):
"""分享存在但视频不存在"""
mock_share_repo.get_by_token.return_value = sample_share
mock_video_repo.get.return_value = None
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
uc = AccessShareUseCase(share_repo, video_repo)
with pytest.raises(VideoNotFoundError):
uc.execute(share.share_token)
# ── ListSharesByVideoUseCase ────────────────────────────────────────────────
use_case.execute(sample_share.share_token)
class TestListSharesByVideoUseCase:
def test_list_success(self):
shares = [_make_share(), _make_share()]
repo = MagicMock()
repo.list_by_video.return_value = shares
"""ListSharesByVideoUseCase 测试"""
uc = ListSharesByVideoUseCase(repo)
result = uc.execute("v1", "u1")
def test_list_by_video(self, mock_share_repo, sample_share):
"""列出某个视频的所有分享"""
mock_share_repo.list_by_video.return_value = [sample_share]
assert len(result) == 2
repo.list_by_video.assert_called_once_with("v1", "u1")
use_case = ListSharesByVideoUseCase(mock_share_repo)
result = use_case.execute("video_001", "user_001")
def test_list_empty(self):
repo = MagicMock()
repo.list_by_video.return_value = []
assert len(result) == 1
mock_share_repo.list_by_video.assert_called_once_with("video_001", "user_001")
def test_list_by_video_empty(self, mock_share_repo):
"""视频没有分享记录时返回空列表"""
mock_share_repo.list_by_video.return_value = []
use_case = ListSharesByVideoUseCase(mock_share_repo)
result = use_case.execute("video_001", "user_001")
uc = ListSharesByVideoUseCase(repo)
result = uc.execute("v1", "u1")
assert result == []
# ── ListSharesByUserUseCase ─────────────────────────────────────────────────
class TestListSharesByUserUseCase:
def test_list_with_pagination(self):
shares = [_make_share() for _ in range(5)]
repo = MagicMock()
repo.list_by_user.return_value = shares
repo.count_by_user.return_value = 20
"""ListSharesByUserUseCase 测试"""
uc = ListSharesByUserUseCase(repo)
items, total = uc.execute("u1", skip=0, limit=5)
def test_list_by_user(self, mock_share_repo, sample_share):
"""列出用户的所有分享"""
mock_share_repo.list_by_user.return_value = [sample_share]
mock_share_repo.count_by_user.return_value = 1
assert len(items) == 5
assert total == 20
repo.list_by_user.assert_called_once_with("u1", skip=0, limit=5)
repo.count_by_user.assert_called_once_with("u1")
use_case = ListSharesByUserUseCase(mock_share_repo)
items, total = use_case.execute("user_001")
def test_list_default_params(self):
repo = MagicMock()
repo.list_by_user.return_value = []
repo.count_by_user.return_value = 0
assert len(items) == 1
assert total == 1
mock_share_repo.list_by_user.assert_called_once_with("user_001", skip=0, limit=20)
uc = ListSharesByUserUseCase(repo)
uc.execute("u1")
def test_list_by_user_with_pagination(self, mock_share_repo):
"""带分页参数查询"""
mock_share_repo.list_by_user.return_value = []
mock_share_repo.count_by_user.return_value = 50
repo.list_by_user.assert_called_once_with("u1", skip=0, limit=20)
use_case = ListSharesByUserUseCase(mock_share_repo)
items, total = use_case.execute("user_001", skip=10, limit=5)
assert total == 50
mock_share_repo.list_by_user.assert_called_once_with("user_001", skip=10, limit=5)
# ── UpdateShareUseCase ──────────────────────────────────────────────────────
def test_list_by_user_empty(self, mock_share_repo):
"""用户没有分享记录"""
mock_share_repo.list_by_user.return_value = []
mock_share_repo.count_by_user.return_value = 0
use_case = ListSharesByUserUseCase(mock_share_repo)
items, total = use_case.execute("user_001")
assert items == []
assert total == 0
class TestUpdateShareUseCase:
def test_update_password(self):
share = _make_share(password="oldpass")
repo = MagicMock()
repo.get_by_id.return_value = share
repo.update.side_effect = lambda s: s
"""UpdateShareUseCase 测试"""
uc = UpdateShareUseCase(repo)
cmd = UpdateShareCommand(share_id=share.id, user_id="u1", password="newpass")
result = uc.execute(cmd)
def test_update_password(self, mock_share_repo, sample_share):
"""更新分享密码"""
mock_share_repo.get_by_id.return_value = sample_share
mock_share_repo.update.side_effect = lambda s: s
assert result is not None
assert share.verify_password("newpass") is True
assert share.verify_password("oldpass") is False
repo.update.assert_called_once()
use_case = UpdateShareUseCase(mock_share_repo)
command = UpdateShareCommand(
share_id=sample_share.id,
user_id="user_001",
password="newpassword",
)
result = use_case.execute(command)
def test_clear_password(self):
share = _make_share(password="oldpass")
repo = MagicMock()
repo.get_by_id.return_value = share
repo.update.side_effect = lambda s: s
assert result.has_password is True
mock_share_repo.update.assert_called_once()
uc = UpdateShareUseCase(repo)
cmd = UpdateShareCommand(share_id=share.id, user_id="u1", password="")
result = uc.execute(cmd)
def test_clear_password(self, mock_share_repo, sample_share_with_password):
"""清除分享密码(空字符串)"""
mock_share_repo.get_by_id.return_value = sample_share_with_password
mock_share_repo.update.side_effect = lambda s: s
use_case = UpdateShareUseCase(mock_share_repo)
command = UpdateShareCommand(
share_id=sample_share_with_password.id,
user_id="user_001",
password="", # 空字符串表示清除
)
result = use_case.execute(command)
assert result.has_password is False
assert result.password_hash is None
def test_update_password_none_no_change(self):
share = _make_share(password="oldpass")
repo = MagicMock()
repo.get_by_id.return_value = share
repo.update.side_effect = lambda s: s
def test_update_password_none_no_change(self, mock_share_repo, sample_share_with_password):
"""password=None 不修改密码"""
original_hash = sample_share_with_password.password_hash
mock_share_repo.get_by_id.return_value = sample_share_with_password
mock_share_repo.update.side_effect = lambda s: s
uc = UpdateShareUseCase(repo)
cmd = UpdateShareCommand(share_id=share.id, user_id="u1", password=None)
result = uc.execute(cmd)
use_case = UpdateShareUseCase(mock_share_repo)
command = UpdateShareCommand(
share_id=sample_share_with_password.id,
user_id="user_001",
password=None, # None表示不修改
)
result = use_case.execute(command)
# password=None 表示不修改
assert result.has_password is True
assert share.verify_password("oldpass") is True
assert result.password_hash == original_hash
def test_update_expires_at(self):
share = _make_share()
new_expiry = datetime.now(timezone.utc) + timedelta(days=30)
repo = MagicMock()
repo.get_by_id.return_value = share
repo.update.side_effect = lambda s: s
def test_update_expires_at(self, mock_share_repo, sample_share):
"""更新有效期"""
mock_share_repo.get_by_id.return_value = sample_share
mock_share_repo.update.side_effect = lambda s: s
uc = UpdateShareUseCase(repo)
cmd = UpdateShareCommand(share_id=share.id, user_id="u1", expires_at=new_expiry)
result = uc.execute(cmd)
future = datetime.now(timezone.utc) + timedelta(days=3)
use_case = UpdateShareUseCase(mock_share_repo)
command = UpdateShareCommand(
share_id=sample_share.id,
user_id="user_001",
expires_at=future,
)
result = use_case.execute(command)
assert result.expires_at == new_expiry
assert result.expires_at == future
def test_update_expires_at_past_raises(self):
share = _make_share()
past = datetime.now(timezone.utc) - timedelta(days=1)
repo = MagicMock()
repo.get_by_id.return_value = share
def test_update_expires_at_past_raises(self, mock_share_repo, sample_share):
"""设置过去的有效期抛出 ValueError"""
mock_share_repo.get_by_id.return_value = sample_share
uc = UpdateShareUseCase(repo)
cmd = UpdateShareCommand(share_id=share.id, user_id="u1", expires_at=past)
past = datetime.now(timezone.utc) - timedelta(hours=1)
use_case = UpdateShareUseCase(mock_share_repo)
command = UpdateShareCommand(
share_id=sample_share.id,
user_id="user_001",
expires_at=past,
)
with pytest.raises(ValueError, match="expires_at cannot be in the past"):
uc.execute(cmd)
use_case.execute(command)
def test_update_not_found_raises(self):
repo = MagicMock()
repo.get_by_id.return_value = None
mock_share_repo.update.assert_not_called()
uc = UpdateShareUseCase(repo)
cmd = UpdateShareCommand(share_id="nonexistent", user_id="u1")
def test_update_share_not_found(self, mock_share_repo):
"""分享不存在抛出 NotFoundError"""
mock_share_repo.get_by_id.return_value = None
use_case = UpdateShareUseCase(mock_share_repo)
command = UpdateShareCommand(
share_id="nonexistent",
user_id="user_001",
password="newpass",
)
with pytest.raises(NotFoundError):
uc.execute(cmd)
use_case.execute(command)
def test_update_wrong_user_not_found(self):
share = _make_share(user_id="u2")
repo = MagicMock()
repo.get_by_id.return_value = None # 仓储层已经按user_id过滤了
uc = UpdateShareUseCase(repo)
cmd = UpdateShareCommand(share_id=share.id, user_id="u1")
with pytest.raises(NotFoundError):
uc.execute(cmd)
# ── RevokeShareUseCase ──────────────────────────────────────────────────────
mock_share_repo.update.assert_not_called()
class TestRevokeShareUseCase:
def test_revoke_success(self):
repo = MagicMock()
repo.get_by_id.return_value = MagicMock()
repo.delete.return_value = True
"""RevokeShareUseCase 测试"""
def test_revoke_success(self, mock_share_repo, sample_share):
"""撤销分享成功"""
mock_share_repo.get_by_id.return_value = sample_share
mock_share_repo.delete.return_value = True
use_case = RevokeShareUseCase(mock_share_repo)
result = use_case.execute(sample_share.id, "user_001")
uc = RevokeShareUseCase(repo)
result = uc.execute("s1", "u1")
assert result is True
repo.delete.assert_called_once_with("s1", "u1")
mock_share_repo.delete.assert_called_once_with(sample_share.id, "user_001")
def test_revoke_not_found_raises(self):
repo = MagicMock()
repo.get_by_id.return_value = None
def test_revoke_not_found(self, mock_share_repo):
"""分享不存在抛出 NotFoundError"""
mock_share_repo.get_by_id.return_value = None
use_case = RevokeShareUseCase(mock_share_repo)
uc = RevokeShareUseCase(repo)
with pytest.raises(NotFoundError):
uc.execute("s1", "u1")
use_case.execute("nonexistent", "user_001")
# ── RecordShareDownloadUseCase ──────────────────────────────────────────────
mock_share_repo.delete.assert_not_called()
class TestRecordShareDownloadUseCase:
def test_record_download_success(self):
share = _make_share(download_count=3)
repo = MagicMock()
repo.get_by_token.return_value = share
"""RecordShareDownloadUseCase 测试"""
uc = RecordShareDownloadUseCase(repo)
uc.execute(share.share_token)
def test_record_download_no_password(self, mock_share_repo, sample_share):
"""无密码分享记录下载"""
mock_share_repo.get_by_token.return_value = sample_share
repo.increment_download.assert_called_once_with(share.id)
use_case = RecordShareDownloadUseCase(mock_share_repo)
use_case.execute(sample_share.share_token)
def test_record_download_with_password(self):
share = _make_share(password="secret")
repo = MagicMock()
repo.get_by_token.return_value = share
mock_share_repo.increment_download.assert_called_once_with(sample_share.id)
uc = RecordShareDownloadUseCase(repo)
uc.execute(share.share_token, password="secret")
repo.increment_download.assert_called_once()
def test_record_download_with_password(self, mock_share_repo, sample_share_with_password):
"""带密码分享正确密码记录下载"""
mock_share_repo.get_by_token.return_value = sample_share_with_password
def test_record_download_wrong_password_raises(self):
share = _make_share(password="secret")
repo = MagicMock()
repo.get_by_token.return_value = share
use_case = RecordShareDownloadUseCase(mock_share_repo)
use_case.execute(sample_share_with_password.share_token, password="secret123")
mock_share_repo.increment_download.assert_called_once()
def test_record_download_wrong_password(self, mock_share_repo, sample_share_with_password):
"""密码错误不记录下载"""
mock_share_repo.get_by_token.return_value = sample_share_with_password
use_case = RecordShareDownloadUseCase(mock_share_repo)
uc = RecordShareDownloadUseCase(repo)
with pytest.raises(InvalidPasswordError):
uc.execute(share.share_token, password="wrong")
use_case.execute(sample_share_with_password.share_token, password="wrong")
def test_record_download_expired_raises(self):
share = _make_share()
share.expires_at = datetime.now(timezone.utc) - timedelta(days=1)
repo = MagicMock()
repo.get_by_token.return_value = share
mock_share_repo.increment_download.assert_not_called()
uc = RecordShareDownloadUseCase(repo)
with pytest.raises(ShareExpiredError):
uc.execute(share.share_token)
def test_record_download_not_found(self, mock_share_repo):
"""分享不存在抛出 NotFoundError"""
mock_share_repo.get_by_token.return_value = None
def test_record_download_not_found_raises(self):
repo = MagicMock()
repo.get_by_token.return_value = None
use_case = RecordShareDownloadUseCase(mock_share_repo)
uc = RecordShareDownloadUseCase(repo)
with pytest.raises(NotFoundError):
uc.execute("nonexistent")
use_case.execute("invalid_token")
def test_record_download_expired(self, mock_share_repo, sample_share_expired):
"""已过期分享不能下载"""
mock_share_repo.get_by_token.return_value = sample_share_expired
use_case = RecordShareDownloadUseCase(mock_share_repo)
with pytest.raises(ShareExpiredError):
use_case.execute(sample_share_expired.share_token)
mock_share_repo.increment_download.assert_not_called()