Compare commits
37 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| e0a24f636d | |||
| b410b65e22 | |||
| e4f8b8d492 | |||
| b1f27b0ef6 | |||
| 16b18dc2fd | |||
| 56707a834c | |||
| c0c049b2e1 | |||
| 1dd612c89c | |||
| b6ead113cf | |||
| 37152a0729 | |||
| 05c9137ece | |||
| bae9b509de | |||
| e6ecce79ee | |||
| 6dbb281acf | |||
| d934e0e322 | |||
| 304354208a | |||
| b5e27a5e9e | |||
| b98c83582b | |||
| afea0a7566 | |||
| 7b08029e2e | |||
| 6b5959bafd | |||
| 7783f36172 | |||
| f9e8d6efdc | |||
| 401a55cfcb | |||
| f80a311082 | |||
| 4c1139b681 | |||
| d2578a9093 | |||
| e723025889 | |||
| 3c6b8aa362 | |||
| 265e49587f | |||
| 70ad41bf94 | |||
| c2a522babf | |||
| 35446732ea | |||
| d2097744e4 | |||
| ea9e536740 | |||
| 8c678b9cd0 | |||
| 8c4c428cdc |
+36
-12
@@ -1,30 +1,54 @@
|
||||
# 生产环境配置模板(实际使用时复制为 .env.production)
|
||||
|
||||
# ==================== 基础配置 ====================
|
||||
APP_ENV=production
|
||||
ENVIRONMENT=production
|
||||
DEBUG=false
|
||||
USE_IN_MEMORY_DB=false
|
||||
LOG_LEVEL=WARNING
|
||||
|
||||
# 数据库(必须修改)
|
||||
# ==================== 数据库(必须修改)====================
|
||||
DATABASE_URL=postgresql://prod_user:CHANGE_THIS_PASSWORD@db-prod:5432/xiaoxia_prod
|
||||
|
||||
# Redis(必须修改)
|
||||
# ==================== Redis(必须修改)====================
|
||||
REDIS_URL=redis://:CHANGE_THIS_PASSWORD@redis-prod:6379/0
|
||||
ENABLE_REDIS_SESSIONS=false
|
||||
|
||||
# JWT(必须修改,至少 32 字符)
|
||||
# ==================== JWT(必须修改,至少 32 字符)====================
|
||||
JWT_SECRET_KEY=CHANGE_THIS_TO_A_RANDOM_SECRET_KEY_AT_LEAST_32_CHARS
|
||||
|
||||
# SMTP(必须配置)
|
||||
# ==================== 邮件(必须配置)====================
|
||||
ENABLE_EMAIL_DELIVERY=false
|
||||
SMTP_HOST=smtp.gmail.com
|
||||
SMTP_PORT=587
|
||||
SMTP_USER=your-email@gmail.com
|
||||
SMTP_PASSWORD=your-app-specific-password
|
||||
SMTP_USER=CHANGE_ME_SMTP_USER
|
||||
SMTP_PASSWORD=CHANGE_ME_SMTP_PASSWORD
|
||||
SMTP_FROM_EMAIL=noreply@yourdomain.com
|
||||
SMTP_FROM_NAME=小虾 SaaS
|
||||
SMTP_USE_TLS=true
|
||||
|
||||
# 应用配置
|
||||
BASE_URL=https://yourdomain.com
|
||||
# ==================== 应用配置 ====================
|
||||
APP_BASE_URL=https://yourdomain.com
|
||||
|
||||
# CORS(修改为实际域名)
|
||||
CORS_ORIGINS=["https://yourdomain.com","https://app.yourdomain.com"]
|
||||
# ==================== CORS(修改为实际域名,逗号分隔)====================
|
||||
CORS_ORIGINS_RAW=https://yourdomain.com,https://app.yourdomain.com
|
||||
|
||||
# 监控(可选)
|
||||
SENTRY_DSN=https://your-sentry-dsn@sentry.io/project-id
|
||||
# ==================== 阿里云 OSS(必须配置)====================
|
||||
OSS_ENDPOINT=oss-cn-hangzhou.aliyuncs.com
|
||||
OSS_ACCESS_KEY_ID=CHANGE_ME_ACCESS_KEY_ID
|
||||
OSS_ACCESS_KEY_SECRET=CHANGE_ME_ACCESS_KEY_SECRET
|
||||
OSS_BUCKET_NAME=xiaoxia-autocut
|
||||
OSS_DIRECT_UPLOAD_MAX_MB=2000
|
||||
OSS_DIRECT_UPLOAD_EXPIRE_SECONDS=900
|
||||
|
||||
# ==================== 生成文件 ====================
|
||||
GENERATED_FILES_DIR=/app/generated
|
||||
GENERATED_FILES_URL_PREFIX=/generated-files
|
||||
PUBLIC_API_BASE_URL=https://api.xiaoxiajianji.com
|
||||
|
||||
# ==================== Celery ====================
|
||||
CELERY_BROKER_URL=redis://:CHANGE_THIS_PASSWORD@redis-prod:6379/0
|
||||
CELERY_RESULT_BACKEND=redis://:CHANGE_THIS_PASSWORD@redis-prod:6379/1
|
||||
|
||||
# ==================== 监控(可选)====================
|
||||
# SENTRY_DSN=https://your-sentry-dsn@sentry.io/project-id
|
||||
|
||||
@@ -169,6 +169,21 @@ jobs:
|
||||
scp -i "$key_path" "dist/release-artifacts/xiaoxia-web-${GITHUB_REF_NAME}.tar" \
|
||||
"$production_user@$production_host:/var/lib/xiaoxia-saas-production/web-${GITHUB_REF_NAME}.tar"
|
||||
|
||||
- name: Cleanup old Docker images
|
||||
if: always()
|
||||
shell: sh
|
||||
run: |
|
||||
set -eu
|
||||
if [ -f scripts/cleanup_old_images.sh ]; then
|
||||
chmod +x scripts/cleanup_old_images.sh
|
||||
scripts/cleanup_old_images.sh
|
||||
else
|
||||
echo "Cleanup script not found, doing basic prune..."
|
||||
docker image prune -f 2>/dev/null || true
|
||||
fi
|
||||
echo "Disk usage after cleanup:"
|
||||
df -h / | tail -1
|
||||
|
||||
deploy-production:
|
||||
name: Deploy Production
|
||||
runs-on: runtime-builder:host
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
"""add generation task extensions
|
||||
|
||||
Revision ID: 015
|
||||
Revises: 014
|
||||
Create Date: 2026-06-29
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import mysql
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers
|
||||
revision = "015"
|
||||
down_revision = "014"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column("generation_tasks", sa.Column("template_id", sa.String(36), nullable=False, server_default=""))
|
||||
op.add_column("generation_tasks", sa.Column("asset_ids", mysql.JSON(), nullable=False, server_default="[]"))
|
||||
op.add_column("generation_tasks", sa.Column("title_ids", mysql.JSON(), nullable=False, server_default="[]"))
|
||||
op.add_column("generation_tasks", sa.Column("voice_ids", mysql.JSON(), nullable=False, server_default="[]"))
|
||||
|
||||
op.create_index(op.f("ix_generation_tasks_template_id"), "generation_tasks", ["template_id"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index(op.f("ix_generation_tasks_template_id"), table_name="generation_tasks")
|
||||
op.drop_column("generation_tasks", "voice_ids")
|
||||
op.drop_column("generation_tasks", "title_ids")
|
||||
op.drop_column("generation_tasks", "asset_ids")
|
||||
op.drop_column("generation_tasks", "template_id")
|
||||
@@ -1,3 +1,4 @@
|
||||
from app.api.routes.dashboard import router as dashboard_router
|
||||
from app.api.routes.asset_diagnosis import router as asset_diagnosis_router
|
||||
from app.api.routes.asset_libraries import router as asset_libraries_router
|
||||
from app.api.routes.assets import router as assets_router
|
||||
@@ -110,3 +111,8 @@ api_router.include_router(
|
||||
prefix="/templates",
|
||||
tags=["Template"],
|
||||
)
|
||||
api_router.include_router(
|
||||
dashboard_router,
|
||||
prefix="/dashboard",
|
||||
tags=["Dashboard"],
|
||||
)
|
||||
|
||||
@@ -202,7 +202,7 @@ def get_project_asset_diagnosis(
|
||||
if project is None:
|
||||
raise HTTPException(status_code=404, detail=f"Project {project_id} not found")
|
||||
|
||||
libraries = asset_library_repository.list_by_project(project_id)
|
||||
libraries = asset_library_repository.find_by_project(project_id)
|
||||
assets: list[Asset] = []
|
||||
for library in libraries:
|
||||
assets.extend(asset_repository.list_by_library(library.id))
|
||||
|
||||
@@ -101,7 +101,7 @@ def _require_project_and_library(
|
||||
if project is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
|
||||
|
||||
libraries = asset_library_repository.list_by_project(project_id)
|
||||
libraries = asset_library_repository.find_by_project(project_id)
|
||||
if not any(item.id == library_id for item in libraries):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Asset library not found")
|
||||
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import (
|
||||
get_asset_repository,
|
||||
get_generation_task_repository,
|
||||
get_project_repository,
|
||||
get_title_library_repository,
|
||||
get_voice_library_repository,
|
||||
)
|
||||
from app.schemas.dashboard import DashboardOverviewResponse, RecentTaskItem, SubscriptionInfo
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _status_value(status) -> str:
|
||||
return status.value if hasattr(status, "value") else str(status)
|
||||
|
||||
|
||||
def _generation_step(status: str) -> str:
|
||||
if status == "pending":
|
||||
return "等待 Worker 执行"
|
||||
if status == "running":
|
||||
return "正在生成成片"
|
||||
if status == "completed":
|
||||
return "生成完成"
|
||||
if status == "failed":
|
||||
return "生成失败"
|
||||
return status
|
||||
|
||||
|
||||
@router.get("/overview", response_model=DashboardOverviewResponse)
|
||||
def get_dashboard_overview(
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
asset_repository: Any = Depends(get_asset_repository),
|
||||
generation_task_repository: Any = Depends(get_generation_task_repository),
|
||||
title_library_repository: Any = Depends(get_title_library_repository),
|
||||
voice_library_repository: Any = Depends(get_voice_library_repository),
|
||||
) -> DashboardOverviewResponse:
|
||||
"""Dashboard 概览:用户级汇总数据。"""
|
||||
user_id = authenticated_user.user.id
|
||||
|
||||
# 获取用户可访问的所有 project
|
||||
projects = project_repository.find_accessible_projects(user_id)
|
||||
project_ids = [p.id for p in projects]
|
||||
|
||||
# 素材统计
|
||||
total_assets = asset_repository.count_by_project_ids(project_ids)
|
||||
used_storage_bytes = asset_repository.sum_storage_by_project_ids(project_ids)
|
||||
|
||||
# 标题库 / 配音库统计
|
||||
total_titles = title_library_repository.count_by_user(user_id)
|
||||
total_voices = voice_library_repository.count_by_user(user_id)
|
||||
|
||||
# 生成任务统计
|
||||
total_tasks = generation_task_repository.count_by_user(user_id)
|
||||
|
||||
# 最近任务(SQL 层 LIMIT 5)
|
||||
recent = generation_task_repository.list_recent_by_user(user_id, limit=5)
|
||||
recent_tasks = []
|
||||
for task in recent:
|
||||
s = _status_value(task.status)
|
||||
recent_tasks.append(
|
||||
RecentTaskItem(
|
||||
id=task.id,
|
||||
task_type="generation",
|
||||
status=s,
|
||||
current_step=_generation_step(s),
|
||||
error_message=task.error_message or "",
|
||||
updated_at=task.completed_at or task.started_at or task.created_at,
|
||||
)
|
||||
)
|
||||
|
||||
# 订阅信息
|
||||
user = authenticated_user.user
|
||||
subscription = SubscriptionInfo(
|
||||
plan=getattr(user, "subscription_plan", "free") or "free",
|
||||
is_active=getattr(user, "subscription_status", "") == "active",
|
||||
)
|
||||
|
||||
return DashboardOverviewResponse(
|
||||
total_assets=total_assets,
|
||||
used_storage_bytes=used_storage_bytes,
|
||||
total_titles=total_titles,
|
||||
total_voices=total_voices,
|
||||
total_tasks=total_tasks,
|
||||
total_products=len(projects),
|
||||
subscription=subscription,
|
||||
recent_tasks=recent_tasks,
|
||||
)
|
||||
@@ -16,6 +16,7 @@ from app.schemas.generated_video import (
|
||||
from app.schemas.generation_task import (
|
||||
CreateGenerationTaskRequest,
|
||||
GenerationTaskResponse,
|
||||
ListGenerationTasksResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
@@ -45,6 +46,10 @@ def _to_generation_task_response(task) -> GenerationTaskResponse:
|
||||
asset_library_id=task.asset_library_id,
|
||||
strategy_id=task.strategy_id,
|
||||
voice_library_id=task.voice_library_id,
|
||||
template_id=task.template_id,
|
||||
asset_ids=task.asset_ids,
|
||||
title_ids=task.title_ids,
|
||||
voice_ids=task.voice_ids,
|
||||
status=task.status,
|
||||
progress=task.progress,
|
||||
result_count=task.result_count,
|
||||
@@ -79,6 +84,43 @@ def _ensure_library_has_ready_video_assets(assets) -> None:
|
||||
)
|
||||
|
||||
|
||||
def _resolve_project_and_library(
|
||||
request: CreateGenerationTaskRequest,
|
||||
project_repository: Any,
|
||||
asset_library_repository: Any,
|
||||
asset_repository: Any,
|
||||
authenticated_user: AuthenticatedUser,
|
||||
) -> tuple[str, str]:
|
||||
"""解析 project_id 和 asset_library_id。
|
||||
|
||||
支持两种模式:
|
||||
- 显式传入(向后兼容)
|
||||
- 从 asset_ids 反查 asset_library(模板模式)
|
||||
返回 (project_id, asset_library_id)。
|
||||
"""
|
||||
project_id = request.project_id.strip()
|
||||
asset_library_id = request.asset_library_id.strip()
|
||||
|
||||
# 模板模式:project_id 未提供时,从 asset_ids 反查所属 project
|
||||
if not project_id and request.asset_ids:
|
||||
first_asset_id = request.asset_ids[0]
|
||||
asset = asset_repository.find_by_id(first_asset_id)
|
||||
if asset is not None:
|
||||
project_id = asset.project_id
|
||||
if not asset_library_id:
|
||||
asset_library_id = asset.library_id
|
||||
|
||||
# 向后兼容校验:project_id 已提供时验证权限
|
||||
if project_id:
|
||||
project = project_repository.find_by_id(project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=404, detail=f"Project {project_id} not found")
|
||||
if not project.can_access(authenticated_user.user.id):
|
||||
raise HTTPException(status_code=403, detail="Access denied to project")
|
||||
|
||||
return project_id, asset_library_id
|
||||
|
||||
|
||||
@router.post("/tasks", response_model=GenerationTaskResponse)
|
||||
def create_generation_task(
|
||||
request: CreateGenerationTaskRequest,
|
||||
@@ -88,26 +130,30 @@ def create_generation_task(
|
||||
asset_library_repository: Any = Depends(get_asset_library_repository),
|
||||
asset_repository: Any = Depends(get_asset_repository),
|
||||
) -> GenerationTaskResponse:
|
||||
project = project_repository.find_by_id(request.project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=404, detail=f"Project {request.project_id} not found")
|
||||
if not project.can_access(authenticated_user.user.id):
|
||||
raise HTTPException(status_code=403, detail="Access denied to project")
|
||||
|
||||
library = asset_library_repository.get(request.asset_library_id)
|
||||
if library is None or library.project_id != request.project_id:
|
||||
raise HTTPException(status_code=404, detail=f"AssetLibrary {request.asset_library_id} not found")
|
||||
|
||||
assets = asset_repository.list_by_library(request.asset_library_id)
|
||||
_ensure_library_has_ready_video_assets(assets)
|
||||
project_id, asset_library_id = _resolve_project_and_library(
|
||||
request, project_repository, asset_library_repository, asset_repository, authenticated_user
|
||||
)
|
||||
|
||||
# asset_library 存在性校验(仅在提供了 asset_library_id 时)
|
||||
if asset_library_id:
|
||||
library = asset_library_repository.get(asset_library_id)
|
||||
if library is None or (project_id and library.project_id != project_id):
|
||||
raise HTTPException(status_code=404, detail=f"AssetLibrary {asset_library_id} not found")
|
||||
|
||||
assets = asset_repository.find_by_library(asset_library_id)
|
||||
_ensure_library_has_ready_video_assets(assets)
|
||||
|
||||
use_case = CreateGenerationTaskUseCase(generation_task_repository)
|
||||
task = use_case.execute(
|
||||
CreateGenerationTaskCommand(
|
||||
project_id=request.project_id,
|
||||
asset_library_id=request.asset_library_id,
|
||||
project_id=project_id,
|
||||
asset_library_id=asset_library_id,
|
||||
strategy_id=request.strategy_id,
|
||||
voice_library_id=request.voice_library_id,
|
||||
template_id=request.template_id,
|
||||
asset_ids=request.asset_ids,
|
||||
title_ids=request.title_ids,
|
||||
voice_ids=request.voice_ids,
|
||||
created_by_user_id=authenticated_user.user.id,
|
||||
)
|
||||
)
|
||||
@@ -115,6 +161,17 @@ def create_generation_task(
|
||||
return _to_generation_task_response(task)
|
||||
|
||||
|
||||
@router.get("/tasks", response_model=ListGenerationTasksResponse)
|
||||
def list_generation_tasks(
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
generation_task_repository: Any = Depends(get_generation_task_repository),
|
||||
) -> ListGenerationTasksResponse:
|
||||
"""用户级生成任务列表(跨 project)。"""
|
||||
tasks = generation_task_repository.list_by_user(authenticated_user.user.id)
|
||||
items = [_to_generation_task_response(task) for task in tasks]
|
||||
return ListGenerationTasksResponse(items=items)
|
||||
|
||||
|
||||
@router.get("/tasks/{task_id}", response_model=GenerationTaskResponse)
|
||||
def get_generation_task(
|
||||
task_id: str,
|
||||
@@ -126,7 +183,8 @@ def get_generation_task(
|
||||
task = use_case.execute(task_id)
|
||||
if task is None:
|
||||
raise HTTPException(status_code=404, detail=f"GenerationTask {task_id} not found")
|
||||
_check_project_access(task.project_id, authenticated_user.user.id, project_repository)
|
||||
if task.project_id:
|
||||
_check_project_access(task.project_id, authenticated_user.user.id, project_repository)
|
||||
return _to_generation_task_response(task)
|
||||
|
||||
|
||||
@@ -141,7 +199,42 @@ def list_generation_results(
|
||||
task = generation_task_repository.get(task_id)
|
||||
if task is None:
|
||||
raise HTTPException(status_code=404, detail=f"GenerationTask {task_id} not found")
|
||||
_check_project_access(task.project_id, authenticated_user.user.id, project_repository)
|
||||
if task.project_id:
|
||||
_check_project_access(task.project_id, authenticated_user.user.id, project_repository)
|
||||
use_case = ListGeneratedVideosByTaskUseCase(generated_video_repository)
|
||||
items = use_case.execute(task_id)
|
||||
return ListGeneratedVideosResponse(items=[_to_generated_video_response(item) for item in items])
|
||||
|
||||
|
||||
@router.post("/tasks/{task_id}/retry", response_model=GenerationTaskResponse)
|
||||
def retry_generation_task(
|
||||
task_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
generation_task_repository: Any = Depends(get_generation_task_repository),
|
||||
) -> GenerationTaskResponse:
|
||||
"""简化重试:通过 task_id 直接重试失败任务。"""
|
||||
task = generation_task_repository.get(task_id)
|
||||
if task is None:
|
||||
raise HTTPException(status_code=404, detail="Generation task not found")
|
||||
if task.created_by_user_id and task.created_by_user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="Access denied to this task")
|
||||
status_val = task.status.value if hasattr(task.status, "value") else str(task.status)
|
||||
if status_val != "failed":
|
||||
raise HTTPException(status_code=409, detail="Only failed tasks can be retried")
|
||||
|
||||
use_case = CreateGenerationTaskUseCase(generation_task_repository)
|
||||
retried = use_case.execute(
|
||||
CreateGenerationTaskCommand(
|
||||
project_id=task.project_id,
|
||||
asset_library_id=task.asset_library_id,
|
||||
strategy_id=task.strategy_id,
|
||||
voice_library_id=task.voice_library_id,
|
||||
template_id=task.template_id,
|
||||
asset_ids=task.asset_ids,
|
||||
title_ids=task.title_ids,
|
||||
voice_ids=task.voice_ids,
|
||||
created_by_user_id=authenticated_user.user.id,
|
||||
)
|
||||
)
|
||||
celery_app.send_task("worker.generate_video", args=[retried.id])
|
||||
return _to_generation_task_response(retried)
|
||||
|
||||
@@ -24,6 +24,7 @@ async def readiness_check():
|
||||
checks = {
|
||||
"database": await _check_database(),
|
||||
"redis": await _check_redis(),
|
||||
"oss": _check_oss(),
|
||||
}
|
||||
all_healthy = all(check["status"] == "healthy" for check in checks.values())
|
||||
response = {
|
||||
@@ -97,6 +98,38 @@ async def _check_redis() -> dict:
|
||||
}
|
||||
|
||||
|
||||
def _check_oss() -> dict:
|
||||
try:
|
||||
from app.core.storage import get_storage_service
|
||||
|
||||
svc = get_storage_service()
|
||||
if not svc.access_key_id or not svc.access_key_secret:
|
||||
return {
|
||||
"status": "unhealthy",
|
||||
"type": "oss",
|
||||
"message": "OSS credentials not configured (OSS_ACCESS_KEY_ID / OSS_ACCESS_KEY_SECRET missing)",
|
||||
}
|
||||
if svc.bucket is None:
|
||||
return {
|
||||
"status": "unhealthy",
|
||||
"type": "oss",
|
||||
"message": "OSS SDK (oss2) not installed or bucket client init failed",
|
||||
}
|
||||
# Try a lightweight OSS API call to verify connectivity & credentials
|
||||
svc.bucket.get_bucket_info()
|
||||
return {
|
||||
"status": "healthy",
|
||||
"type": "oss",
|
||||
"message": f"OSS connected: endpoint={svc.endpoint} bucket={svc.bucket_name}",
|
||||
}
|
||||
except Exception as error:
|
||||
return {
|
||||
"status": "unhealthy",
|
||||
"type": "oss",
|
||||
"message": f"OSS check failed: {type(error).__name__}: {error}",
|
||||
}
|
||||
|
||||
|
||||
async def _check_migrations() -> dict:
|
||||
if settings.USE_IN_MEMORY_DB:
|
||||
return {
|
||||
|
||||
@@ -22,8 +22,10 @@ router = APIRouter()
|
||||
def _to_project_response(item) -> ProjectResponse:
|
||||
return ProjectResponse(
|
||||
id=item.id,
|
||||
owner_user_id=item.owner_user_id,
|
||||
name=item.name,
|
||||
description=item.description,
|
||||
shared_users=item.shared_users,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -7,7 +7,12 @@ from app.dependencies import (
|
||||
get_ingest_job_repository,
|
||||
get_project_repository,
|
||||
)
|
||||
from app.schemas.task_center import ListProjectTasksResponse, ProjectTaskResponse
|
||||
from app.schemas.task_center import (
|
||||
ListProjectTasksResponse,
|
||||
ListTasksResponse,
|
||||
ProjectTaskResponse,
|
||||
UserTaskResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from packages.application import (
|
||||
@@ -34,28 +39,136 @@ def _humanize_task_error(error_message: str) -> str:
|
||||
return f"任务失败:{raw}"
|
||||
|
||||
|
||||
def _status_value(status) -> str:
|
||||
"""安全获取状态值(兼容 StrEnum 和 plain string)。"""
|
||||
return status.value if hasattr(status, "value") else str(status)
|
||||
|
||||
|
||||
def _generation_step(task) -> str:
|
||||
if task.status.value == "pending":
|
||||
s = _status_value(task.status)
|
||||
if s == "pending":
|
||||
return "等待 Worker 执行"
|
||||
if task.status.value == "running":
|
||||
if s == "running":
|
||||
return "正在生成成片"
|
||||
if task.status.value == "completed":
|
||||
if s == "completed":
|
||||
return "生成完成"
|
||||
if task.status.value == "failed":
|
||||
if s == "failed":
|
||||
return "生成失败"
|
||||
return task.status.value
|
||||
return s
|
||||
|
||||
|
||||
def _ingest_step(job) -> str:
|
||||
if job.status.value == "pending":
|
||||
s = _status_value(job.status)
|
||||
if s == "pending":
|
||||
return "等待导入"
|
||||
if job.status.value == "processing":
|
||||
if s == "processing":
|
||||
return "正在分析素材"
|
||||
if job.status.value == "completed":
|
||||
if s == "completed":
|
||||
return "导入完成"
|
||||
if job.status.value == "failed":
|
||||
if s == "failed":
|
||||
return "导入失败"
|
||||
return job.status.value
|
||||
return s
|
||||
|
||||
|
||||
def _generation_task_to_project_response(task) -> ProjectTaskResponse:
|
||||
return ProjectTaskResponse(
|
||||
id=f"generation:{task.id}",
|
||||
task_type="generation",
|
||||
project_id=task.project_id,
|
||||
status=_status_value(task.status),
|
||||
progress=task.progress,
|
||||
current_step=_generation_step(task),
|
||||
error_message=task.error_message,
|
||||
user_message=_humanize_task_error(task.error_message),
|
||||
retryable=_status_value(task.status) == "failed",
|
||||
source_id=task.id,
|
||||
template_id=task.template_id,
|
||||
created_at=task.created_at,
|
||||
updated_at=task.completed_at or task.started_at or task.created_at,
|
||||
)
|
||||
|
||||
|
||||
# ── 用户级端点(放在项目级端点之前,避免路由冲突) ──
|
||||
|
||||
|
||||
@router.get("/tasks", response_model=ListTasksResponse)
|
||||
def list_user_tasks(
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
ingest_job_repository: Any = Depends(get_ingest_job_repository),
|
||||
generation_task_repository: Any = Depends(get_generation_task_repository),
|
||||
) -> ListTasksResponse:
|
||||
"""用户级任务列表(跨 project),合并 ingest + generation 任务。"""
|
||||
user_id = authenticated_user.user.id
|
||||
items: list[UserTaskResponse] = []
|
||||
|
||||
for task in generation_task_repository.list_by_user(user_id):
|
||||
items.append(
|
||||
UserTaskResponse(
|
||||
id=f"generation:{task.id}",
|
||||
task_type="generation",
|
||||
project_id=task.project_id,
|
||||
template_id=task.template_id,
|
||||
status=_status_value(task.status),
|
||||
progress=task.progress,
|
||||
current_step=_generation_step(task),
|
||||
error_message=task.error_message,
|
||||
user_message=_humanize_task_error(task.error_message),
|
||||
retryable=_status_value(task.status) == "failed",
|
||||
source_id=task.id,
|
||||
created_at=task.created_at,
|
||||
updated_at=task.completed_at or task.started_at or task.created_at,
|
||||
)
|
||||
)
|
||||
|
||||
items.sort(key=lambda item: item.updated_at or item.created_at or "", reverse=True)
|
||||
return ListTasksResponse(items=items)
|
||||
|
||||
|
||||
@router.post("/tasks/{task_id}/retry", response_model=UserTaskResponse)
|
||||
def retry_task_by_id(
|
||||
task_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
generation_task_repository: Any = Depends(get_generation_task_repository),
|
||||
) -> UserTaskResponse:
|
||||
"""简化重试:通过 task_id 直接重试失败的生成任务。"""
|
||||
task = generation_task_repository.get(task_id)
|
||||
if task is None:
|
||||
raise HTTPException(status_code=404, detail="Generation task not found")
|
||||
if task.created_by_user_id and task.created_by_user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="Access denied to this task")
|
||||
if _status_value(task.status) != "failed":
|
||||
raise HTTPException(status_code=409, detail="Only failed tasks can be retried")
|
||||
|
||||
use_case = CreateGenerationTaskUseCase(generation_task_repository)
|
||||
retried = use_case.execute(
|
||||
CreateGenerationTaskCommand(
|
||||
project_id=task.project_id,
|
||||
asset_library_id=task.asset_library_id,
|
||||
strategy_id=task.strategy_id,
|
||||
voice_library_id=task.voice_library_id,
|
||||
template_id=task.template_id,
|
||||
asset_ids=task.asset_ids,
|
||||
title_ids=task.title_ids,
|
||||
voice_ids=task.voice_ids,
|
||||
created_by_user_id=authenticated_user.user.id,
|
||||
)
|
||||
)
|
||||
celery_app.send_task("worker.generate_video", args=[retried.id])
|
||||
return UserTaskResponse(
|
||||
id=f"generation:{retried.id}",
|
||||
task_type="generation",
|
||||
project_id=retried.project_id,
|
||||
template_id=retried.template_id,
|
||||
status=_status_value(retried.status),
|
||||
progress=retried.progress,
|
||||
current_step=_generation_step(retried),
|
||||
source_id=retried.id,
|
||||
created_at=retried.created_at,
|
||||
updated_at=retried.created_at,
|
||||
)
|
||||
|
||||
|
||||
# ── 项目级端点 ──
|
||||
|
||||
|
||||
@router.get("/projects/{project_id}/tasks", response_model=ListProjectTasksResponse)
|
||||
@@ -77,34 +190,19 @@ def list_project_tasks(
|
||||
id=f"ingest:{job.id}",
|
||||
task_type="ingest",
|
||||
project_id=job.project_id,
|
||||
status=job.status.value,
|
||||
progress=100.0 if job.status.value == "completed" else 0.0,
|
||||
status=_status_value(job.status),
|
||||
progress=100.0 if _status_value(job.status) == "completed" else 0.0,
|
||||
current_step=_ingest_step(job),
|
||||
error_message=job.error_message,
|
||||
user_message=_humanize_task_error(job.error_message),
|
||||
retryable=job.status.value == "failed",
|
||||
retryable=_status_value(job.status) == "failed",
|
||||
source_id=job.id,
|
||||
created_at=job.created_at,
|
||||
updated_at=job.updated_at,
|
||||
)
|
||||
)
|
||||
for task in generation_task_repository.list_by_project(project_id):
|
||||
items.append(
|
||||
ProjectTaskResponse(
|
||||
id=f"generation:{task.id}",
|
||||
task_type="generation",
|
||||
project_id=task.project_id,
|
||||
status=task.status.value,
|
||||
progress=task.progress,
|
||||
current_step=_generation_step(task),
|
||||
error_message=task.error_message,
|
||||
user_message=_humanize_task_error(task.error_message),
|
||||
retryable=task.status.value == "failed",
|
||||
source_id=task.id,
|
||||
created_at=task.created_at,
|
||||
updated_at=task.completed_at or task.started_at or task.created_at,
|
||||
)
|
||||
)
|
||||
items.append(_generation_task_to_project_response(task))
|
||||
items.sort(key=lambda item: item.updated_at or item.created_at or "", reverse=True)
|
||||
return ListProjectTasksResponse(items=items)
|
||||
|
||||
@@ -121,7 +219,7 @@ def retry_project_task(
|
||||
task = generation_task_repository.get(source_id)
|
||||
if task is None:
|
||||
raise HTTPException(status_code=404, detail="Generation task not found")
|
||||
if task.status.value != "failed":
|
||||
if _status_value(task.status) != "failed":
|
||||
raise HTTPException(status_code=409, detail="Only failed tasks can be retried")
|
||||
use_case = CreateGenerationTaskUseCase(generation_task_repository)
|
||||
retried = use_case.execute(
|
||||
@@ -130,27 +228,20 @@ def retry_project_task(
|
||||
asset_library_id=task.asset_library_id,
|
||||
strategy_id=task.strategy_id,
|
||||
voice_library_id=task.voice_library_id,
|
||||
edit_plan_id=task.edit_plan_id,
|
||||
template_id=task.template_id,
|
||||
asset_ids=task.asset_ids,
|
||||
title_ids=task.title_ids,
|
||||
voice_ids=task.voice_ids,
|
||||
created_by_user_id=authenticated_user.user.id,
|
||||
)
|
||||
)
|
||||
celery_app.send_task("worker.generate_video", args=[retried.id])
|
||||
return ProjectTaskResponse(
|
||||
id=f"generation:{retried.id}",
|
||||
task_type="generation",
|
||||
project_id=retried.project_id,
|
||||
status=retried.status.value,
|
||||
progress=retried.progress,
|
||||
current_step=_generation_step(retried),
|
||||
source_id=retried.id,
|
||||
created_at=retried.created_at,
|
||||
updated_at=retried.created_at,
|
||||
)
|
||||
return _generation_task_to_project_response(retried)
|
||||
if task_type == "ingest":
|
||||
job = ingest_job_repository.get(source_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail="Ingest job not found")
|
||||
if job.status.value != "failed":
|
||||
if _status_value(job.status) != "failed":
|
||||
raise HTTPException(status_code=409, detail="Only failed tasks can be retried")
|
||||
use_case = SubmitIngestJobUseCase(ingest_job_repository)
|
||||
retried = use_case.execute(
|
||||
@@ -165,7 +256,7 @@ def retry_project_task(
|
||||
id=f"ingest:{retried.id}",
|
||||
task_type="ingest",
|
||||
project_id=retried.project_id,
|
||||
status=retried.status.value,
|
||||
status=_status_value(retried.status),
|
||||
progress=0,
|
||||
current_step=_ingest_step(retried),
|
||||
source_id=retried.id,
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import logging
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
@@ -24,6 +25,8 @@ from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile, s
|
||||
|
||||
from packages.application import GetProjectUseCase, SubmitIngestJobCommand, SubmitIngestJobUseCase
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# 允许上传的文件 MIME 类型
|
||||
@@ -70,7 +73,7 @@ def _require_project_and_library(
|
||||
if project is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
|
||||
|
||||
libraries = asset_library_repository.list_by_project(project_id)
|
||||
libraries = asset_library_repository.find_by_project(project_id)
|
||||
if not any(item.id == library_id for item in libraries):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Asset library not found")
|
||||
|
||||
@@ -131,7 +134,14 @@ async def prepare_direct_upload(
|
||||
expires_seconds=settings.OSS_DIRECT_UPLOAD_EXPIRE_SECONDS,
|
||||
)
|
||||
except RuntimeError as error:
|
||||
logger.error("OSS not configured for direct upload prepare: %s", error)
|
||||
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=str(error)) from error
|
||||
except Exception as error:
|
||||
logger.exception("Unexpected error in direct upload prepare: %s", error)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to prepare upload: {type(error).__name__}",
|
||||
) from error
|
||||
|
||||
return DirectUploadPrepareResponse(
|
||||
upload_url=str(payload["url"]),
|
||||
@@ -162,7 +172,15 @@ async def complete_direct_upload(
|
||||
normalized_key = storage_service._normalize_storage_key(request.storage_key)
|
||||
if not normalized_key.startswith("uploads/"):
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid upload key")
|
||||
if not storage_service.file_exists(normalized_key):
|
||||
try:
|
||||
file_exists = storage_service.file_exists(normalized_key)
|
||||
except Exception as error:
|
||||
logger.exception("OSS error checking file existence for key=%s: %s", normalized_key, error)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Storage service unavailable",
|
||||
) from error
|
||||
if not file_exists:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Uploaded file not found")
|
||||
|
||||
job = _submit_ingest_job(
|
||||
@@ -181,7 +199,8 @@ async def complete_direct_upload(
|
||||
description="上传素材文件(multipart/form-data),支持视频、音频、图片。触发导入流水线自动处理。",
|
||||
)
|
||||
async def upload_asset(
|
||||
form_data: Annotated[UploadAssetRequest, Form()],
|
||||
project_id: str = Form(..., min_length=1, description="项目 ID"),
|
||||
library_id: str = Form(..., min_length=1, description="素材库 ID"),
|
||||
file: UploadFile = File(..., description="要上传的文件(视频、音频、图片等)"),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
ingest_job_repository: Any = Depends(get_ingest_job_repository),
|
||||
@@ -190,8 +209,6 @@ async def upload_asset(
|
||||
storage_service: OSSStorageService = Depends(get_storage_service),
|
||||
) -> UploadAssetResponse:
|
||||
"""上传素材文件并触发导入流水线。"""
|
||||
project_id = form_data.project_id
|
||||
library_id = form_data.library_id
|
||||
_require_project_and_library(project_id, library_id, project_repository, asset_library_repository)
|
||||
|
||||
# P2-5: 服务端验证 MIME 类型
|
||||
@@ -201,11 +218,21 @@ async def upload_asset(
|
||||
safe_filename = file.filename.replace("/", "_").replace("\\", "_") if file.filename else "unknown"
|
||||
storage_key = f"uploads/{file_id}/{safe_filename}"
|
||||
|
||||
file_url = storage_service.upload_file(
|
||||
file.file,
|
||||
storage_key,
|
||||
content_type=validated_content_type,
|
||||
)
|
||||
try:
|
||||
file_url = storage_service.upload_file(
|
||||
file.file,
|
||||
storage_key,
|
||||
content_type=validated_content_type,
|
||||
)
|
||||
except RuntimeError as error:
|
||||
logger.error("OSS not configured for upload: %s", error)
|
||||
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=str(error)) from error
|
||||
except Exception as error:
|
||||
logger.exception("Unexpected error uploading file to OSS: %s", error)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to upload file: {type(error).__name__}",
|
||||
) from error
|
||||
|
||||
job = _submit_ingest_job(
|
||||
project_id=project_id,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import field_validator
|
||||
from pydantic import AliasChoices, Field, field_validator
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
@@ -83,8 +83,11 @@ class Settings(BaseSettings):
|
||||
OSS_ACCESS_KEY_ID: str = ""
|
||||
OSS_ACCESS_KEY_SECRET: str = ""
|
||||
OSS_BUCKET_NAME: str = "xiaoxia-autocut"
|
||||
OSS_DIRECT_UPLOAD_MAX_MB: int = 800
|
||||
OSS_DIRECT_UPLOAD_EXPRESS_SECRET: int = 900
|
||||
OSS_DIRECT_UPLOAD_MAX_MB: int = Field(
|
||||
default=2000,
|
||||
validation_alias=AliasChoices("OSS_DIRECT_UPLOAD_MAX_MB", "MAX_UPLOAD_SIZE_MB"),
|
||||
)
|
||||
OSS_DIRECT_UPLOAD_EXPIRE_SECONDS: int = 900
|
||||
|
||||
LOG_LEVEL: str = "INFO"
|
||||
CORS_ORIGINS_RAW: str = (
|
||||
|
||||
@@ -28,17 +28,38 @@ class OSSStorageService:
|
||||
self.local_url_prefix = os.getenv("GENERATED_FILES_URL_PREFIX", "/generated-files")
|
||||
self.bucket = None
|
||||
|
||||
if settings.OSS_ACCESS_KEY_ID and settings.OSS_ACCESS_KEY_SECRET:
|
||||
has_key_id = bool(settings.OSS_ACCESS_KEY_ID)
|
||||
has_key_secret = bool(settings.OSS_ACCESS_KEY_SECRET)
|
||||
|
||||
if has_key_id and has_key_secret:
|
||||
if oss2 is not None:
|
||||
auth = oss2.Auth(
|
||||
settings.OSS_ACCESS_KEY_ID,
|
||||
settings.OSS_ACCESS_KEY_SECRET,
|
||||
)
|
||||
self.bucket = oss2.Bucket(
|
||||
auth,
|
||||
settings.OSS_ENDPOINT,
|
||||
settings.OSS_BUCKET_NAME,
|
||||
)
|
||||
try:
|
||||
auth = oss2.Auth(
|
||||
settings.OSS_ACCESS_KEY_ID,
|
||||
settings.OSS_ACCESS_KEY_SECRET,
|
||||
)
|
||||
self.bucket = oss2.Bucket(
|
||||
auth,
|
||||
settings.OSS_ENDPOINT,
|
||||
settings.OSS_BUCKET_NAME,
|
||||
)
|
||||
logger.info(
|
||||
"OSS initialized: endpoint=%s bucket=%s",
|
||||
settings.OSS_ENDPOINT,
|
||||
settings.OSS_BUCKET_NAME,
|
||||
)
|
||||
except Exception as error:
|
||||
logger.error("Failed to initialize OSS bucket client: %s", error)
|
||||
else:
|
||||
logger.error("oss2 SDK is not installed — OSS operations will fail")
|
||||
else:
|
||||
missing = []
|
||||
if not has_key_id:
|
||||
missing.append("OSS_ACCESS_KEY_ID")
|
||||
if not has_key_secret:
|
||||
missing.append("OSS_ACCESS_KEY_SECRET")
|
||||
logger.error("OSS credentials not configured — missing: %s", ", ".join(missing))
|
||||
|
||||
self.access_key_id = settings.OSS_ACCESS_KEY_ID
|
||||
self.access_key_secret = settings.OSS_ACCESS_KEY_SECRET
|
||||
self.endpoint = settings.OSS_ENDPOINT
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class RecentTaskItem(BaseModel):
|
||||
id: str
|
||||
task_type: str = "generation"
|
||||
status: str
|
||||
current_step: str = ""
|
||||
error_message: str = ""
|
||||
updated_at: datetime | None = None
|
||||
|
||||
|
||||
class SubscriptionInfo(BaseModel):
|
||||
"""用户订阅信息。"""
|
||||
plan: str = "free"
|
||||
is_active: bool = False
|
||||
|
||||
|
||||
class DashboardOverviewResponse(BaseModel):
|
||||
"""Dashboard 概览数据。"""
|
||||
total_assets: int = 0
|
||||
used_storage_bytes: int = 0
|
||||
total_titles: int = 0
|
||||
total_voices: int = 0
|
||||
total_tasks: int = 0
|
||||
total_products: int = 0
|
||||
subscription: SubscriptionInfo = Field(default_factory=SubscriptionInfo)
|
||||
recent_tasks: list[RecentTaskItem] = Field(default_factory=list)
|
||||
@@ -1,12 +1,35 @@
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
|
||||
class CreateGenerationTaskRequest(BaseModel):
|
||||
project_id: str = Field(..., min_length=1)
|
||||
asset_library_id: str = Field(..., min_length=1)
|
||||
"""创建生成任务请求。
|
||||
|
||||
支持两种模式(至少提供一种):
|
||||
- 项目模式:project_id + asset_library_id(向后兼容)
|
||||
- 模板模式:template_id + asset_ids / title_ids / voice_ids
|
||||
"""
|
||||
project_id: str = ""
|
||||
asset_library_id: str = ""
|
||||
strategy_id: str = ""
|
||||
voice_library_id: str = ""
|
||||
created_by_user_id: str = ""
|
||||
# ── 模板模式新增字段 ──
|
||||
template_id: str = ""
|
||||
asset_ids: list[str] = Field(default_factory=list)
|
||||
title_ids: list[str] = Field(default_factory=list)
|
||||
voice_ids: list[str] = Field(default_factory=list)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_at_least_one_mode(self) -> "CreateGenerationTaskRequest":
|
||||
has_project = bool(self.project_id.strip())
|
||||
has_template = bool(self.template_id.strip())
|
||||
if not has_project and not has_template:
|
||||
raise ValueError("project_id 或 template_id 至少需要提供一个")
|
||||
has_library = bool(self.asset_library_id.strip())
|
||||
has_assets = bool(self.asset_ids or self.title_ids or self.voice_ids)
|
||||
if not has_library and not has_assets:
|
||||
raise ValueError("asset_library_id 或 asset_ids/title_ids/voice_ids 至少需要提供一个")
|
||||
return self
|
||||
|
||||
|
||||
class GenerationTaskResponse(BaseModel):
|
||||
@@ -15,7 +38,16 @@ class GenerationTaskResponse(BaseModel):
|
||||
asset_library_id: str
|
||||
strategy_id: str
|
||||
voice_library_id: str
|
||||
template_id: str = ""
|
||||
asset_ids: list[str] = Field(default_factory=list)
|
||||
title_ids: list[str] = Field(default_factory=list)
|
||||
voice_ids: list[str] = Field(default_factory=list)
|
||||
status: str
|
||||
progress: float
|
||||
result_count: int
|
||||
error_message: str
|
||||
|
||||
|
||||
class ListGenerationTasksResponse(BaseModel):
|
||||
"""用户级生成任务列表响应(跨 project)。"""
|
||||
items: list[GenerationTaskResponse]
|
||||
|
||||
@@ -14,9 +14,32 @@ class ProjectTaskResponse(BaseModel):
|
||||
user_message: str = ""
|
||||
retryable: bool = False
|
||||
source_id: str = ""
|
||||
template_id: str = ""
|
||||
created_at: datetime | None = None
|
||||
updated_at: datetime | None = None
|
||||
|
||||
|
||||
class ListProjectTasksResponse(BaseModel):
|
||||
items: list[ProjectTaskResponse] = Field(default_factory=list)
|
||||
|
||||
|
||||
class UserTaskResponse(BaseModel):
|
||||
"""用户级任务响应(跨 project,用于模板模式)。"""
|
||||
id: str
|
||||
task_type: str
|
||||
project_id: str = ""
|
||||
template_id: str = ""
|
||||
status: str
|
||||
progress: float
|
||||
current_step: str
|
||||
error_message: str = ""
|
||||
user_message: str = ""
|
||||
retryable: bool = False
|
||||
source_id: str = ""
|
||||
created_at: datetime | None = None
|
||||
updated_at: datetime | None = None
|
||||
|
||||
|
||||
class ListTasksResponse(BaseModel):
|
||||
"""用户级任务列表响应(GET /api/v1/tasks)。"""
|
||||
items: list[UserTaskResponse] = Field(default_factory=list)
|
||||
|
||||
@@ -154,6 +154,7 @@ export const uploadAsset = async (
|
||||
|
||||
/** 预签名直传准备 */
|
||||
export const prepareDirectUpload = async (data: {
|
||||
project_id: string;
|
||||
library_id: string;
|
||||
filename: string;
|
||||
content_type: string;
|
||||
@@ -172,6 +173,7 @@ export const prepareDirectUpload = async (data: {
|
||||
|
||||
/** 直传完成确认 */
|
||||
export const completeDirectUpload = async (data: {
|
||||
project_id: string;
|
||||
library_id: string;
|
||||
storage_key: string;
|
||||
}): Promise<{ storage_key: string; ingest_job_id: string }> => {
|
||||
@@ -184,7 +186,11 @@ export const uploadAssetDirect = async (data: {
|
||||
file: File;
|
||||
library_id: string;
|
||||
}): Promise<{ storage_key: string; ingest_job_id: string }> => {
|
||||
// 后端要求 project_id,前端自动获取默认项目
|
||||
const project = await getOrCreateDefaultProject();
|
||||
|
||||
const prepared = await prepareDirectUpload({
|
||||
project_id: project.id,
|
||||
library_id: data.library_id,
|
||||
filename: data.file.name,
|
||||
content_type: data.file.type || 'application/octet-stream',
|
||||
@@ -206,6 +212,7 @@ export const uploadAssetDirect = async (data: {
|
||||
}
|
||||
|
||||
return completeDirectUpload({
|
||||
project_id: project.id,
|
||||
library_id: data.library_id,
|
||||
storage_key: prepared.storage_key,
|
||||
});
|
||||
|
||||
@@ -1,98 +0,0 @@
|
||||
/**
|
||||
* 编辑计划 API
|
||||
* Phase 1 重构:去掉 projectId,编辑计划直接归属用户
|
||||
*/
|
||||
import apiClient from './client';
|
||||
|
||||
/** 编辑模式 */
|
||||
export type EditingMode = 'one-take' | 'pip' | 'voiceover' | 'voice_pip';
|
||||
|
||||
/** 编辑模板 */
|
||||
export interface EditTemplateItem {
|
||||
id: string;
|
||||
name: string;
|
||||
description: string;
|
||||
target_duration: number;
|
||||
clip_count: number;
|
||||
is_active: boolean;
|
||||
category?: string;
|
||||
thumbnail_url?: string;
|
||||
}
|
||||
|
||||
/** 编辑计划片段 */
|
||||
export interface EditPlanClipItem {
|
||||
id: string;
|
||||
asset_id: string;
|
||||
asset_name: string;
|
||||
sequence: number;
|
||||
start_time: number;
|
||||
duration: number;
|
||||
reason: string;
|
||||
layer?: 'main' | 'pip' | 'broll';
|
||||
thumbnail_url?: string;
|
||||
}
|
||||
|
||||
/** 编辑计划 */
|
||||
export interface EditPlanItem {
|
||||
id: string;
|
||||
template_id: string;
|
||||
asset_library_id: string;
|
||||
title_id: string;
|
||||
status: string;
|
||||
editing_mode?: EditingMode;
|
||||
summary: string;
|
||||
clips: EditPlanClipItem[];
|
||||
created_at?: string;
|
||||
updated_at?: string;
|
||||
}
|
||||
|
||||
// ─── 编辑计划 ──────────────────────────────────────────────
|
||||
|
||||
/** 获取当前用户的编辑计划列表 */
|
||||
export const getEditPlans = async (): Promise<EditPlanItem[]> => {
|
||||
const response = await apiClient.get('/edit-plans');
|
||||
return response.data.items || [];
|
||||
};
|
||||
|
||||
/** 获取单个编辑计划 */
|
||||
export const getEditPlan = async (planId: string): Promise<EditPlanItem> => {
|
||||
const response = await apiClient.get(`/edit-plans/${planId}`);
|
||||
return response.data;
|
||||
};
|
||||
|
||||
/** 创建编辑计划 */
|
||||
export const createEditPlan = async (data: {
|
||||
asset_library_id: string;
|
||||
template_id?: string;
|
||||
title_id?: string;
|
||||
}): Promise<EditPlanItem> => {
|
||||
const response = await apiClient.post('/edit-plans', data);
|
||||
return response.data;
|
||||
};
|
||||
|
||||
/** 智能编排 - 自动生成编辑计划 */
|
||||
export const autoGenerateEditPlan = async (params: {
|
||||
template_id: string;
|
||||
asset_ids?: string[];
|
||||
title_ids?: string[];
|
||||
voice_ids?: string[];
|
||||
editing_mode?: EditingMode;
|
||||
target_duration?: number;
|
||||
}): Promise<EditPlanItem> => {
|
||||
const response = await apiClient.post('/edit-plans/auto-generate', params);
|
||||
return response.data;
|
||||
};
|
||||
|
||||
/** 更新编辑计划 */
|
||||
export const updateEditPlan = async (
|
||||
planId: string,
|
||||
data: Partial<EditPlanItem>
|
||||
): Promise<EditPlanItem> => {
|
||||
const response = await apiClient.patch(`/edit-plans/${planId}`, data);
|
||||
return response.data;
|
||||
};
|
||||
|
||||
/** 删除编辑计划 */
|
||||
export const deleteEditPlan = async (planId: string): Promise<void> => {
|
||||
await apiClient.delete(`/edit-plans/${planId}`);
|
||||
};
|
||||
@@ -1,8 +1,8 @@
|
||||
/**
|
||||
* 剪辑计划编辑器 API
|
||||
* 当前使用 mock 数据,后端 API 就绪后替换
|
||||
* 对接后端 /api/v1/templates 路由
|
||||
*/
|
||||
// import apiClient from './client'; // TODO: 后端 API 就绪后启用
|
||||
import apiClient from './client';
|
||||
|
||||
/* ──────────── 类型定义 ──────────── */
|
||||
|
||||
@@ -53,11 +53,11 @@ export interface BgmConfig {
|
||||
|
||||
/** 模板片段 */
|
||||
export interface TemplateSegment {
|
||||
id: string;
|
||||
id?: string;
|
||||
segment_order: number;
|
||||
duration_min: number;
|
||||
duration_max: number;
|
||||
material_type: string | null; // 仅 口播+混剪 模式:人物/场景
|
||||
material_type: string | null;
|
||||
}
|
||||
|
||||
/** 剪辑模板 */
|
||||
@@ -72,6 +72,7 @@ export interface EditingTemplate {
|
||||
bgm_config: BgmConfig;
|
||||
estimated_duration: number;
|
||||
segments: TemplateSegment[];
|
||||
is_active?: boolean;
|
||||
created_at: string;
|
||||
updated_at: string;
|
||||
}
|
||||
@@ -80,6 +81,7 @@ export interface EditingTemplate {
|
||||
export interface TemplateCategory {
|
||||
id: string;
|
||||
name: string;
|
||||
created_at?: string;
|
||||
}
|
||||
|
||||
/** 创建/更新模板请求体 */
|
||||
@@ -100,124 +102,46 @@ export interface GenerateFromTemplatePayload {
|
||||
voiceover_duration: number;
|
||||
}
|
||||
|
||||
/** 使用模板生成响应 */
|
||||
export interface GenerateFromTemplateResponse {
|
||||
task_id: string;
|
||||
warning?: string;
|
||||
/** 验证/生成响应 */
|
||||
export interface ValidateWarning {
|
||||
code: string;
|
||||
message: string;
|
||||
details?: Record<string, unknown>;
|
||||
}
|
||||
|
||||
/* ──────────── Mock 数据 ──────────── */
|
||||
/** 使用模板生成响应 */
|
||||
export interface GenerateFromTemplateResponse {
|
||||
template: EditingTemplate;
|
||||
warnings: ValidateWarning[];
|
||||
}
|
||||
|
||||
let _nextId = 100;
|
||||
const nextId = () => String(++_nextId);
|
||||
/** 列表响应(带分页) */
|
||||
export interface ListTemplatesResponse {
|
||||
items: EditingTemplate[];
|
||||
total: number;
|
||||
}
|
||||
|
||||
const MOCK_CATEGORIES: TemplateCategory[] = [
|
||||
{ id: 'cat-1', name: '生活' },
|
||||
{ id: 'cat-2', name: '美食' },
|
||||
{ id: 'cat-3', name: '旅行' },
|
||||
{ id: 'cat-4', name: '知识' },
|
||||
];
|
||||
/** 分类列表响应 */
|
||||
export interface ListCategoriesResponse {
|
||||
items: TemplateCategory[];
|
||||
}
|
||||
|
||||
const MOCK_TEMPLATES: EditingTemplate[] = [
|
||||
{
|
||||
id: 'tpl-1',
|
||||
name: '生活 Vlog 模板',
|
||||
mode: 'pip',
|
||||
category: '生活',
|
||||
tags: ['vlog', '日常'],
|
||||
title_config: {
|
||||
ai_auto_select: true,
|
||||
content: '',
|
||||
font_preset: '思源黑体',
|
||||
font_color: '#ffffff',
|
||||
font_size: 32,
|
||||
position: 'top',
|
||||
},
|
||||
subtitle_config: {
|
||||
enabled: true,
|
||||
position: 'bottom',
|
||||
font: '思源黑体',
|
||||
color: '#ffffff',
|
||||
size: 24,
|
||||
animation: 'fade',
|
||||
},
|
||||
bgm_config: { enabled: true, music_id: 'bgm-1' },
|
||||
estimated_duration: 30,
|
||||
segments: [
|
||||
{ id: 'seg-1', segment_order: 1, duration_min: 5, duration_max: 15, material_type: null },
|
||||
{ id: 'seg-2', segment_order: 2, duration_min: 10, duration_max: 20, material_type: null },
|
||||
],
|
||||
created_at: '2026-06-20T10:00:00Z',
|
||||
updated_at: '2026-06-20T10:00:00Z',
|
||||
},
|
||||
{
|
||||
id: 'tpl-2',
|
||||
name: '知识分享口播',
|
||||
mode: 'voice_over',
|
||||
category: '知识',
|
||||
tags: ['口播', '分享'],
|
||||
title_config: {
|
||||
ai_auto_select: false,
|
||||
content: '每日知识分享',
|
||||
font_preset: '站酷快乐体',
|
||||
font_color: '#ffdd00',
|
||||
font_size: 36,
|
||||
position: 'top',
|
||||
},
|
||||
subtitle_config: {
|
||||
enabled: true,
|
||||
position: 'bottom',
|
||||
font: '思源黑体',
|
||||
color: '#ffffff',
|
||||
size: 28,
|
||||
animation: 'typewriter',
|
||||
},
|
||||
bgm_config: { enabled: false, music_id: '' },
|
||||
estimated_duration: 60,
|
||||
segments: [
|
||||
{ id: 'seg-3', segment_order: 1, duration_min: 10, duration_max: 30, material_type: null },
|
||||
{ id: 'seg-4', segment_order: 2, duration_min: 20, duration_max: 40, material_type: null },
|
||||
{ id: 'seg-5', segment_order: 3, duration_min: 10, duration_max: 20, material_type: null },
|
||||
],
|
||||
created_at: '2026-06-21T10:00:00Z',
|
||||
updated_at: '2026-06-21T10:00:00Z',
|
||||
},
|
||||
{
|
||||
id: 'tpl-3',
|
||||
name: '一镜到底展示',
|
||||
mode: 'one_take',
|
||||
category: '生活',
|
||||
tags: ['一镜到底'],
|
||||
title_config: {
|
||||
ai_auto_select: true,
|
||||
content: '',
|
||||
font_preset: '思源黑体',
|
||||
font_color: '#ffffff',
|
||||
font_size: 32,
|
||||
position: 'center',
|
||||
},
|
||||
subtitle_config: { enabled: false, position: 'bottom', font: '思源黑体', color: '#ffffff', size: 24, animation: 'fade' },
|
||||
bgm_config: { enabled: true, music_id: 'bgm-2' },
|
||||
estimated_duration: 15,
|
||||
segments: [
|
||||
{ id: 'seg-6', segment_order: 1, duration_min: 10, duration_max: 20, material_type: null },
|
||||
],
|
||||
created_at: '2026-06-22T10:00:00Z',
|
||||
updated_at: '2026-06-22T10:00:00Z',
|
||||
},
|
||||
];
|
||||
|
||||
/* ──────────── Mock API 函数 ──────────── */
|
||||
|
||||
const delay = (ms = 200) => new Promise((r) => setTimeout(r, ms));
|
||||
// ============ API 函数 ============
|
||||
|
||||
/** 获取模板列表 */
|
||||
export const getEditingTemplates = async (params?: {
|
||||
category?: string;
|
||||
tag?: string;
|
||||
skip?: number;
|
||||
limit?: number;
|
||||
}): Promise<EditingTemplate[]> => {
|
||||
await delay();
|
||||
let list = [...MOCK_TEMPLATES];
|
||||
const response = await apiClient.get<ListTemplatesResponse>('/templates', {
|
||||
params: {
|
||||
skip: params?.skip ?? 0,
|
||||
limit: params?.limit ?? 50,
|
||||
},
|
||||
});
|
||||
let list = response.data.items;
|
||||
if (params?.category) list = list.filter((t) => t.category === params.category);
|
||||
if (params?.tag) list = list.filter((t) => t.tags.includes(params.tag!));
|
||||
return list;
|
||||
@@ -225,38 +149,16 @@ export const getEditingTemplates = async (params?: {
|
||||
|
||||
/** 获取模板详情 */
|
||||
export const getEditingTemplate = async (id: string): Promise<EditingTemplate> => {
|
||||
await delay();
|
||||
const tpl = MOCK_TEMPLATES.find((t) => t.id === id);
|
||||
if (!tpl) throw new Error('模板不存在');
|
||||
return { ...tpl };
|
||||
const response = await apiClient.get<EditingTemplate>(`/templates/${id}`);
|
||||
return response.data;
|
||||
};
|
||||
|
||||
/** 创建模板 */
|
||||
export const createEditingTemplate = async (
|
||||
data: SaveTemplatePayload,
|
||||
): Promise<EditingTemplate> => {
|
||||
await delay(300);
|
||||
const now = new Date().toISOString();
|
||||
const tpl: EditingTemplate = {
|
||||
id: nextId(),
|
||||
name: data.name,
|
||||
mode: data.mode,
|
||||
category: data.category,
|
||||
tags: data.tags,
|
||||
title_config: data.title_config,
|
||||
subtitle_config: data.subtitle_config,
|
||||
bgm_config: data.bgm_config,
|
||||
estimated_duration: data.estimated_duration,
|
||||
segments: data.segments.map((s, i) => ({
|
||||
...s,
|
||||
id: nextId(),
|
||||
segment_order: i + 1,
|
||||
})),
|
||||
created_at: now,
|
||||
updated_at: now,
|
||||
};
|
||||
MOCK_TEMPLATES.push(tpl);
|
||||
return tpl;
|
||||
const response = await apiClient.post<EditingTemplate>('/templates', data);
|
||||
return response.data;
|
||||
};
|
||||
|
||||
/** 更新模板 */
|
||||
@@ -264,51 +166,29 @@ export const updateEditingTemplate = async (
|
||||
id: string,
|
||||
data: SaveTemplatePayload,
|
||||
): Promise<EditingTemplate> => {
|
||||
await delay(300);
|
||||
const idx = MOCK_TEMPLATES.findIndex((t) => t.id === id);
|
||||
if (idx === -1) throw new Error('模板不存在');
|
||||
const updated: EditingTemplate = {
|
||||
...MOCK_TEMPLATES[idx],
|
||||
name: data.name,
|
||||
mode: data.mode,
|
||||
category: data.category,
|
||||
tags: data.tags,
|
||||
title_config: data.title_config,
|
||||
subtitle_config: data.subtitle_config,
|
||||
bgm_config: data.bgm_config,
|
||||
estimated_duration: data.estimated_duration,
|
||||
segments: data.segments.map((s, i) => ({
|
||||
...s,
|
||||
id: nextId(),
|
||||
segment_order: i + 1,
|
||||
})),
|
||||
updated_at: new Date().toISOString(),
|
||||
};
|
||||
MOCK_TEMPLATES[idx] = updated;
|
||||
return updated;
|
||||
const response = await apiClient.patch<EditingTemplate>(`/templates/${id}`, data);
|
||||
return response.data;
|
||||
};
|
||||
|
||||
/** 删除模板 */
|
||||
export const deleteEditingTemplate = async (id: string): Promise<void> => {
|
||||
await delay();
|
||||
const idx = MOCK_TEMPLATES.findIndex((t) => t.id === id);
|
||||
if (idx !== -1) MOCK_TEMPLATES.splice(idx, 1);
|
||||
await apiClient.delete(`/templates/${id}`);
|
||||
};
|
||||
|
||||
/** 获取模板分类列表 */
|
||||
export const getTemplateCategories = async (): Promise<TemplateCategory[]> => {
|
||||
await delay();
|
||||
return [...MOCK_CATEGORIES];
|
||||
const response = await apiClient.get<ListCategoriesResponse>('/templates/categories/list');
|
||||
return response.data.items;
|
||||
};
|
||||
|
||||
/** 使用模板生成视频 */
|
||||
/** 使用模板生成视频(调用 validate 端点) */
|
||||
export const generateFromTemplate = async (
|
||||
_templateId: string,
|
||||
_data: GenerateFromTemplatePayload,
|
||||
templateId: string,
|
||||
data: GenerateFromTemplatePayload,
|
||||
): Promise<GenerateFromTemplateResponse> => {
|
||||
await delay(500);
|
||||
return {
|
||||
task_id: nextId(),
|
||||
warning: undefined,
|
||||
};
|
||||
const response = await apiClient.post<GenerateFromTemplateResponse>(
|
||||
`/templates/${templateId}/validate`,
|
||||
data,
|
||||
);
|
||||
return response.data;
|
||||
};
|
||||
|
||||
+37
-29
@@ -1,13 +1,20 @@
|
||||
/**
|
||||
* 任务相关 API
|
||||
* Phase 1 重构:去掉 projectId,任务直接归属用户
|
||||
* 对接后端方案 A 扩展后的端点(PR #109)
|
||||
* - POST /api/v1/generation/tasks — 创建生成任务(template_id + asset_ids 细粒度模式)
|
||||
* - GET /api/v1/tasks — 用户级任务列表(跨 project)
|
||||
* - POST /api/v1/tasks/{task_id}/retry — 简化重试
|
||||
*/
|
||||
import apiClient from './client';
|
||||
|
||||
/** 任务条目 */
|
||||
/* ──────────── 类型定义 ──────────── */
|
||||
|
||||
/** 任务条目(对应用户级 UserTaskResponse) */
|
||||
export interface TaskItem {
|
||||
id: string;
|
||||
task_type: 'ingest' | 'generation' | string;
|
||||
project_id: string;
|
||||
template_id: string;
|
||||
status: string;
|
||||
progress: number;
|
||||
current_step: string;
|
||||
@@ -19,18 +26,6 @@ export interface TaskItem {
|
||||
updated_at?: string | null;
|
||||
}
|
||||
|
||||
/** 获取当前用户的所有任务(生成记录) */
|
||||
export const getUserTasks = async (): Promise<TaskItem[]> => {
|
||||
const response = await apiClient.get('/tasks');
|
||||
return response.data.items || [];
|
||||
};
|
||||
|
||||
/** 重试失败的任务 */
|
||||
export const retryTask = async (taskId: string): Promise<TaskItem> => {
|
||||
const response = await apiClient.post(`/tasks/${taskId}/retry`);
|
||||
return response.data;
|
||||
};
|
||||
|
||||
/** 创建生成任务请求参数 */
|
||||
export interface CreateGenerationTaskRequest {
|
||||
template_id: string;
|
||||
@@ -39,31 +34,44 @@ export interface CreateGenerationTaskRequest {
|
||||
voice_ids: string[];
|
||||
}
|
||||
|
||||
/** 创建生成任务响应 */
|
||||
/** 创建生成任务响应(对齐后端 GenerationTaskResponse) */
|
||||
export interface CreateGenerationTaskResponse {
|
||||
task_id: string;
|
||||
id: string;
|
||||
project_id: string;
|
||||
asset_library_id: string;
|
||||
strategy_id: string;
|
||||
voice_library_id: string;
|
||||
template_id: string;
|
||||
asset_ids: string[];
|
||||
title_ids: string[];
|
||||
voice_ids: string[];
|
||||
status: string;
|
||||
message: string;
|
||||
progress: number;
|
||||
result_count: number;
|
||||
error_message: string;
|
||||
}
|
||||
|
||||
// TODO: 后端生成接口适配扁平化架构后切换为 false
|
||||
const USE_MOCK = true;
|
||||
/* ──────────── API 函数 ──────────── */
|
||||
|
||||
/** 创建生成任务(一键生成) */
|
||||
export const createGenerationTask = async (
|
||||
params: CreateGenerationTaskRequest,
|
||||
): Promise<CreateGenerationTaskResponse> => {
|
||||
if (USE_MOCK) {
|
||||
await new Promise((r) => setTimeout(r, 800));
|
||||
return {
|
||||
task_id: `task_${Date.now()}`,
|
||||
status: 'pending',
|
||||
message: '生成任务已创建',
|
||||
};
|
||||
}
|
||||
const response = await apiClient.post<CreateGenerationTaskResponse>(
|
||||
const { data } = await apiClient.post<CreateGenerationTaskResponse>(
|
||||
'/generation/tasks',
|
||||
params,
|
||||
);
|
||||
return response.data;
|
||||
return data;
|
||||
};
|
||||
|
||||
/** 获取当前用户的所有任务(跨 project) */
|
||||
export const getUserTasks = async (): Promise<TaskItem[]> => {
|
||||
const { data } = await apiClient.get('/tasks');
|
||||
return data.items || [];
|
||||
};
|
||||
|
||||
/** 重试失败的任务 */
|
||||
export const retryTask = async (taskId: string): Promise<TaskItem> => {
|
||||
const { data } = await apiClient.post(`/tasks/${taskId}/retry`);
|
||||
return data;
|
||||
};
|
||||
|
||||
@@ -37,12 +37,18 @@ import {
|
||||
getAssets,
|
||||
deleteAsset,
|
||||
uploadAsset,
|
||||
uploadAssetDirect,
|
||||
} from '@/api/assets';
|
||||
import { getOrCreateDefaultProject } from '@/api/projects';
|
||||
|
||||
const { Title, Text } = Typography;
|
||||
const { Dragger } = Upload;
|
||||
|
||||
/** 最大文件大小:2048MB(与后端 OSS_DIRECT_UPLOAD_MAX_MB=2000 对齐,留少量余量) */
|
||||
const MAX_FILE_SIZE = 2048 * 1024 * 1024; // 2147483648 bytes
|
||||
/** 大文件阈值:100MB,超过此值走 OSS 直传通道 */
|
||||
const LARGE_FILE_THRESHOLD = 100 * 1024 * 1024; // 104857600 bytes
|
||||
|
||||
/** 素材库类型图标 */
|
||||
const KindIcon: React.FC<{ kind: string }> = ({ kind }) => {
|
||||
switch (kind) {
|
||||
@@ -99,7 +105,7 @@ const AssetLibrary: React.FC = () => {
|
||||
onError: (err: any) => { if (!err?.__msgShown) message.error('创建失败') },
|
||||
});
|
||||
|
||||
// 上传素材
|
||||
// 上传素材(小文件表单上传)
|
||||
const uploadMutation = useMutation({
|
||||
mutationFn: uploadAsset,
|
||||
onSuccess: () => {
|
||||
@@ -110,6 +116,17 @@ const AssetLibrary: React.FC = () => {
|
||||
onError: (err: any) => { if (!err?.__msgShown) message.error('上传失败') },
|
||||
});
|
||||
|
||||
// 大文件直传(OSS 预签名)
|
||||
const directUploadMutation = useMutation({
|
||||
mutationFn: uploadAssetDirect,
|
||||
onSuccess: () => {
|
||||
message.success('上传成功');
|
||||
queryClient.invalidateQueries({ queryKey: ['assets', activeLibrary] });
|
||||
queryClient.invalidateQueries({ queryKey: ['asset-libraries'] });
|
||||
},
|
||||
onError: (err: any) => { if (!err?.__msgShown) message.error('上传失败') },
|
||||
});
|
||||
|
||||
// 删除素材
|
||||
const deleteMutation = useMutation({
|
||||
mutationFn: deleteAsset,
|
||||
@@ -127,16 +144,32 @@ const AssetLibrary: React.FC = () => {
|
||||
message.warning('请先选择素材库');
|
||||
return false;
|
||||
}
|
||||
|
||||
// 文件大小校验:上限 2GB
|
||||
if (file.size > MAX_FILE_SIZE) {
|
||||
message.error('文件大小不能超过 2GB');
|
||||
return false;
|
||||
}
|
||||
|
||||
setUploading(true);
|
||||
try {
|
||||
const project = await getOrCreateDefaultProject();
|
||||
const formData = new FormData();
|
||||
formData.append('file', file);
|
||||
formData.append('library_id', activeLibrary);
|
||||
formData.append('project_id', project.id);
|
||||
await uploadMutation.mutateAsync(formData);
|
||||
if (file.size > LARGE_FILE_THRESHOLD) {
|
||||
// 大文件(>100MB)走 OSS 预签名直传通道
|
||||
await directUploadMutation.mutateAsync({
|
||||
file,
|
||||
library_id: activeLibrary,
|
||||
});
|
||||
} else {
|
||||
// 小文件走表单上传
|
||||
const project = await getOrCreateDefaultProject();
|
||||
const formData = new FormData();
|
||||
formData.append('file', file);
|
||||
formData.append('library_id', activeLibrary);
|
||||
formData.append('project_id', project.id);
|
||||
await uploadMutation.mutateAsync(formData);
|
||||
}
|
||||
} catch {
|
||||
// uploadMutation.onError 已处理错误提示
|
||||
// mutation.onError 已处理错误提示
|
||||
} finally {
|
||||
setUploading(false);
|
||||
}
|
||||
@@ -216,7 +249,7 @@ const AssetLibrary: React.FC = () => {
|
||||
beforeUpload={handleUpload}
|
||||
showUploadList={false}
|
||||
multiple
|
||||
disabled={uploading || uploadMutation.isPending}
|
||||
disabled={uploading || uploadMutation.isPending || directUploadMutation.isPending}
|
||||
style={{ marginBottom: 24 }}
|
||||
>
|
||||
<p className="ant-upload-drag-icon">
|
||||
@@ -226,7 +259,7 @@ const AssetLibrary: React.FC = () => {
|
||||
{uploading ? '正在上传,请稍候...' : '点击或拖拽文件到此区域上传'}
|
||||
</p>
|
||||
<p className="ant-upload-hint">
|
||||
支持 {kindLabel[currentLib?.kind || 'video']} 格式文件
|
||||
支持 {kindLabel[currentLib?.kind || 'video']} 格式文件,单文件不超过 2GB
|
||||
</p>
|
||||
</Dragger>
|
||||
)}
|
||||
|
||||
@@ -168,7 +168,7 @@ const EditingPlanner: React.FC = () => {
|
||||
mutationFn: ({ templateId, duration }: { templateId: string; duration: number }) =>
|
||||
generateFromTemplate(templateId, { voiceover_duration: duration }),
|
||||
onSuccess: (data) => {
|
||||
const msg = data.warning ? `生成任务已提交(${data.warning})` : '生成任务已提交';
|
||||
const msg = data.warnings && data.warnings.length > 0 ? `生成任务已提交(${data.warnings.map(w => w.message).join('; ')})` : '生成任务已提交';
|
||||
message.success(msg);
|
||||
setGenerateModalOpen(false);
|
||||
},
|
||||
|
||||
@@ -136,7 +136,7 @@ const TimelinePanel: React.FC<TimelinePanelProps> = ({
|
||||
size="small"
|
||||
danger
|
||||
icon={<DeleteOutlined />}
|
||||
onClick={() => onRemoveSegment(seg.id)}
|
||||
onClick={() => onRemoveSegment(seg.id!)}
|
||||
disabled={isOneShot}
|
||||
style={{ marginLeft: 'auto' }}
|
||||
/>
|
||||
@@ -154,7 +154,7 @@ const TimelinePanel: React.FC<TimelinePanelProps> = ({
|
||||
min={1}
|
||||
max={seg.duration_max}
|
||||
value={seg.duration_min}
|
||||
onChange={(v) => onUpdateSegment(seg.id, { duration_min: v })}
|
||||
onChange={(v) => onUpdateSegment(seg.id!, { duration_min: v })}
|
||||
/>
|
||||
</div>
|
||||
<div>
|
||||
@@ -163,7 +163,7 @@ const TimelinePanel: React.FC<TimelinePanelProps> = ({
|
||||
min={seg.duration_min}
|
||||
max={60}
|
||||
value={seg.duration_max}
|
||||
onChange={(v) => onUpdateSegment(seg.id, { duration_max: v })}
|
||||
onChange={(v) => onUpdateSegment(seg.id!, { duration_max: v })}
|
||||
/>
|
||||
</div>
|
||||
</>
|
||||
@@ -175,7 +175,7 @@ const TimelinePanel: React.FC<TimelinePanelProps> = ({
|
||||
<Select
|
||||
size="small"
|
||||
value={seg.material_type || '人物'}
|
||||
onChange={(v) => onUpdateSegment(seg.id, { material_type: v })}
|
||||
onChange={(v) => onUpdateSegment(seg.id!, { material_type: v })}
|
||||
style={{ width: '100%', marginTop: 4 }}
|
||||
options={[
|
||||
{ value: '人物', label: '人物' },
|
||||
|
||||
@@ -8,7 +8,7 @@ server {
|
||||
gzip_vary on;
|
||||
gzip_min_length 1024;
|
||||
gzip_types text/plain text/css text/xml text/javascript application/javascript application/json application/xml+rss;
|
||||
client_max_body_size 800m;
|
||||
client_max_body_size 2g;
|
||||
|
||||
location /api/ {
|
||||
proxy_pass http://xiaoxia-api-production:8000/api/;
|
||||
|
||||
@@ -28,7 +28,7 @@ server {
|
||||
include /etc/letsencrypt/options-ssl-nginx.conf;
|
||||
ssl_dhparam /etc/letsencrypt/ssl-dhparams.pem;
|
||||
|
||||
client_max_body_size 800m;
|
||||
client_max_body_size 2g;
|
||||
|
||||
location / {
|
||||
proxy_pass http://127.0.0.1:3002/;
|
||||
@@ -51,7 +51,7 @@ server {
|
||||
include /etc/letsencrypt/options-ssl-nginx.conf;
|
||||
ssl_dhparam /etc/letsencrypt/ssl-dhparams.pem;
|
||||
|
||||
client_max_body_size 800m;
|
||||
client_max_body_size 2g;
|
||||
|
||||
location /api/ {
|
||||
proxy_pass http://127.0.0.1:8001;
|
||||
@@ -112,7 +112,7 @@ server {
|
||||
include /etc/letsencrypt/options-ssl-nginx.conf;
|
||||
ssl_dhparam /etc/letsencrypt/ssl-dhparams.pem;
|
||||
|
||||
client_max_body_size 800m;
|
||||
client_max_body_size 2g;
|
||||
|
||||
location / {
|
||||
proxy_pass http://127.0.0.1:3000/;
|
||||
|
||||
@@ -16,7 +16,7 @@ class InMemoryAssetLibraryRepository:
|
||||
def get(self, library_id: str) -> AssetLibrary | None:
|
||||
return self._libraries.get(library_id)
|
||||
|
||||
def list_by_project(self, project_id: str, kind: AssetLibraryKind | None = None) -> list[AssetLibrary]:
|
||||
def find_by_project(self, project_id: str, kind: AssetLibraryKind | None = None) -> list[AssetLibrary]:
|
||||
items = [library for library in self._libraries.values() if library.project_id == project_id]
|
||||
if kind is not None:
|
||||
items = [library for library in items if library.kind == kind]
|
||||
|
||||
@@ -11,7 +11,7 @@ class SQLAlchemyAssetRepository:
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
async def find_by_library(
|
||||
def find_by_library(
|
||||
self,
|
||||
library_id: str,
|
||||
skip: int = 0,
|
||||
@@ -22,7 +22,7 @@ class SQLAlchemyAssetRepository:
|
||||
).offset(skip).limit(limit).all()
|
||||
return [self._to_domain(model) for model in models]
|
||||
|
||||
async def find_by_project(
|
||||
def find_by_project(
|
||||
self,
|
||||
project_id: str,
|
||||
skip: int = 0,
|
||||
@@ -33,7 +33,7 @@ class SQLAlchemyAssetRepository:
|
||||
).offset(skip).limit(limit).all()
|
||||
return [self._to_domain(model) for model in models]
|
||||
|
||||
async def find_by_id(self, asset_id: str) -> Asset | None:
|
||||
def find_by_id(self, asset_id: str) -> Asset | None:
|
||||
model = self.session.query(AssetModel).filter(AssetModel.id == asset_id).first()
|
||||
if model is None:
|
||||
return None
|
||||
@@ -42,7 +42,7 @@ class SQLAlchemyAssetRepository:
|
||||
def get(self, asset_id: str) -> Asset | None:
|
||||
return self.find_by_id(asset_id)
|
||||
|
||||
async def create(self, asset: Asset) -> Asset:
|
||||
def create(self, asset: Asset) -> Asset:
|
||||
now = datetime.now(timezone.utc)
|
||||
model = AssetModel(
|
||||
id=asset.id,
|
||||
@@ -70,7 +70,7 @@ class SQLAlchemyAssetRepository:
|
||||
self.session.commit()
|
||||
return asset
|
||||
|
||||
async def update(self, asset: Asset) -> Asset:
|
||||
def update(self, asset: Asset) -> Asset:
|
||||
model = self.session.query(AssetModel).filter(AssetModel.id == asset.id).first()
|
||||
if model is None:
|
||||
raise ValueError(f"Asset {asset.id} not found")
|
||||
@@ -92,7 +92,7 @@ class SQLAlchemyAssetRepository:
|
||||
self.session.commit()
|
||||
return asset
|
||||
|
||||
async def delete(self, asset_id: str) -> bool:
|
||||
def delete(self, asset_id: str) -> bool:
|
||||
model = self.session.query(AssetModel).filter(AssetModel.id == asset_id).first()
|
||||
if model:
|
||||
self.session.delete(model)
|
||||
@@ -100,11 +100,27 @@ class SQLAlchemyAssetRepository:
|
||||
return True
|
||||
return False
|
||||
|
||||
async def count_by_project(self, project_id: str) -> int:
|
||||
def count_by_project(self, project_id: str) -> int:
|
||||
return self.session.query(AssetModel).filter(
|
||||
AssetModel.project_id == project_id
|
||||
).count()
|
||||
|
||||
def count_by_project_ids(self, project_ids: list[str]) -> int:
|
||||
if not project_ids:
|
||||
return 0
|
||||
return self.session.query(AssetModel).filter(
|
||||
AssetModel.project_id.in_(project_ids)
|
||||
).count()
|
||||
|
||||
def sum_storage_by_project_ids(self, project_ids: list[str]) -> int:
|
||||
if not project_ids:
|
||||
return 0
|
||||
from sqlalchemy import func
|
||||
result = self.session.query(func.coalesce(func.sum(AssetModel.file_size), 0)).filter(
|
||||
AssetModel.project_id.in_(project_ids)
|
||||
).scalar()
|
||||
return int(result or 0)
|
||||
|
||||
def _to_domain(self, model: AssetModel) -> Asset:
|
||||
metadata = {}
|
||||
if model.classification_result:
|
||||
|
||||
@@ -4,6 +4,29 @@ from packages.adapters.sqlalchemy_impl.models import GenerationTaskModel
|
||||
from packages.domain import GenerationTask
|
||||
|
||||
|
||||
def _to_domain(model: GenerationTaskModel) -> GenerationTask:
|
||||
"""Convert ORM model to domain entity."""
|
||||
return GenerationTask(
|
||||
id=model.id,
|
||||
project_id=model.project_id,
|
||||
strategy_id=model.strategy_id,
|
||||
asset_library_id=model.asset_library_id,
|
||||
voice_library_id=model.voice_library_id,
|
||||
template_id=model.template_id,
|
||||
asset_ids=list(model.asset_ids or []),
|
||||
title_ids=list(model.title_ids or []),
|
||||
voice_ids=list(model.voice_ids or []),
|
||||
status=model.status,
|
||||
progress=model.progress,
|
||||
result_count=int(model.result_count or 0),
|
||||
error_message=model.error_message,
|
||||
started_at=model.started_at,
|
||||
completed_at=model.completed_at,
|
||||
created_by_user_id=model.created_by_user_id,
|
||||
created_at=model.created_at,
|
||||
)
|
||||
|
||||
|
||||
class SQLAlchemyGenerationTaskRepository:
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
@@ -15,6 +38,10 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
strategy_id=task.strategy_id,
|
||||
asset_library_id=task.asset_library_id,
|
||||
voice_library_id=task.voice_library_id,
|
||||
template_id=task.template_id,
|
||||
asset_ids=task.asset_ids,
|
||||
title_ids=task.title_ids,
|
||||
voice_ids=task.voice_ids,
|
||||
status=task.status,
|
||||
progress=task.progress,
|
||||
result_count=task.result_count,
|
||||
@@ -32,31 +59,55 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
model = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.id == task_id).first()
|
||||
if model is None:
|
||||
return None
|
||||
return GenerationTask(
|
||||
id=model.id,
|
||||
project_id=model.project_id,
|
||||
strategy_id=model.strategy_id,
|
||||
asset_library_id=model.asset_library_id,
|
||||
voice_library_id=model.voice_library_id,
|
||||
status=model.status,
|
||||
progress=model.progress,
|
||||
result_count=int(model.result_count or 0),
|
||||
error_message=model.error_message,
|
||||
started_at=model.started_at,
|
||||
completed_at=model.completed_at,
|
||||
created_by_user_id=model.created_by_user_id,
|
||||
created_at=model.created_at,
|
||||
)
|
||||
return _to_domain(model)
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[GenerationTask]:
|
||||
models = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.project_id == project_id).all()
|
||||
return [self.get(model.id) for model in models if self.get(model.id) is not None]
|
||||
models = (
|
||||
self.session.query(GenerationTaskModel)
|
||||
.filter(GenerationTaskModel.project_id == project_id)
|
||||
.order_by(GenerationTaskModel.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
return [_to_domain(m) for m in models]
|
||||
|
||||
def list_by_user(self, user_id: str) -> list[GenerationTask]:
|
||||
models = (
|
||||
self.session.query(GenerationTaskModel)
|
||||
.filter(GenerationTaskModel.created_by_user_id == user_id)
|
||||
.order_by(GenerationTaskModel.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
return [_to_domain(m) for m in models]
|
||||
|
||||
def count_by_user(self, user_id: str) -> int:
|
||||
return (
|
||||
self.session.query(GenerationTaskModel)
|
||||
.filter(GenerationTaskModel.created_by_user_id == user_id)
|
||||
.count()
|
||||
)
|
||||
|
||||
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]:
|
||||
models = (
|
||||
self.session.query(GenerationTaskModel)
|
||||
.filter(GenerationTaskModel.created_by_user_id == user_id)
|
||||
.order_by(GenerationTaskModel.created_at.desc())
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
return [_to_domain(m) for m in models]
|
||||
|
||||
def update(self, task: GenerationTask) -> GenerationTask:
|
||||
model = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.id == task.id).first()
|
||||
if model is None:
|
||||
raise ValueError(f"GenerationTask {task.id} not found")
|
||||
model.project_id = task.project_id
|
||||
model.asset_library_id = task.asset_library_id
|
||||
model.strategy_id = task.strategy_id
|
||||
model.voice_library_id = task.voice_library_id
|
||||
model.template_id = task.template_id
|
||||
model.asset_ids = task.asset_ids
|
||||
model.title_ids = task.title_ids
|
||||
model.voice_ids = task.voice_ids
|
||||
model.status = task.status
|
||||
model.progress = task.progress
|
||||
model.result_count = task.result_count
|
||||
|
||||
@@ -65,8 +65,6 @@ class AssetModel(Base):
|
||||
asset_library_id = Column(String(36), nullable=False, index=True)
|
||||
name = Column(String(500), nullable=False)
|
||||
file_type = Column(String(20), nullable=False, index=True)
|
||||
# storage_key: OSS 对象键(相对路径),用于内部存储和操作
|
||||
storage_key = Column(String(255), nullable=False)
|
||||
# file_size: 文件大小(字节),使用 Integer 类型以确保精确性
|
||||
file_size = Column(Integer, nullable=False)
|
||||
# file_url: 完整可访问的 URL,用于客户端直接访问文件
|
||||
@@ -139,10 +137,14 @@ class GenerationTaskModel(Base):
|
||||
__tablename__ = "generation_tasks"
|
||||
|
||||
id = Column(String(32), primary_key=True)
|
||||
project_id = Column(String(32), nullable=False, index=True)
|
||||
project_id = Column(String(32), nullable=False, default="", index=True)
|
||||
strategy_id = Column(String(32), nullable=False, default="")
|
||||
asset_library_id = Column(String(32), nullable=False, index=True)
|
||||
asset_library_id = Column(String(32), nullable=False, default="", index=True)
|
||||
voice_library_id = Column(String(32), nullable=False, default="")
|
||||
template_id = Column(String(36), nullable=False, default="", index=True)
|
||||
asset_ids = Column(JSON, nullable=False, default=list)
|
||||
title_ids = Column(JSON, nullable=False, default=list)
|
||||
voice_ids = Column(JSON, nullable=False, default=list)
|
||||
editing_mode = Column(String(20), nullable=False, default="one_take", index=True) # 剪辑模式: one_take, pip, voice_over, voice_pip
|
||||
status = Column(String(20), nullable=False, default="pending", index=True)
|
||||
progress = Column(Float, nullable=False, default=0.0)
|
||||
@@ -150,7 +152,7 @@ class GenerationTaskModel(Base):
|
||||
error_message = Column(Text, nullable=False, default="")
|
||||
started_at = Column(DateTime, nullable=True)
|
||||
completed_at = Column(DateTime, nullable=True)
|
||||
created_by_user_id = Column(String(32), nullable=False, default="")
|
||||
created_by_user_id = Column(String(32), nullable=False, default="", index=True)
|
||||
extra_meta = Column('metadata', JSON, nullable=False, default=dict)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from uuid import uuid4
|
||||
|
||||
from packages.domain import GenerationTask
|
||||
@@ -9,11 +9,14 @@ from packages.ports.generation_task_repository import GenerationTaskRepository
|
||||
|
||||
@dataclass(slots=True)
|
||||
class CreateGenerationTaskCommand:
|
||||
project_id: str
|
||||
asset_library_id: str
|
||||
project_id: str = ""
|
||||
asset_library_id: str = ""
|
||||
strategy_id: str = ""
|
||||
voice_library_id: str = ""
|
||||
edit_plan_id: str = ""
|
||||
template_id: str = ""
|
||||
asset_ids: list[str] = field(default_factory=list)
|
||||
title_ids: list[str] = field(default_factory=list)
|
||||
voice_ids: list[str] = field(default_factory=list)
|
||||
created_by_user_id: str = ""
|
||||
|
||||
|
||||
@@ -28,7 +31,10 @@ class CreateGenerationTaskUseCase:
|
||||
asset_library_id=command.asset_library_id,
|
||||
strategy_id=command.strategy_id,
|
||||
voice_library_id=command.voice_library_id,
|
||||
edit_plan_id=command.edit_plan_id,
|
||||
template_id=command.template_id,
|
||||
asset_ids=command.asset_ids,
|
||||
title_ids=command.title_ids,
|
||||
voice_ids=command.voice_ids,
|
||||
status="pending",
|
||||
progress=0.0,
|
||||
result_count=0,
|
||||
|
||||
@@ -250,71 +250,6 @@ class IngestJob:
|
||||
)
|
||||
|
||||
|
||||
# 继续读取其他实体定义 - 生成任务、生成视频等
|
||||
@dataclass(slots=True)
|
||||
class ClassificationJob:
|
||||
id: str
|
||||
project_id: str
|
||||
asset_id: str
|
||||
status: str = "pending"
|
||||
classification: str = ""
|
||||
confidence: float = 0.0
|
||||
error_message: str = ""
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class GenerationTask:
|
||||
id: str
|
||||
project_id: str
|
||||
asset_library_id: str
|
||||
strategy_id: str = ""
|
||||
voice_library_id: str = ""
|
||||
edit_plan_id: str = ""
|
||||
status: str = "pending"
|
||||
progress: float = 0.0
|
||||
result_count: int = 0
|
||||
error_message: str = ""
|
||||
started_at: datetime | None = None
|
||||
completed_at: datetime | None = None
|
||||
created_by_user_id: str = ""
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class GeneratedVideo:
|
||||
id: str
|
||||
project_id: str
|
||||
generation_task_id: str
|
||||
name: str
|
||||
file_url: str
|
||||
file_size: float = 0.0
|
||||
duration: float = 0.0
|
||||
thumbnail_url: str | None = None
|
||||
width: int = 0
|
||||
height: int = 0
|
||||
fps: float = 0.0
|
||||
status: str = "completed"
|
||||
review_status: str = "pending_review"
|
||||
generation_params: dict[str, Any] = field(default_factory=dict)
|
||||
generated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime | None = None
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class EditTemplate:
|
||||
id: str
|
||||
project_id: str
|
||||
name: str
|
||||
description: str = ""
|
||||
target_duration: float = 30.0
|
||||
clip_count: int = 3
|
||||
is_active: bool = True
|
||||
created_by_user_id: str = ""
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -21,6 +21,10 @@ class GenerationTask:
|
||||
asset_library_id: str
|
||||
strategy_id: str = ""
|
||||
voice_library_id: str = ""
|
||||
template_id: str = ""
|
||||
asset_ids: list[str] = field(default_factory=list)
|
||||
title_ids: list[str] = field(default_factory=list)
|
||||
voice_ids: list[str] = field(default_factory=list)
|
||||
status: GenerationTaskStatus = GenerationTaskStatus.PENDING
|
||||
progress: float = 0.0
|
||||
result_count: int = 0
|
||||
@@ -38,17 +42,25 @@ class GenerationTask:
|
||||
*,
|
||||
strategy_id: str = "",
|
||||
voice_library_id: str = "",
|
||||
template_id: str = "",
|
||||
asset_ids: list[str] | None = None,
|
||||
title_ids: list[str] | None = None,
|
||||
voice_ids: list[str] | None = None,
|
||||
created_by_user_id: str = "",
|
||||
) -> "GenerationTask":
|
||||
if not project_id.strip():
|
||||
raise ValueError("project_id 不能为空")
|
||||
if not asset_library_id.strip():
|
||||
raise ValueError("asset_library_id 不能为空")
|
||||
if not project_id.strip() and not template_id.strip():
|
||||
raise ValueError("project_id 或 template_id 至少需要提供一个")
|
||||
if not asset_library_id.strip() and not (asset_ids or title_ids or voice_ids):
|
||||
raise ValueError("asset_library_id 或 asset_ids/title_ids/voice_ids 至少需要提供一个")
|
||||
return cls(
|
||||
id=uuid4().hex,
|
||||
project_id=project_id.strip(),
|
||||
asset_library_id=asset_library_id.strip(),
|
||||
strategy_id=strategy_id.strip(),
|
||||
voice_library_id=voice_library_id.strip(),
|
||||
template_id=template_id.strip(),
|
||||
asset_ids=list(asset_ids) if asset_ids else [],
|
||||
title_ids=list(title_ids) if title_ids else [],
|
||||
voice_ids=list(voice_ids) if voice_ids else [],
|
||||
created_by_user_id=created_by_user_id.strip(),
|
||||
)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""兼容层:旧素材仓储接口定义,保留给遗留异步适配器使用。"""
|
||||
"""素材仓储接口定义。"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
@@ -7,15 +7,15 @@ from packages.domain import Asset
|
||||
|
||||
class AssetRepository(ABC):
|
||||
@abstractmethod
|
||||
async def create(self, asset: Asset) -> Asset:
|
||||
def create(self, asset: Asset) -> Asset:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def find_by_id(self, asset_id: str) -> Asset | None:
|
||||
def find_by_id(self, asset_id: str) -> Asset | None:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def find_by_project(
|
||||
def find_by_project(
|
||||
self,
|
||||
project_id: str,
|
||||
skip: int = 0,
|
||||
@@ -24,7 +24,7 @@ class AssetRepository(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def find_by_library(
|
||||
def find_by_library(
|
||||
self,
|
||||
library_id: str,
|
||||
skip: int = 0,
|
||||
@@ -33,13 +33,21 @@ class AssetRepository(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def update(self, asset: Asset) -> Asset:
|
||||
def update(self, asset: Asset) -> Asset:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def delete(self, asset_id: str) -> bool:
|
||||
def delete(self, asset_id: str) -> bool:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def count_by_project(self, project_id: str) -> int:
|
||||
def count_by_project(self, project_id: str) -> int:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def count_by_project_ids(self, project_ids: list[str]) -> int:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def sum_storage_by_project_ids(self, project_ids: list[str]) -> int:
|
||||
pass
|
||||
|
||||
@@ -12,4 +12,10 @@ class GenerationTaskRepository(Protocol):
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[GenerationTask]: ...
|
||||
|
||||
def list_by_user(self, user_id: str) -> list[GenerationTask]: ...
|
||||
|
||||
def count_by_user(self, user_id: str) -> int: ...
|
||||
|
||||
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]: ...
|
||||
|
||||
def update(self, task: GenerationTask) -> GenerationTask: ...
|
||||
|
||||
Executable
+63
@@ -0,0 +1,63 @@
|
||||
#!/bin/sh
|
||||
# cleanup_old_images.sh
|
||||
# 清理构建服务器上的旧 Docker 镜像和 Registry 旧版本
|
||||
# 保留最近 KEEP_VERSIONS 个版本(默认 2)
|
||||
# 在 CI 构建完成后调用,防止磁盘空间耗尽
|
||||
set -eu
|
||||
|
||||
KEEP_VERSIONS="${KEEP_VERSIONS:-2}"
|
||||
REGISTRY_HOST="${REGISTRY_HOST:-172.30.18.198:5000}"
|
||||
SERVICES="xiaoxia-saas-api xiaoxia-saas-worker xiaoxia-saas-web"
|
||||
|
||||
echo "=== Docker Image Cleanup ==="
|
||||
echo "Keeping last ${KEEP_VERSIONS} versions per service"
|
||||
echo ""
|
||||
|
||||
for svc in $SERVICES; do
|
||||
# 收集所有 v0.N.N 格式的版本号(去重,按版本号排序)
|
||||
versions=$(docker images --format "{{.Repository}}:{{.Tag}}" | \
|
||||
grep "${svc}" | \
|
||||
grep -E "v[0-9]+\.[0-9]+\.[0-9]+" | \
|
||||
sed -E "s/.*:${svc}:(v[0-9]+\.[0-9]+\.[0-9]+).*/\1/" | \
|
||||
sed -E "s/.*:v([0-9]+\.[0-9]+\.[0-9]+).*/v\1/" | \
|
||||
sort -t. -k1,1V -k2,2n -k3,3n | \
|
||||
uniq)
|
||||
|
||||
total=$(echo "$versions" | grep -c "^v" || true)
|
||||
|
||||
if [ "$total" -gt "$KEEP_VERSIONS" ]; then
|
||||
remove_count=$((total - KEEP_VERSIONS))
|
||||
to_remove=$(echo "$versions" | head -n "$remove_count")
|
||||
|
||||
echo "[Registry] Cleaning old blobs for ${svc}..."
|
||||
for ver in $to_remove; do
|
||||
repo_name=$(echo "$svc" | sed "s/xiaoxia-saas-//")
|
||||
manifest_url="http://admin:Xiaoxia2026@localhost:5000/v2/${repo_name}/manifests/${ver}"
|
||||
digest=$(curl -sI -H "Accept: application/vnd.docker.distribution.manifest.v2+json" "$manifest_url" | grep -i "^docker-content-digest:" | tr -d "\r" | awk "{print \$2}")
|
||||
if [ -n "$digest" ]; then
|
||||
curl -s -X DELETE -H "Accept: application/vnd.docker.distribution.manifest.v2+json" "http://admin:Xiaoxia2026@localhost:5000/v2/${repo_name}/manifests/${digest}" > /dev/null 2>&1 || true
|
||||
echo " Deleted registry tag: ${svc}:${ver}"
|
||||
fi
|
||||
done
|
||||
|
||||
echo "[Local] Removing old local images for ${svc}..."
|
||||
for ver in $to_remove; do
|
||||
docker rmi "${svc}:${ver}" 2>/dev/null && echo " Removed local: ${svc}:${ver}" || true
|
||||
docker rmi "${REGISTRY_HOST}/${svc}:${ver}" 2>/dev/null && echo " Removed registry-ref: ${REGISTRY_HOST}/${svc}:${ver}" || true
|
||||
done
|
||||
else
|
||||
echo "[${svc}] ${total} version(s) found, within keep limit (${KEEP_VERSIONS})"
|
||||
fi
|
||||
echo ""
|
||||
done
|
||||
|
||||
# 清理悬空镜像(构建中间层)
|
||||
echo "=== Pruning dangling images ==="
|
||||
pruned=$(docker image prune -f 2>&1)
|
||||
echo "$pruned" | tail -1
|
||||
|
||||
# 清理 dev 标签的旧层(dev 标签会被下次构建覆盖)
|
||||
echo ""
|
||||
echo "=== Cleanup complete ==="
|
||||
echo "Current images:"
|
||||
docker images --format "table {{.Repository}}\t{{.Tag}}\t{{.Size}}" | grep -E "(xiaoxia|REPOSITORY)" || true
|
||||
@@ -0,0 +1,140 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
冒烟测试脚本 - 部署后自动验证核心端点可用性
|
||||
用法: python3 smoke_test.py <API_BASE_URL> [--email EMAIL] [--password PASSWORD] [--json]
|
||||
示例: python3 smoke_test.py https://saas-api.xiaoxiajianji.com --email test@example.com --password test123 --json
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
import time
|
||||
import urllib.request
|
||||
import urllib.error
|
||||
import ssl
|
||||
|
||||
CORE_ENDPOINTS = [
|
||||
{"name": "upload/direct/prepare", "method": "POST", "path": "/api/v1/upload/direct/prepare",
|
||||
"body": {"project_id": "smoke-test", "file_name": "test.mp4", "file_size": 1024, "content_type": "video/mp4"},
|
||||
"expect": [200, 401, 422]},
|
||||
{"name": "upload/chunk/init", "method": "POST", "path": "/api/v1/upload/chunk/init",
|
||||
"body": {"project_id": "smoke-test", "file_name": "test.mp4", "file_size": 1024000, "total_chunks": 2},
|
||||
"expect": [200, 401, 422]},
|
||||
{"name": "dashboard/overview", "method": "GET", "path": "/api/v1/dashboard/overview",
|
||||
"expect": [200, 401]},
|
||||
{"name": "assets", "method": "GET", "path": "/api/v1/assets?library_id=smoke-test",
|
||||
"expect": [200, 401]},
|
||||
{"name": "generation/tasks", "method": "POST", "path": "/api/v1/generation/tasks",
|
||||
"body": {},
|
||||
"expect": [200, 401, 422]},
|
||||
]
|
||||
|
||||
|
||||
def make_request(base_url, endpoint, token=None):
|
||||
url = f"{base_url}{endpoint['path']}"
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if token:
|
||||
headers["Authorization"] = f"Bearer {token}"
|
||||
|
||||
data = json.dumps(endpoint.get("body", {})).encode() if endpoint.get("body") is not None else None
|
||||
req = urllib.request.Request(url, data=data, headers=headers, method=endpoint["method"])
|
||||
|
||||
ctx = ssl.create_default_context()
|
||||
ctx.check_hostname = False
|
||||
ctx.verify_mode = ssl.CERT_NONE
|
||||
|
||||
try:
|
||||
start = time.time()
|
||||
resp = urllib.request.urlopen(req, timeout=15, context=ctx)
|
||||
elapsed = round((time.time() - start) * 1000)
|
||||
body = resp.read().decode()
|
||||
return {"status": resp.status, "elapsed_ms": elapsed, "body": body[:200], "error": None}
|
||||
except urllib.error.HTTPError as e:
|
||||
elapsed = round((time.time() - start) * 1000)
|
||||
body = ""
|
||||
try:
|
||||
body = e.read().decode()[:200]
|
||||
except:
|
||||
pass
|
||||
return {"status": e.code, "elapsed_ms": elapsed, "body": body, "error": None}
|
||||
except Exception as e:
|
||||
return {"status": 0, "elapsed_ms": 0, "body": "", "error": str(e)}
|
||||
|
||||
|
||||
def login(base_url, email, password):
|
||||
url = f"{base_url}/api/v1/auth/login"
|
||||
data = json.dumps({"email": email, "password": password}).encode()
|
||||
req = urllib.request.Request(url, data=data, headers={"Content-Type": "application/json"}, method="POST")
|
||||
ctx = ssl.create_default_context()
|
||||
ctx.check_hostname = False
|
||||
ctx.verify_mode = ssl.CERT_NONE
|
||||
try:
|
||||
resp = urllib.request.urlopen(req, timeout=10, context=ctx)
|
||||
body = json.loads(resp.read().decode())
|
||||
return body.get("token") or body.get("data", {}).get("token") or body.get("access_token")
|
||||
except:
|
||||
return None
|
||||
|
||||
|
||||
def run_smoke_test(base_url, email=None, password=None, output_json=False):
|
||||
base_url = base_url.rstrip("/")
|
||||
token = None
|
||||
|
||||
if email and password:
|
||||
token = login(base_url, email, password)
|
||||
if not output_json:
|
||||
print(f"{'✅ 登录成功' if token else '⚠️ 登录失败,将以未认证模式测试'}")
|
||||
|
||||
results = []
|
||||
all_passed = True
|
||||
|
||||
for ep in CORE_ENDPOINTS:
|
||||
result = make_request(base_url, ep, token)
|
||||
passed = result["status"] in ep["expect"] and result["error"] is None
|
||||
is_5xx = 500 <= result["status"] < 600
|
||||
if is_5xx:
|
||||
passed = False
|
||||
all_passed = False
|
||||
|
||||
results.append({
|
||||
"name": ep["name"],
|
||||
"path": ep["path"],
|
||||
"status": result["status"],
|
||||
"elapsed_ms": result["elapsed_ms"],
|
||||
"passed": passed,
|
||||
"error": result["error"],
|
||||
"is_5xx": is_5xx
|
||||
})
|
||||
|
||||
if not output_json:
|
||||
icon = "✅" if passed else "❌"
|
||||
print(f" {icon} {ep['name']}: {result['status']} ({result['elapsed_ms']}ms)")
|
||||
|
||||
if output_json:
|
||||
print(json.dumps({"success": all_passed, "results": results, "base_url": base_url}, indent=2))
|
||||
|
||||
return 0 if all_passed else 1
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="冒烟测试 - 部署后核心端点验证")
|
||||
parser.add_argument("base_url", help="API 基础地址,如 https://saas-api.xiaoxiajianji.com")
|
||||
parser.add_argument("--email", help="登录邮箱")
|
||||
parser.add_argument("--password", help="登录密码")
|
||||
parser.add_argument("--json", action="store_true", help="JSON 格式输出")
|
||||
args = parser.parse_args()
|
||||
|
||||
if not args.json:
|
||||
print(f"\n🔍 冒烟测试: {args.base_url}")
|
||||
print("-" * 50)
|
||||
|
||||
exit_code = run_smoke_test(args.base_url, args.email, args.password, args.json)
|
||||
|
||||
if not args.json:
|
||||
print("-" * 50)
|
||||
print(f"{'✅ 全部通过' if exit_code == 0 else '❌ 存在失败端点'}\n")
|
||||
|
||||
sys.exit(exit_code)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,11 +1,28 @@
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from uuid import uuid4
|
||||
|
||||
# 设置必要环境变量(必须在导入 app 模块之前)
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app.api.routes.asset_diagnosis import _build_diagnosis
|
||||
|
||||
from packages.domain import Asset, AssetStatus, ClassificationStatus
|
||||
from packages.domain import (
|
||||
Asset,
|
||||
AssetLibrary,
|
||||
AssetLibraryKind,
|
||||
AssetStatus,
|
||||
ClassificationStatus,
|
||||
Project,
|
||||
)
|
||||
|
||||
|
||||
def _asset(name: str, mime_type: str, *, status=AssetStatus.READY, duration=None, quality_score=None):
|
||||
@@ -25,7 +42,6 @@ def _asset(name: str, mime_type: str, *, status=AssetStatus.READY, duration=None
|
||||
|
||||
def test_asset_diagnosis_reports_missing_video_gap():
|
||||
diagnosis = _build_diagnosis(
|
||||
"workspace-1",
|
||||
"project-1",
|
||||
[_asset("voice.mp3", "audio/mpeg"), _asset("image.jpg", "image/jpeg")],
|
||||
)
|
||||
@@ -39,7 +55,6 @@ def test_asset_diagnosis_scores_ready_video_assets():
|
||||
used_asset = _asset("video-1.mp4", "video/mp4", duration=8)
|
||||
used_asset.metadata = {"generation_use_count": 1, "review_status": "pending_review"}
|
||||
diagnosis = _build_diagnosis(
|
||||
"workspace-1",
|
||||
"project-1",
|
||||
[
|
||||
used_asset,
|
||||
@@ -64,7 +79,6 @@ def test_asset_diagnosis_scores_ready_video_assets():
|
||||
|
||||
def test_asset_diagnosis_flags_unready_and_low_quality_assets():
|
||||
diagnosis = _build_diagnosis(
|
||||
"workspace-1",
|
||||
"project-1",
|
||||
[
|
||||
_asset("video.mp4", "video/mp4", duration=10, quality_score=40),
|
||||
@@ -78,3 +92,127 @@ def test_asset_diagnosis_flags_unready_and_low_quality_assets():
|
||||
smart_view_counts = {item.key: item.count for item in diagnosis.smart_views}
|
||||
assert smart_view_counts["needs_attention"] == 2
|
||||
assert smart_view_counts["high_risk"] == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 路由层测试 — 验证 find_by_project 调用正确性
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _StubProjectRepository:
|
||||
def __init__(self, projects: dict[str, Project] | None = None):
|
||||
self._projects = projects or {}
|
||||
|
||||
def get(self, project_id: str) -> Project | None:
|
||||
return self._projects.get(project_id)
|
||||
|
||||
def find_by_id(self, project_id: str) -> Project | None:
|
||||
return self._projects.get(project_id)
|
||||
|
||||
|
||||
class _StubAssetLibraryRepository:
|
||||
def __init__(self, libraries: dict[str, AssetLibrary] | None = None):
|
||||
self._libraries = libraries or {}
|
||||
self.find_by_project_called_with: list[str] = []
|
||||
|
||||
def find_by_project(self, project_id: str, kind=None) -> list[AssetLibrary]:
|
||||
self.find_by_project_called_with.append(project_id)
|
||||
return [lib for lib in self._libraries.values() if lib.project_id == project_id]
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[AssetLibrary]:
|
||||
raise AssertionError("路由不应调用 list_by_project,应调用 find_by_project")
|
||||
|
||||
|
||||
class _StubAssetRepository:
|
||||
def __init__(self, assets: dict[str, Asset] | None = None):
|
||||
self._assets = assets or {}
|
||||
self.list_by_library_called_with: list[str] = []
|
||||
|
||||
def list_by_library(self, library_id: str, skip: int = 0, limit: int = 50) -> list[Asset]:
|
||||
self.list_by_library_called_with.append(library_id)
|
||||
return [a for a in self._assets.values() if a.library_id == library_id]
|
||||
|
||||
def count_by_library(self, library_id: str) -> int:
|
||||
return len([a for a in self._assets.values() if a.library_id == library_id])
|
||||
|
||||
|
||||
def _dep(name: str):
|
||||
from app import dependencies
|
||||
|
||||
return getattr(dependencies, name)
|
||||
|
||||
|
||||
def _build_route_test_app(project_repo, library_repo, asset_repo):
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from app.api.routes.asset_diagnosis import router
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(router, prefix="/api/v1")
|
||||
|
||||
mock_user = MagicMock(spec=AuthenticatedUser)
|
||||
mock_user.id = "user-1"
|
||||
mock_user.email = "test@example.com"
|
||||
app.dependency_overrides[get_current_user] = lambda: mock_user
|
||||
app.dependency_overrides[_dep("get_project_repository")] = lambda: project_repo
|
||||
app.dependency_overrides[_dep("get_asset_library_repository")] = lambda: library_repo
|
||||
app.dependency_overrides[_dep("get_asset_repository")] = lambda: asset_repo
|
||||
return app
|
||||
|
||||
|
||||
class TestAssetDiagnosisRoute:
|
||||
"""路由层测试 — 验证 find_by_project 调用正确性。"""
|
||||
|
||||
def test_returns_404_when_project_not_found(self):
|
||||
project_repo = _StubProjectRepository()
|
||||
library_repo = _StubAssetLibraryRepository()
|
||||
asset_repo = _StubAssetRepository()
|
||||
app = _build_route_test_app(project_repo, library_repo, asset_repo)
|
||||
client = TestClient(app)
|
||||
|
||||
resp = client.get("/api/v1/projects/nonexistent/asset-diagnosis")
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_find_by_project_called_with_correct_project_id(self):
|
||||
project = Project(id="proj-123", name="Test", owner_user_id="user-1")
|
||||
library = AssetLibrary(
|
||||
id="lib-1", name="Lib", project_id="proj-123", kind=AssetLibraryKind.VIDEO
|
||||
)
|
||||
project_repo = _StubProjectRepository({"proj-123": project})
|
||||
library_repo = _StubAssetLibraryRepository({"lib-1": library})
|
||||
asset_repo = _StubAssetRepository()
|
||||
app = _build_route_test_app(project_repo, library_repo, asset_repo)
|
||||
client = TestClient(app)
|
||||
|
||||
resp = client.get("/api/v1/projects/proj-123/asset-diagnosis")
|
||||
assert resp.status_code == 200
|
||||
assert library_repo.find_by_project_called_with == ["proj-123"]
|
||||
|
||||
def test_list_by_library_called_for_each_library(self):
|
||||
project = Project(id="proj-1", name="Test", owner_user_id="user-1")
|
||||
lib1 = AssetLibrary(id="lib-1", name="Lib1", project_id="proj-1", kind=AssetLibraryKind.VIDEO)
|
||||
lib2 = AssetLibrary(id="lib-2", name="Lib2", project_id="proj-1", kind=AssetLibraryKind.VOICE)
|
||||
project_repo = _StubProjectRepository({"proj-1": project})
|
||||
library_repo = _StubAssetLibraryRepository({"lib-1": lib1, "lib-2": lib2})
|
||||
asset_repo = _StubAssetRepository()
|
||||
app = _build_route_test_app(project_repo, library_repo, asset_repo)
|
||||
client = TestClient(app)
|
||||
|
||||
resp = client.get("/api/v1/projects/proj-1/asset-diagnosis")
|
||||
assert resp.status_code == 200
|
||||
assert set(asset_repo.list_by_library_called_with) == {"lib-1", "lib-2"}
|
||||
|
||||
def test_find_by_project_not_list_by_project(self):
|
||||
"""路由调用 find_by_project 而非 list_by_project(否则会触发 AssertionError)。"""
|
||||
project = Project(id="proj-1", name="Test", owner_user_id="user-1")
|
||||
library = AssetLibrary(id="lib-1", name="Lib", project_id="proj-1", kind=AssetLibraryKind.VIDEO)
|
||||
project_repo = _StubProjectRepository({"proj-1": project})
|
||||
library_repo = _StubAssetLibraryRepository({"lib-1": library})
|
||||
asset_repo = _StubAssetRepository()
|
||||
app = _build_route_test_app(project_repo, library_repo, asset_repo)
|
||||
client = TestClient(app)
|
||||
|
||||
resp = client.get("/api/v1/projects/proj-1/asset-diagnosis")
|
||||
# 如果调用了 list_by_project,会抛 AssertionError → 500
|
||||
assert resp.status_code == 200
|
||||
|
||||
@@ -0,0 +1,316 @@
|
||||
"""
|
||||
chunked_upload.py 路由单元测试
|
||||
|
||||
覆盖:
|
||||
- init_chunked_upload 端点正常路径
|
||||
- init_chunked_upload 项目/素材库不存在时返回 404
|
||||
- find_by_project 调用正确性
|
||||
- 文件大小校验
|
||||
- OSS 凭证校验
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# 设置必要环境变量(必须在导入 app 模块之前)
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
import pytest
|
||||
|
||||
# 确保 app 模块可导入
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from packages.domain import AssetLibrary, AssetLibraryKind, Project
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Stub 实现(不继承 Port ABC)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class StubProjectRepository:
|
||||
def __init__(self, projects: dict[str, Project] | None = None):
|
||||
self._projects = projects or {}
|
||||
|
||||
def get(self, project_id: str) -> Project | None:
|
||||
return self._projects.get(project_id)
|
||||
|
||||
def find_by_id(self, project_id: str) -> Project | None:
|
||||
return self._projects.get(project_id)
|
||||
|
||||
|
||||
class StubAssetLibraryRepository:
|
||||
def __init__(self, libraries: dict[str, AssetLibrary] | None = None):
|
||||
self._libraries = libraries or {}
|
||||
|
||||
def find_by_project(self, project_id: str, kind=None) -> list[AssetLibrary]:
|
||||
items = [lib for lib in self._libraries.values() if lib.project_id == project_id]
|
||||
if kind is not None:
|
||||
items = [lib for lib in items if lib.kind == kind]
|
||||
return items
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[AssetLibrary]:
|
||||
raise AssertionError("路由不应调用 list_by_project,应调用 find_by_project")
|
||||
|
||||
|
||||
class StubChunkedUploadRepository:
|
||||
def __init__(self):
|
||||
self._uploads = {}
|
||||
|
||||
def add(self, upload) -> None:
|
||||
self._uploads[upload.upload_id] = upload
|
||||
|
||||
def get(self, upload_id: str):
|
||||
return self._uploads.get(upload_id)
|
||||
|
||||
def update(self, upload) -> None:
|
||||
self._uploads[upload.upload_id] = upload
|
||||
|
||||
def find_by_project(self, project_id: str, skip: int = 0, limit: int = 50):
|
||||
return [u for u in self._uploads.values() if u.project_id == project_id]
|
||||
|
||||
def count_by_project(self, project_id: str) -> int:
|
||||
return len([u for u in self._uploads.values() if u.project_id == project_id])
|
||||
|
||||
|
||||
class StubIngestJobRepository:
|
||||
def add(self, job) -> None:
|
||||
pass
|
||||
|
||||
def get(self, job_id: str):
|
||||
return None
|
||||
|
||||
def update_status(self, job_id, status, **kwargs):
|
||||
pass
|
||||
|
||||
def list_by_library(self, library_id: str, skip: int = 0, limit: int = 50):
|
||||
return []
|
||||
|
||||
def count_by_library(self, library_id: str) -> int:
|
||||
return 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_project(id: str = "proj-1", owner_user_id: str = "user-1") -> Project:
|
||||
return Project(id=id, name="Test Project", owner_user_id=owner_user_id)
|
||||
|
||||
|
||||
def _make_library(id: str = "lib-1", project_id: str = "proj-1") -> AssetLibrary:
|
||||
return AssetLibrary(id=id, name="Test Library", project_id=project_id, kind=AssetLibraryKind.VIDEO)
|
||||
|
||||
|
||||
def _dep(name: str):
|
||||
from app import dependencies
|
||||
|
||||
return getattr(dependencies, name)
|
||||
|
||||
|
||||
def _build_app(
|
||||
project_repo=None,
|
||||
library_repo=None,
|
||||
storage=None,
|
||||
chunked_repo=None,
|
||||
ingest_repo=None,
|
||||
) -> FastAPI:
|
||||
from app.api.routes.chunked_upload import router
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import get_storage_service
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(router, prefix="/api/v1")
|
||||
|
||||
project_repo = project_repo or StubProjectRepository()
|
||||
library_repo = library_repo or StubAssetLibraryRepository()
|
||||
storage = storage or MagicMock()
|
||||
storage.is_configured = True
|
||||
chunked_repo = chunked_repo or StubChunkedUploadRepository()
|
||||
ingest_repo = ingest_repo or StubIngestJobRepository()
|
||||
|
||||
# Mock auth
|
||||
mock_user = MagicMock(spec=AuthenticatedUser)
|
||||
mock_user.id = "user-1"
|
||||
mock_user.email = "test@example.com"
|
||||
|
||||
app.dependency_overrides[get_current_user] = lambda: mock_user
|
||||
app.dependency_overrides[_dep("get_project_repository")] = lambda: project_repo
|
||||
app.dependency_overrides[_dep("get_asset_library_repository")] = lambda: library_repo
|
||||
app.dependency_overrides[get_storage_service] = lambda: storage
|
||||
app.dependency_overrides[_dep("get_ingest_job_repository")] = lambda: ingest_repo
|
||||
|
||||
return app
|
||||
|
||||
|
||||
def _client(**kwargs) -> TestClient:
|
||||
return TestClient(_build_app(**kwargs))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 测试用例
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestInitChunkedUpload:
|
||||
"""init_chunked_upload 端点测试。"""
|
||||
|
||||
def test_returns_upload_record_on_success(self):
|
||||
"""正常初始化分片上传。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(project_repo=project_repo, library_repo=library_repo)
|
||||
|
||||
file_size = 100 * 1024 * 1024 # 100MB
|
||||
chunk_size = 5 * 1024 * 1024 # 5MB
|
||||
total_chunks = (file_size + chunk_size - 1) // chunk_size
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/init",
|
||||
json={
|
||||
"project_id": project.id,
|
||||
"library_id": library.id,
|
||||
"filename": "large-video.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": file_size,
|
||||
"total_chunks": total_chunks,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["filename"] == "large-video.mp4"
|
||||
assert "upload_id" in data
|
||||
|
||||
def test_returns_404_when_project_not_found(self):
|
||||
"""项目不存在时返回 404。"""
|
||||
library = _make_library()
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(
|
||||
project_repo=StubProjectRepository(),
|
||||
library_repo=library_repo,
|
||||
)
|
||||
|
||||
file_size = 100 * 1024 * 1024
|
||||
total_chunks = (file_size + 5 * 1024 * 1024 - 1) // (5 * 1024 * 1024)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/init",
|
||||
json={
|
||||
"project_id": "nonexistent",
|
||||
"library_id": library.id,
|
||||
"filename": "test.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": file_size,
|
||||
"total_chunks": total_chunks,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
assert "Project not found" in resp.json()["detail"]
|
||||
|
||||
def test_returns_404_when_library_not_in_project(self):
|
||||
"""素材库不属于该项目时返回 404。"""
|
||||
project = _make_project()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
|
||||
# 素材库属于另一个项目
|
||||
other_library = _make_library(project_id="other-project")
|
||||
library_repo = StubAssetLibraryRepository({other_library.id: other_library})
|
||||
|
||||
client = _client(project_repo=project_repo, library_repo=library_repo)
|
||||
|
||||
file_size = 100 * 1024 * 1024
|
||||
total_chunks = (file_size + 5 * 1024 * 1024 - 1) // (5 * 1024 * 1024)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/init",
|
||||
json={
|
||||
"project_id": project.id,
|
||||
"library_id": "nonexistent-lib",
|
||||
"filename": "test.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": file_size,
|
||||
"total_chunks": total_chunks,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
assert "Asset library not found" in resp.json()["detail"]
|
||||
|
||||
def test_rejects_file_exceeding_max_size(self):
|
||||
"""超过 2GB 限制的文件被拒绝(schema 层 le=2GB 会返回 422)。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(project_repo=project_repo, library_repo=library_repo)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/init",
|
||||
json={
|
||||
"project_id": project.id,
|
||||
"library_id": library.id,
|
||||
"filename": "huge.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": 3 * 1024 * 1024 * 1024, # 3GB,超过 2GB 限制
|
||||
"total_chunks": 600,
|
||||
},
|
||||
)
|
||||
# schema le=2147483648 → 422; route-level check → 413
|
||||
assert resp.status_code in (400, 413, 422)
|
||||
|
||||
def test_find_by_project_is_called_not_list_by_project(self):
|
||||
"""验证路由调用 find_by_project 而非 list_by_project。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(project_repo=project_repo, library_repo=library_repo)
|
||||
|
||||
file_size = 100 * 1024 * 1024
|
||||
total_chunks = (file_size + 5 * 1024 * 1024 - 1) // (5 * 1024 * 1024)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/init",
|
||||
json={
|
||||
"project_id": project.id,
|
||||
"library_id": library.id,
|
||||
"filename": "test.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": file_size,
|
||||
"total_chunks": total_chunks,
|
||||
},
|
||||
)
|
||||
# 如果调用了 list_by_project,StubAssetLibraryRepository 会抛 AssertionError
|
||||
assert resp.status_code == 200
|
||||
|
||||
|
||||
class TestChunkedUploadConstants:
|
||||
"""分片上传常量测试。"""
|
||||
|
||||
def test_max_file_size_is_2gb(self):
|
||||
from app.api.routes.chunked_upload import MAX_FILE_SIZE
|
||||
|
||||
assert MAX_FILE_SIZE == 2 * 1024 * 1024 * 1024
|
||||
|
||||
def test_default_chunk_size_is_5mb(self):
|
||||
from app.api.routes.chunked_upload import DEFAULT_CHUNK_SIZE
|
||||
|
||||
assert DEFAULT_CHUNK_SIZE == 5 * 1024 * 1024
|
||||
|
||||
def test_chunk_expiry_hours_is_24(self):
|
||||
from app.api.routes.chunked_upload import CHUNK_EXPIRY_HOURS
|
||||
|
||||
assert CHUNK_EXPIRY_HOURS == 24
|
||||
@@ -0,0 +1,198 @@
|
||||
"""
|
||||
config.py OSS 配置字段单元测试
|
||||
|
||||
覆盖:
|
||||
- OSS 相关字段默认值
|
||||
- 环境变量覆盖
|
||||
- 字段名与代码引用一致
|
||||
- pydantic_settings 加载行为
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _load_settings_class():
|
||||
"""
|
||||
直接加载 config.py 模块,绕过 apps/api/__init__.py 的副作用。
|
||||
apps/api/__init__.py 会导入 main.py,而 main.py 依赖 app 模块。
|
||||
"""
|
||||
config_path = Path(__file__).resolve().parents[2] / "apps" / "api" / "app" / "config.py"
|
||||
spec = importlib.util.spec_from_file_location("config_module", config_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module.Settings
|
||||
|
||||
|
||||
def _fresh_settings(**env_overrides: dict[str, str]):
|
||||
"""
|
||||
每次创建一个全新的 Settings 实例。
|
||||
env_overrides 会注入到 os.environ。
|
||||
"""
|
||||
env = {
|
||||
"JWT_SECRET_KEY": "unit-test-secret-key-12345",
|
||||
**env_overrides,
|
||||
}
|
||||
with patch.dict(os.environ, env, clear=False):
|
||||
Settings = _load_settings_class()
|
||||
return Settings()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 默认值测试
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestOSSConfigDefaults:
|
||||
"""OSS 配置字段默认值必须与代码引用一致。"""
|
||||
|
||||
def test_oss_endpoint_default(self):
|
||||
settings = _fresh_settings()
|
||||
assert settings.OSS_ENDPOINT == "oss-cn-hangzhou.aliiyuncs.com"
|
||||
|
||||
def test_oss_access_key_id_default_empty(self):
|
||||
settings = _fresh_settings()
|
||||
assert settings.OSS_ACCESS_KEY_ID == ""
|
||||
|
||||
def test_oss_access_key_secret_default_empty(self):
|
||||
settings = _fresh_settings()
|
||||
assert settings.OSS_ACCESS_KEY_SECRET == ""
|
||||
|
||||
def test_oss_bucket_name_default(self):
|
||||
settings = _fresh_settings()
|
||||
assert settings.OSS_BUCKET_NAME == "xiaoxia-autocut"
|
||||
|
||||
def test_oss_direct_upload_max_mb_default_is_2000(self):
|
||||
"""PR #117 修复:默认值从 800 改为 2000。"""
|
||||
settings = _fresh_settings()
|
||||
assert settings.OSS_DIRECT_UPLOAD_MAX_MB == 2000
|
||||
|
||||
def test_oss_direct_upload_expire_seconds_default(self):
|
||||
settings = _fresh_settings()
|
||||
assert settings.OSS_DIRECT_UPLOAD_EXPIRE_SECONDS == 900
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 环境变量覆盖测试
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestOSSConfigEnvOverride:
|
||||
"""环境变量能正确覆盖 OSS 配置字段。"""
|
||||
|
||||
def test_oss_endpoint_override(self):
|
||||
settings = _fresh_settings(OSS_ENDPOINT="oss-cn-shanghai.aliiyuncs.com")
|
||||
assert settings.OSS_ENDPOINT == "oss-cn-shanghai.aliiyuncs.com"
|
||||
|
||||
def test_oss_access_key_id_override(self):
|
||||
settings = _fresh_settings(OSS_ACCESS_KEY_ID="test-key-id")
|
||||
assert settings.OSS_ACCESS_KEY_ID == "test-key-id"
|
||||
|
||||
def test_oss_access_key_secret_override(self):
|
||||
settings = _fresh_settings(OSS_ACCESS_KEY_SECRET="test-key-secret")
|
||||
assert settings.OSS_ACCESS_KEY_SECRET == "test-key-secret"
|
||||
|
||||
def test_oss_bucket_name_override(self):
|
||||
settings = _fresh_settings(OSS_BUCKET_NAME="test-bucket")
|
||||
assert settings.OSS_BUCKET_NAME == "test-bucket"
|
||||
|
||||
def test_oss_direct_upload_max_mb_override(self):
|
||||
settings = _fresh_settings(OSS_DIRECT_UPLOAD_MAX_MB="4096")
|
||||
assert settings.OSS_DIRECT_UPLOAD_MAX_MB == 4096
|
||||
|
||||
def test_oss_direct_upload_expire_seconds_override(self):
|
||||
settings = _fresh_settings(OSS_DIRECT_UPLOAD_EXPIRE_SECONDS="1800")
|
||||
assert settings.OSS_DIRECT_UPLOAD_EXPIRE_SECONDS == 1800
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 字段名一致性测试(防止再次出现字段名拼写错误导致 500)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestOSSConfigFieldNameConsistency:
|
||||
"""
|
||||
确保 Settings 类包含代码中实际引用的所有字段。
|
||||
防止类似 OSS_DIRECT_UPLOAD_EXPRESS_SECRET 的拼写错误再次发生。
|
||||
"""
|
||||
|
||||
def test_settings_has_oss_endpoint_field(self):
|
||||
settings = _fresh_settings()
|
||||
assert hasattr(settings, "OSS_ENDPOINT")
|
||||
|
||||
def test_settings_has_oss_access_key_id_field(self):
|
||||
settings = _fresh_settings()
|
||||
assert hasattr(settings, "OSS_ACCESS_KEY_ID")
|
||||
|
||||
def test_settings_has_oss_access_key_secret_field(self):
|
||||
settings = _fresh_settings()
|
||||
assert hasattr(settings, "OSS_ACCESS_KEY_SECRET")
|
||||
|
||||
def test_settings_has_oss_bucket_name_field(self):
|
||||
settings = _fresh_settings()
|
||||
assert hasattr(settings, "OSS_BUCKET_NAME")
|
||||
|
||||
def test_settings_has_oss_direct_upload_max_mb_field(self):
|
||||
settings = _fresh_settings()
|
||||
assert hasattr(settings, "OSS_DIRECT_UPLOAD_MAX_MB")
|
||||
|
||||
def test_settings_has_oss_direct_upload_expire_seconds_field(self):
|
||||
settings = _fresh_settings()
|
||||
assert hasattr(settings, "OSS_DIRECT_UPLOAD_EXPIRE_SECONDS")
|
||||
|
||||
def test_property_aliases_match_field_names(self):
|
||||
"""确保 property 访问器与字段值一致。"""
|
||||
settings = _fresh_settings(
|
||||
OSS_ENDPOINT="ep",
|
||||
OSS_ACCESS_KEY_ID="kid",
|
||||
OSS_ACCESS_KEY_SECRET="ksec",
|
||||
OSS_BUCKET_NAME="bkt",
|
||||
)
|
||||
assert settings.oss_endpoint == "ep"
|
||||
assert settings.oss_access_key_id == "kid"
|
||||
assert settings.oss_access_key_secret == "ksec"
|
||||
assert settings.oss_bucket_name == "bkt"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# extra="ignore" 行为测试
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSettingsExtraIgnore:
|
||||
"""extra="ignore" 确保未知环境变量不会导致启动失败。"""
|
||||
|
||||
def test_unknown_env_var_is_ignored(self):
|
||||
settings = _fresh_settings(UNKNOWN_RANDOM_VAR="whatever")
|
||||
assert not hasattr(settings, "UNKNOWN_RANDOM_VAR")
|
||||
|
||||
def test_settings_loads_without_error(self):
|
||||
settings = _fresh_settings()
|
||||
assert settings.APP_NAME == "xiaoxia-saas"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 别名测试(MAX_UPLOAD_SIZE_MB 兼容旧配置)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestOSSConfigAliases:
|
||||
"""OSS_DIRECT_UPLOAD_MAX_MB 支持 MAX_UPLOAD_SIZE_MB 别名。"""
|
||||
|
||||
def test_max_upload_size_mb_alias_works(self):
|
||||
"""旧环境变量 MAX_UPLOAD_SIZE_MB 仍能生效。"""
|
||||
env = {
|
||||
"JWT_SECRET_KEY": "unit-test-secret-key-12345",
|
||||
"MAX_UPLOAD_SIZE_MB": "3000",
|
||||
}
|
||||
with patch.dict(os.environ, env, clear=False):
|
||||
os.environ.pop("OSS_DIRECT_UPLOAD_MAX_MB", None)
|
||||
Settings = _load_settings_class()
|
||||
settings = Settings()
|
||||
assert settings.OSS_DIRECT_UPLOAD_MAX_MB == 3000
|
||||
@@ -0,0 +1,650 @@
|
||||
"""
|
||||
upload.py 表单上传端点单元测试
|
||||
|
||||
覆盖(24个测试用例):
|
||||
- 正常路径(5):视频/音频/图片上传成功、创建导入任务、响应包含URL
|
||||
- 文件名校验(5):路径遍历防护、反斜杠处理、空文件名、特殊字符、无扩展名
|
||||
- 异常路径(8):不支持文件类型、项目/素材库不存在、存储服务错误、缺少参数
|
||||
- 多格式支持(3):多种视频(4种)/音频(5种)/图片(6种)格式
|
||||
- MIME验证(6):有效类型、空类型(400)、不支持类型(415)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# 设置必要环境变量(必须在导入 app 模块之前)
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
import pytest
|
||||
|
||||
# 确保 app 模块可导入
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from packages.domain import AssetLibrary, AssetLibraryKind, Project
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 测试用 Stub(不继承 Port ABC,因为 Port 定义 async 方法,路由实际使用同步 duck-type)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class StubProjectRepository:
|
||||
def __init__(self, projects: dict[str, Project] | None = None):
|
||||
self._projects = projects or {}
|
||||
|
||||
def get(self, project_id: str) -> Project | None:
|
||||
return self._projects.get(project_id)
|
||||
|
||||
def find_by_id(self, project_id: str) -> Project | None:
|
||||
return self._projects.get(project_id)
|
||||
|
||||
|
||||
class StubAssetLibraryRepository:
|
||||
def __init__(self, libraries: dict[str, AssetLibrary] | None = None):
|
||||
self._libraries = libraries or {}
|
||||
|
||||
def find_by_project(self, project_id: str, kind=None) -> list[AssetLibrary]:
|
||||
items = [lib for lib in self._libraries.values() if lib.project_id == project_id]
|
||||
if kind is not None:
|
||||
items = [lib for lib in items if lib.kind == kind]
|
||||
return items
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[AssetLibrary]:
|
||||
raise AssertionError("路由不应调用 list_by_project,应调用 find_by_project")
|
||||
|
||||
|
||||
class StubIngestJobRepository:
|
||||
def __init__(self):
|
||||
self._jobs = {}
|
||||
self._counter = 0
|
||||
|
||||
def add(self, job) -> None:
|
||||
self._jobs[job.id] = job
|
||||
|
||||
def get(self, job_id: str):
|
||||
return self._jobs.get(job_id)
|
||||
|
||||
def update_status(self, job_id, status, **kwargs):
|
||||
if job_id in self._jobs:
|
||||
self._jobs[job_id].status = status
|
||||
|
||||
def list_by_library(self, library_id: str, skip: int = 0, limit: int = 50):
|
||||
return []
|
||||
|
||||
def count_by_library(self, library_id: str) -> int:
|
||||
return 0
|
||||
|
||||
def create(self, job):
|
||||
"""创建一个模拟的 ingest job"""
|
||||
self._jobs[job.id] = job
|
||||
return job
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 测试 Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_project(id: str = "proj-1", owner_user_id: str = "user-1") -> Project:
|
||||
return Project(id=id, name="Test Project", owner_user_id=owner_user_id)
|
||||
|
||||
|
||||
def _make_library(
|
||||
id: str = "lib-1",
|
||||
project_id: str = "proj-1",
|
||||
kind: AssetLibraryKind = AssetLibraryKind.VIDEO,
|
||||
) -> AssetLibrary:
|
||||
return AssetLibrary(id=id, name="Test Library", project_id=project_id, kind=kind)
|
||||
|
||||
|
||||
def _build_app(
|
||||
project_repo: StubProjectRepository | None = None,
|
||||
library_repo: StubAssetLibraryRepository | None = None,
|
||||
storage: MagicMock | None = None,
|
||||
ingest_repo: StubIngestJobRepository | None = None,
|
||||
) -> FastAPI:
|
||||
"""构建一个最小化的 FastAPI app,只注册 upload 路由。"""
|
||||
from app.api.routes.upload import router
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import (
|
||||
get_asset_library_repository,
|
||||
get_ingest_job_repository,
|
||||
get_project_repository,
|
||||
)
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(router, prefix="/api/v1")
|
||||
|
||||
project_repo = project_repo or StubProjectRepository()
|
||||
library_repo = library_repo or StubAssetLibraryRepository()
|
||||
storage = storage or MagicMock()
|
||||
storage.is_configured = True
|
||||
storage.upload_file.return_value = "https://bucket.oss.example.com/uploads/test.mp4"
|
||||
ingest_repo = ingest_repo or StubIngestJobRepository()
|
||||
|
||||
# Mock auth
|
||||
mock_user = MagicMock(spec=AuthenticatedUser)
|
||||
mock_user.id = "user-1"
|
||||
mock_user.email = "test@example.com"
|
||||
|
||||
app.dependency_overrides[get_current_user] = lambda: mock_user
|
||||
app.dependency_overrides[get_project_repository] = lambda: project_repo
|
||||
app.dependency_overrides[get_asset_library_repository] = lambda: library_repo
|
||||
app.dependency_overrides[get_storage_service] = lambda: storage
|
||||
app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo
|
||||
|
||||
return app
|
||||
|
||||
|
||||
def _client(**kwargs) -> TestClient:
|
||||
app = _build_app(**kwargs)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 测试用例
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFormUploadSuccess:
|
||||
"""表单上传正常成功路径测试。"""
|
||||
|
||||
def test_upload_video_file_successfully(self):
|
||||
"""上传视频文件成功。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
ingest_repo = StubIngestJobRepository()
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://bucket.oss.example.com/uploads/video.mp4"
|
||||
|
||||
client = _client(
|
||||
project_repo=project_repo,
|
||||
library_repo=library_repo,
|
||||
storage=storage,
|
||||
ingest_repo=ingest_repo,
|
||||
)
|
||||
|
||||
# 模拟一个视频文件
|
||||
file_content = b"fake video content" * 100
|
||||
files = {"file": ("test-video.mp4", io.BytesIO(file_content), "video/mp4")}
|
||||
data = {"project_id": project.id, "library_id": library.id}
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
|
||||
assert resp.status_code == 200
|
||||
result = resp.json()
|
||||
assert "storage_key" in result
|
||||
assert "ingest_job_id" in result
|
||||
assert "url" in result
|
||||
assert storage.upload_file.called
|
||||
|
||||
def test_upload_image_file_successfully(self):
|
||||
"""上传图片文件成功。"""
|
||||
project = _make_project()
|
||||
library = _make_library(kind=AssetLibraryKind.IMAGE)
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://bucket.oss.example.com/uploads/image.jpg"
|
||||
|
||||
client = _client(
|
||||
project_repo=project_repo,
|
||||
library_repo=library_repo,
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
# 模拟一个图片文件
|
||||
file_content = b"fake image content" * 50
|
||||
files = {"file": ("test-image.jpg", io.BytesIO(file_content), "image/jpeg")}
|
||||
data = {"project_id": project.id, "library_id": library.id}
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert storage.upload_file.called
|
||||
|
||||
def test_upload_audio_file_successfully(self):
|
||||
"""上传音频文件成功。"""
|
||||
project = _make_project()
|
||||
library = _make_library(kind=AssetLibraryKind.VOICE)
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://bucket.oss.example.com/uploads/audio.mp3"
|
||||
|
||||
client = _client(
|
||||
project_repo=project_repo,
|
||||
library_repo=library_repo,
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
file_content = b"fake audio content" * 50
|
||||
files = {"file": ("test-audio.mp3", io.BytesIO(file_content), "audio/mpeg")}
|
||||
data = {"project_id": project.id, "library_id": library.id}
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert storage.upload_file.called
|
||||
|
||||
def test_ingest_job_is_created(self):
|
||||
"""验证上传成功后创建导入任务。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
ingest_repo = StubIngestJobRepository()
|
||||
|
||||
client = _client(
|
||||
project_repo=project_repo,
|
||||
library_repo=library_repo,
|
||||
ingest_repo=ingest_repo,
|
||||
)
|
||||
|
||||
file_content = b"fake video content"
|
||||
files = {"file": ("test.mp4", io.BytesIO(file_content), "video/mp4")}
|
||||
data = {"project_id": project.id, "library_id": library.id}
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
|
||||
assert resp.status_code == 200
|
||||
result = resp.json()
|
||||
assert "ingest_job_id" in result
|
||||
# 验证 ingest job 被创建
|
||||
job = ingest_repo.get(result["ingest_job_id"])
|
||||
assert job is not None
|
||||
|
||||
def test_response_contains_url(self):
|
||||
"""验证响应包含文件 URL。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
storage = MagicMock()
|
||||
expected_url = "https://bucket.oss.example.com/uploads/my-video.mp4"
|
||||
storage.upload_file.return_value = expected_url
|
||||
|
||||
client = _client(
|
||||
project_repo=project_repo,
|
||||
library_repo=library_repo,
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
file_content = b"fake video content"
|
||||
files = {"file": ("test.mp4", io.BytesIO(file_content), "video/mp4")}
|
||||
data = {"project_id": project.id, "library_id": library.id}
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
|
||||
assert resp.status_code == 200
|
||||
result = resp.json()
|
||||
assert "url" in result
|
||||
assert result["url"].startswith("https://")
|
||||
|
||||
|
||||
class TestFormUploadFilenameSanitization:
|
||||
"""文件名校验测试。"""
|
||||
|
||||
def _upload_with_filename(self, filename: str):
|
||||
"""辅助方法:使用指定文件名上传文件。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://bucket.oss.example.com/uploads/test.mp4"
|
||||
|
||||
client = _client(
|
||||
project_repo=project_repo,
|
||||
library_repo=library_repo,
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
file_content = b"fake video content"
|
||||
files = {"file": (filename, io.BytesIO(file_content), "video/mp4")}
|
||||
data = {"project_id": project.id, "library_id": library.id}
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
return resp, storage
|
||||
|
||||
def test_path_traversal_prevented(self):
|
||||
"""路径遍历防护:../ 被替换为 __。"""
|
||||
resp, storage = self._upload_with_filename("../../../etc/passwd")
|
||||
|
||||
assert resp.status_code == 200
|
||||
# 验证文件名被清理
|
||||
call_args = storage.upload_file.call_args
|
||||
storage_key = call_args[0][1] # 第二个位置参数是 storage_key
|
||||
# 提取文件名部分(去掉 uploads/{file_id}/ 前缀)
|
||||
filename_part = storage_key.split("/", 2)[-1]
|
||||
# 文件名部分不应包含 / 或 \
|
||||
assert "/" not in filename_part
|
||||
assert "\\" not in filename_part
|
||||
|
||||
def test_backslash_replaced(self):
|
||||
"""反斜杠被替换为下划线。"""
|
||||
resp, storage = self._upload_with_filename(r"folder\subfolder\video.mp4")
|
||||
|
||||
assert resp.status_code == 200
|
||||
call_args = storage.upload_file.call_args
|
||||
storage_key = call_args[0][1]
|
||||
assert "\\" not in storage_key
|
||||
|
||||
def test_empty_filename_becomes_unknown(self):
|
||||
"""空文件名被 FastAPI 拒绝(422)。"""
|
||||
resp, storage = self._upload_with_filename("")
|
||||
|
||||
# FastAPI 验证文件名不能为空
|
||||
assert resp.status_code == 422
|
||||
|
||||
def test_special_characters_in_filename(self):
|
||||
"""特殊字符文件名正常处理。"""
|
||||
resp, storage = self._upload_with_filename("my-video (2024) [HD].mp4")
|
||||
|
||||
assert resp.status_code == 200
|
||||
call_args = storage.upload_file.call_args
|
||||
storage_key = call_args[0][1]
|
||||
assert "my-video (2024) [HD].mp4" in storage_key
|
||||
|
||||
def test_filename_without_extension(self):
|
||||
"""无扩展名文件名正常处理。"""
|
||||
resp, storage = self._upload_with_filename("no-extension-file")
|
||||
|
||||
assert resp.status_code == 200
|
||||
call_args = storage.upload_file.call_args
|
||||
storage_key = call_args[0][1]
|
||||
assert "no-extension-file" in storage_key
|
||||
|
||||
|
||||
class TestFormUploadMissingFile:
|
||||
"""缺少文件字段测试。"""
|
||||
|
||||
def test_missing_file_field_returns_422(self):
|
||||
"""缺少文件字段时返回 422。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(project_repo=project_repo, library_repo=library_repo)
|
||||
|
||||
# 只提交表单数据,不上传文件
|
||||
data = {"project_id": project.id, "library_id": library.id}
|
||||
resp = client.post("/api/v1", data=data)
|
||||
|
||||
assert resp.status_code == 422
|
||||
|
||||
|
||||
class TestFormUploadMissingParameters:
|
||||
"""缺少必填参数测试。"""
|
||||
|
||||
def test_missing_project_id_returns_422(self):
|
||||
"""缺少 project_id 时返回 422。"""
|
||||
library = _make_library()
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(library_repo=library_repo)
|
||||
|
||||
file_content = b"fake video content"
|
||||
files = {"file": ("test.mp4", io.BytesIO(file_content), "video/mp4")}
|
||||
data = {"library_id": library.id} # 缺少 project_id
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
|
||||
assert resp.status_code == 422
|
||||
|
||||
def test_missing_library_id_returns_422(self):
|
||||
"""缺少 library_id 时返回 422。"""
|
||||
project = _make_project()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
|
||||
client = _client(project_repo=project_repo)
|
||||
|
||||
file_content = b"fake video content"
|
||||
files = {"file": ("test.mp4", io.BytesIO(file_content), "video/mp4")}
|
||||
data = {"project_id": project.id} # 缺少 library_id
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
|
||||
assert resp.status_code == 422
|
||||
|
||||
def test_empty_project_id_returns_422(self):
|
||||
"""project_id 为空字符串时返回 422(min_length=1)。"""
|
||||
library = _make_library()
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(library_repo=library_repo)
|
||||
|
||||
file_content = b"fake video content"
|
||||
files = {"file": ("test.mp4", io.BytesIO(file_content), "video/mp4")}
|
||||
data = {"project_id": "", "library_id": library.id}
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
|
||||
assert resp.status_code == 422
|
||||
|
||||
|
||||
class TestFormUploadProjectNotFound:
|
||||
"""项目/素材库不存在测试。"""
|
||||
|
||||
def test_returns_404_when_project_not_found(self):
|
||||
"""项目不存在时返回 404。"""
|
||||
library = _make_library()
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(
|
||||
project_repo=StubProjectRepository(),
|
||||
library_repo=library_repo,
|
||||
)
|
||||
|
||||
file_content = b"fake video content"
|
||||
files = {"file": ("test.mp4", io.BytesIO(file_content), "video/mp4")}
|
||||
data = {"project_id": "nonexistent", "library_id": library.id}
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
|
||||
assert resp.status_code == 404
|
||||
assert "Project not found" in resp.json()["detail"]
|
||||
|
||||
def test_returns_404_when_library_not_in_project(self):
|
||||
"""素材库不属于该项目时返回 404。"""
|
||||
project = _make_project()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
|
||||
# 素材库属于另一个项目
|
||||
other_library = _make_library(project_id="other-project")
|
||||
library_repo = StubAssetLibraryRepository({other_library.id: other_library})
|
||||
|
||||
client = _client(project_repo=project_repo, library_repo=library_repo)
|
||||
|
||||
file_content = b"fake video content"
|
||||
files = {"file": ("test.mp4", io.BytesIO(file_content), "video/mp4")}
|
||||
data = {"project_id": project.id, "library_id": "nonexistent-lib"}
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
|
||||
assert resp.status_code == 404
|
||||
assert "Asset library not found" in resp.json()["detail"]
|
||||
|
||||
|
||||
class TestFormUploadOSSNotConfigured:
|
||||
"""OSS 未配置测试。"""
|
||||
|
||||
def test_returns_503_when_oss_not_configured(self):
|
||||
"""OSS 未配置时返回 503(RuntimeError)。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.side_effect = RuntimeError("OSS 未配置")
|
||||
|
||||
client = _client(
|
||||
project_repo=project_repo,
|
||||
library_repo=library_repo,
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
file_content = b"fake video content"
|
||||
files = {"file": ("test.mp4", io.BytesIO(file_content), "video/mp4")}
|
||||
data = {"project_id": project.id, "library_id": library.id}
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
|
||||
assert resp.status_code == 503
|
||||
|
||||
|
||||
class TestFormUploadStorageErrors:
|
||||
"""存储服务错误测试。"""
|
||||
|
||||
def test_returns_500_on_generic_storage_error(self):
|
||||
"""存储服务通用错误返回 500。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.side_effect = Exception("Network error")
|
||||
|
||||
client = _client(
|
||||
project_repo=project_repo,
|
||||
library_repo=library_repo,
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
file_content = b"fake video content"
|
||||
files = {"file": ("test.mp4", io.BytesIO(file_content), "video/mp4")}
|
||||
data = {"project_id": project.id, "library_id": library.id}
|
||||
|
||||
resp = client.post("/api/v1", files=files, data=data)
|
||||
|
||||
assert resp.status_code == 500
|
||||
assert "Failed to upload file" in resp.json()["detail"]
|
||||
|
||||
|
||||
class TestFormUploadMultipleFormats:
|
||||
"""多格式支持测试。"""
|
||||
|
||||
def _upload_with_mime(self, mime_type: str, filename: str = "test.mp4"):
|
||||
"""辅助方法:使用指定 MIME 类型上传文件。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://bucket.oss.example.com/uploads/test.mp4"
|
||||
|
||||
client = _client(
|
||||
project_repo=project_repo,
|
||||
library_repo=library_repo,
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
file_content = b"fake content"
|
||||
files = {"file": (filename, io.BytesIO(file_content), mime_type)}
|
||||
data = {"project_id": project.id, "library_id": library.id}
|
||||
|
||||
return client.post("/api/v1", files=files, data=data)
|
||||
|
||||
def test_multiple_video_formats(self):
|
||||
"""支持多种视频格式(4种)。"""
|
||||
video_formats = [
|
||||
("video/mp4", "test.mp4"),
|
||||
("video/mpeg", "test.mpeg"),
|
||||
("video/quicktime", "test.mov"),
|
||||
("video/x-msvideo", "test.avi"),
|
||||
]
|
||||
for mime_type, filename in video_formats:
|
||||
resp = self._upload_with_mime(mime_type, filename)
|
||||
assert resp.status_code == 200, f"Failed for {mime_type}"
|
||||
|
||||
def test_multiple_audio_formats(self):
|
||||
"""支持多种音频格式(5种)。"""
|
||||
audio_formats = [
|
||||
("audio/mpeg", "test.mp3"),
|
||||
("audio/wav", "test.wav"),
|
||||
("audio/ogg", "test.ogg"),
|
||||
("audio/flac", "test.flac"),
|
||||
("audio/aac", "test.aac"),
|
||||
]
|
||||
for mime_type, filename in audio_formats:
|
||||
resp = self._upload_with_mime(mime_type, filename)
|
||||
assert resp.status_code == 200, f"Failed for {mime_type}"
|
||||
|
||||
def test_multiple_image_formats(self):
|
||||
"""支持多种图片格式(6种)。"""
|
||||
image_formats = [
|
||||
("image/jpeg", "test.jpg"),
|
||||
("image/png", "test.png"),
|
||||
("image/gif", "test.gif"),
|
||||
("image/webp", "test.webp"),
|
||||
("image/bmp", "test.bmp"),
|
||||
("image/svg+xml", "test.svg"),
|
||||
]
|
||||
for mime_type, filename in image_formats:
|
||||
resp = self._upload_with_mime(mime_type, filename)
|
||||
assert resp.status_code == 200, f"Failed for {mime_type}"
|
||||
|
||||
|
||||
class TestFormUploadMIMEValidation:
|
||||
"""MIME 类型验证测试。"""
|
||||
|
||||
def _upload_with_content_type(self, content_type: str | None):
|
||||
"""辅助方法:使用指定 Content-Type 上传文件。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://bucket.oss.example.com/uploads/test.mp4"
|
||||
|
||||
client = _client(
|
||||
project_repo=project_repo,
|
||||
library_repo=library_repo,
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
file_content = b"fake content"
|
||||
# 使用元组形式明确指定 content_type
|
||||
files = {"file": ("test.mp4", io.BytesIO(file_content), content_type)}
|
||||
data = {"project_id": project.id, "library_id": library.id}
|
||||
|
||||
return client.post("/api/v1", files=files, data=data)
|
||||
|
||||
def test_valid_mime_type_accepted(self):
|
||||
"""有效 MIME 类型被接受。"""
|
||||
resp = self._upload_with_content_type("video/mp4")
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_empty_content_type_returns_400(self):
|
||||
"""空 Content-Type 返回 400。"""
|
||||
# 使用空字符串作为 content type
|
||||
resp = self._upload_with_content_type("")
|
||||
assert resp.status_code == 400
|
||||
assert "Content-Type header is required" in resp.json()["detail"]
|
||||
|
||||
def test_unsupported_content_type_returns_415(self):
|
||||
"""不支持的 Content-Type 返回 415。"""
|
||||
resp = self._upload_with_content_type("text/plain")
|
||||
assert resp.status_code == 415
|
||||
assert "not supported" in resp.json()["detail"]
|
||||
@@ -0,0 +1,390 @@
|
||||
"""
|
||||
upload.py 路由单元测试
|
||||
|
||||
覆盖:
|
||||
- _require_project_and_library 中 find_by_project 调用正确性
|
||||
- OSS 凭证校验(未配置时返回 503)
|
||||
- prepare_direct_upload 正常路径
|
||||
- 文件类型校验
|
||||
- 异常处理路径
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# 设置必要环境变量(必须在导入 app 模块之前)
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
import pytest
|
||||
|
||||
# 确保 app 模块可导入
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from packages.domain import AssetLibrary, AssetLibraryKind, Project
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 测试用 Stub(不继承 Port ABC,因为 Port 定义 async 方法,路由实际使用同步 duck-type)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class StubProjectRepository:
|
||||
def __init__(self, projects: dict[str, Project] | None = None):
|
||||
self._projects = projects or {}
|
||||
|
||||
def get(self, project_id: str) -> Project | None:
|
||||
return self._projects.get(project_id)
|
||||
|
||||
def find_by_id(self, project_id: str) -> Project | None:
|
||||
return self._projects.get(project_id)
|
||||
|
||||
|
||||
class StubAssetLibraryRepository:
|
||||
def __init__(self, libraries: dict[str, AssetLibrary] | None = None):
|
||||
self._libraries = libraries or {}
|
||||
|
||||
def find_by_project(self, project_id: str, kind=None) -> list[AssetLibrary]:
|
||||
items = [lib for lib in self._libraries.values() if lib.project_id == project_id]
|
||||
if kind is not None:
|
||||
items = [lib for lib in items if lib.kind == kind]
|
||||
return items
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[AssetLibrary]:
|
||||
"""故意保留旧方法名,验证路由不会调用它。"""
|
||||
raise AssertionError("路由不应调用 list_by_project,应调用 find_by_project")
|
||||
|
||||
|
||||
class StubIngestJobRepository:
|
||||
def add(self, job) -> None:
|
||||
pass
|
||||
|
||||
def get(self, job_id: str):
|
||||
return None
|
||||
|
||||
def update_status(self, job_id, status, **kwargs):
|
||||
pass
|
||||
|
||||
def list_by_library(self, library_id: str, skip: int = 0, limit: int = 50):
|
||||
return []
|
||||
|
||||
def count_by_library(self, library_id: str) -> int:
|
||||
return 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 测试 Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_project(id: str = "proj-1", owner_user_id: str = "user-1") -> Project:
|
||||
return Project(id=id, name="Test Project", owner_user_id=owner_user_id)
|
||||
|
||||
|
||||
def _make_library(
|
||||
id: str = "lib-1",
|
||||
project_id: str = "proj-1",
|
||||
kind: AssetLibraryKind = AssetLibraryKind.VIDEO,
|
||||
) -> AssetLibrary:
|
||||
return AssetLibrary(id=id, name="Test Library", project_id=project_id, kind=kind)
|
||||
|
||||
|
||||
def _build_app(
|
||||
project_repo: StubProjectRepository | None = None,
|
||||
library_repo: StubAssetLibraryRepository | None = None,
|
||||
storage: MagicMock | None = None,
|
||||
ingest_repo: StubIngestJobRepository | None = None,
|
||||
) -> FastAPI:
|
||||
"""构建一个最小化的 FastAPI app,只注册 upload 路由。"""
|
||||
from app.api.routes.upload import router
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import (
|
||||
get_asset_library_repository,
|
||||
get_ingest_job_repository,
|
||||
get_project_repository,
|
||||
)
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(router, prefix="/api/v1")
|
||||
|
||||
project_repo = project_repo or StubProjectRepository()
|
||||
library_repo = library_repo or StubAssetLibraryRepository()
|
||||
storage = storage or MagicMock()
|
||||
storage.is_configured = True
|
||||
storage.create_direct_upload_post.return_value = {
|
||||
"url": "https://bucket.oss.example.com",
|
||||
"method": "POST",
|
||||
"storage_key": "uploads/abc/test.mp4",
|
||||
"expires_at": "2026-01-01T00:00:00Z",
|
||||
"fields": {"key": "uploads/abc/test.mp4"},
|
||||
}
|
||||
ingest_repo = ingest_repo or StubIngestJobRepository()
|
||||
|
||||
# Mock auth
|
||||
mock_user = MagicMock(spec=AuthenticatedUser)
|
||||
mock_user.id = "user-1"
|
||||
mock_user.email = "test@example.com"
|
||||
|
||||
app.dependency_overrides[get_current_user] = lambda: mock_user
|
||||
app.dependency_overrides[get_project_repository] = lambda: project_repo
|
||||
app.dependency_overrides[get_asset_library_repository] = lambda: library_repo
|
||||
app.dependency_overrides[get_storage_service] = lambda: storage
|
||||
app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo
|
||||
|
||||
return app
|
||||
|
||||
|
||||
def _client(**kwargs) -> TestClient:
|
||||
app = _build_app(**kwargs)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 测试用例
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRequireProjectAndLibrary:
|
||||
"""_require_project_and_library 辅助函数测试。"""
|
||||
|
||||
def test_returns_200_when_project_and_library_exist(self):
|
||||
"""项目和素材库都存在时,正常返回。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(project_repo=project_repo, library_repo=library_repo)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/direct/prepare",
|
||||
json={
|
||||
"project_id": project.id,
|
||||
"library_id": library.id,
|
||||
"filename": "test.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": 1024,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_returns_404_when_project_not_found(self):
|
||||
"""项目不存在时返回 404。"""
|
||||
library = _make_library()
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(
|
||||
project_repo=StubProjectRepository(),
|
||||
library_repo=library_repo,
|
||||
)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/direct/prepare",
|
||||
json={
|
||||
"project_id": "nonexistent",
|
||||
"library_id": library.id,
|
||||
"filename": "test.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": 1024,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
assert "Project not found" in resp.json()["detail"]
|
||||
|
||||
def test_returns_404_when_library_not_found(self):
|
||||
"""素材库不属于该项目时返回 404。"""
|
||||
project = _make_project()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
|
||||
other_library = _make_library(project_id="other-project")
|
||||
library_repo = StubAssetLibraryRepository({other_library.id: other_library})
|
||||
|
||||
client = _client(project_repo=project_repo, library_repo=library_repo)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/direct/prepare",
|
||||
json={
|
||||
"project_id": project.id,
|
||||
"library_id": "nonexistent-lib",
|
||||
"filename": "test.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": 1024,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
assert "Asset library not found" in resp.json()["detail"]
|
||||
|
||||
def test_find_by_project_is_called_not_list_by_project(self):
|
||||
"""
|
||||
验证路由调用的是 find_by_project 而不是 list_by_project。
|
||||
StubAssetLibraryRepository.list_by_project 会抛出 AssertionError。
|
||||
"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(project_repo=project_repo, library_repo=library_repo)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/direct/prepare",
|
||||
json={
|
||||
"project_id": project.id,
|
||||
"library_id": library.id,
|
||||
"filename": "test.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": 1024,
|
||||
},
|
||||
)
|
||||
# 如果调用了 list_by_project,会抛 AssertionError 导致 500
|
||||
assert resp.status_code == 200
|
||||
|
||||
|
||||
class TestPrepareDirectUpload:
|
||||
"""prepare_direct_upload 端点测试。"""
|
||||
|
||||
def test_returns_upload_credentials_when_configured(self):
|
||||
"""OSS 已配置时,返回上传凭证。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
storage = MagicMock()
|
||||
storage.is_configured = True
|
||||
storage.create_direct_upload_post.return_value = {
|
||||
"url": "https://bucket.oss.example.com",
|
||||
"method": "POST",
|
||||
"storage_key": "uploads/abc/test-video.mp4",
|
||||
"expires_at": "2026-01-01T00:00:00Z",
|
||||
"fields": {"key": "uploads/abc/test-video.mp4"},
|
||||
}
|
||||
|
||||
client = _client(
|
||||
project_repo=project_repo,
|
||||
library_repo=library_repo,
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/direct/prepare",
|
||||
json={
|
||||
"project_id": project.id,
|
||||
"library_id": library.id,
|
||||
"filename": "test-video.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": 1024 * 1024,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert "upload_url" in data
|
||||
assert "storage_key" in data
|
||||
|
||||
def test_returns_503_when_oss_not_configured(self):
|
||||
"""OSS 未配置时,返回 503。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
storage = MagicMock()
|
||||
storage.create_direct_upload_post.side_effect = RuntimeError("OSS 未配置")
|
||||
|
||||
client = _client(
|
||||
project_repo=project_repo,
|
||||
library_repo=library_repo,
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/direct/prepare",
|
||||
json={
|
||||
"project_id": project.id,
|
||||
"library_id": library.id,
|
||||
"filename": "test.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": 1024,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 503
|
||||
|
||||
|
||||
class TestCompleteDirectUpload:
|
||||
"""complete_direct_upload 端点测试。"""
|
||||
|
||||
def test_returns_404_when_project_not_found(self):
|
||||
"""项目不存在时返回 404。"""
|
||||
library = _make_library()
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(
|
||||
project_repo=StubProjectRepository(),
|
||||
library_repo=library_repo,
|
||||
)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/direct/complete",
|
||||
json={
|
||||
"project_id": "nonexistent",
|
||||
"library_id": library.id,
|
||||
"storage_key": "uploads/abc/test.mp4",
|
||||
"filename": "test.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": 1024,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
class TestMimeTypeValidation:
|
||||
"""文件类型校验测试。"""
|
||||
|
||||
def test_accepts_video_mp4(self):
|
||||
"""video/mp4 是合法类型。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(project_repo=project_repo, library_repo=library_repo)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/direct/prepare",
|
||||
json={
|
||||
"project_id": project.id,
|
||||
"library_id": library.id,
|
||||
"filename": "test.mp4",
|
||||
"content_type": "video/mp4",
|
||||
"file_size": 1024,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_rejects_invalid_mime_type(self):
|
||||
"""非法文件类型被拒绝。"""
|
||||
project = _make_project()
|
||||
library = _make_library()
|
||||
project_repo = StubProjectRepository({project.id: project})
|
||||
library_repo = StubAssetLibraryRepository({library.id: library})
|
||||
|
||||
client = _client(project_repo=project_repo, library_repo=library_repo)
|
||||
|
||||
resp = client.post(
|
||||
"/api/v1/direct/prepare",
|
||||
json={
|
||||
"project_id": project.id,
|
||||
"library_id": library.id,
|
||||
"filename": "malware.exe",
|
||||
"content_type": "application/x-executable",
|
||||
"file_size": 1024,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 415
|
||||
Reference in New Issue
Block a user