refactor(#774): 统一配置管理,消除重复定义和默认值不一致 #784
+221
-83
@@ -1,50 +1,62 @@
|
||||
"""API 服务配置 — 继承 SharedSettings,只追加 API 特有字段。
|
||||
|
||||
通用配置(DB/Redis/OSS/Celery/CosyVoice/Doubao 等)统一在
|
||||
packages/shared/config.py 的 SharedSettings 中定义,这里不重复。
|
||||
|
||||
历史上 API 端使用 UPPER_CASE 命名风格的字段,目前通过 property
|
||||
别名向后兼容。新代码统一使用 snake_case(继承自 SharedSettings)。
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import AliasChoices, Field, field_validator
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
from pydantic_settings import SettingsConfigDict
|
||||
|
||||
from packages.shared.config import SharedSettings
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
APP_NAME: str = "xiaoxia-saas"
|
||||
APP_VERSION: str = "0.1.61"
|
||||
ENVIRONMENT: str = "development"
|
||||
DEBUG: bool = True
|
||||
class Settings(SharedSettings):
|
||||
"""API 服务专用配置。
|
||||
|
||||
通用配置继承自 SharedSettings,这里只定义 API 独有字段。
|
||||
"""
|
||||
|
||||
# ── 应用基本信息 ────────────────────────────────────────────────────
|
||||
app_name: str = "xiaoxia-saas"
|
||||
app_version: str = "0.1.61"
|
||||
|
||||
# 应用基础 URL,用于生成认证邮件中的链接
|
||||
# 开发环境默认 http://localhost:3000
|
||||
# 生产环境应通过环境变量 APP_BASE_URL 设置
|
||||
APP_BASE_URL: str = "http://localhost:3000"
|
||||
app_base_url: str = "http://localhost:3000"
|
||||
|
||||
# Container bind address; external expose is controlled by Docker/Nginx.
|
||||
API_HOST: str = "0.0.0.0" # nosec: B104
|
||||
API_PORT: int = 8000
|
||||
api_host: str = "0.0.0.0" # nosec: B104
|
||||
api_port: int = 8000
|
||||
|
||||
DATABASE_URL: str = "postgresql+psycopg://postgres:postgres@localhost:5432/xiaoxia_saas"
|
||||
DATABASE_POOL_SIZE: int = 20
|
||||
DATABASE_MAX_OVERFLOW: int = 10 # 调整为合理值:pool_size(20) + max_overflow(10) = 最大30连接
|
||||
DATABASE_POOL_TIMEOUT: int = 30
|
||||
DATABASE_POOL_RECYCLE: int = 3600
|
||||
USE_IN_MEMORY_DB: bool = False
|
||||
AUTO_CREATE_SCHEMA: bool = False
|
||||
# ── 数据库特有 ──────────────────────────────────────────────────────
|
||||
use_in_memory_db: bool = False
|
||||
|
||||
REDIS_URL: str = "redis://localhost:6379/0"
|
||||
ENABLE_REDIS_SESSIONS: bool = False
|
||||
# ── Redis 特有 ──────────────────────────────────────────────────────
|
||||
enable_redis_sessions: bool = False
|
||||
|
||||
# ── JWT ────────────────────────────────────────────────────────────
|
||||
# JWT secret key - MUST be set via environment variable, no default allowed
|
||||
JWT_SECRET_KEY: Optional[str] = None
|
||||
jwt_secret_key: Optional[str] = None
|
||||
|
||||
# JWT 算法与过期时间(与 .env.example 对齐)
|
||||
JWT_ALGORITHM: str = "HS256"
|
||||
JWT_ACCESS_TOKEN_EXPIRE_MINUTES: int = 30
|
||||
JWT_REFRESH_TOKEN_EXPIRE_DAYS: int = 30
|
||||
# JWT 算法与过期时间
|
||||
jwt_algorithm: str = "HS256"
|
||||
jwt_access_token_expire_minutes: int = 30
|
||||
jwt_refresh_token_expire_days: int = 30
|
||||
|
||||
@field_validator("JWT_SECRET_KEY", mode="before")
|
||||
@field_validator("jwt_secret_key", mode="before")
|
||||
@classmethod
|
||||
def validate_jwt_secret_key(cls, v):
|
||||
if v is None or v == "":
|
||||
raise ValueError(
|
||||
"JWT_SECRET_KEY must be set via environment variable. " "Do not use default value in production!"
|
||||
"JWT_SECRET_KEY must be set via environment variable. "
|
||||
"Do not use default value in production!"
|
||||
)
|
||||
# Block known insecure default values
|
||||
insecure_defaults = [
|
||||
@@ -56,29 +68,23 @@ class Settings(BaseSettings):
|
||||
]
|
||||
if v.lower() in [d.lower() for d in insecure_defaults]:
|
||||
raise ValueError(
|
||||
f"JWT_SECRET_KEY '{v}' is insecure. " "Please set a strong random secret via environment variable."
|
||||
f"JWT_SECRET_KEY '{v}' is insecure. "
|
||||
"Please set a strong random secret via environment variable."
|
||||
)
|
||||
return v
|
||||
|
||||
ENABLE_EMAIL_DELIVERY: bool = False
|
||||
SMTP_HOST: str = "smtp.gmail.com"
|
||||
SMTP_PORT: int = 587
|
||||
SMTP_USER: str = ""
|
||||
SMTP_PASSWORD: str = ""
|
||||
SMTP_FROM_EMAIL: str = ""
|
||||
SMTP_FROM_NAME: str = "小虾 SaaS"
|
||||
SMTP_USE_TLS: bool = True
|
||||
# ── 邮件 ────────────────────────────────────────────────────────────
|
||||
enable_email_delivery: bool = False
|
||||
smtp_host: str = "smtp.gmail.com"
|
||||
smtp_port: int = 587
|
||||
smtp_user: str = ""
|
||||
smtp_password: str = ""
|
||||
smtp_from_email: str = ""
|
||||
smtp_from_name: str = "小虾 SaaS"
|
||||
smtp_use_tls: bool = True
|
||||
|
||||
CELERY_BROKER_URL: str = "redis://localhost:6379/0"
|
||||
CELERY_RESULT_BACKEND: str = "redis://localhost:6379/1"
|
||||
|
||||
# OSS 七牛云相关
|
||||
OSS_ENDPOINT: str = "oss-cn-hangzhou.aliyuncs.com"
|
||||
OSS_ACCESS_KEY_ID: str = ""
|
||||
OSS_ACCESS_KEY_SECRET: str = ""
|
||||
OSS_BUCKET_NAME: str = "xiaoxia-autocut"
|
||||
|
||||
@field_validator("OSS_ACCESS_KEY_ID", mode="before")
|
||||
# ── OSS 特有校验 ────────────────────────────────────────────────────
|
||||
@field_validator("oss_access_key_id", mode="before")
|
||||
@classmethod
|
||||
def validate_oss_access_key_id(cls, v):
|
||||
if (v is None or v == "") and os.getenv("APP_ENV", "development") != "development":
|
||||
@@ -88,7 +94,7 @@ class Settings(BaseSettings):
|
||||
)
|
||||
return v or ""
|
||||
|
||||
@field_validator("OSS_ACCESS_KEY_SECRET", mode="before")
|
||||
@field_validator("oss_access_key_secret", mode="before")
|
||||
@classmethod
|
||||
def validate_oss_access_key_secret(cls, v):
|
||||
if (v is None or v == "") and os.getenv("APP_ENV", "development") != "development":
|
||||
@@ -98,16 +104,17 @@ class Settings(BaseSettings):
|
||||
)
|
||||
return v or ""
|
||||
|
||||
OSS_DIRECT_UPLOAD_MAX_MB: int = Field(
|
||||
oss_direct_upload_max_mb: int = Field(
|
||||
default=2000,
|
||||
validation_alias=AliasChoices("OSS_DIRECT_UPLOAD_MAX_MB", "MAX_UPLOAD_SIZE_MB"),
|
||||
validation_alias=AliasChoices("oss_direct_upload_max_mb", "max_upload_size_mb"),
|
||||
)
|
||||
OSS_DIRECT_UPLOAD_EXPIRE_SECONDS: int = 900
|
||||
|
||||
CORS_ORIGINS_RAW: str = "http://localhost:3000,http://localhost:5173,http://localhost:8000"
|
||||
# ── CORS ────────────────────────────────────────────────────────────
|
||||
cors_origins_raw: str = "http://localhost:3000,http://localhost:5173,http://localhost:8000"
|
||||
|
||||
# ── 渲染引擎 ────────────────────────────────────────────────────────
|
||||
# 渲染引擎选择:legacy=旧VideoComposeService,unified=新UnifiedRenderService
|
||||
RENDER_ENGINE: str = "legacy"
|
||||
render_engine: str = "legacy"
|
||||
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
@@ -116,47 +123,178 @@ class Settings(BaseSettings):
|
||||
extra="ignore",
|
||||
)
|
||||
|
||||
@property
|
||||
def cors_origins(self) -> list[str]:
|
||||
return [origin.strip() for origin in self.cors_origins_raw.split(",") if origin.strip()]
|
||||
|
||||
# ── 向后兼容:UPPER_CASE property 别名 ──────────────────────────────
|
||||
# 新代码请使用 snake_case(继承的字段名),以下别名仅用于兼容旧代码
|
||||
|
||||
@property
|
||||
def APP_NAME(self) -> str:
|
||||
return self.app_name
|
||||
|
||||
@property
|
||||
def APP_VERSION(self) -> str:
|
||||
return self.app_version
|
||||
|
||||
@property
|
||||
def ENVIRONMENT(self) -> str:
|
||||
return self.environment
|
||||
|
||||
@property
|
||||
def DEBUG(self) -> bool:
|
||||
return self.debug
|
||||
|
||||
@property
|
||||
def APP_BASE_URL(self) -> str:
|
||||
return self.app_base_url
|
||||
|
||||
@property
|
||||
def API_HOST(self) -> str:
|
||||
return self.api_host
|
||||
|
||||
@property
|
||||
def API_PORT(self) -> int:
|
||||
return self.api_port
|
||||
|
||||
@property
|
||||
def DATABASE_URL(self) -> str:
|
||||
return self.database_url
|
||||
|
||||
@property
|
||||
def DATABASE_POOL_SIZE(self) -> int:
|
||||
return self.database_pool_size
|
||||
|
||||
@property
|
||||
def DATABASE_MAX_OVERFLOW(self) -> int:
|
||||
return self.database_max_overflow
|
||||
|
||||
@property
|
||||
def DATABASE_POOL_TIMEOUT(self) -> int:
|
||||
return self.database_pool_timeout
|
||||
|
||||
@property
|
||||
def DATABASE_POOL_RECYCLE(self) -> int:
|
||||
return self.database_pool_recycle
|
||||
|
||||
@property
|
||||
def USE_IN_MEMORY_DB(self) -> bool:
|
||||
return self.use_in_memory_db
|
||||
|
||||
@property
|
||||
def AUTO_CREATE_SCHEMA(self) -> bool:
|
||||
return self.auto_create_schema
|
||||
|
||||
@property
|
||||
def REDIS_URL(self) -> str:
|
||||
return self.redis_url
|
||||
|
||||
@property
|
||||
def ENABLE_REDIS_SESSIONS(self) -> bool:
|
||||
return self.enable_redis_sessions
|
||||
|
||||
@property
|
||||
def JWT_SECRET_KEY(self) -> Optional[str]:
|
||||
return self.jwt_secret_key
|
||||
|
||||
@property
|
||||
def JWT_ALGORITHM(self) -> str:
|
||||
return self.jwt_algorithm
|
||||
|
||||
@property
|
||||
def JWT_ACCESS_TOKEN_EXPIRE_MINUTES(self) -> int:
|
||||
return self.jwt_access_token_expire_minutes
|
||||
|
||||
@property
|
||||
def JWT_REFRESH_TOKEN_EXPIRE_DAYS(self) -> int:
|
||||
return self.jwt_refresh_token_expire_days
|
||||
|
||||
@property
|
||||
def ENABLE_EMAIL_DELIVERY(self) -> bool:
|
||||
return self.enable_email_delivery
|
||||
|
||||
@property
|
||||
def SMTP_HOST(self) -> str:
|
||||
return self.smtp_host
|
||||
|
||||
@property
|
||||
def SMTP_PORT(self) -> int:
|
||||
return self.smtp_port
|
||||
|
||||
@property
|
||||
def SMTP_USER(self) -> str:
|
||||
return self.smtp_user
|
||||
|
||||
@property
|
||||
def SMTP_PASSWORD(self) -> str:
|
||||
return self.smtp_password
|
||||
|
||||
@property
|
||||
def SMTP_FROM_EMAIL(self) -> str:
|
||||
return self.smtp_from_email
|
||||
|
||||
@property
|
||||
def SMTP_FROM_NAME(self) -> str:
|
||||
return self.smtp_from_name
|
||||
|
||||
@property
|
||||
def SMTP_USE_TLS(self) -> bool:
|
||||
return self.smtp_use_tls
|
||||
|
||||
@property
|
||||
def CELERY_BROKER_URL(self) -> str:
|
||||
return self.celery_broker_url
|
||||
|
||||
@property
|
||||
def CELERY_RESULT_BACKEND(self) -> str:
|
||||
return self.celery_result_backend
|
||||
|
||||
@property
|
||||
def OSS_ENDPOINT(self) -> str:
|
||||
return self.oss_endpoint
|
||||
|
||||
@property
|
||||
def OSS_ACCESS_KEY_ID(self) -> str:
|
||||
return self.oss_access_key_id
|
||||
|
||||
@property
|
||||
def OSS_ACCESS_KEY_SECRET(self) -> str:
|
||||
return self.oss_access_key_secret
|
||||
|
||||
@property
|
||||
def OSS_BUCKET_NAME(self) -> str:
|
||||
return self.oss_bucket_name
|
||||
|
||||
@property
|
||||
def OSS_DIRECT_UPLOAD_MAX_MB(self) -> int:
|
||||
return self.oss_direct_upload_max_mb
|
||||
|
||||
@property
|
||||
def OSS_DIRECT_UPLOAD_EXPIRE_SECONDS(self) -> int:
|
||||
return self.oss_direct_upload_expire_seconds
|
||||
|
||||
@property
|
||||
def CORS_ORIGINS_RAW(self) -> str:
|
||||
return self.cors_origins_raw
|
||||
|
||||
@property
|
||||
def CORS_ORIGINS(self) -> list[str]:
|
||||
return [origin.strip() for origin in self.CORS_ORIGINS_RAW.split(",") if origin.strip()]
|
||||
return self.cors_origins
|
||||
|
||||
@property
|
||||
def database_url(self) -> str:
|
||||
return self.DATABASE_URL
|
||||
|
||||
@property
|
||||
def redis_url(self) -> str:
|
||||
return self.REDIS_URL
|
||||
|
||||
@property
|
||||
def celery_broker_url(self) -> str:
|
||||
return self.CELERY_BROKER_URL
|
||||
|
||||
@property
|
||||
def celery_result_backend(self) -> str:
|
||||
return self.CELERY_RESULT_BACKEND
|
||||
|
||||
@property
|
||||
def oss_endpoint(self) -> str:
|
||||
return self.OSS_ENDPOINT
|
||||
|
||||
@property
|
||||
def oss_access_key_id(self) -> str:
|
||||
return self.OSS_ACCESS_KEY_ID
|
||||
|
||||
@property
|
||||
def oss_access_key_secret(self) -> str:
|
||||
return self.OSS_ACCESS_KEY_SECRET
|
||||
|
||||
@property
|
||||
def oss_bucket_name(self) -> str:
|
||||
return self.OSS_BUCKET_NAME
|
||||
def RENDER_ENGINE(self) -> str:
|
||||
return self.render_engine
|
||||
|
||||
|
||||
_settings: Optional[Settings] = None
|
||||
_settings: Optional["Settings"] = None
|
||||
|
||||
|
||||
def get_settings() -> Settings:
|
||||
def get_settings() -> "Settings":
|
||||
"""获取 API 配置单例。
|
||||
|
||||
优先读取 APP_ENV 指定的环境文件(.env.{env}),不存在则读 .env。
|
||||
"""
|
||||
global _settings
|
||||
if _settings is None:
|
||||
env = os.getenv("APP_ENV", "development")
|
||||
|
||||
@@ -1,23 +1,29 @@
|
||||
"""Worker 服务配置 — 继承 SharedSettings,只追加 Worker 特有字段。
|
||||
|
||||
通用配置(DB/Redis/Celery/OSS/CosyVoice/Doubao 等)统一在
|
||||
packages/shared/config.py 的 SharedSettings 中定义,这里不重复。
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
from pydantic_settings import SettingsConfigDict
|
||||
|
||||
from packages.shared.config import SharedSettings
|
||||
|
||||
|
||||
class WorkerSettings(BaseSettings):
|
||||
class WorkerSettings(SharedSettings):
|
||||
"""Worker 服务专用配置。
|
||||
|
||||
通用配置继承自 SharedSettings,这里只定义 Worker 独有字段。
|
||||
Celery broker/backend 使用继承的 celery_broker_url / celery_result_backend;
|
||||
历史上 Worker 使用 broker_url / result_backend 字段名,通过 property 别名兼容。
|
||||
"""
|
||||
|
||||
# ── Worker 特有 ────────────────────────────────────────────────────
|
||||
worker_name: str = "xiaoxia-saas-worker"
|
||||
broker_url: str = "redis://redis:6379/0"
|
||||
result_backend: str = "redis://redis:6379/1"
|
||||
worker_concurrency: int = 4
|
||||
worker_max_tasks_per_child: int = 1000
|
||||
database_url: str = "postgresql+psycopg://postgres:postgres@postgres:5432/xiaoxia_saas"
|
||||
database_pool_size: int = 20
|
||||
database_max_overflow: int = 40
|
||||
database_pool_timeout: int = 30
|
||||
database_pool_recycle: int = 3600
|
||||
environment: str = "development"
|
||||
auto_create_schema: bool = False
|
||||
redis_url: str = "redis://redis:6379/0"
|
||||
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
@@ -26,11 +32,24 @@ class WorkerSettings(BaseSettings):
|
||||
extra="ignore",
|
||||
)
|
||||
|
||||
# ── 向后兼容:Celery 字段名别名 ──────────────────────────────────
|
||||
@property
|
||||
def broker_url(self) -> str:
|
||||
return self.celery_broker_url
|
||||
|
||||
_settings: Optional[WorkerSettings] = None
|
||||
@property
|
||||
def result_backend(self) -> str:
|
||||
return self.celery_result_backend
|
||||
|
||||
|
||||
def get_settings() -> WorkerSettings:
|
||||
_settings: Optional["WorkerSettings"] = None
|
||||
|
||||
|
||||
def get_settings() -> "WorkerSettings":
|
||||
"""获取 Worker 配置单例。
|
||||
|
||||
优先读取 APP_ENV 指定的环境文件(.env.{env}),不存在则读 .env。
|
||||
"""
|
||||
global _settings
|
||||
if _settings is None:
|
||||
env = os.getenv("APP_ENV", "development")
|
||||
|
||||
+30
-14
@@ -1,4 +1,8 @@
|
||||
"""Shared settings for API and Worker services."""
|
||||
"""统一配置入口 — API 和 Worker 共享的基础配置。
|
||||
|
||||
所有服务通用配置定义在这里,两端各自的 Settings 类继承本类,
|
||||
只追加服务特有字段。彻底消除重复定义和默认值不一致问题。
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import Optional
|
||||
@@ -7,29 +11,42 @@ from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
class SharedSettings(BaseSettings):
|
||||
"""Settings shared between API and Worker services."""
|
||||
"""所有服务共享的基础配置。
|
||||
|
||||
# Database
|
||||
API 和 Worker 都继承本类,确保:
|
||||
1. 数据库/Redis/OSS/Celery 等核心配置默认值一致
|
||||
2. 环境变量命名统一(小写风格,pydantic-settings 自动兼容大写)
|
||||
3. 单例模式和 env 文件加载逻辑只实现一次
|
||||
"""
|
||||
|
||||
# ── 环境 ──────────────────────────────────────────────────────────────
|
||||
environment: str = "development"
|
||||
debug: bool = True
|
||||
auto_create_schema: bool = False
|
||||
|
||||
# ── 数据库 ────────────────────────────────────────────────────────────
|
||||
database_url: str = "postgresql+psycopg://postgres:postgres@localhost:5432/xiaoxia_saas"
|
||||
database_pool_size: int = 20
|
||||
database_max_overflow: int = 40
|
||||
database_max_overflow: int = 10 # pool_size(20) + max_overflow(10) = 最大30连接
|
||||
database_pool_timeout: int = 30
|
||||
database_pool_recycle: int = 3600
|
||||
|
||||
# Redis
|
||||
# ── Redis ────────────────────────────────────────────────────────────
|
||||
redis_url: str = "redis://localhost:6379/0"
|
||||
|
||||
# Celery
|
||||
# ── Celery ───────────────────────────────────────────────────────────
|
||||
celery_broker_url: str = "redis://localhost:6379/0"
|
||||
celery_result_backend: str = "redis://localhost:6379/1"
|
||||
|
||||
# OSS Aliyun
|
||||
# ── OSS 阿里云 ──────────────────────────────────────────────────────
|
||||
oss_endpoint: str = "oss-cn-hangzhou.aliyuncs.com"
|
||||
oss_access_key_id: str = ""
|
||||
oss_access_key_secret: str = ""
|
||||
oss_bucket_name: str = "xiaoxia-autocut"
|
||||
oss_direct_upload_max_mb: int = 2000
|
||||
oss_direct_upload_expire_seconds: int = 900
|
||||
|
||||
# CosyVoice (阿里云百炼语音合成)
|
||||
# ── CosyVoice (阿里云百炼语音合成) ───────────────────────────────────
|
||||
cosyvoice_api_key: str = ""
|
||||
cosyvoice_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
|
||||
cosyvoice_model: str = "cosyvoice-v3-flash"
|
||||
@@ -39,17 +56,13 @@ class SharedSettings(BaseSettings):
|
||||
# 音色克隆模型名(固定为 voice-enrollment)
|
||||
cosyvoice_clone_model: str = "voice-enrollment"
|
||||
|
||||
# 豆包大模型(火山引擎方舟)
|
||||
# ── 豆包大模型(火山引擎方舟) ────────────────────────────────────────
|
||||
doubao_api_key: str = ""
|
||||
doubao_model: str = "doubao-seed-1-6-250615"
|
||||
doubao_base_url: str = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
doubao_timeout: int = 30
|
||||
doubao_max_retries: int = 2
|
||||
|
||||
# Environment
|
||||
environment: str = "development"
|
||||
auto_create_schema: bool = False
|
||||
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
env_file_encoding="utf-8",
|
||||
@@ -62,7 +75,10 @@ _settings: Optional[SharedSettings] = None
|
||||
|
||||
|
||||
def get_shared_settings() -> SharedSettings:
|
||||
"""Get shared settings instance (global singleton)."""
|
||||
"""获取共享配置单例。
|
||||
|
||||
优先读取 APP_ENV 指定的环境文件(.env.{env}),不存在则读 .env。
|
||||
"""
|
||||
global _settings
|
||||
if _settings is None:
|
||||
env = os.getenv("APP_ENV", "development")
|
||||
|
||||
Executable
+342
@@ -0,0 +1,342 @@
|
||||
"""素材库 + 素材 UseCase 单元测试(P3-1 第31波)。"""
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.asset_libraries import (
|
||||
CreateAssetLibraryCommand,
|
||||
CreateAssetLibraryUseCase,
|
||||
ListAssetLibrariesUseCase,
|
||||
)
|
||||
from packages.application.assets import (
|
||||
CreateAssetCommand,
|
||||
CreateAssetUseCase,
|
||||
ListAssetsUseCase,
|
||||
)
|
||||
from packages.domain import AssetLibraryKind, AssetStatus, ClassificationStatus
|
||||
|
||||
|
||||
# ── ListAssetLibrariesUseCase ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestListAssetLibrariesUseCase:
|
||||
"""ListAssetLibrariesUseCase 单元测试。"""
|
||||
|
||||
def test_list_libraries_success(self):
|
||||
"""正常列出项目素材库。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_project.return_value = ["lib1", "lib2"]
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("proj_001")
|
||||
|
||||
assert result == ["lib1", "lib2"]
|
||||
mock_repo.find_by_project.assert_called_once_with("proj_001")
|
||||
|
||||
def test_list_libraries_empty(self):
|
||||
"""项目无素材库返回空列表。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_project.return_value = []
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("proj_001")
|
||||
|
||||
assert result == []
|
||||
|
||||
def test_list_libraries_empty_project_id_raises(self):
|
||||
"""空 project_id 抛 ValueError。"""
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
use_case.execute("")
|
||||
|
||||
def test_list_libraries_whitespace_project_id_raises(self):
|
||||
"""全空格 project_id 抛 ValueError。"""
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
use_case.execute(" ")
|
||||
|
||||
def test_list_libraries_project_id_stripped(self):
|
||||
"""project_id 首尾空格被清理。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_project.return_value = []
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
use_case.execute(" proj_001 ")
|
||||
|
||||
mock_repo.find_by_project.assert_called_once_with("proj_001")
|
||||
|
||||
|
||||
# ── CreateAssetLibraryUseCase ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateAssetLibraryUseCase:
|
||||
"""CreateAssetLibraryUseCase 单元测试。"""
|
||||
|
||||
def test_create_video_library_success(self):
|
||||
"""正常创建视频素材库。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateAssetLibraryUseCase(mock_repo)
|
||||
command = CreateAssetLibraryCommand(
|
||||
project_id="proj_001",
|
||||
name="我的视频",
|
||||
kind=AssetLibraryKind.VIDEO,
|
||||
)
|
||||
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.project_id == "proj_001"
|
||||
assert result.name == "我的视频"
|
||||
assert result.kind == AssetLibraryKind.VIDEO
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_create_audio_library_success(self):
|
||||
"""创建音频素材库。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateAssetLibraryUseCase(mock_repo)
|
||||
command = CreateAssetLibraryCommand(
|
||||
project_id="proj_001",
|
||||
name="配乐库",
|
||||
kind=AssetLibraryKind.VOICE,
|
||||
)
|
||||
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.kind == AssetLibraryKind.VOICE
|
||||
assert result.asset_count == 0
|
||||
assert result.total_size == 0
|
||||
|
||||
def test_create_image_library_success(self):
|
||||
"""创建图片素材库。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateAssetLibraryUseCase(mock_repo)
|
||||
command = CreateAssetLibraryCommand(
|
||||
project_id="proj_001",
|
||||
name="图片素材",
|
||||
kind=AssetLibraryKind.IMAGE,
|
||||
)
|
||||
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.kind == AssetLibraryKind.IMAGE
|
||||
|
||||
def test_create_library_name_stripped(self):
|
||||
"""素材库名称首尾空格被清理。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateAssetLibraryUseCase(mock_repo)
|
||||
command = CreateAssetLibraryCommand(
|
||||
project_id="proj_001",
|
||||
name=" 我的素材库 ",
|
||||
kind=AssetLibraryKind.VIDEO,
|
||||
)
|
||||
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.name == "我的素材库"
|
||||
|
||||
|
||||
# ── ListAssetsUseCase ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestListAssetsUseCase:
|
||||
"""ListAssetsUseCase 单元测试。"""
|
||||
|
||||
def test_list_assets_success(self):
|
||||
"""正常列出素材库素材。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_library.return_value = ["asset1", "asset2"]
|
||||
use_case = ListAssetsUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("lib_001")
|
||||
|
||||
assert result == ["asset1", "asset2"]
|
||||
mock_repo.find_by_library.assert_called_once_with("lib_001")
|
||||
|
||||
def test_list_assets_empty(self):
|
||||
"""素材库为空返回空列表。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_library.return_value = []
|
||||
use_case = ListAssetsUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("lib_001")
|
||||
|
||||
assert result == []
|
||||
|
||||
def test_list_assets_empty_library_id_raises(self):
|
||||
"""空 library_id 抛 ValueError。"""
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListAssetsUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="library_id 不能为空"):
|
||||
use_case.execute("")
|
||||
|
||||
def test_list_assets_whitespace_library_id_raises(self):
|
||||
"""全空格 library_id 抛 ValueError。"""
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListAssetsUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="library_id 不能为空"):
|
||||
use_case.execute(" ")
|
||||
|
||||
def test_list_assets_library_id_stripped(self):
|
||||
"""library_id 首尾空格被清理。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_library.return_value = []
|
||||
use_case = ListAssetsUseCase(mock_repo)
|
||||
|
||||
use_case.execute(" lib_001 ")
|
||||
|
||||
mock_repo.find_by_library.assert_called_once_with("lib_001")
|
||||
|
||||
|
||||
# ── CreateAssetUseCase ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateAssetUseCase:
|
||||
"""CreateAssetUseCase 单元测试。"""
|
||||
|
||||
def test_create_video_asset_success(self):
|
||||
"""正常创建视频素材。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateAssetUseCase(mock_repo)
|
||||
command = CreateAssetCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
name="video.mp4",
|
||||
storage_key="assets/video.mp4",
|
||||
mime_type="video/mp4",
|
||||
file_size=1024000,
|
||||
duration=30.5,
|
||||
width=1920,
|
||||
height=1080,
|
||||
fps=25.0,
|
||||
codec="h264",
|
||||
)
|
||||
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.name == "video.mp4"
|
||||
assert result.project_id == "proj_001"
|
||||
assert result.library_id == "lib_001"
|
||||
assert result.storage_key == "assets/video.mp4"
|
||||
assert result.mime_type == "video/mp4"
|
||||
assert result.file_size == 1024000
|
||||
assert result.duration == 30.5
|
||||
assert result.width == 1920
|
||||
assert result.height == 1080
|
||||
assert result.fps == 25.0
|
||||
assert result.codec == "h264"
|
||||
assert result.status == AssetStatus.UPLOADING
|
||||
assert result.classification_status == ClassificationStatus.PENDING
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_create_audio_asset_success(self):
|
||||
"""创建音频素材。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateAssetUseCase(mock_repo)
|
||||
command = CreateAssetCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_002",
|
||||
name="bgm.mp3",
|
||||
storage_key="audio/bgm.mp3",
|
||||
mime_type="audio/mpeg",
|
||||
file_size=512000,
|
||||
duration=180.0,
|
||||
)
|
||||
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.mime_type == "audio/mpeg"
|
||||
assert result.duration == 180.0
|
||||
assert result.width is None
|
||||
assert result.height is None
|
||||
|
||||
def test_create_asset_default_status(self):
|
||||
"""默认状态为 UPLOADING + PENDING。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateAssetUseCase(mock_repo)
|
||||
command = CreateAssetCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
name="test.mp4",
|
||||
storage_key="test.mp4",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.status == AssetStatus.UPLOADING
|
||||
assert result.classification_status == ClassificationStatus.PENDING
|
||||
|
||||
def test_create_asset_custom_status(self):
|
||||
"""可以指定自定义状态。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateAssetUseCase(mock_repo)
|
||||
command = CreateAssetCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
name="test.mp4",
|
||||
storage_key="test.mp4",
|
||||
mime_type="video/mp4",
|
||||
status=AssetStatus.READY,
|
||||
classification_status=ClassificationStatus.COMPLETED,
|
||||
quality_score=85.5,
|
||||
uploaded_by_user_id="user_001",
|
||||
)
|
||||
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.status == AssetStatus.READY
|
||||
assert result.classification_status == ClassificationStatus.COMPLETED
|
||||
assert result.quality_score == 85.5
|
||||
assert result.uploaded_by_user_id == "user_001"
|
||||
|
||||
def test_create_asset_with_metadata(self):
|
||||
"""带 metadata 创建素材。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateAssetUseCase(mock_repo)
|
||||
metadata = {"location": "beijing", "camera": "sony"}
|
||||
command = CreateAssetCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
name="test.mp4",
|
||||
storage_key="test.mp4",
|
||||
mime_type="video/mp4",
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.metadata == metadata
|
||||
|
||||
def test_create_asset_with_thumbnail(self):
|
||||
"""带缩略图创建素材。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateAssetUseCase(mock_repo)
|
||||
command = CreateAssetCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
name="test.mp4",
|
||||
storage_key="test.mp4",
|
||||
mime_type="video/mp4",
|
||||
thumbnail_url="http://cdn.com/thumb.jpg",
|
||||
)
|
||||
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.thumbnail_url == "http://cdn.com/thumb.jpg"
|
||||
Executable
+452
@@ -0,0 +1,452 @@
|
||||
"""Project 领域模型 + UseCase 单元测试(P3-1 第30波)。"""
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.projects import (
|
||||
CreateProjectCommand,
|
||||
CreateProjectUseCase,
|
||||
DeleteProjectUseCase,
|
||||
GetProjectUseCase,
|
||||
ListProjectsUseCase,
|
||||
ShareProjectUseCase,
|
||||
UnshareProjectUseCase,
|
||||
)
|
||||
from packages.domain.entities import Project
|
||||
|
||||
|
||||
# ── Project 领域模型 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestProjectDomain:
|
||||
"""Project 领域模型单元测试。"""
|
||||
|
||||
def test_create_project_basic(self):
|
||||
"""创建项目基本字段正确。"""
|
||||
project = Project.create(owner_user_id="user_001", name="我的项目")
|
||||
|
||||
assert project.id
|
||||
assert project.owner_user_id == "user_001"
|
||||
assert project.name == "我的项目"
|
||||
assert project.description == ""
|
||||
assert project.shared_users == []
|
||||
assert project.created_at
|
||||
|
||||
def test_create_project_with_description(self):
|
||||
"""创建项目带描述。"""
|
||||
project = Project.create(
|
||||
owner_user_id="user_001",
|
||||
name="测试项目",
|
||||
description="这是一个测试项目",
|
||||
)
|
||||
|
||||
assert project.name == "测试项目"
|
||||
assert project.description == "这是一个测试项目"
|
||||
|
||||
def test_create_project_name_stripped(self):
|
||||
"""项目名称首尾空格被清理。"""
|
||||
project = Project.create(owner_user_id="user_001", name=" 我的项目 ")
|
||||
|
||||
assert project.name == "我的项目"
|
||||
|
||||
def test_create_project_description_stripped(self):
|
||||
"""项目描述首尾空格被清理。"""
|
||||
project = Project.create(
|
||||
owner_user_id="user_001",
|
||||
name="测试",
|
||||
description=" 描述内容 ",
|
||||
)
|
||||
|
||||
assert project.description == "描述内容"
|
||||
|
||||
def test_create_project_empty_name_raises(self):
|
||||
"""空项目名抛 ValueError。"""
|
||||
with pytest.raises(ValueError, match="项目名称不能为空"):
|
||||
Project.create(owner_user_id="user_001", name="")
|
||||
|
||||
def test_create_project_whitespace_name_raises(self):
|
||||
"""全空格项目名抛 ValueError。"""
|
||||
with pytest.raises(ValueError, match="项目名称不能为空"):
|
||||
Project.create(owner_user_id="user_001", name=" ")
|
||||
|
||||
def test_is_owner_true(self):
|
||||
"""is_owner 所有者返回 True。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
|
||||
assert project.is_owner("user_001") is True
|
||||
|
||||
def test_is_owner_false(self):
|
||||
"""is_owner 非所有者返回 False。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
|
||||
assert project.is_owner("user_002") is False
|
||||
|
||||
def test_is_shared_with_true(self):
|
||||
"""is_shared_with 已共享用户返回 True。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
project.shared_users = ["user_002", "user_003"]
|
||||
|
||||
assert project.is_shared_with("user_002") is True
|
||||
assert project.is_shared_with("user_003") is True
|
||||
|
||||
def test_is_shared_with_false(self):
|
||||
"""is_shared_with 未共享用户返回 False。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
project.shared_users = ["user_002"]
|
||||
|
||||
assert project.is_shared_with("user_004") is False
|
||||
|
||||
def test_is_shared_with_empty_list(self):
|
||||
"""共享列表为空时返回 False。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
|
||||
assert project.is_shared_with("user_002") is False
|
||||
|
||||
def test_can_access_owner(self):
|
||||
"""所有者可以访问。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
|
||||
assert project.can_access("user_001") is True
|
||||
|
||||
def test_can_access_shared_user(self):
|
||||
"""共享用户可以访问。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
project.shared_users = ["user_002"]
|
||||
|
||||
assert project.can_access("user_002") is True
|
||||
|
||||
def test_can_access_other_user(self):
|
||||
"""其他用户不能访问。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
project.shared_users = ["user_002"]
|
||||
|
||||
assert project.can_access("user_003") is False
|
||||
|
||||
|
||||
# ── ListProjectsUseCase ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestListProjectsUseCase:
|
||||
"""ListProjectsUseCase 单元测试。"""
|
||||
|
||||
def test_list_projects_success(self):
|
||||
"""正常列出用户项目。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_accessible_projects.return_value = [
|
||||
Project.create(owner_user_id="user_001", name="项目1"),
|
||||
Project.create(owner_user_id="user_001", name="项目2"),
|
||||
]
|
||||
use_case = ListProjectsUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("user_001")
|
||||
|
||||
assert len(result) == 2
|
||||
mock_repo.find_accessible_projects.assert_called_once_with("user_001")
|
||||
|
||||
def test_list_projects_empty(self):
|
||||
"""用户没有项目返回空列表。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_accessible_projects.return_value = []
|
||||
use_case = ListProjectsUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("user_001")
|
||||
|
||||
assert result == []
|
||||
|
||||
def test_list_projects_empty_user_id_raises(self):
|
||||
"""空 user_id 抛 ValueError。"""
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListProjectsUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="user_id 不能为空"):
|
||||
use_case.execute("")
|
||||
|
||||
def test_list_projects_whitespace_user_id_raises(self):
|
||||
"""全空格 user_id 抛 ValueError。"""
|
||||
mock_repo = MagicMock()
|
||||
use_case = ListProjectsUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="user_id 不能为空"):
|
||||
use_case.execute(" ")
|
||||
|
||||
def test_list_projects_user_id_stripped(self):
|
||||
"""user_id 首尾空格被清理。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_accessible_projects.return_value = []
|
||||
use_case = ListProjectsUseCase(mock_repo)
|
||||
|
||||
use_case.execute(" user_001 ")
|
||||
|
||||
mock_repo.find_accessible_projects.assert_called_once_with("user_001")
|
||||
|
||||
|
||||
# ── GetProjectUseCase ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGetProjectUseCase:
|
||||
"""GetProjectUseCase 单元测试。"""
|
||||
|
||||
def test_get_project_success(self):
|
||||
"""正常获取项目。"""
|
||||
project = Project.create(owner_user_id="user_001", name="测试项目")
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = project
|
||||
use_case = GetProjectUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("proj_001")
|
||||
|
||||
assert result == project
|
||||
mock_repo.find_by_id.assert_called_once_with("proj_001")
|
||||
|
||||
def test_get_project_not_found(self):
|
||||
"""项目不存在返回 None。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = None
|
||||
use_case = GetProjectUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("proj_001")
|
||||
|
||||
assert result is None
|
||||
|
||||
def test_get_project_empty_id_raises(self):
|
||||
"""空 project_id 抛 ValueError。"""
|
||||
mock_repo = MagicMock()
|
||||
use_case = GetProjectUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
use_case.execute("")
|
||||
|
||||
def test_get_project_whitespace_id_raises(self):
|
||||
"""全空格 project_id 抛 ValueError。"""
|
||||
mock_repo = MagicMock()
|
||||
use_case = GetProjectUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
use_case.execute(" ")
|
||||
|
||||
def test_get_project_id_stripped(self):
|
||||
"""project_id 首尾空格被清理。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = None
|
||||
use_case = GetProjectUseCase(mock_repo)
|
||||
|
||||
use_case.execute(" proj_001 ")
|
||||
|
||||
mock_repo.find_by_id.assert_called_once_with("proj_001")
|
||||
|
||||
|
||||
# ── CreateProjectUseCase ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateProjectUseCase:
|
||||
"""CreateProjectUseCase 单元测试。"""
|
||||
|
||||
def test_create_project_success(self):
|
||||
"""正常创建项目。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.save.side_effect = lambda p: p
|
||||
use_case = CreateProjectUseCase(mock_repo)
|
||||
command = CreateProjectCommand(name="新项目", description="项目描述")
|
||||
|
||||
result = use_case.execute(command, "user_001")
|
||||
|
||||
assert result.name == "新项目"
|
||||
assert result.description == "项目描述"
|
||||
assert result.owner_user_id == "user_001"
|
||||
mock_repo.save.assert_called_once()
|
||||
|
||||
def test_create_project_default_description(self):
|
||||
"""不传描述默认为空。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.save.side_effect = lambda p: p
|
||||
use_case = CreateProjectUseCase(mock_repo)
|
||||
command = CreateProjectCommand(name="新项目")
|
||||
|
||||
result = use_case.execute(command, "user_001")
|
||||
|
||||
assert result.description == ""
|
||||
|
||||
def test_create_project_empty_name_raises(self):
|
||||
"""空项目名抛 ValueError(领域层校验)。"""
|
||||
mock_repo = MagicMock()
|
||||
use_case = CreateProjectUseCase(mock_repo)
|
||||
command = CreateProjectCommand(name="")
|
||||
|
||||
with pytest.raises(ValueError, match="项目名称不能为空"):
|
||||
use_case.execute(command, "user_001")
|
||||
|
||||
mock_repo.save.assert_not_called()
|
||||
|
||||
|
||||
# ── ShareProjectUseCase ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestShareProjectUseCase:
|
||||
"""ShareProjectUseCase 单元测试。"""
|
||||
|
||||
def test_share_project_success(self):
|
||||
"""正常共享项目给用户。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = project
|
||||
mock_repo.save.side_effect = lambda p: p
|
||||
use_case = ShareProjectUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute(project.id, "user_001", "user_002")
|
||||
|
||||
assert "user_002" in result.shared_users
|
||||
mock_repo.save.assert_called_once()
|
||||
|
||||
def test_share_project_already_shared_no_duplicate(self):
|
||||
"""已共享用户再次共享不重复添加。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
project.shared_users = ["user_002"]
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = project
|
||||
mock_repo.save.side_effect = lambda p: p
|
||||
use_case = ShareProjectUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute(project.id, "user_001", "user_002")
|
||||
|
||||
assert result.shared_users == ["user_002"]
|
||||
mock_repo.save.assert_not_called()
|
||||
|
||||
def test_share_project_not_found_raises(self):
|
||||
"""项目不存在抛 ValueError。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = None
|
||||
use_case = ShareProjectUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="项目不存在"):
|
||||
use_case.execute("proj_xxx", "user_001", "user_002")
|
||||
|
||||
mock_repo.save.assert_not_called()
|
||||
|
||||
def test_share_project_not_owner_raises(self):
|
||||
"""非所有者不能共享项目。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = project
|
||||
use_case = ShareProjectUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="只有项目所有者可以共享项目"):
|
||||
use_case.execute(project.id, "user_002", "user_003")
|
||||
|
||||
mock_repo.save.assert_not_called()
|
||||
|
||||
|
||||
# ── UnshareProjectUseCase ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestUnshareProjectUseCase:
|
||||
"""UnshareProjectUseCase 单元测试。"""
|
||||
|
||||
def test_unshare_project_success(self):
|
||||
"""正常取消共享。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
project.shared_users = ["user_002", "user_003"]
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = project
|
||||
mock_repo.save.side_effect = lambda p: p
|
||||
use_case = UnshareProjectUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute(project.id, "user_001", "user_002")
|
||||
|
||||
assert result.shared_users == ["user_003"]
|
||||
mock_repo.save.assert_called_once()
|
||||
|
||||
def test_unshare_project_not_shared_noop(self):
|
||||
"""未共享的用户取消共享不操作。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
project.shared_users = ["user_003"]
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = project
|
||||
mock_repo.save.side_effect = lambda p: p
|
||||
use_case = UnshareProjectUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute(project.id, "user_001", "user_002")
|
||||
|
||||
assert result.shared_users == ["user_003"]
|
||||
mock_repo.save.assert_not_called()
|
||||
|
||||
def test_unshare_project_not_found_raises(self):
|
||||
"""项目不存在抛 ValueError。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = None
|
||||
use_case = UnshareProjectUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="项目不存在"):
|
||||
use_case.execute("proj_xxx", "user_001", "user_002")
|
||||
|
||||
mock_repo.save.assert_not_called()
|
||||
|
||||
def test_unshare_project_not_owner_raises(self):
|
||||
"""非所有者不能取消共享。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
project.shared_users = ["user_002"]
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = project
|
||||
use_case = UnshareProjectUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="只有项目所有者可以取消共享"):
|
||||
use_case.execute(project.id, "user_002", "user_003")
|
||||
|
||||
mock_repo.save.assert_not_called()
|
||||
|
||||
|
||||
# ── DeleteProjectUseCase ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestDeleteProjectUseCase:
|
||||
"""DeleteProjectUseCase 单元测试。"""
|
||||
|
||||
def test_delete_project_success(self):
|
||||
"""所有者正常删除项目。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = project
|
||||
mock_repo.delete.return_value = True
|
||||
use_case = DeleteProjectUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute(project.id, "user_001")
|
||||
|
||||
assert result is True
|
||||
mock_repo.delete.assert_called_once_with(project.id)
|
||||
|
||||
def test_delete_project_not_found(self):
|
||||
"""项目不存在返回 False。"""
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = None
|
||||
use_case = DeleteProjectUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("proj_xxx", "user_001")
|
||||
|
||||
assert result is False
|
||||
mock_repo.delete.assert_not_called()
|
||||
|
||||
def test_delete_project_not_owner_raises(self):
|
||||
"""非所有者删除抛 PermissionError。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = project
|
||||
use_case = DeleteProjectUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(PermissionError, match="只有项目所有者可以删除项目"):
|
||||
use_case.execute(project.id, "user_002")
|
||||
|
||||
mock_repo.delete.assert_not_called()
|
||||
|
||||
def test_delete_project_shared_user_raises(self):
|
||||
"""共享用户不能删除项目。"""
|
||||
project = Project.create(owner_user_id="user_001", name="项目")
|
||||
project.shared_users = ["user_002"]
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.find_by_id.return_value = project
|
||||
use_case = DeleteProjectUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(PermissionError, match="只有项目所有者可以删除项目"):
|
||||
use_case.execute(project.id, "user_002")
|
||||
|
||||
mock_repo.delete.assert_not_called()
|
||||
Reference in New Issue
Block a user