Compare commits
20 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 4642130a04 | |||
| 66e7faa06a | |||
| 2f92398212 | |||
| c9449a1dbd | |||
| 675945f4fb | |||
| 5e74d0b565 | |||
| ebe68429bc | |||
| 18f534bbd6 | |||
| f36aaea374 | |||
| 80bc63d58b | |||
| 73a566c621 | |||
| 9057ba25c8 | |||
| 35b38ed48c | |||
| 15142e3168 | |||
| 953dd9a6e6 | |||
| 060307197c | |||
| 6930d4543f | |||
| 7b2f35ad3d | |||
| 5256bd0d8b | |||
| 33270dd026 |
@@ -1,188 +0,0 @@
|
||||
/**
|
||||
* 认证相关 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
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
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)
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
/**
|
||||
* 认证相关 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"
|
||||
@@ -0,0 +1,37 @@
|
||||
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")
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
/**
|
||||
* 认证相关类型定义
|
||||
*/
|
||||
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
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,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
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>
|
||||
</>
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
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>
|
||||
)}
|
||||
</>
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
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"
|
||||
@@ -0,0 +1,41 @@
|
||||
/** 格式化文件大小 */
|
||||
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: "🎵 音频" },
|
||||
]
|
||||
@@ -1,182 +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 "./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
@@ -159,3 +159,4 @@
|
||||
flex-wrap: wrap;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
/** 路由 → 标题映射(用于自动生成面包屑) */
|
||||
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": "查重结果",
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
/**
|
||||
* 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"
|
||||
@@ -0,0 +1,23 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
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>
|
||||
)
|
||||
Regular → Executable
+5
-75
@@ -2,11 +2,9 @@
|
||||
* 选中贴纸的属性编辑器
|
||||
*/
|
||||
import React from "react"
|
||||
import type { StickerItem, TextStickerPreset } from "@/pages/editing-planner/types"
|
||||
import {
|
||||
TEXT_PRESET_STYLES,
|
||||
TEXT_STICKER_PRESET_LABELS,
|
||||
} from "@/pages/editing-planner/constants/sticker"
|
||||
import type { StickerItem } from "@/pages/editing-planner/types"
|
||||
import { StickerPreview } from "./StickerPreview"
|
||||
import { TextStickerPropsEditor } from "./TextStickerPropsEditor"
|
||||
|
||||
interface StickerPropsEditorProps {
|
||||
sticker: StickerItem
|
||||
@@ -121,78 +119,10 @@ const StickerPropsEditor: React.FC<StickerPropsEditorProps> = ({
|
||||
</div>
|
||||
|
||||
{/* 文字贴纸特有属性 */}
|
||||
{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>
|
||||
</>
|
||||
)}
|
||||
<TextStickerPropsEditor sticker={sticker} onUpdate={onUpdate} />
|
||||
|
||||
{/* 预览 */}
|
||||
<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>
|
||||
<StickerPreview sticker={sticker} />
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
+57
@@ -0,0 +1,57 @@
|
||||
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,3 +1,7 @@
|
||||
/**
|
||||
* Auth API 测试
|
||||
* 对应 api/auth/ 目录化后的模块
|
||||
*/
|
||||
import { describe, expect, it, vi, beforeEach } from "vitest"
|
||||
import {
|
||||
normalizeUser,
|
||||
@@ -10,7 +14,6 @@ import {
|
||||
resetPassword,
|
||||
verifyEmail,
|
||||
} from "@/api/auth"
|
||||
|
||||
const mockPost = vi.fn()
|
||||
const mockGet = vi.fn()
|
||||
const mockAxiosPost = vi.fn()
|
||||
|
||||
Regular → Executable
+2
@@ -61,6 +61,8 @@ 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"
|
||||
|
||||
Executable
+281
@@ -0,0 +1,281 @@
|
||||
"""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
|
||||
@@ -1,58 +1,55 @@
|
||||
"""video_share 视频分享领域实体单测."""
|
||||
"""视频分享领域模型单元测试 — wave215"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import pytest
|
||||
from domain.video_share import (
|
||||
|
||||
from packages.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_none_password_returns_empty(self):
|
||||
assert _hash_password(None) == ""
|
||||
|
||||
def test_same_password_same_hash(self):
|
||||
h1 = _hash_password("mypassword")
|
||||
h2 = _hash_password("mypassword")
|
||||
h1 = _hash_password("secret123")
|
||||
h2 = _hash_password("secret123")
|
||||
assert h1 == h2
|
||||
assert h1 != ""
|
||||
|
||||
def test_different_passwords_different_hashes(self):
|
||||
h1 = _hash_password("password1")
|
||||
h2 = _hash_password("password2")
|
||||
def test_different_password_different_hash(self):
|
||||
h1 = _hash_password("pass1")
|
||||
h2 = _hash_password("pass2")
|
||||
assert h1 != h2
|
||||
|
||||
def test_hash_is_hex_string(self):
|
||||
def test_hash_is_sha256_hex(self):
|
||||
h = _hash_password("test")
|
||||
assert isinstance(h, str)
|
||||
assert len(h) == 64 # SHA-256 hex
|
||||
int(h, 16) # 应该能被解析为16进制
|
||||
assert len(h) == 64
|
||||
assert re.match(r"^[0-9a-f]{64}$", h)
|
||||
|
||||
def test_hash_contains_salt(self):
|
||||
# 直接的 SHA-256(password) 应该不等于加盐后的
|
||||
from hashlib import sha256
|
||||
# 直接SHA-256("test") vs 加盐后的结果应该不同
|
||||
import hashlib
|
||||
|
||||
raw = sha256("mypass".encode()).hexdigest()
|
||||
salted = _hash_password("mypass")
|
||||
assert raw != salted
|
||||
direct = hashlib.sha256(b"test").hexdigest()
|
||||
salted = _hash_password("test")
|
||||
assert direct != salted
|
||||
|
||||
|
||||
# ── generate_share_token ─────────────────────────────────────────────────────
|
||||
# ── Token 生成 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGenerateShareToken:
|
||||
"""generate_share_token 函数"""
|
||||
|
||||
def test_default_length(self):
|
||||
def test_default_length_12(self):
|
||||
token = generate_share_token()
|
||||
assert len(token) == 12
|
||||
|
||||
@@ -60,231 +57,228 @@ class TestGenerateShareToken:
|
||||
token = generate_share_token(20)
|
||||
assert len(token) == 20
|
||||
|
||||
def test_short_token(self):
|
||||
token = generate_share_token(6)
|
||||
assert len(token) == 6
|
||||
|
||||
def test_url_friendly_chars(self):
|
||||
def test_url_friendly_no_ambiguous_chars(self):
|
||||
# 不应包含容易混淆的字符:i, l, o, I, L, O, 0, 1
|
||||
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
|
||||
for ch in "ilO01":
|
||||
assert ch not in token
|
||||
|
||||
def test_unique_tokens(self):
|
||||
tokens = {generate_share_token() for _ in range(100)}
|
||||
assert len(tokens) == 100 # 应该都是唯一的
|
||||
|
||||
def test_alphanumeric(self):
|
||||
def test_alphanumeric_only(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:
|
||||
"""VideoShare.create 工厂方法"""
|
||||
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
|
||||
|
||||
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_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_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_empty_password_no_hash(self):
|
||||
share = VideoShare.create(video_id="v1", user_id="u1", password="")
|
||||
assert share.password_hash is None
|
||||
|
||||
def test_with_expiry(self):
|
||||
def test_create_with_expires_at(self):
|
||||
future = datetime.now(timezone.utc) + timedelta(days=7)
|
||||
s = VideoShare.create(video_id="v1", user_id="u1", expires_at=future)
|
||||
assert s.expires_at == future
|
||||
share = VideoShare.create(video_id="v1", user_id="u1", expires_at=future)
|
||||
assert share.expires_at == future
|
||||
|
||||
def test_empty_video_id_raises(self):
|
||||
with pytest.raises(ValueError, match="video_id"):
|
||||
VideoShare.create(video_id="", user_id="u1")
|
||||
|
||||
def test_whitespace_video_id_raises(self):
|
||||
with pytest.raises(ValueError):
|
||||
VideoShare.create(video_id=" ", user_id="u1")
|
||||
|
||||
def test_empty_user_id_raises(self):
|
||||
with pytest.raises(ValueError, match="user_id"):
|
||||
VideoShare.create(video_id="v1", user_id="")
|
||||
|
||||
def test_past_expiry_raises(self):
|
||||
past = datetime.now(timezone.utc) - timedelta(hours=1)
|
||||
with pytest.raises(ValueError, match="past"):
|
||||
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_video_id_stripped(self):
|
||||
s = VideoShare.create(video_id=" vid_123 ", user_id="u1")
|
||||
assert s.video_id == "vid_123"
|
||||
def test_create_empty_video_id_raises(self):
|
||||
with pytest.raises(ValueError, match="video_id cannot be empty"):
|
||||
VideoShare.create(video_id="", user_id="u1")
|
||||
|
||||
def test_user_id_stripped(self):
|
||||
s = VideoShare.create(video_id="v1", user_id=" user_456 ")
|
||||
assert s.user_id == "user_456"
|
||||
def test_create_whitespace_video_id_raises(self):
|
||||
with pytest.raises(ValueError, match="video_id cannot be empty"):
|
||||
VideoShare.create(video_id=" ", user_id="u1")
|
||||
|
||||
def test_unique_ids(self):
|
||||
def test_create_empty_user_id_raises(self):
|
||||
with pytest.raises(ValueError, match="user_id cannot be empty"):
|
||||
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_create_unique_id_each_time(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_unique_tokens(self):
|
||||
def test_create_unique_token_each_time(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
|
||||
|
||||
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
|
||||
|
||||
# ── has_password ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
# ── VideoShare 属性方法 ─────────────────────────────────────────────────────
|
||||
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
|
||||
|
||||
|
||||
class TestVideoShareProperties:
|
||||
"""VideoShare 属性方法"""
|
||||
# ── is_expired ──────────────────────────────────────────────────────────────
|
||||
|
||||
def test_has_password_true(self):
|
||||
s = VideoShare.create(video_id="v1", user_id="u1", password="pass")
|
||||
assert s.has_password is True
|
||||
|
||||
def test_has_password_false(self):
|
||||
s = VideoShare.create(video_id="v1", user_id="u1")
|
||||
assert s.has_password is False
|
||||
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_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):
|
||||
def test_future_expiry_not_expired(self):
|
||||
future = datetime.now(timezone.utc) + timedelta(hours=1)
|
||||
s = VideoShare.create(video_id="v1", user_id="u1", expires_at=future)
|
||||
assert s.is_expired is False
|
||||
share = VideoShare.create(video_id="v1", user_id="u1", expires_at=future)
|
||||
assert share.is_expired is False
|
||||
|
||||
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
|
||||
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
|
||||
|
||||
|
||||
# ── VideoShare 方法 ─────────────────────────────────────────────────────────
|
||||
# ── is_accessible ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestVideoShareMethods:
|
||||
"""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
|
||||
|
||||
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_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_correct(self):
|
||||
s = VideoShare.create(video_id="v1", user_id="u1", password="mypass")
|
||||
assert s.verify_password("mypass") 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_wrong(self):
|
||||
s = VideoShare.create(video_id="v1", user_id="u1", password="mypass")
|
||||
assert s.verify_password("wrongpass") is False
|
||||
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_empty_false(self):
|
||||
s = VideoShare.create(video_id="v1", user_id="u1", password="mypass")
|
||||
assert s.verify_password("") is False
|
||||
|
||||
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
|
||||
# ── verify_password ─────────────────────────────────────────────────────────
|
||||
|
||||
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
|
||||
|
||||
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
|
||||
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_makes_inaccessible(self):
|
||||
s = VideoShare.create(video_id="v1", user_id="u1")
|
||||
assert s.is_accessible is True
|
||||
s.revoke()
|
||||
assert s.is_accessible is False
|
||||
share = VideoShare.create(video_id="v1", user_id="u1")
|
||||
share.revoke()
|
||||
assert share.is_accessible 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
|
||||
def test_revoke_idempotent(self):
|
||||
share = VideoShare.create(video_id="v1", user_id="u1")
|
||||
share.revoke()
|
||||
share.revoke() # 第二次也不报错
|
||||
assert share.is_active is False
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+393
-304
@@ -1,63 +1,68 @@
|
||||
"""JWT 服务单元测试 — wave130."""
|
||||
|
||||
from __future__ import annotations
|
||||
"""JWT 服务与处理器单元测试."""
|
||||
|
||||
import time
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import jwt as pyjwt
|
||||
import jwt
|
||||
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 # 满足长度要求的测试密钥
|
||||
|
||||
|
||||
TEST_SECRET = "test-secret-key-for-unit-testing-only-1234567890"
|
||||
TEST_ALGORITHM = "HS256"
|
||||
|
||||
|
||||
# ── JWTConfig 配置 ──────────────────────────────────────────────────────────
|
||||
# ── JWTConfig 测试 ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestJWTConfig:
|
||||
def test_normal_config(self):
|
||||
config = JWTConfig(secret_key=TEST_SECRET)
|
||||
assert config.SECRET_KEY == TEST_SECRET
|
||||
"""JWTConfig 配置类测试"""
|
||||
|
||||
def test_init_with_valid_secret(self):
|
||||
config = JWTConfig(secret_key=STRONG_SECRET)
|
||||
assert config.SECRET_KEY == STRONG_SECRET
|
||||
assert config.ALGORITHM == "HS256"
|
||||
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 15
|
||||
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 7
|
||||
|
||||
def test_custom_config(self):
|
||||
def test_init_custom_values(self):
|
||||
config = JWTConfig(
|
||||
secret_key=TEST_SECRET,
|
||||
secret_key=STRONG_SECRET,
|
||||
algorithm="HS384",
|
||||
access_token_expire_minutes=60,
|
||||
refresh_token_expire_days=30,
|
||||
refresh_token_expire_days=14,
|
||||
)
|
||||
assert config.SECRET_KEY == STRONG_SECRET
|
||||
assert config.ALGORITHM == "HS384"
|
||||
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 60
|
||||
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 30
|
||||
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 14
|
||||
|
||||
def test_empty_secret_raises(self):
|
||||
with pytest.raises(ValueError, match="secret_key must be provided"):
|
||||
JWTConfig(secret_key="")
|
||||
|
||||
def test_whitespace_secret_raises(self):
|
||||
with pytest.raises(ValueError):
|
||||
def test_whitespace_only_secret_raises(self):
|
||||
with pytest.raises(ValueError, match="secret_key must be provided"):
|
||||
JWTConfig(secret_key=" ")
|
||||
|
||||
def test_none_secret_raises(self):
|
||||
with pytest.raises(ValueError):
|
||||
JWTConfig(secret_key=None) # type: ignore
|
||||
with pytest.raises(ValueError, match="secret_key must be provided"):
|
||||
JWTConfig(secret_key=None)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_secret",
|
||||
"insecure_secret",
|
||||
[
|
||||
"your-secret-key-change-in-production",
|
||||
"your-secret-key",
|
||||
@@ -68,323 +73,407 @@ class TestJWTConfig:
|
||||
"Your-Secret-Key",
|
||||
],
|
||||
)
|
||||
def test_insecure_defaults_rejected(self, bad_secret):
|
||||
def test_insecure_default_secret_raises(self, insecure_secret):
|
||||
with pytest.raises(ValueError, match="insecure"):
|
||||
JWTConfig(secret_key=bad_secret)
|
||||
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
|
||||
|
||||
|
||||
# ── JWTService 初始化 ──────────────────────────────────────────────────────
|
||||
# ── JWTService 初始化测试 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestJWTServiceInit:
|
||||
def test_with_config_works(self):
|
||||
config = JWTConfig(secret_key=TEST_SECRET)
|
||||
"""JWTService 初始化测试"""
|
||||
|
||||
def test_init_with_config(self):
|
||||
config = JWTConfig(secret_key=STRONG_SECRET)
|
||||
service = JWTService(config)
|
||||
assert service.config is config
|
||||
|
||||
def test_none_config_raises(self):
|
||||
with pytest.raises(ValueError, match="JWTService requires"):
|
||||
def test_init_none_config_raises(self):
|
||||
with pytest.raises(ValueError, match="JWTService requires a JWTConfig"):
|
||||
JWTService(None)
|
||||
|
||||
|
||||
# ── 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 常量 ──────────────────────────────────────────────────────────
|
||||
# ── TokenType 测试 ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTokenType:
|
||||
"""TokenType 常量测试"""
|
||||
|
||||
def test_access_value(self):
|
||||
assert TokenType.ACCESS == "access"
|
||||
|
||||
def test_refresh_value(self):
|
||||
assert TokenType.REFRESH == "refresh"
|
||||
|
||||
def test_different_types(self):
|
||||
def test_access_and_refresh_different(self):
|
||||
assert TokenType.ACCESS != TokenType.REFRESH
|
||||
|
||||
|
||||
# ── 多算法支持 ──────────────────────────────────────────────────────────────
|
||||
# ── JWTService create_access_token 测试 ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestDifferentAlgorithms:
|
||||
def test_hs384_works(self):
|
||||
config = JWTConfig(secret_key=TEST_SECRET * 2, algorithm="HS384")
|
||||
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)
|
||||
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")
|
||||
payload = service.verify_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
|
||||
def test_hs512_works(self):
|
||||
config = JWTConfig(secret_key=TEST_SECRET * 3, algorithm="HS512")
|
||||
service = JWTService(config)
|
||||
token = service.create_access_token(user_id="u1")
|
||||
payload = service.verify_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)
|
||||
|
||||
token = service_256.create_access_token(user_id="u1")
|
||||
# 用 HS256 解码应该失败
|
||||
with pytest.raises(InvalidTokenError):
|
||||
service_384.verify_token(token)
|
||||
jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
# 用 HS384 解码应该成功
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS384"])
|
||||
assert payload["sub"] == "u1"
|
||||
|
||||
|
||||
# ── 边界:空用户ID等 ────────────────────────────────────────────────────────
|
||||
# ── JWTService create_refresh_token 测试 ────────────────────────────────────
|
||||
|
||||
|
||||
class TestEdgeCases:
|
||||
def setup_method(self):
|
||||
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
|
||||
class TestCreateRefreshToken:
|
||||
"""创建 refresh_token 测试"""
|
||||
|
||||
def test_empty_user_id(self):
|
||||
token = self.service.create_access_token(user_id="")
|
||||
payload = self.service.verify_access_token(token)
|
||||
assert payload["sub"] == ""
|
||||
@pytest.fixture
|
||||
def service(self):
|
||||
return JWTService(JWTConfig(secret_key=STRONG_SECRET))
|
||||
|
||||
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_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_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
|
||||
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_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_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_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}"
|
||||
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)
|
||||
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)
|
||||
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_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
|
||||
|
||||
|
||||
# ── 全局 JWT handler 测试 ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGlobalJWTHandler:
|
||||
"""全局 JWT Handler 配置与获取测试"""
|
||||
|
||||
def test_configure_creates_handler(self):
|
||||
handler = configure_jwt_handler(secret_key=STRONG_SECRET)
|
||||
assert isinstance(handler, JWTHandler)
|
||||
|
||||
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_get_before_configure_raises(self):
|
||||
# 重置全局状态(通过设置 None 模拟未配置)
|
||||
import packages.application.auth.jwt_handler as mod
|
||||
|
||||
mod._default_handler = None
|
||||
with pytest.raises(RuntimeError, match="JWT handler not configured"):
|
||||
get_jwt_handler()
|
||||
|
||||
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
|
||||
|
||||
+601
-412
File diff suppressed because it is too large
Load Diff
@@ -1,6 +1,4 @@
|
||||
"""密码重置 UseCase 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
"""密码重置 Use Case 单元测试."""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
@@ -15,285 +13,477 @@ 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():
|
||||
return MagicMock()
|
||||
"""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
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_email_service():
|
||||
"""mock 邮件服务."""
|
||||
svc = MagicMock()
|
||||
svc.send_password_reset_email.return_value = (True, None)
|
||||
return svc
|
||||
|
||||
|
||||
@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
|
||||
# ── RequestPasswordResetUseCase 测试 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRequestPasswordResetRequest:
|
||||
"""RequestPasswordResetRequest 测试"""
|
||||
class TestRequestPasswordReset:
|
||||
"""请求密码重置用例测试"""
|
||||
|
||||
def test_email_lowercased_and_stripped(self):
|
||||
"""邮箱转小写并去空格"""
|
||||
req = RequestPasswordResetRequest(" User@Example.COM ")
|
||||
assert req.email == "user@example.com"
|
||||
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
|
||||
|
||||
def test_empty_email(self):
|
||||
"""空邮箱"""
|
||||
req = RequestPasswordResetRequest("")
|
||||
assert req.email == ""
|
||||
|
||||
|
||||
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,
|
||||
uc = RequestPasswordResetUseCase(
|
||||
user_repository=mock_user_repo,
|
||||
base_url="https://app.example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RequestPasswordResetRequest("user@example.com")
|
||||
use_case.execute(request)
|
||||
req = RequestPasswordResetRequest(email="user@example.com")
|
||||
ok, error = uc.execute(req)
|
||||
|
||||
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
|
||||
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_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
|
||||
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")
|
||||
|
||||
use_case = RequestPasswordResetUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
uc = RequestPasswordResetUseCase(
|
||||
user_repository=mock_user_repo,
|
||||
base_url="https://app.example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RequestPasswordResetRequest("user@example.com")
|
||||
success, error = use_case.execute(request)
|
||||
req = RequestPasswordResetRequest(email="user@example.com")
|
||||
ok, error = uc.execute(req)
|
||||
|
||||
assert success is True
|
||||
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_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
|
||||
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
|
||||
|
||||
use_case = RequestPasswordResetUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
uc = RequestPasswordResetUseCase(
|
||||
user_repository=mock_user_repo,
|
||||
base_url="https://app.example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RequestPasswordResetRequest("user@example.com")
|
||||
req = RequestPasswordResetRequest(email="user@example.com")
|
||||
uc.execute(req)
|
||||
|
||||
use_case.execute(request)
|
||||
token1 = sample_user.password_reset_token
|
||||
call_kwargs = mock_email_service.send_password_reset_email.call_args.kwargs
|
||||
assert call_kwargs["username"] == "Display Name"
|
||||
|
||||
use_case.execute(request)
|
||||
token2 = sample_user.password_reset_token
|
||||
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
|
||||
|
||||
assert token1 != token2
|
||||
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 测试 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRequestPasswordResetRequest:
|
||||
"""请求数据类测试"""
|
||||
|
||||
def test_email_stripped_and_lowercased(self):
|
||||
req = RequestPasswordResetRequest(email=" 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"
|
||||
|
||||
|
||||
# ── ResetPasswordRequest 测试 ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestResetPasswordRequest:
|
||||
"""ResetPasswordRequest 测试"""
|
||||
"""重置密码请求数据类测试"""
|
||||
|
||||
def test_stores_token_and_password(self):
|
||||
"""正确存储 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
|
||||
req = ResetPasswordRequest(token="token123", new_password="password123")
|
||||
assert req.token == "token123"
|
||||
assert req.new_password == "password123"
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,100 +1,295 @@
|
||||
"""text_splitter 单元测试."""
|
||||
"""TTS 文本分段工具单元测试."""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.tts_job.text_splitter import split_text
|
||||
|
||||
|
||||
class TestSplitText:
|
||||
def test_empty_text_returns_empty(self):
|
||||
class TestSplitTextEmpty:
|
||||
"""空文本测试"""
|
||||
|
||||
def test_empty_string(self):
|
||||
"""空字符串返回空列表."""
|
||||
assert split_text("") == []
|
||||
|
||||
def test_whitespace_only(self):
|
||||
assert split_text(" \n\t ") == []
|
||||
def test_only_whitespace(self):
|
||||
"""纯空白文本返回空列表."""
|
||||
assert split_text(" \n \t ") == []
|
||||
|
||||
def test_short_text_single_segment(self):
|
||||
text = "你好世界。"
|
||||
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 = "你好,世界。"
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
assert result[0] == text
|
||||
|
||||
def test_exact_max_chars(self):
|
||||
def test_exactly_max_chars(self):
|
||||
"""刚好等于 max_chars 的文本返回一个段落."""
|
||||
text = "a" * 500
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
assert len(result[0]) == 500
|
||||
|
||||
def test_splits_on_sentence_boundary(self):
|
||||
# 两个长句子,各300字左右,超过50字阈值
|
||||
sent1 = "你" * 300 + "。"
|
||||
sent2 = "我" * 300 + "。"
|
||||
text = sent1 + sent2
|
||||
def test_one_under_max(self):
|
||||
"""max_chars-1 的文本返回一个段落."""
|
||||
text = "a" * 499
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 2
|
||||
assert result[0] == sent1
|
||||
assert result[1] == sent2
|
||||
assert len(result) == 1
|
||||
assert len(result[0]) == 499
|
||||
|
||||
def test_long_sentence_hard_cut(self):
|
||||
# 一个超长句子,没有句末标点,会被硬切
|
||||
text = "长" * 800
|
||||
result = split_text(text, max_chars=500)
|
||||
|
||||
class TestSplitTextSentenceBoundary:
|
||||
"""句子边界分段测试"""
|
||||
|
||||
def test_split_at_period(self):
|
||||
"""在句号处拆分."""
|
||||
text = "第一句。第二句。第三句。"
|
||||
# 每句5字符,max_chars=10,每次两句就接近10
|
||||
result = split_text(text, max_chars=10)
|
||||
assert len(result) >= 2
|
||||
assert all(len(seg) <= 500 for seg in result)
|
||||
# 合起来应该等于原文本
|
||||
assert "".join(result) == text
|
||||
# 所有段落都不超过 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:
|
||||
"""短段落合并测试"""
|
||||
|
||||
def test_short_segments_merged(self):
|
||||
# 多个短句应该被合并
|
||||
sentences = [f"第{i}句。" for i in range(10)]
|
||||
text = "".join(sentences)
|
||||
result = split_text(text, max_chars=200)
|
||||
# 每句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
|
||||
|
||||
def test_multiple_punctuation_types(self):
|
||||
# 构造足够长的文本触发分段
|
||||
text = "第一" * 30 + "。" + "第二" * 30 + "!" + "第三" * 30 + "?" + "第四" * 30 + ";"
|
||||
"""多个短段合并为一个."""
|
||||
# 生成5个短句,每句5字符,max_chars=100,应该合并成一段
|
||||
text = "一。二。三。四。五。"
|
||||
result = split_text(text, max_chars=100)
|
||||
assert len(result) >= 2
|
||||
assert "".join(result) == text
|
||||
assert len(result) == 1
|
||||
assert len(result[0]) <= 100
|
||||
|
||||
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_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_newline_as_sentence_end(self):
|
||||
text = "第一段\n第二段\n第三段"
|
||||
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) >= 1
|
||||
assert "".join(result) == text.strip()
|
||||
assert len(result) >= 2
|
||||
for seg in result:
|
||||
assert len(seg) <= 50
|
||||
# 重新拼回应该等于原文本(除了可能的空格处理)
|
||||
combined = "".join(result)
|
||||
assert combined == text.replace(" ", "") # strip 不影响中文字符
|
||||
|
||||
def test_minimum_segment_length(self):
|
||||
# 句子太短(<50字)不会立即分段
|
||||
text = "短句一。短句二。短句三。"
|
||||
def test_chinese_long_paragraph(self):
|
||||
"""长中文段落."""
|
||||
text = "这是一个测试句子。" * 100 # 100个句子
|
||||
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)
|
||||
|
||||
|
||||
class TestSplitTextCustomMaxChars:
|
||||
"""自定义 max_chars 测试"""
|
||||
|
||||
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_large_max_chars(self):
|
||||
"""很大的 max_chars(不拆分)."""
|
||||
text = "这是一段测试文本。" * 10
|
||||
result = split_text(text, max_chars=10000)
|
||||
assert len(result) == 1
|
||||
|
||||
def test_trailing_content_added(self):
|
||||
# 最后一段不完整的句子也要加上
|
||||
text = "完整的句子。剩余内容"
|
||||
result = split_text(text, max_chars=50)
|
||||
assert "".join(result) == text
|
||||
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_no_empty_segments(self):
|
||||
text = "。。。。。" # 全是标点
|
||||
result = split_text(text, max_chars=2)
|
||||
assert all(len(seg) > 0 for seg in result)
|
||||
"""没有空字符串段落."""
|
||||
text = "句子。。。双标点。"
|
||||
result = split_text(text, max_chars=10)
|
||||
assert all(seg for seg in result) # 所有段非空
|
||||
|
||||
def test_chinese_and_english_mixed(self):
|
||||
text = "Hello世界。这是测试Test文本。Mixed混合。"
|
||||
result = split_text(text, max_chars=20)
|
||||
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 "".join(result) == text
|
||||
assert all(len(s) <= 10 for s in result)
|
||||
|
||||
def test_consecutive_punctuation(self):
|
||||
"""连续标点."""
|
||||
text = "真的吗!?不对。。。好吧。"
|
||||
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
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""标题库 UseCase 单元测试."""
|
||||
"""标题库 Use Cases 单元测试 — wave217"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -24,384 +24,447 @@ 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(id: str, name: str, text: str, usage_count: int = 0, category: str = "default") -> TitleLibraryItem:
|
||||
|
||||
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,
|
||||
):
|
||||
return TitleLibraryItem(
|
||||
id=id,
|
||||
user_id="user_1",
|
||||
id=item_id,
|
||||
user_id=user_id,
|
||||
name=name,
|
||||
text=text,
|
||||
category=category,
|
||||
description="",
|
||||
tags=[],
|
||||
description=description,
|
||||
tags=tags or [],
|
||||
usage_count=usage_count,
|
||||
is_active=True,
|
||||
metadata_={},
|
||||
is_active=is_active,
|
||||
metadata_=metadata_ or {},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_item():
|
||||
return _make_item("title_1", "爆款标题", "这是一个爆款标题文案", usage_count=5)
|
||||
# ── ListTitleLibraryUseCase ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestListTitleLibraryUseCase:
|
||||
"""ListTitleLibraryUseCase 测试"""
|
||||
def test_list_default_params(self):
|
||||
items = [_make_item("t1"), _make_item("t2")]
|
||||
repo = MagicMock()
|
||||
repo.list_by_user.return_value = items
|
||||
|
||||
def test_list_returns_results(self, mock_repo, sample_item):
|
||||
"""正常返回标题列表"""
|
||||
mock_repo.list_by_user.return_value = [sample_item]
|
||||
use_case = ListTitleLibraryUseCase(mock_repo)
|
||||
uc = ListTitleLibraryUseCase(repo)
|
||||
result = uc.execute("u1")
|
||||
|
||||
result = use_case.execute("user_1")
|
||||
assert len(result) == 2
|
||||
repo.list_by_user.assert_called_once_with("u1", category=None, skip=0, limit=50)
|
||||
|
||||
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)
|
||||
def test_list_with_category(self):
|
||||
repo = MagicMock()
|
||||
repo.list_by_user.return_value = []
|
||||
|
||||
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")
|
||||
uc = ListTitleLibraryUseCase(repo)
|
||||
result = uc.execute("u1", category="marketing")
|
||||
|
||||
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:
|
||||
"""GetTitleLibraryUseCase 测试"""
|
||||
def test_get_found(self):
|
||||
item = _make_item()
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = item
|
||||
|
||||
def test_get_existing(self, mock_repo, sample_item):
|
||||
"""获取存在的标题"""
|
||||
mock_repo.get.return_value = sample_item
|
||||
use_case = GetTitleLibraryUseCase(mock_repo)
|
||||
uc = GetTitleLibraryUseCase(repo)
|
||||
result = uc.execute("t1", "u1")
|
||||
assert result.id == "t1"
|
||||
repo.get.assert_called_once_with("t1", "u1")
|
||||
|
||||
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")
|
||||
def test_get_not_found(self):
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = None
|
||||
|
||||
uc = GetTitleLibraryUseCase(repo)
|
||||
result = uc.execute("t999", "u1")
|
||||
assert result is None
|
||||
|
||||
|
||||
# ── CreateTitleLibraryUseCase ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateTitleLibraryUseCase:
|
||||
"""CreateTitleLibraryUseCase 测试"""
|
||||
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
|
||||
|
||||
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)
|
||||
uc = CreateTitleLibraryUseCase(repo)
|
||||
cmd = CreateTitleLibraryCommand(user_id="u1", name="好标题", text="这是一个好标题的内容")
|
||||
|
||||
command = CreateTitleLibraryCommand(
|
||||
user_id="user_1",
|
||||
name="新标题",
|
||||
text="新标题文案",
|
||||
category="default",
|
||||
description="",
|
||||
tags=[],
|
||||
metadata_={},
|
||||
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"},
|
||||
)
|
||||
result = use_case.execute(command, plan_name="free")
|
||||
|
||||
assert result.id == "title_1"
|
||||
mock_repo.count_by_user.assert_called_once_with("user_1")
|
||||
mock_repo.create.assert_called_once()
|
||||
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
|
||||
|
||||
def test_create_quota_exceeded(self, mock_repo):
|
||||
"""超过配额时抛出 QuotaExceededError"""
|
||||
mock_repo.count_by_user.return_value = 9999
|
||||
use_case = CreateTitleLibraryUseCase(mock_repo)
|
||||
result = uc.execute(cmd)
|
||||
|
||||
command = CreateTitleLibraryCommand(
|
||||
user_id="user_1",
|
||||
name="新标题",
|
||||
text="文案",
|
||||
category="default",
|
||||
description="",
|
||||
tags=[],
|
||||
metadata_={},
|
||||
)
|
||||
with pytest.raises(QuotaExceededError):
|
||||
use_case.execute(command, plan_name="free")
|
||||
assert result.category == "marketing"
|
||||
assert result.description == "描述"
|
||||
assert result.tags == ["tag1", "tag2"]
|
||||
assert result.metadata_ == {"key": "value"}
|
||||
|
||||
mock_repo.create.assert_not_called()
|
||||
|
||||
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"}
|
||||
# ── UpdateTitleLibraryUseCase ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestUpdateTitleLibraryUseCase:
|
||||
"""UpdateTitleLibraryUseCase 测试"""
|
||||
def test_update_name(self):
|
||||
existing = _make_item(name="old")
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = existing
|
||||
repo.update.return_value = existing
|
||||
|
||||
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)
|
||||
uc = UpdateTitleLibraryUseCase(repo)
|
||||
cmd = UpdateTitleLibraryCommand(title_id="t1", user_id="u1", name="new")
|
||||
result = uc.execute(cmd)
|
||||
|
||||
command = UpdateTitleLibraryCommand(title_id="title_1", user_id="user_1", name="新名称")
|
||||
result = use_case.execute(command)
|
||||
assert result.name == "new"
|
||||
repo.update.assert_called_once()
|
||||
|
||||
assert result.name == "新名称"
|
||||
# 其他字段不变
|
||||
assert result.text == "这是一个爆款标题文案"
|
||||
mock_repo.get.assert_called_once_with("title_1", "user_1")
|
||||
mock_repo.update.assert_called_once()
|
||||
def test_update_text(self):
|
||||
existing = _make_item(text="old")
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = existing
|
||||
repo.update.return_value = existing
|
||||
|
||||
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)
|
||||
uc = UpdateTitleLibraryUseCase(repo)
|
||||
cmd = UpdateTitleLibraryCommand(title_id="t1", user_id="u1", text="new text")
|
||||
result = uc.execute(cmd)
|
||||
assert result.text == "new text"
|
||||
|
||||
command = UpdateTitleLibraryCommand(
|
||||
title_id="title_1",
|
||||
user_id="user_1",
|
||||
text="新文案内容",
|
||||
category="美食",
|
||||
is_active=False,
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
def test_update_category(self):
|
||||
existing = _make_item(category="old")
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = existing
|
||||
repo.update.return_value = existing
|
||||
|
||||
assert result.text == "新文案内容"
|
||||
assert result.category == "美食"
|
||||
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.is_active is False
|
||||
|
||||
def test_update_nonexistent_raises(self, mock_repo):
|
||||
"""更新不存在的标题抛出 NotFoundError"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = UpdateTitleLibraryUseCase(mock_repo)
|
||||
def test_update_not_found_raises(self):
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = None
|
||||
|
||||
command = UpdateTitleLibraryCommand(title_id="noexist", user_id="user_1", name="新名称")
|
||||
with pytest.raises(NotFoundError, match="not found"):
|
||||
use_case.execute(command)
|
||||
uc = UpdateTitleLibraryUseCase(repo)
|
||||
cmd = UpdateTitleLibraryCommand(title_id="t999", user_id="u1", name="x")
|
||||
with pytest.raises(NotFoundError):
|
||||
uc.execute(cmd)
|
||||
|
||||
mock_repo.update.assert_not_called()
|
||||
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 ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestDeleteTitleLibraryUseCase:
|
||||
"""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")
|
||||
def test_delete_success(self):
|
||||
repo = MagicMock()
|
||||
repo.delete.return_value = True
|
||||
|
||||
uc = DeleteTitleLibraryUseCase(repo)
|
||||
result = uc.execute("t1", "u1")
|
||||
assert result is True
|
||||
mock_repo.delete.assert_called_once_with("title_1", "user_1")
|
||||
repo.delete.assert_called_once_with("t1", "u1")
|
||||
|
||||
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")
|
||||
def test_delete_not_found(self):
|
||||
repo = MagicMock()
|
||||
repo.delete.return_value = False
|
||||
|
||||
uc = DeleteTitleLibraryUseCase(repo)
|
||||
result = uc.execute("t999", "u1")
|
||||
assert result is False
|
||||
|
||||
|
||||
# ── IncrementTitleUsageUseCase ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestIncrementTitleUsageUseCase:
|
||||
"""IncrementTitleUsageUseCase 测试"""
|
||||
def test_increment_default_1(self):
|
||||
repo = MagicMock()
|
||||
repo.increment_usage_count.return_value = True
|
||||
|
||||
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)
|
||||
uc = IncrementTitleUsageUseCase(repo)
|
||||
cmd = IncrementTitleUsageCommand(title_id="t1", user_id="u1")
|
||||
result = uc.execute(cmd)
|
||||
|
||||
assert result is True
|
||||
mock_repo.increment_usage_count.assert_called_once_with("title_1", "user_1", increment=1)
|
||||
repo.increment_usage_count.assert_called_once_with("t1", "u1", increment=1)
|
||||
|
||||
def test_increment_zero_returns_false(self, mock_repo):
|
||||
"""增量为0返回False,不调用repository"""
|
||||
use_case = IncrementTitleUsageUseCase(mock_repo)
|
||||
def test_increment_custom_amount(self):
|
||||
repo = MagicMock()
|
||||
repo.increment_usage_count.return_value = True
|
||||
|
||||
command = IncrementTitleUsageCommand(title_id="title_1", user_id="user_1", increment=0)
|
||||
result = use_case.execute(command)
|
||||
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)
|
||||
|
||||
assert result is False
|
||||
mock_repo.increment_usage_count.assert_not_called()
|
||||
repo.increment_usage_count.assert_not_called()
|
||||
|
||||
def test_increment_negative_returns_false(self, mock_repo):
|
||||
"""负增量返回False"""
|
||||
use_case = IncrementTitleUsageUseCase(mock_repo)
|
||||
def test_increment_negative_returns_false(self):
|
||||
repo = MagicMock()
|
||||
|
||||
command = IncrementTitleUsageCommand(title_id="title_1", user_id="user_1", increment=-1)
|
||||
result = use_case.execute(command)
|
||||
uc = IncrementTitleUsageUseCase(repo)
|
||||
cmd = IncrementTitleUsageCommand(title_id="t1", user_id="u1", increment=-1)
|
||||
result = uc.execute(cmd)
|
||||
|
||||
assert result is False
|
||||
mock_repo.increment_usage_count.assert_not_called()
|
||||
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)
|
||||
|
||||
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)
|
||||
# ── PickTitleUseCase ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPickTitleUseCase:
|
||||
"""PickTitleUseCase 智能选标题测试"""
|
||||
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
|
||||
|
||||
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)
|
||||
uc = PickTitleUseCase(repo)
|
||||
cmd = PickTitleCommand(user_id="u1")
|
||||
|
||||
command = PickTitleCommand(user_id="user_1")
|
||||
result = use_case.execute(command)
|
||||
# 由于有随机性,多次验证都在候选池(最少使用的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
|
||||
|
||||
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()
|
||||
# 验证查询参数
|
||||
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
|
||||
|
||||
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)
|
||||
def test_pick_empty_returns_none(self):
|
||||
repo = MagicMock()
|
||||
repo.list_by_user.return_value = []
|
||||
|
||||
uc = PickTitleUseCase(repo)
|
||||
cmd = PickTitleCommand(user_id="u1")
|
||||
result = uc.execute(cmd)
|
||||
assert result is None
|
||||
|
||||
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)
|
||||
def test_pick_single_item(self):
|
||||
item = _make_item("t1")
|
||||
repo = MagicMock()
|
||||
repo.list_by_user.return_value = [item]
|
||||
|
||||
command = PickTitleCommand(user_id="user_1", category="美食")
|
||||
result = use_case.execute(command)
|
||||
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)
|
||||
|
||||
assert result is not None
|
||||
call_kwargs = mock_repo.list_by_user.call_args[1]
|
||||
assert call_kwargs["category"] == "美食"
|
||||
assert call_kwargs["is_active"] is True
|
||||
repo.list_by_user.assert_called_once()
|
||||
assert repo.list_by_user.call_args[1]["category"] == "marketing"
|
||||
|
||||
def test_pick_exclude_ids(self, mock_repo):
|
||||
"""排除指定ID"""
|
||||
def test_pick_exclude_ids(self):
|
||||
items = [
|
||||
_make_item("t1", "标题1", "文案1", usage_count=1),
|
||||
_make_item("t2", "标题2", "文案2", usage_count=2),
|
||||
_make_item("t3", "标题3", "文案3", usage_count=3),
|
||||
_make_item("t1", usage_count=1),
|
||||
_make_item("t2", usage_count=2),
|
||||
_make_item("t3", usage_count=3),
|
||||
]
|
||||
mock_repo.list_by_user.return_value = items
|
||||
use_case = PickTitleUseCase(mock_repo)
|
||||
repo = MagicMock()
|
||||
repo.list_by_user.return_value = items
|
||||
|
||||
command = PickTitleCommand(user_id="user_1", exclude_ids=["t1", "t2"])
|
||||
result = use_case.execute(command)
|
||||
uc = PickTitleUseCase(repo)
|
||||
cmd = PickTitleCommand(user_id="u1", exclude_ids=["t1", "t2"])
|
||||
|
||||
# 排除两个后只剩t3
|
||||
# 排除 t1, t2 后只剩 t3
|
||||
result = uc.execute(cmd)
|
||||
assert result.id == "t3"
|
||||
|
||||
def test_pick_exclude_all_falls_back(self, mock_repo):
|
||||
"""排除全部时从所有标题中选"""
|
||||
def test_pick_exclude_all_fallback_to_all(self):
|
||||
items = [
|
||||
_make_item("t1", "标题1", "文案1", usage_count=1),
|
||||
_make_item("t2", "标题2", "文案2", usage_count=2),
|
||||
_make_item("t1", usage_count=1),
|
||||
_make_item("t2", usage_count=2),
|
||||
]
|
||||
mock_repo.list_by_user.return_value = items
|
||||
use_case = PickTitleUseCase(mock_repo)
|
||||
repo = MagicMock()
|
||||
repo.list_by_user.return_value = items
|
||||
|
||||
command = PickTitleCommand(user_id="user_1", exclude_ids=["t1", "t2"])
|
||||
result = use_case.execute(command)
|
||||
uc = PickTitleUseCase(repo)
|
||||
cmd = PickTitleCommand(user_id="u1", exclude_ids=["t1", "t2"])
|
||||
|
||||
# 排除全部后fallback到全部,所以还是能选出一个
|
||||
# 排除后没了,回退到从全部选
|
||||
result = uc.execute(cmd)
|
||||
assert result is not None
|
||||
assert result.id in ("t1", "t2")
|
||||
assert result.id in {"t1", "t2"}
|
||||
|
||||
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)
|
||||
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
|
||||
|
||||
command = PickTitleCommand(user_id="user_1")
|
||||
result = use_case.execute(command)
|
||||
uc = PickTitleUseCase(repo)
|
||||
cmd = PickTitleCommand(user_id="u1")
|
||||
|
||||
assert result.id == "only"
|
||||
|
||||
def test_pick_prefers_less_used(self, mock_repo):
|
||||
"""倾向于选择使用次数少的"""
|
||||
items = [
|
||||
_make_item("t_used", "常用", "常用", usage_count=100),
|
||||
_make_item("t_fresh", "新的", "新的", usage_count=0),
|
||||
]
|
||||
mock_repo.list_by_user.return_value = items
|
||||
use_case = PickTitleUseCase(mock_repo)
|
||||
|
||||
# 跑多次,验证使用少的出现在候选池里
|
||||
results = set()
|
||||
# 运行多次,确保选中的都在前5个使用最少的里(t1~t5, usage 1~5)
|
||||
for _ in range(20):
|
||||
command = PickTitleCommand(user_id="user_1")
|
||||
r = use_case.execute(command)
|
||||
if r:
|
||||
results.add(r.id)
|
||||
result = uc.execute(cmd)
|
||||
assert int(result.id[1:]) <= 5 # 只从前5个里选
|
||||
|
||||
# 两个都在候选池(少于5个),所以都可能被选中
|
||||
assert "t_used" in results or "t_fresh" in results
|
||||
def test_pick_fewer_than_pool_size(self):
|
||||
# 只有3个标题,不足5个池大小
|
||||
items = [
|
||||
_make_item("t1", usage_count=3),
|
||||
_make_item("t2", usage_count=1),
|
||||
_make_item("t3", usage_count=2),
|
||||
]
|
||||
repo = MagicMock()
|
||||
repo.list_by_user.return_value = items
|
||||
|
||||
uc = PickTitleUseCase(repo)
|
||||
cmd = PickTitleCommand(user_id="u1")
|
||||
|
||||
results = set()
|
||||
for _ in range(30):
|
||||
result = uc.execute(cmd)
|
||||
results.add(result.id)
|
||||
|
||||
# 3个都有可能被选中(随机性+少量样本,大概率至少出现2个)
|
||||
assert len(results) >= 1
|
||||
assert results.issubset({"t1", "t2", "t3"})
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
"""验证码服务单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
@@ -9,9 +8,9 @@ import pytest
|
||||
|
||||
from packages.application.auth.verification_code_service import (
|
||||
CODE_TYPE_EMAIL_BIND,
|
||||
CODE_TYPE_EMAIL_LOGIN,
|
||||
CODE_TYPE_PHONE_BIND,
|
||||
DAILY_LIMIT,
|
||||
DEFAULT_TTL_SECONDS,
|
||||
MAX_ATTEMPTS,
|
||||
RESEND_COOLDOWN_SECONDS,
|
||||
VerificationCodeService,
|
||||
@@ -21,298 +20,560 @@ from packages.application.auth.verification_code_service import (
|
||||
)
|
||||
from packages.domain.verification_code import VerificationCode
|
||||
|
||||
# ── Test Fixtures ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
return MagicMock()
|
||||
"""mock 验证码仓储."""
|
||||
repo = MagicMock()
|
||||
repo.find_latest.return_value = None
|
||||
repo.count_today.return_value = 0
|
||||
return repo
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def code_service(mock_repo):
|
||||
return VerificationCodeService(mock_repo)
|
||||
def service(mock_repo):
|
||||
"""验证码服务实例."""
|
||||
return VerificationCodeService(repo=mock_repo)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_code():
|
||||
code = VerificationCode.create(
|
||||
recipient="test@example.com",
|
||||
code_type=CODE_TYPE_EMAIL_BIND,
|
||||
ttl_seconds=300,
|
||||
def _make_code(
|
||||
recipient="test@example.com",
|
||||
code_type=CODE_TYPE_EMAIL_BIND,
|
||||
code="123456",
|
||||
ttl=300,
|
||||
used=False,
|
||||
attempts=0,
|
||||
created_at=None,
|
||||
):
|
||||
"""创建一个测试用验证码实体."""
|
||||
now = created_at or datetime.now(timezone.utc)
|
||||
vc = VerificationCode(
|
||||
id="test-code-id",
|
||||
recipient=recipient,
|
||||
code=code,
|
||||
code_type=code_type,
|
||||
expires_at=now + timedelta(seconds=ttl),
|
||||
used_at=now if used else None,
|
||||
attempts=attempts,
|
||||
created_at=now,
|
||||
)
|
||||
return code
|
||||
return vc
|
||||
|
||||
|
||||
class TestVerificationCodeServiceGenerate:
|
||||
# ── generate 方法测试 ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGenerate:
|
||||
"""generate 方法测试"""
|
||||
|
||||
def test_generate_success(self, code_service, mock_repo, sample_code):
|
||||
"""生成验证码成功"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
mock_repo.save.return_value = None
|
||||
|
||||
code, error = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
def test_generate_success(self, service, mock_repo):
|
||||
"""成功生成验证码."""
|
||||
code, error = service.generate("user@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert error is None
|
||||
assert code is not None
|
||||
assert code.recipient == "test@example.com"
|
||||
assert code.recipient == "user@example.com"
|
||||
assert code.code_type == CODE_TYPE_EMAIL_BIND
|
||||
assert len(code.code) == 6
|
||||
assert code.code.isdigit()
|
||||
assert not code.is_used
|
||||
mock_repo.save.assert_called_once()
|
||||
|
||||
def test_generate_empty_recipient(self, code_service):
|
||||
"""空接收方返回错误"""
|
||||
code, error = code_service.generate("", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "接收方不能为空" in error
|
||||
def test_generate_with_custom_code(self, service, mock_repo):
|
||||
"""使用自定义验证码."""
|
||||
code, error = service.generate("user@example.com", CODE_TYPE_EMAIL_LOGIN, custom_code="999999")
|
||||
|
||||
def test_generate_invalid_type(self, code_service):
|
||||
"""无效验证码类型返回错误"""
|
||||
code, error = code_service.generate("test@example.com", "invalid_type")
|
||||
assert error is None
|
||||
assert code.code == "999999"
|
||||
|
||||
def test_generate_custom_ttl(self, service, mock_repo):
|
||||
"""自定义 TTL."""
|
||||
code, _ = service.generate("u@e.com", CODE_TYPE_EMAIL_BIND, ttl_seconds=600)
|
||||
delta = code.expires_at - code.created_at
|
||||
assert delta.total_seconds() == 600
|
||||
|
||||
def test_generate_default_ttl(self, service, mock_repo):
|
||||
"""默认 TTL."""
|
||||
code, _ = service.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
delta = code.expires_at - code.created_at
|
||||
assert delta.total_seconds() == 300 # 默认5分钟
|
||||
|
||||
def test_generate_empty_recipient(self, service):
|
||||
"""空接收方."""
|
||||
code, error = service.generate("", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "不能为空" in error
|
||||
|
||||
def test_generate_whitespace_recipient(self, service):
|
||||
"""全空白接收方."""
|
||||
code, error = service.generate(" ", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "不能为空" in error
|
||||
|
||||
def test_generate_invalid_type(self, service):
|
||||
"""无效验证码类型."""
|
||||
code, error = service.generate("u@e.com", "invalid_type")
|
||||
assert code is None
|
||||
assert "无效的验证码类型" in error
|
||||
|
||||
def test_generate_cooldown(self, code_service, mock_repo, sample_code):
|
||||
"""冷却期内返回频控错误"""
|
||||
# 最新的验证码刚创建10秒前
|
||||
sample_code.created_at = datetime.now(timezone.utc) - timedelta(seconds=10)
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
mock_repo.count_today.return_value = 1
|
||||
def test_generate_recipient_stripped(self, service, mock_repo):
|
||||
"""接收方前后空格会被清理."""
|
||||
code, _ = service.generate(" user@e.com ", CODE_TYPE_EMAIL_BIND)
|
||||
assert code.recipient == "user@e.com"
|
||||
|
||||
code, error = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
def test_generate_phone_code(self, service, mock_repo):
|
||||
"""手机验证码生成."""
|
||||
code, error = service.generate("13800138000", CODE_TYPE_PHONE_BIND)
|
||||
assert error is None
|
||||
assert code.code_type == CODE_TYPE_PHONE_BIND
|
||||
assert len(code.code) == 6
|
||||
|
||||
|
||||
# ── generate 频控测试 ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGenerateRateLimit:
|
||||
"""generate 频控测试"""
|
||||
|
||||
def test_cooldown_active_rejects(self, service, mock_repo):
|
||||
"""冷却期内拒绝重发."""
|
||||
recent = _make_code(created_at=datetime.now(timezone.utc) - timedelta(seconds=10))
|
||||
mock_repo.find_latest.return_value = recent
|
||||
|
||||
code, error = service.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "发送太频繁" in error
|
||||
assert "秒后再试" in error
|
||||
# 等待时间应该接近 50 秒 (60-10)
|
||||
match = re.search(r"(\d+)\s*秒", error)
|
||||
assert match
|
||||
wait = int(match.group(1))
|
||||
assert 45 <= wait <= 55
|
||||
|
||||
def test_generate_daily_limit_exceeded(self, code_service, mock_repo):
|
||||
"""超过每日上限返回错误"""
|
||||
mock_repo.find_latest.return_value = None # 没有冷却期问题
|
||||
def test_cooldown_expired_allows(self, service, mock_repo):
|
||||
"""冷却期过后允许重发."""
|
||||
old = _make_code(created_at=datetime.now(timezone.utc) - timedelta(seconds=120))
|
||||
mock_repo.find_latest.return_value = old
|
||||
mock_repo.count_today.return_value = 1
|
||||
|
||||
code, error = service.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert error is None
|
||||
assert code is not None
|
||||
|
||||
def test_daily_limit_reached(self, service, mock_repo):
|
||||
"""达到每日上限."""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = DAILY_LIMIT
|
||||
|
||||
code, error = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
code, error = service.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "今日发送次数已达上限" in error
|
||||
|
||||
def test_generate_recipient_stripped(self, code_service, mock_repo, sample_code):
|
||||
"""recipient 会被 strip"""
|
||||
def test_daily_limit_one_below_allows(self, service, mock_repo):
|
||||
"""未达到上限时允许."""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = DAILY_LIMIT - 1
|
||||
|
||||
code, error = service.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert error is None
|
||||
assert code is not None
|
||||
|
||||
def test_custom_daily_limit(self, mock_repo):
|
||||
"""自定义每日上限."""
|
||||
svc = VerificationCodeService(repo=mock_repo, daily_limit=3)
|
||||
mock_repo.count_today.return_value = 3
|
||||
|
||||
code, error = svc.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "已达上限" in error
|
||||
|
||||
def test_custom_cooldown(self, mock_repo):
|
||||
"""自定义冷却时间."""
|
||||
svc = VerificationCodeService(repo=mock_repo, resend_cooldown=30)
|
||||
recent = _make_code(created_at=datetime.now(timezone.utc) - timedelta(seconds=10))
|
||||
mock_repo.find_latest.return_value = recent
|
||||
|
||||
code, error = svc.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
match = re.search(r"(\d+)\s*秒", error)
|
||||
assert match
|
||||
wait = int(match.group(1))
|
||||
assert 15 <= wait <= 25
|
||||
|
||||
def test_cooldown_different_types_independent(self, service, mock_repo):
|
||||
"""不同类型的验证码冷却独立."""
|
||||
# email_bind 类型有一个近期验证码
|
||||
recent = _make_code(code_type=CODE_TYPE_EMAIL_BIND)
|
||||
mock_repo.find_latest.side_effect = lambda r, t: recent if t == CODE_TYPE_EMAIL_BIND else None
|
||||
mock_repo.count_today.return_value = 0
|
||||
mock_repo.save.return_value = None
|
||||
|
||||
code_service.generate(" test@example.com ", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
# 传给 repo 的应该是 strip 后的值
|
||||
save_call = mock_repo.save.call_args[0][0]
|
||||
assert save_call.recipient == "test@example.com"
|
||||
|
||||
def test_generate_custom_code(self, code_service, mock_repo):
|
||||
"""使用自定义验证码"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
mock_repo.save.return_value = None
|
||||
|
||||
code, _ = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND, custom_code="123456")
|
||||
assert code.code == "123456"
|
||||
|
||||
def test_generate_custom_ttl(self, code_service, mock_repo):
|
||||
"""自定义 TTL"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
mock_repo.save.return_value = None
|
||||
|
||||
code, _ = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND, ttl_seconds=600)
|
||||
# email_login 类型应该可以正常发送
|
||||
code, error = service.generate("u@e.com", CODE_TYPE_EMAIL_LOGIN)
|
||||
assert error is None
|
||||
assert code is not None
|
||||
|
||||
|
||||
class TestVerificationCodeServiceVerify:
|
||||
# ── verify 方法测试 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestVerify:
|
||||
"""verify 方法测试"""
|
||||
|
||||
def test_verify_success(self, code_service, mock_repo, sample_code):
|
||||
"""验证成功"""
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
def test_verify_success(self, service, mock_repo):
|
||||
"""验证码正确."""
|
||||
code = _make_code(code="654321")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, sample_code.code)
|
||||
|
||||
assert success is True
|
||||
ok, error = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "654321")
|
||||
assert ok is True
|
||||
assert error is None
|
||||
assert sample_code.is_used is True
|
||||
assert code.is_used # 标记为已使用
|
||||
assert mock_repo.save.call_count >= 2 # increment + mark_used
|
||||
|
||||
def test_verify_wrong_code(self, code_service, mock_repo, sample_code):
|
||||
"""验证码错误"""
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
def test_verify_wrong_code(self, service, mock_repo):
|
||||
"""验证码错误."""
|
||||
code = _make_code(code="123456")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "wrongcode")
|
||||
|
||||
assert success is False
|
||||
ok, error = service.verify("test@e.com", CODE_TYPE_EMAIL_BIND, "000000")
|
||||
assert ok is False
|
||||
assert "验证码错误" in error
|
||||
assert not code.is_used # 不标记为已使用
|
||||
assert code.attempts == 1 # 尝试次数+1
|
||||
|
||||
def test_verify_not_found(self, code_service, mock_repo):
|
||||
"""验证码不存在"""
|
||||
def test_verify_no_code_found(self, service, mock_repo):
|
||||
"""找不到验证码."""
|
||||
mock_repo.find_latest.return_value = None
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
|
||||
assert success is False
|
||||
ok, error = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
assert ok is False
|
||||
assert "不存在或已过期" in error
|
||||
|
||||
def test_verify_expired(self, code_service, mock_repo):
|
||||
"""验证码已过期"""
|
||||
expired_code = VerificationCode.create(
|
||||
recipient="test@example.com",
|
||||
code_type=CODE_TYPE_EMAIL_BIND,
|
||||
ttl_seconds=1, # 1秒过期
|
||||
)
|
||||
# 手动设置过期时间
|
||||
expired_code.expires_at = datetime.now(timezone.utc) - timedelta(seconds=10)
|
||||
mock_repo.find_latest.return_value = expired_code
|
||||
def test_verify_empty_params(self, service):
|
||||
"""参数为空."""
|
||||
ok, error = service.verify("", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
assert ok is False
|
||||
assert "参数不完整" in error
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, expired_code.code)
|
||||
ok2, error2 = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, "")
|
||||
assert ok2 is False
|
||||
assert "参数不完整" in error2
|
||||
|
||||
assert success is False
|
||||
assert "已过期" in error
|
||||
def test_verify_whitespace_params(self, service, mock_repo):
|
||||
"""参数前后空格会被清理."""
|
||||
code = _make_code(recipient="u@e.com", code="111111")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
def test_verify_already_used(self, code_service, mock_repo, sample_code):
|
||||
"""验证码已使用"""
|
||||
sample_code.mark_used()
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
ok, error = service.verify(" u@e.com ", CODE_TYPE_EMAIL_BIND, " 111111 ")
|
||||
assert ok is True
|
||||
assert error is None
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, sample_code.code)
|
||||
def test_verify_already_used(self, service, mock_repo):
|
||||
"""验证码已使用."""
|
||||
code = _make_code(used=True)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
assert success is False
|
||||
ok, error = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, code.code)
|
||||
assert ok is False
|
||||
assert "已使用" in error
|
||||
|
||||
def test_verify_max_attempts_exceeded(self, code_service, mock_repo, sample_code):
|
||||
"""尝试次数过多"""
|
||||
# 先把尝试次数加到超过上限
|
||||
for _ in range(MAX_ATTEMPTS + 1):
|
||||
sample_code.increment_attempts()
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
def test_verify_expired(self, service, mock_repo):
|
||||
"""验证码已过期."""
|
||||
code = _make_code(ttl=-60) # 已过期
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, sample_code.code)
|
||||
ok, error = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, code.code)
|
||||
assert ok is False
|
||||
assert "已过期" in error
|
||||
|
||||
assert success is False
|
||||
def test_verify_too_many_attempts(self, service, mock_repo):
|
||||
"""尝试次数过多."""
|
||||
code = _make_code(attempts=MAX_ATTEMPTS + 1)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
ok, error = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, code.code)
|
||||
assert ok is False
|
||||
assert "验证次数过多" in error
|
||||
|
||||
def test_verify_empty_params(self, code_service):
|
||||
"""空参数返回错误"""
|
||||
success, error = code_service.verify("", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
assert success is False
|
||||
assert "参数不完整" in error
|
||||
def test_verify_attempts_increment_each_time(self, service, mock_repo):
|
||||
"""每次错误尝试都增加尝试次数."""
|
||||
code = _make_code(code="123456", attempts=0)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "")
|
||||
assert success is False
|
||||
assert "参数不完整" in error
|
||||
for _ in range(3):
|
||||
service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, "wrong")
|
||||
|
||||
def test_verify_increments_attempts(self, code_service, mock_repo, sample_code):
|
||||
"""验证会增加尝试次数"""
|
||||
initial_attempts = sample_code.attempts
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
assert code.attempts == 3
|
||||
|
||||
code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "wrong")
|
||||
def test_verify_without_consume(self, service, mock_repo):
|
||||
"""验证成功但不标记为已使用(consume=False)."""
|
||||
code = _make_code(code="999999")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
assert sample_code.attempts == initial_attempts + 1
|
||||
|
||||
def test_verify_no_consume(self, code_service, mock_repo, sample_code):
|
||||
"""consume=False 时不标记为已使用"""
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
|
||||
success, _ = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, sample_code.code, consume=False)
|
||||
|
||||
assert success is True
|
||||
assert sample_code.is_used is False
|
||||
|
||||
|
||||
class TestVerifyPhone:
|
||||
"""validate_phone 函数测试"""
|
||||
|
||||
def test_valid_phone(self):
|
||||
"""有效手机号"""
|
||||
ok, err = validate_phone("13800000001")
|
||||
ok, error = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, "999999", consume=False)
|
||||
assert ok is True
|
||||
assert err == ""
|
||||
assert error is None
|
||||
assert not code.is_used # 不标记为已使用
|
||||
|
||||
def test_valid_phone_with_plus86(self):
|
||||
"""带 +86 前缀的手机号"""
|
||||
ok, err = validate_phone("+8613800000001")
|
||||
def test_verify_consume_default_true(self, service, mock_repo):
|
||||
"""默认 consume=True."""
|
||||
code = _make_code(code="123456")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
assert code.is_used
|
||||
|
||||
def test_verify_used_checked_before_attempts(self, service, mock_repo):
|
||||
"""已使用优先于其他检查."""
|
||||
code = _make_code(used=True, attempts=0)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
ok, error = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, code.code)
|
||||
assert ok is False
|
||||
assert "已使用" in error
|
||||
# attempts 会被 increment,但错误原因是已使用
|
||||
assert code.attempts == 1
|
||||
|
||||
def test_custom_max_attempts(self, mock_repo):
|
||||
"""自定义最大尝试次数."""
|
||||
svc = VerificationCodeService(repo=mock_repo, max_attempts=2)
|
||||
code = _make_code(attempts=2)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
ok, error = svc.verify("u@e.com", CODE_TYPE_EMAIL_BIND, code.code)
|
||||
assert ok is False
|
||||
assert "验证次数过多" in error
|
||||
|
||||
|
||||
# ── validate_phone 测试 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestValidatePhone:
|
||||
"""手机号格式校验测试"""
|
||||
|
||||
def test_valid_11_digit(self):
|
||||
"""标准11位手机号."""
|
||||
ok, msg = validate_phone("13800138000")
|
||||
assert ok is True
|
||||
assert msg == ""
|
||||
|
||||
def test_valid_with_plus_86(self):
|
||||
"""带+86前缀."""
|
||||
ok, msg = validate_phone("+8613800138000")
|
||||
assert ok is True
|
||||
|
||||
def test_invalid_phone_short(self):
|
||||
"""太短的手机号"""
|
||||
ok, err = validate_phone("123")
|
||||
def test_invalid_too_short(self):
|
||||
"""位数不足."""
|
||||
ok, msg = validate_phone("1380013800")
|
||||
assert ok is False
|
||||
assert "格式不正确" in err
|
||||
assert "格式不正确" in msg
|
||||
|
||||
def test_invalid_phone_wrong_prefix(self):
|
||||
"""号段不对的手机号"""
|
||||
ok, err = validate_phone("11000000000")
|
||||
def test_invalid_too_long(self):
|
||||
"""位数过多."""
|
||||
ok, msg = validate_phone("138001380001")
|
||||
assert ok is False
|
||||
|
||||
def test_empty_phone(self):
|
||||
"""空手机号"""
|
||||
ok, err = validate_phone("")
|
||||
def test_invalid_starts_with_2(self):
|
||||
"""开头不是1."""
|
||||
ok, msg = validate_phone("23800138000")
|
||||
assert ok is False
|
||||
assert "不能为空" in err
|
||||
|
||||
def test_phone_with_spaces(self):
|
||||
"""带空格的手机号会被 strip"""
|
||||
ok, _ = validate_phone(" 13800000001 ")
|
||||
def test_invalid_starts_with_12(self):
|
||||
"""第二位不在3-9."""
|
||||
ok, msg = validate_phone("12800138000")
|
||||
assert ok is False
|
||||
|
||||
def test_invalid_empty(self):
|
||||
"""空字符串."""
|
||||
ok, msg = validate_phone("")
|
||||
assert ok is False
|
||||
assert "不能为空" in msg
|
||||
|
||||
def test_invalid_whitespace_only(self):
|
||||
"""仅空白."""
|
||||
ok, msg = validate_phone(" ")
|
||||
assert ok is False
|
||||
assert "不能为空" in msg
|
||||
|
||||
def test_valid_all_prefixes_3_to_9(self):
|
||||
"""第二位3-9都有效."""
|
||||
for n in range(3, 10):
|
||||
ok, _ = validate_phone(f"1{n}800138000")
|
||||
assert ok is True, f"1{n} prefix should be valid"
|
||||
|
||||
def test_invalid_contains_letters(self):
|
||||
"""包含字母."""
|
||||
ok, msg = validate_phone("13800abc000")
|
||||
assert ok is False
|
||||
|
||||
def test_strips_whitespace(self):
|
||||
"""前后空格会被清理."""
|
||||
ok, msg = validate_phone(" 13800138000 ")
|
||||
assert ok is True
|
||||
|
||||
|
||||
# ── normalize_phone 测试 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestNormalizePhone:
|
||||
"""normalize_phone 函数测试"""
|
||||
"""手机号标准化测试"""
|
||||
|
||||
def test_removes_plus86(self):
|
||||
"""去掉 +86 前缀"""
|
||||
assert normalize_phone("+8613800000001") == "13800000001"
|
||||
def test_strip_plus_86(self):
|
||||
"""去掉+86前缀."""
|
||||
assert normalize_phone("+8613800138000") == "13800138000"
|
||||
|
||||
def test_no_prefix_stays_same(self):
|
||||
"""没有前缀保持不变"""
|
||||
assert normalize_phone("13800000001") == "13800000001"
|
||||
"""无前缀保持不变."""
|
||||
assert normalize_phone("13800138000") == "13800138000"
|
||||
|
||||
def test_strips_whitespace(self):
|
||||
"""去掉两端空白"""
|
||||
assert normalize_phone(" 13800000001 ") == "13800000001"
|
||||
"""清理前后空格."""
|
||||
assert normalize_phone(" 13800138000 ") == "13800138000"
|
||||
|
||||
def test_plus_86_with_spaces(self):
|
||||
"""带空格的+86."""
|
||||
assert normalize_phone(" +8613800138000 ") == "13800138000"
|
||||
|
||||
|
||||
# ── validate_email 测试 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestValidateEmail:
|
||||
"""validate_email 函数测试"""
|
||||
"""邮箱格式校验测试"""
|
||||
|
||||
def test_valid_email(self):
|
||||
"""有效邮箱"""
|
||||
ok, err = validate_email("test@example.com")
|
||||
def test_valid_simple(self):
|
||||
"""标准邮箱."""
|
||||
ok, msg = validate_email("user@example.com")
|
||||
assert ok is True
|
||||
assert err == ""
|
||||
assert msg == ""
|
||||
|
||||
def test_valid_email_with_subdomain(self):
|
||||
"""带子域名的邮箱"""
|
||||
ok, _ = validate_email("user@mail.example.com")
|
||||
def test_valid_with_dots(self):
|
||||
"""带点号的用户名."""
|
||||
ok, _ = validate_email("user.name@example.com")
|
||||
assert ok is True
|
||||
|
||||
def test_valid_email_with_plus(self):
|
||||
"""带 + 号的邮箱"""
|
||||
def test_valid_with_plus(self):
|
||||
"""带加号的邮箱."""
|
||||
ok, _ = validate_email("user+tag@example.com")
|
||||
assert ok is True
|
||||
|
||||
def test_invalid_email_no_at(self):
|
||||
"""没有 @ 的邮箱"""
|
||||
ok, err = validate_email("notanemail")
|
||||
assert ok is False
|
||||
assert "格式不正确" in err
|
||||
|
||||
def test_invalid_email_no_domain(self):
|
||||
"""没有域名的邮箱"""
|
||||
ok, err = validate_email("user@")
|
||||
assert ok is False
|
||||
|
||||
def test_empty_email(self):
|
||||
"""空邮箱"""
|
||||
ok, err = validate_email("")
|
||||
assert ok is False
|
||||
assert "不能为空" in err
|
||||
|
||||
def test_email_with_spaces(self):
|
||||
"""带空格的邮箱会被 strip"""
|
||||
ok, _ = validate_email(" test@example.com ")
|
||||
def test_valid_with_underscore(self):
|
||||
"""带下划线."""
|
||||
ok, _ = validate_email("user_name@example.com")
|
||||
assert ok is True
|
||||
|
||||
def test_valid_subdomain(self):
|
||||
"""多级域名."""
|
||||
ok, _ = validate_email("user@mail.example.com")
|
||||
assert ok is True
|
||||
|
||||
def test_invalid_no_at(self):
|
||||
"""没有@."""
|
||||
ok, msg = validate_email("userexample.com")
|
||||
assert ok is False
|
||||
assert "格式不正确" in msg
|
||||
|
||||
def test_invalid_empty_local(self):
|
||||
"""@前为空."""
|
||||
ok, _ = validate_email("@example.com")
|
||||
assert ok is False
|
||||
|
||||
def test_invalid_empty_domain(self):
|
||||
"""@后为空."""
|
||||
ok, _ = validate_email("user@")
|
||||
assert ok is False
|
||||
|
||||
def test_invalid_no_tld(self):
|
||||
"""没有顶级域名."""
|
||||
ok, _ = validate_email("user@example")
|
||||
assert ok is False
|
||||
|
||||
def test_invalid_empty(self):
|
||||
"""空字符串."""
|
||||
ok, msg = validate_email("")
|
||||
assert ok is False
|
||||
assert "不能为空" in msg
|
||||
|
||||
def test_invalid_spaces_only(self):
|
||||
"""仅空白."""
|
||||
ok, msg = validate_email(" ")
|
||||
assert ok is False
|
||||
assert "不能为空" in msg
|
||||
|
||||
def test_strips_whitespace(self):
|
||||
"""前后空格会被清理."""
|
||||
ok, msg = validate_email(" user@e.com ")
|
||||
assert ok is True
|
||||
|
||||
def test_invalid_special_chars(self):
|
||||
"""特殊字符."""
|
||||
ok, _ = validate_email("user name@e.com")
|
||||
assert ok is False
|
||||
|
||||
def test_valid_numbers(self):
|
||||
"""数字邮箱."""
|
||||
ok, _ = validate_email("12345@example.com")
|
||||
assert ok is True
|
||||
|
||||
|
||||
# ── VerificationCode 实体辅助验证 ──────────────────────────────────────────
|
||||
|
||||
|
||||
class TestVerificationCodeEntity:
|
||||
"""VerificationCode 实体属性测试"""
|
||||
|
||||
def test_is_expired_false_when_fresh(self):
|
||||
code = _make_code(ttl=300)
|
||||
assert code.is_expired is False
|
||||
|
||||
def test_is_expired_true_when_past(self):
|
||||
code = _make_code(ttl=-1)
|
||||
assert code.is_expired is True
|
||||
|
||||
def test_is_used_false_initially(self):
|
||||
code = _make_code()
|
||||
assert code.is_used is False
|
||||
|
||||
def test_is_used_after_mark_used(self):
|
||||
code = _make_code()
|
||||
code.mark_used()
|
||||
assert code.is_used is True
|
||||
assert code.used_at is not None
|
||||
|
||||
def test_is_valid_fresh(self):
|
||||
code = _make_code()
|
||||
assert code.is_valid is True
|
||||
|
||||
def test_is_valid_when_expired(self):
|
||||
code = _make_code(ttl=-100)
|
||||
assert code.is_valid is False
|
||||
|
||||
def test_is_valid_when_used(self):
|
||||
code = _make_code(used=True)
|
||||
assert code.is_valid is False
|
||||
|
||||
def test_increment_attempts(self):
|
||||
code = _make_code(attempts=0)
|
||||
code.increment_attempts()
|
||||
assert code.attempts == 1
|
||||
code.increment_attempts()
|
||||
assert code.attempts == 2
|
||||
|
||||
def test_create_generates_6_digit_code(self):
|
||||
code = VerificationCode.create("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert len(code.code) == 6
|
||||
assert code.code.isdigit()
|
||||
|
||||
def test_create_custom_code(self):
|
||||
code = VerificationCode.create("u@e.com", CODE_TYPE_EMAIL_BIND, custom_code="555555")
|
||||
assert code.code == "555555"
|
||||
|
||||
def test_create_strips_recipient(self):
|
||||
code = VerificationCode.create(" u@e.com ", CODE_TYPE_EMAIL_BIND)
|
||||
assert code.recipient == "u@e.com"
|
||||
|
||||
def test_create_sets_expiry(self):
|
||||
code = VerificationCode.create("u@e.com", CODE_TYPE_EMAIL_BIND, ttl_seconds=120)
|
||||
delta = code.expires_at - code.created_at
|
||||
assert delta.total_seconds() == 120
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""视频分享 UseCase 单元测试."""
|
||||
"""视频分享 Use Cases 单元测试 — wave215"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -22,6 +22,7 @@ from packages.application.video_share.use_cases import (
|
||||
PasswordRequiredError,
|
||||
RecordShareDownloadUseCase,
|
||||
RevokeShareUseCase,
|
||||
ShareAccessResult,
|
||||
ShareExpiredError,
|
||||
UpdateShareUseCase,
|
||||
VideoNotFoundError,
|
||||
@@ -29,486 +30,459 @@ from packages.application.video_share.use_cases import (
|
||||
from packages.domain.generated_video import GeneratedVideo
|
||||
from packages.domain.video_share import VideoShare
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_share_repo():
|
||||
return MagicMock()
|
||||
# ── helpers ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@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():
|
||||
def _make_share(
|
||||
video_id="v1",
|
||||
user_id="u1",
|
||||
password=None,
|
||||
expires_at=None,
|
||||
is_active=True,
|
||||
view_count=0,
|
||||
download_count=0,
|
||||
):
|
||||
share = VideoShare.create(
|
||||
video_id="video_001",
|
||||
user_id="user_001",
|
||||
video_id=video_id,
|
||||
user_id=user_id,
|
||||
password=password,
|
||||
expires_at=expires_at,
|
||||
)
|
||||
share.is_active = is_active
|
||||
share.view_count = view_count
|
||||
share.download_count = download_count
|
||||
return share
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_share_with_password():
|
||||
share = VideoShare.create(
|
||||
video_id="video_001",
|
||||
user_id="user_001",
|
||||
password="secret123",
|
||||
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,
|
||||
)
|
||||
return share
|
||||
|
||||
|
||||
@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
|
||||
# ── CreateShareUseCase ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateShareUseCase:
|
||||
"""CreateShareUseCase 测试"""
|
||||
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
|
||||
|
||||
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
|
||||
uc = CreateShareUseCase(share_repo, video_repo)
|
||||
cmd = CreateShareCommand(video_id="v1", user_id="u1")
|
||||
result = uc.execute(cmd)
|
||||
|
||||
use_case = CreateShareUseCase(mock_share_repo, mock_video_repo)
|
||||
command = CreateShareCommand(video_id="video_001", user_id="user_001")
|
||||
result = use_case.execute(command)
|
||||
assert result.video_id == "v1"
|
||||
assert result.user_id == "u1"
|
||||
video_repo.get.assert_called_once_with("v1")
|
||||
share_repo.create.assert_called_once()
|
||||
|
||||
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()
|
||||
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
|
||||
|
||||
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)
|
||||
uc = CreateShareUseCase(share_repo, video_repo)
|
||||
cmd = CreateShareCommand(video_id="v1", user_id="u1", password="secret")
|
||||
result = uc.execute(cmd)
|
||||
|
||||
assert result.has_password is True
|
||||
assert result.password_hash is not 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
|
||||
def test_video_not_found_raises(self):
|
||||
share_repo = MagicMock()
|
||||
video_repo = MagicMock()
|
||||
video_repo.get.return_value = None
|
||||
|
||||
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")
|
||||
uc = CreateShareUseCase(share_repo, video_repo)
|
||||
cmd = CreateShareCommand(video_id="v999", user_id="u1")
|
||||
|
||||
with pytest.raises(VideoNotFoundError):
|
||||
use_case.execute(command)
|
||||
uc.execute(cmd)
|
||||
|
||||
mock_share_repo.create.assert_not_called()
|
||||
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
|
||||
|
||||
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")
|
||||
uc = CreateShareUseCase(share_repo, video_repo)
|
||||
cmd = CreateShareCommand(video_id="v1", user_id="u1")
|
||||
|
||||
with pytest.raises(VideoNotFoundError):
|
||||
use_case.execute(command)
|
||||
uc.execute(cmd)
|
||||
|
||||
mock_share_repo.create.assert_not_called()
|
||||
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 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGetShareByTokenUseCase:
|
||||
"""GetShareByTokenUseCase 测试"""
|
||||
def test_get_success(self):
|
||||
share = _make_share()
|
||||
repo = MagicMock()
|
||||
repo.get_by_token.return_value = share
|
||||
|
||||
def test_get_share_success(self, mock_share_repo, sample_share):
|
||||
"""通过 token 正常获取分享信息"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share
|
||||
uc = GetShareByTokenUseCase(repo)
|
||||
result = uc.execute(share.share_token)
|
||||
assert result.id == share.id
|
||||
|
||||
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)
|
||||
def test_not_found_raises(self):
|
||||
repo = MagicMock()
|
||||
repo.get_by_token.return_value = None
|
||||
|
||||
uc = GetShareByTokenUseCase(repo)
|
||||
with pytest.raises(NotFoundError):
|
||||
use_case.execute("invalid_token")
|
||||
uc.execute("nonexistent")
|
||||
|
||||
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)
|
||||
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
|
||||
|
||||
uc = GetShareByTokenUseCase(repo)
|
||||
with pytest.raises(ShareExpiredError):
|
||||
use_case.execute(sample_share_expired.share_token)
|
||||
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 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAccessShareUseCase:
|
||||
"""AccessShareUseCase 测试"""
|
||||
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
|
||||
|
||||
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
|
||||
uc = AccessShareUseCase(share_repo, video_repo)
|
||||
result = uc.execute(share.share_token)
|
||||
|
||||
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 isinstance(result, ShareAccessResult)
|
||||
assert result.share.id == share.id
|
||||
assert result.video.id == video.id
|
||||
assert result.password_verified is True
|
||||
mock_share_repo.increment_view.assert_called_once_with(sample_share.id)
|
||||
assert sample_share.view_count == 1
|
||||
assert share.view_count == 1
|
||||
share_repo.increment_view.assert_called_once_with(share.id)
|
||||
|
||||
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")
|
||||
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
|
||||
|
||||
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, 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)
|
||||
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
|
||||
|
||||
uc = AccessShareUseCase(share_repo, video_repo)
|
||||
with pytest.raises(PasswordRequiredError):
|
||||
use_case.execute(sample_share_with_password.share_token)
|
||||
uc.execute(share.share_token)
|
||||
|
||||
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)
|
||||
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
|
||||
|
||||
uc = AccessShareUseCase(share_repo, video_repo)
|
||||
with pytest.raises(InvalidPasswordError):
|
||||
use_case.execute(sample_share_with_password.share_token, password="wrongpass")
|
||||
uc.execute(share.share_token, password="wrong")
|
||||
|
||||
mock_share_repo.increment_view.assert_not_called()
|
||||
|
||||
def test_access_share_not_found(self, mock_share_repo, mock_video_repo):
|
||||
"""分享不存在抛出 NotFoundError"""
|
||||
mock_share_repo.get_by_token.return_value = None
|
||||
|
||||
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
use_case.execute("invalid_token")
|
||||
|
||||
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)
|
||||
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
|
||||
|
||||
uc = AccessShareUseCase(share_repo, MagicMock())
|
||||
with pytest.raises(ShareExpiredError):
|
||||
use_case.execute(sample_share_expired.share_token)
|
||||
uc.execute(share.share_token)
|
||||
|
||||
mock_share_repo.increment_view.assert_not_called()
|
||||
def test_access_share_not_found(self):
|
||||
share_repo = MagicMock()
|
||||
share_repo.get_by_token.return_value = None
|
||||
|
||||
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
|
||||
uc = AccessShareUseCase(share_repo, MagicMock())
|
||||
with pytest.raises(NotFoundError):
|
||||
uc.execute("nonexistent")
|
||||
|
||||
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
|
||||
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
|
||||
|
||||
uc = AccessShareUseCase(share_repo, video_repo)
|
||||
with pytest.raises(VideoNotFoundError):
|
||||
use_case.execute(sample_share.share_token)
|
||||
uc.execute(share.share_token)
|
||||
|
||||
|
||||
# ── ListSharesByVideoUseCase ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestListSharesByVideoUseCase:
|
||||
"""ListSharesByVideoUseCase 测试"""
|
||||
def test_list_success(self):
|
||||
shares = [_make_share(), _make_share()]
|
||||
repo = MagicMock()
|
||||
repo.list_by_video.return_value = shares
|
||||
|
||||
def test_list_by_video(self, mock_share_repo, sample_share):
|
||||
"""列出某个视频的所有分享"""
|
||||
mock_share_repo.list_by_video.return_value = [sample_share]
|
||||
uc = ListSharesByVideoUseCase(repo)
|
||||
result = uc.execute("v1", "u1")
|
||||
|
||||
use_case = ListSharesByVideoUseCase(mock_share_repo)
|
||||
result = use_case.execute("video_001", "user_001")
|
||||
assert len(result) == 2
|
||||
repo.list_by_video.assert_called_once_with("v1", "u1")
|
||||
|
||||
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")
|
||||
def test_list_empty(self):
|
||||
repo = MagicMock()
|
||||
repo.list_by_video.return_value = []
|
||||
|
||||
uc = ListSharesByVideoUseCase(repo)
|
||||
result = uc.execute("v1", "u1")
|
||||
assert result == []
|
||||
|
||||
|
||||
# ── ListSharesByUserUseCase ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestListSharesByUserUseCase:
|
||||
"""ListSharesByUserUseCase 测试"""
|
||||
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
|
||||
|
||||
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
|
||||
uc = ListSharesByUserUseCase(repo)
|
||||
items, total = uc.execute("u1", skip=0, limit=5)
|
||||
|
||||
use_case = ListSharesByUserUseCase(mock_share_repo)
|
||||
items, total = use_case.execute("user_001")
|
||||
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")
|
||||
|
||||
assert len(items) == 1
|
||||
assert total == 1
|
||||
mock_share_repo.list_by_user.assert_called_once_with("user_001", skip=0, limit=20)
|
||||
def test_list_default_params(self):
|
||||
repo = MagicMock()
|
||||
repo.list_by_user.return_value = []
|
||||
repo.count_by_user.return_value = 0
|
||||
|
||||
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
|
||||
uc = ListSharesByUserUseCase(repo)
|
||||
uc.execute("u1")
|
||||
|
||||
use_case = ListSharesByUserUseCase(mock_share_repo)
|
||||
items, total = use_case.execute("user_001", skip=10, limit=5)
|
||||
repo.list_by_user.assert_called_once_with("u1", skip=0, limit=20)
|
||||
|
||||
assert total == 50
|
||||
mock_share_repo.list_by_user.assert_called_once_with("user_001", skip=10, limit=5)
|
||||
|
||||
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
|
||||
# ── UpdateShareUseCase ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestUpdateShareUseCase:
|
||||
"""UpdateShareUseCase 测试"""
|
||||
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
|
||||
|
||||
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
|
||||
uc = UpdateShareUseCase(repo)
|
||||
cmd = UpdateShareCommand(share_id=share.id, user_id="u1", password="newpass")
|
||||
result = uc.execute(cmd)
|
||||
|
||||
use_case = UpdateShareUseCase(mock_share_repo)
|
||||
command = UpdateShareCommand(
|
||||
share_id=sample_share.id,
|
||||
user_id="user_001",
|
||||
password="newpassword",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
assert result is not None
|
||||
assert share.verify_password("newpass") is True
|
||||
assert share.verify_password("oldpass") is False
|
||||
repo.update.assert_called_once()
|
||||
|
||||
assert result.has_password is True
|
||||
mock_share_repo.update.assert_called_once()
|
||||
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
|
||||
|
||||
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)
|
||||
uc = UpdateShareUseCase(repo)
|
||||
cmd = UpdateShareCommand(share_id=share.id, user_id="u1", password="")
|
||||
result = uc.execute(cmd)
|
||||
|
||||
assert result.has_password is False
|
||||
assert result.password_hash is None
|
||||
|
||||
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
|
||||
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
|
||||
|
||||
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)
|
||||
uc = UpdateShareUseCase(repo)
|
||||
cmd = UpdateShareCommand(share_id=share.id, user_id="u1", password=None)
|
||||
result = uc.execute(cmd)
|
||||
|
||||
assert result.password_hash == original_hash
|
||||
# password=None 表示不修改
|
||||
assert result.has_password is True
|
||||
assert share.verify_password("oldpass") is True
|
||||
|
||||
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
|
||||
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
|
||||
|
||||
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)
|
||||
uc = UpdateShareUseCase(repo)
|
||||
cmd = UpdateShareCommand(share_id=share.id, user_id="u1", expires_at=new_expiry)
|
||||
result = uc.execute(cmd)
|
||||
|
||||
assert result.expires_at == future
|
||||
assert result.expires_at == new_expiry
|
||||
|
||||
def test_update_expires_at_past_raises(self, mock_share_repo, sample_share):
|
||||
"""设置过去的有效期抛出 ValueError"""
|
||||
mock_share_repo.get_by_id.return_value = sample_share
|
||||
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
|
||||
|
||||
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,
|
||||
)
|
||||
uc = UpdateShareUseCase(repo)
|
||||
cmd = UpdateShareCommand(share_id=share.id, user_id="u1", expires_at=past)
|
||||
|
||||
with pytest.raises(ValueError, match="expires_at cannot be in the past"):
|
||||
use_case.execute(command)
|
||||
uc.execute(cmd)
|
||||
|
||||
mock_share_repo.update.assert_not_called()
|
||||
def test_update_not_found_raises(self):
|
||||
repo = MagicMock()
|
||||
repo.get_by_id.return_value = None
|
||||
|
||||
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",
|
||||
)
|
||||
uc = UpdateShareUseCase(repo)
|
||||
cmd = UpdateShareCommand(share_id="nonexistent", user_id="u1")
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
use_case.execute(command)
|
||||
uc.execute(cmd)
|
||||
|
||||
mock_share_repo.update.assert_not_called()
|
||||
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 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRevokeShareUseCase:
|
||||
"""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")
|
||||
def test_revoke_success(self):
|
||||
repo = MagicMock()
|
||||
repo.get_by_id.return_value = MagicMock()
|
||||
repo.delete.return_value = True
|
||||
|
||||
uc = RevokeShareUseCase(repo)
|
||||
result = uc.execute("s1", "u1")
|
||||
assert result is True
|
||||
mock_share_repo.delete.assert_called_once_with(sample_share.id, "user_001")
|
||||
repo.delete.assert_called_once_with("s1", "u1")
|
||||
|
||||
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)
|
||||
def test_revoke_not_found_raises(self):
|
||||
repo = MagicMock()
|
||||
repo.get_by_id.return_value = None
|
||||
|
||||
uc = RevokeShareUseCase(repo)
|
||||
with pytest.raises(NotFoundError):
|
||||
use_case.execute("nonexistent", "user_001")
|
||||
uc.execute("s1", "u1")
|
||||
|
||||
mock_share_repo.delete.assert_not_called()
|
||||
|
||||
# ── RecordShareDownloadUseCase ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRecordShareDownloadUseCase:
|
||||
"""RecordShareDownloadUseCase 测试"""
|
||||
def test_record_download_success(self):
|
||||
share = _make_share(download_count=3)
|
||||
repo = MagicMock()
|
||||
repo.get_by_token.return_value = share
|
||||
|
||||
def test_record_download_no_password(self, mock_share_repo, sample_share):
|
||||
"""无密码分享记录下载"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share
|
||||
uc = RecordShareDownloadUseCase(repo)
|
||||
uc.execute(share.share_token)
|
||||
|
||||
use_case = RecordShareDownloadUseCase(mock_share_repo)
|
||||
use_case.execute(sample_share.share_token)
|
||||
repo.increment_download.assert_called_once_with(share.id)
|
||||
|
||||
mock_share_repo.increment_download.assert_called_once_with(sample_share.id)
|
||||
def test_record_download_with_password(self):
|
||||
share = _make_share(password="secret")
|
||||
repo = MagicMock()
|
||||
repo.get_by_token.return_value = share
|
||||
|
||||
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
|
||||
uc = RecordShareDownloadUseCase(repo)
|
||||
uc.execute(share.share_token, password="secret")
|
||||
repo.increment_download.assert_called_once()
|
||||
|
||||
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)
|
||||
def test_record_download_wrong_password_raises(self):
|
||||
share = _make_share(password="secret")
|
||||
repo = MagicMock()
|
||||
repo.get_by_token.return_value = share
|
||||
|
||||
uc = RecordShareDownloadUseCase(repo)
|
||||
with pytest.raises(InvalidPasswordError):
|
||||
use_case.execute(sample_share_with_password.share_token, password="wrong")
|
||||
uc.execute(share.share_token, password="wrong")
|
||||
|
||||
mock_share_repo.increment_download.assert_not_called()
|
||||
|
||||
def test_record_download_not_found(self, mock_share_repo):
|
||||
"""分享不存在抛出 NotFoundError"""
|
||||
mock_share_repo.get_by_token.return_value = None
|
||||
|
||||
use_case = RecordShareDownloadUseCase(mock_share_repo)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
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)
|
||||
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
|
||||
|
||||
uc = RecordShareDownloadUseCase(repo)
|
||||
with pytest.raises(ShareExpiredError):
|
||||
use_case.execute(sample_share_expired.share_token)
|
||||
uc.execute(share.share_token)
|
||||
|
||||
mock_share_repo.increment_download.assert_not_called()
|
||||
def test_record_download_not_found_raises(self):
|
||||
repo = MagicMock()
|
||||
repo.get_by_token.return_value = None
|
||||
|
||||
uc = RecordShareDownloadUseCase(repo)
|
||||
with pytest.raises(NotFoundError):
|
||||
uc.execute("nonexistent")
|
||||
|
||||
Reference in New Issue
Block a user