Files

924 lines
35 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import json
import logging
import os
from datetime import timezone, timedelta
from types import SimpleNamespace
CST = timezone(timedelta(hours=8))
from fastapi import APIRouter, Depends, HTTPException, Query, UploadFile, File, status
from fastapi.responses import RedirectResponse
from sqlalchemy import select, func
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_db, get_current_user
from app.models.user import User
from app.models.project import Project
from app.models.generation_record import GenerationRecord
from app.schemas.generation import (
OptimizeParams,
GenerationRecordOut,
GenerationRecordPageListOut,
OptimizeResult,
UpdatePromptRequest,
)
from app.services.generation.pipeline.db_lock_service import (
DatabaseRowLockBusy,
execute_with_lock_timeout,
)
from app.services.video_url import validate_and_get_record_id, get_video_stream_url
from app.services.private_portrait.reference_resolver import (
batch_resolve_private_portrait_reference_display_urls,
resolve_private_portrait_reference_display_urls,
resolve_private_portrait_references,
)
from app.services.resource_signed_url_service import build_resource_signed_url
from app.services.resource_capacity_service import assert_user_resource_capacity_available
from app.services.upload_resource import delete_unbound_upload_resource, upload_reference_file, cleanup_upload_resource_files_after_commit
from app.services.upload_resource.log_service import log_upload_resource_exception, safe_rollback_with_log
from app.enums.upload_resource import UploadResourceEventEnum, UploadResourceModuleEnum, UploadResourceTypeEnum
from app.enums.generation_status import (
ASPECT_RATIOS,
DURATIONS,
IMAGE_SIZES,
RESOLUTIONS,
GenerationRecordPipelineStage,
GenerationType,
)
from app.enums.common import LogEventStatusEnum
from app.enums.generation_record import (
GenerationRecordConfigSourceEnum,
GenerationRecordEventTypeEnum,
)
from app.services.generation.billing_service import (
OWNER_GENERATION_RECORD,
charge_generation_media_for_record,
get_next_credit_attempt_no,
)
from app.services.generation.ai.engine_service import (
get_image_engine,
get_video_engine,
)
from app.services.generation.pipeline.generation_record_config_service import (
ensure_generation_record_config_frozen,
frozen_generation_record_engine_view,
generation_record_config_fallback_hint,
generation_record_engine_snapshot,
is_generation_record_config_complete,
is_generation_record_config_recoverable,
log_generation_record_config_event,
)
from app.services.generation.media_reference_service import (
calculate_media_reference_usage,
validate_media_reference_usage_for_engine,
)
from app.services.generation.prompt_optimize_service import optimize_generation_prompt
from app.enums.audio_reference import (
AUDIO_ALLOWED_EXTENSIONS,
AUDIO_ALLOWED_MIME_TYPES,
AUDIO_MAX_FILE_SIZE_MB,
)
from app.utils.exceptions import RecordNotFoundError, InvalidStatusError
router = APIRouter(prefix="/generation-records", tags=["generation"])
logger = logging.getLogger("videogen")
def _engine_snapshot(record: GenerationRecord) -> dict | None:
return generation_record_engine_snapshot(record)
def _record_config_complete(record: GenerationRecord) -> bool:
return is_generation_record_config_complete(record)
def _record_status_view(record: GenerationRecord) -> dict[str, object]:
config_complete = _record_config_complete(record)
config_recoverable = is_generation_record_config_recoverable(record)
prompt_failure = record.status == "failed" and record.resource_generation_started_at is None
resource_failure = record.status == "failed" and record.resource_generation_started_at is not None
if record.status in {"pending", "optimizing", "settlement_pending"}:
client_status = "prompt_processing"
operation_phase = "prompt"
elif record.status == "prompt_optimized":
client_status = "ready"
operation_phase = "prompt"
elif record.status == "generating":
client_status = "generating"
operation_phase = "resource"
elif record.status == "completed":
client_status = "success"
operation_phase = "resource"
else:
client_status = "failure"
operation_phase = "prompt" if prompt_failure else "resource"
return {
"config_complete": config_complete,
"config_recoverable": config_recoverable,
"config_fallback_hint": generation_record_config_fallback_hint(record),
"can_generate": record.status == "prompt_optimized" and (config_complete or config_recoverable),
"can_retry": resource_failure and config_complete and record.pipeline_stage != GenerationRecordPipelineStage.UPSCALE_FAILED.value,
"should_poll": record.status in {"optimizing", "settlement_pending", "generating"},
"client_status": client_status,
"operation_phase": operation_phase,
}
def _frozen_engine_view(record: GenerationRecord) -> SimpleNamespace:
return frozen_generation_record_engine_view(record)
def _record_to_out(record: GenerationRecord, project_name: str, refs_override: list[dict] | None = None) -> GenerationRecordOut:
refs = refs_override
if refs is None and record.media_references:
try:
refs = json.loads(record.media_references)
except (json.JSONDecodeError, TypeError):
refs = None
error_message = record.error_message
if error_message:
from app.services.error_codes import ARK_ERRORS
import re
match = re.search(r"code='([^']+)'", error_message)
if match:
code = match.group(1)
if code in ARK_ERRORS:
error_message = ARK_ERRORS[code]
else:
parts = error_message.split(":")
if len(parts) >= 2 and parts[1].strip() in ARK_ERRORS:
error_message = ARK_ERRORS[parts[1].strip()]
return GenerationRecordOut(
id=record.id,
project_id=record.project_id,
project_name=project_name,
original_prompt=record.original_prompt,
optimized_prompt=record.optimized_prompt,
gen_type=record.gen_type,
duration=record.duration,
aspect_ratio=record.aspect_ratio,
resolution=record.resolution,
image_size=record.image_size,
image_proportion=record.image_proportion,
image_px=record.image_px,
status=record.status,
pipeline_stage=record.pipeline_stage,
video_upscale_enabled=bool(record.video_upscale_enabled_snapshot),
video_url=build_resource_signed_url(record.video_url) if record.video_url else '',
video_cover_url=build_resource_signed_url(record.video_cover_url) if record.video_cover_url else '',
image_url=build_resource_signed_url(record.image_url) if record.image_url else '',
references=refs,
engine_id=record.engine_id,
engine_name=(_engine_snapshot(record) or {}).get("name"),
engine_snapshot=_engine_snapshot(record),
include_media_references=bool(record.include_media_references),
**_record_status_view(record),
text_credits_cost=round(record.text_credits_cost or 0.00, 2),
text_tokens_used=record.text_tokens_used or 0,
credits_cost=round(record.credits_cost or 0.00, 2),
video_tokens_used=record.video_tokens_used or 0,
image_tokens_used=record.image_tokens_used or 0,
error_message=error_message,
created_at=record.created_at,
generated_at=record.generated_at,
)
@router.get("", response_model=GenerationRecordPageListOut)
async def list_records(
project_id: str | None = Query(
None,
alias="project_id",
description="查询单个项目的生成记录",
),
status: str | None = Query(
None,
description="查询状态,可以不传。prompt_optimized:待生成 | generating:生成中 | failed:失败 | completed:成功",
examples=["completed"],
),
page: int = Query(
1,
ge=1,
description="分页页码,从1开始",
examples=[1],
),
page_size: int = Query(
10,
ge=1,
le=100,
description="每页返回的生成记录数量,范围 1~100",
examples=[10],
),
record_ids: list[str] | None = Query(
None,
description="对应记录ID数组",
examples=[["0019e8c55ddd1429b86", "0019e8c54dfc1e13262"]],
),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
allowed_statuses = {
"optimizing",
"settlement_pending",
"prompt_optimized",
"generating",
"failed",
"completed",
}
if status and status not in allowed_statuses:
raise HTTPException(
status_code=400,
detail="状态参数错误,仅支持:optimizing、settlement_pending、prompt_optimized、generating、failed、completed",
)
offset = (page - 1) * page_size
conditions = [
GenerationRecord.user_id == current_user.id,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
]
if project_id:
conditions.append(GenerationRecord.project_id == project_id)
if status:
conditions.append(GenerationRecord.status == status)
if record_ids:
conditions.append(GenerationRecord.id.in_(record_ids))
total_result = await db.execute(
select(func.count(GenerationRecord.id))
.join(Project, GenerationRecord.project_id == Project.id)
.where(*conditions)
)
total = total_result.scalar_one() or 0
query = (
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(*conditions)
.order_by(GenerationRecord.created_at.desc())
.offset(offset)
.limit(page_size)
)
result = await db.execute(query)
rows = result.all()
refs_map = await batch_resolve_private_portrait_reference_display_urls(
db,
{record.id: json.loads(record.media_references) if record.media_references else None for record, _project_name in rows},
user_id=current_user.id,
)
return {
"total": int(total),
"page": page,
"page_size": page_size,
"items": [
_record_to_out(record, project_name, refs_override=refs_map.get(record.id))
for record, project_name in rows
],
}
@router.post("/optimize", response_model=OptimizeResult)
async def optimize(
req: OptimizeParams,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
user_id_snapshot = str(current_user.id)
service_result = await optimize_generation_prompt(
db,
req=req,
user_id=user_id_snapshot,
)
refreshed = await db.execute(
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.id == service_result.record_id,
GenerationRecord.user_id == user_id_snapshot,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
)
.limit(1)
)
refreshed_row = refreshed.first()
if not refreshed_row:
raise RecordNotFoundError()
record, project_name = refreshed_row
refs = await resolve_private_portrait_reference_display_urls(
db,
json.loads(record.media_references) if record.media_references else None,
user_id=user_id_snapshot,
)
return OptimizeResult(
optimized_prompt=record.optimized_prompt or "",
text_credits_cost=round(record.text_credits_cost or 0, 2),
text_tokens_used=record.text_tokens_used or 0,
record=_record_to_out(record, project_name, refs_override=refs),
)
@router.post("/{record_id}/generate")
async def generate_record_resource(
record_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
user_id_snapshot = str(current_user.id)
try:
result = await execute_with_lock_timeout(
db,
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.id == record_id,
GenerationRecord.user_id == user_id_snapshot,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
)
.with_for_update(),
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
row = result.first()
if not row:
raise RecordNotFoundError()
record, project_name = row
if record.status != "prompt_optimized":
raise InvalidStatusError("只有提词已完成的记录可以发起资源生成")
log_generation_record_config_event(
event_type=GenerationRecordEventTypeEnum.GENERATION_SUBMIT_START,
event_status=LogEventStatusEnum.STARTED,
source=GenerationRecordConfigSourceEnum.LEGACY_GENERATE_FALLBACK,
record=record,
detail={
"record_id": record.id,
"project_id": record.project_id,
"gen_type": record.gen_type,
"config_complete_before": _record_config_complete(record),
"config_recoverable": is_generation_record_config_recoverable(record),
},
)
await ensure_generation_record_config_frozen(
db,
record,
source=GenerationRecordConfigSourceEnum.LEGACY_GENERATE_FALLBACK,
)
if not _record_config_complete(record):
raise InvalidStatusError("该记录缺少冻结的生成配置,请重新生成提词")
log_generation_record_config_event(
event_type=GenerationRecordEventTypeEnum.GENERATION_SUBMIT_CONFIG_READY,
event_status=LogEventStatusEnum.SUCCESS,
source=GenerationRecordConfigSourceEnum.LEGACY_GENERATE_FALLBACK,
record=record,
detail={
"record_id": record.id,
"project_id": record.project_id,
"gen_type": record.gen_type,
"engine_id": record.engine_id,
"duration": record.duration,
"aspect_ratio": record.aspect_ratio,
"resolution": record.resolution,
"provider_generation_resolution": record.provider_generation_resolution,
"image_size": record.image_size,
"image_proportion": record.image_proportion,
"image_px": record.image_px,
"include_media_references": bool(record.include_media_references),
},
)
await assert_user_resource_capacity_available(db, user_id_snapshot)
attempt_no = await get_next_credit_attempt_no(
db,
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
)
# Confirm that the bound engine still exists and is active, but never rebuild
# the snapshot or replace the user's frozen parameters with current defaults.
if record.gen_type == GenerationType.video.value:
await get_video_engine(db, record.engine_id)
else:
await get_image_engine(db, record.engine_id)
frozen_engine = _frozen_engine_view(record)
# 触发生成前兜底:私域素材在 GenerationRecord 创建时可能未走 resolve_private_portrait_references
# 导致入库 url 存的是前端预览地址而非供应商需要的 asset://。这里强制重新解析,
# 确保供应商侧拿到正确的 remote_asset_id / asset:// URI。
try:
raw_refs = json.loads(record.media_references) if record.media_references else None
except (TypeError, ValueError):
raw_refs = None
if raw_refs:
resolved_refs = await resolve_private_portrait_references(
db,
user_id=user_id_snapshot,
media_references=raw_refs,
gen_type=record.gen_type,
)
if resolved_refs is not None:
record.media_references = json.dumps(resolved_refs, ensure_ascii=False)
reference_usage = calculate_media_reference_usage(
record.media_references,
include=bool(record.include_media_references),
)
validate_media_reference_usage_for_engine(
reference_usage,
gen_type=record.gen_type,
engine=frozen_engine,
)
billing = await charge_generation_media_for_record(
db,
record=record,
project_name=project_name,
description_prefix=project_name + "-",
attempt_no=attempt_no,
engine_id=record.engine_id,
)
record.credits_cost = round(
float(record.credits_cost or 0) + float(billing.total_charged or 0),
2,
)
log_generation_record_config_event(
event_type=GenerationRecordEventTypeEnum.GENERATION_SUBMIT_BILLING_SUCCESS,
event_status=LogEventStatusEnum.SUCCESS,
source=GenerationRecordConfigSourceEnum.LEGACY_GENERATE_FALLBACK,
record=record,
detail={
"record_id": record.id,
"attempt_no": attempt_no,
"engine_id": record.engine_id,
"charged": float(billing.total_charged or 0),
"credits_cost_total": record.credits_cost,
},
)
from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record,
prepare_generation_record_execution,
)
prepare_generation_record_execution(record, attempt_no=attempt_no)
await db.flush()
record_id_snapshot = str(record.id)
enqueue_log_record = GenerationRecord(
id=record_id_snapshot,
user_id=user_id_snapshot,
project_id=str(record.project_id),
original_prompt=record.original_prompt or "",
gen_type=record.gen_type,
duration=record.duration,
aspect_ratio=record.aspect_ratio,
resolution=record.resolution,
image_size=record.image_size,
image_proportion=record.image_proportion,
image_px=record.image_px,
engine_id=record.engine_id,
include_media_references=bool(record.include_media_references),
media_references=record.media_references,
)
enqueue_log_detail = {
"record_id": record_id_snapshot,
"attempt_no": attempt_no,
"reason": "generation_record_api_generate",
}
await commit_and_enqueue_generation_record(
db,
record,
reason="generation_record_api_generate",
)
log_generation_record_config_event(
event_type=GenerationRecordEventTypeEnum.GENERATION_SUBMIT_ENQUEUE_SUCCESS,
event_status=LogEventStatusEnum.SUCCESS,
source=GenerationRecordConfigSourceEnum.LEGACY_GENERATE_FALLBACK,
record=enqueue_log_record,
detail=enqueue_log_detail,
)
refreshed = await db.execute(
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(GenerationRecord.id == record_id_snapshot)
.limit(1)
)
refreshed_row = refreshed.first()
if not refreshed_row:
raise RecordNotFoundError()
record, project_name = refreshed_row
refs = await resolve_private_portrait_reference_display_urls(
db,
json.loads(record.media_references) if record.media_references else None,
user_id=user_id_snapshot,
)
return _record_to_out(record, project_name, refs_override=refs)
@router.post("/{record_id}/retry")
async def retry_generation(
record_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
user_id_snapshot = str(current_user.id)
try:
result = await execute_with_lock_timeout(
db,
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.id == record_id,
GenerationRecord.user_id == user_id_snapshot,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
)
.with_for_update(),
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
row = result.first()
if not row:
raise RecordNotFoundError()
record, project_name = row
if record.status != "failed":
raise InvalidStatusError("只有失败的记录可以重试")
if record.resource_generation_started_at is None:
raise InvalidStatusError("该记录失败于提词阶段,请重新生成提词")
if record.pipeline_stage == GenerationRecordPipelineStage.UPSCALE_FAILED.value:
raise InvalidStatusError("该任务生成失败,请联系客服进行修复")
if not _record_config_complete(record):
raise InvalidStatusError("该记录缺少冻结的生成配置,请重新生成提词")
await assert_user_resource_capacity_available(db, user_id_snapshot)
attempt_no = await get_next_credit_attempt_no(
db,
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
)
if record.gen_type == GenerationType.video.value:
await get_video_engine(db, record.engine_id)
else:
await get_image_engine(db, record.engine_id)
frozen_engine = _frozen_engine_view(record)
# 重试前兜底:私域素材 url 可能仍然是预览地址,重新解析确保供应商拿到 asset://
try:
raw_refs = json.loads(record.media_references) if record.media_references else None
except (TypeError, ValueError):
raw_refs = None
if raw_refs:
resolved_refs = await resolve_private_portrait_references(
db,
user_id=user_id_snapshot,
media_references=raw_refs,
gen_type=record.gen_type,
)
if resolved_refs is not None:
record.media_references = json.dumps(resolved_refs, ensure_ascii=False)
reference_usage = calculate_media_reference_usage(
record.media_references,
include=bool(record.include_media_references),
)
validate_media_reference_usage_for_engine(
reference_usage,
gen_type=record.gen_type,
engine=frozen_engine,
)
billing = await charge_generation_media_for_record(
db,
record=record,
project_name=project_name,
description_prefix="资源生成重试-",
attempt_no=attempt_no,
engine_id=record.engine_id,
)
record.credits_cost = round(
float(record.credits_cost or 0) + float(billing.total_charged or 0),
2,
)
record.manual_retry_count = int(record.manual_retry_count or 0) + 1
record.retry_count = int(record.manual_retry_count or 0)
from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record,
prepare_generation_record_execution,
)
prepare_generation_record_execution(record, attempt_no=attempt_no)
await db.flush()
record_id_snapshot = str(record.id)
await commit_and_enqueue_generation_record(
db,
record,
reason="generation_record_api_retry",
)
refreshed = await db.execute(
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(GenerationRecord.id == record_id_snapshot)
.limit(1)
)
refreshed_row = refreshed.first()
if not refreshed_row:
raise RecordNotFoundError()
record, project_name = refreshed_row
refs = await resolve_private_portrait_reference_display_urls(
db,
json.loads(record.media_references) if record.media_references else None,
user_id=user_id_snapshot,
)
return _record_to_out(record, project_name, refs_override=refs)
@router.put("/{record_id}/prompt")
async def update_prompt(
record_id: str,
req: UpdatePromptRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(GenerationRecord).where(
GenerationRecord.id == record_id,
GenerationRecord.user_id == current_user.id,
GenerationRecord.deleted_at.is_(None),
)
.limit(1)
)
record = result.scalar_one_or_none()
if not record:
raise RecordNotFoundError()
if record.status != "prompt_optimized":
raise InvalidStatusError("只有待生成状态可以修改提示词")
record.optimized_prompt = req.optimized_prompt
await db.flush()
return {"message": "ok"}
@router.get("/{record_id}/video")
async def get_video(
record_id: str,
token: str = Query(...),
db: AsyncSession = Depends(get_db),
):
"""Validate temp token and redirect to video URL."""
validated_id = await validate_and_get_record_id(token)
if validated_id != record_id:
raise HTTPException(status_code=403, detail="无效的视频链接")
video_url = await get_video_stream_url(db, record_id)
if not video_url:
raise HTTPException(status_code=404, detail="视频不存在")
return RedirectResponse(url=video_url)
@router.get("/{record_id}/queue-status")
async def get_queue_status(
record_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""Get queue position, estimated wait time, and current status for a generation record."""
from sqlalchemy import func
result = await db.execute(
select(GenerationRecord).where(
GenerationRecord.id == record_id,
GenerationRecord.user_id == current_user.id,
GenerationRecord.deleted_at.is_(None),
)
.limit(1)
)
record = result.scalar_one_or_none()
if not record:
raise RecordNotFoundError()
queue_position = None
estimated_wait_seconds = None
if record.status == "generating":
resource_started_at = record.resource_generation_started_at or record.created_at
ahead_result = await db.execute(
select(func.count(GenerationRecord.id)).where(
GenerationRecord.status == "generating",
GenerationRecord.deleted_at.is_(None),
func.coalesce(
GenerationRecord.resource_generation_started_at,
GenerationRecord.created_at,
) < resource_started_at,
)
)
ahead = ahead_result.scalar() or 0
queue_position = ahead + 1
estimated_wait_seconds = ahead * 60
return {
"record_id": record.id,
"status": record.status,
"pipeline_stage": record.pipeline_stage,
"video_upscale_enabled": bool(record.video_upscale_enabled_snapshot),
"queue_position": queue_position,
"estimated_wait_seconds": estimated_wait_seconds,
}
@router.post(
"/upload-image",
summary="上传 AI 创作普通参考图片",
description=(
"上传普通 AI 创作参考图片,写入 UploadResource 资源账本并纳入用户上传容量统计。"
"返回 resource_id 和 url。该文件在未绑定业务记录前可通过 /generation-records/delete-file 单独删除,"
"也会出现在 /upload-resources/history 历史素材中供 AI 创作复用。"
),
responses={400: {"description": "文件类型、大小或容量校验失败"}, 401: {"description": "未登录或 Token 无效"}},
)
async def upload_image(
file: UploadFile = File(...),
current_user: User = Depends(get_current_user),
gen_type: str = Query("video", description="生成类型:video-视频,image-图片"),
db: AsyncSession = Depends(get_db),
):
"""上传普通参考图片,记录 UploadResource 并纳入用户容量统计。"""
result = await upload_reference_file(
db,
file=file,
current_user=current_user,
module=UploadResourceModuleEnum.COMMON.value,
resource_type=UploadResourceTypeEnum.IMAGE.value,
gen_type=gen_type,
)
await db.commit()
return {
"url": result.url,
"filename": result.filename,
"type": "image",
"gen_type": gen_type,
"resource_id": result.resource_id,
"file_size_bytes": result.file_size_bytes,
}
@router.post(
"/upload-video",
summary="上传 AI 创作普通参考视频",
description=(
"上传普通 AI 创作参考视频,写入 UploadResource 资源账本并纳入用户上传容量统计。"
"duration_seconds 为前端识别的视频秒数,用于历史复用和 AI 创作视频总时长校验。"
"未绑定业务记录前可单独删除,并会出现在 /upload-resources/history 历史素材中。"
),
responses={400: {"description": "文件类型、大小、容量或视频参数校验失败"}, 401: {"description": "未登录或 Token 无效"}},
)
async def upload_video(
file: UploadFile = File(...),
current_user: User = Depends(get_current_user),
duration_seconds: float | None = Query(None, description="前端识别的视频时长秒数,可选"),
db: AsyncSession = Depends(get_db),
):
"""上传普通参考视频,记录 UploadResource 并纳入用户容量统计。"""
result = await upload_reference_file(
db,
file=file,
current_user=current_user,
module=UploadResourceModuleEnum.COMMON.value,
resource_type=UploadResourceTypeEnum.VIDEO.value,
duration_seconds=duration_seconds,
)
await db.commit()
return {
"url": result.url,
"filename": result.filename,
"type": "video",
"resource_id": result.resource_id,
"file_size_bytes": result.file_size_bytes,
"duration_seconds": result.duration_seconds,
}
@router.post(
"/upload-audio",
summary="上传 AI 创作普通参考音频",
description=(
"上传普通 AI 创作参考音频,写入 UploadResource 资源账本并纳入用户上传容量统计。"
"当前仅支持 mp3、wav;单文件大小受 AUDIO_MAX_FILE_SIZE_MB 限制。"
"duration_seconds 为前端识别的音频秒数,用于 AI 创作音频总时长校验。"
"未绑定业务记录前可单独删除,并会出现在 /upload-resources/history 历史素材中。"
),
responses={400: {"description": "音频格式、MIME、大小或容量校验失败"}, 401: {"description": "未登录或 Token 无效"}},
)
async def upload_audio(
file: UploadFile = File(...),
current_user: User = Depends(get_current_user),
duration_seconds: float | None = Query(None, description="前端识别的音频时长秒数,可选"),
db: AsyncSession = Depends(get_db),
):
"""上传普通参考音频,记录 UploadResource 并纳入用户容量统计。"""
import os
ext = os.path.splitext(file.filename or "")[1].lower().lstrip(".")
if ext not in AUDIO_ALLOWED_EXTENSIONS:
raise HTTPException(status_code=400, detail="仅支持 mp3、wav 音频文件")
expected_mime = AUDIO_ALLOWED_MIME_TYPES.get(ext)
if expected_mime and file.content_type and file.content_type != expected_mime:
raise HTTPException(status_code=400, detail=f"音频 MIME 类型错误,{ext} 必须为 {expected_mime}")
result = await upload_reference_file(
db,
file=file,
current_user=current_user,
module=UploadResourceModuleEnum.COMMON.value,
resource_type=UploadResourceTypeEnum.AUDIO.value,
duration_seconds=duration_seconds,
max_bytes=AUDIO_MAX_FILE_SIZE_MB * 1024 * 1024,
)
await db.commit()
return {
"url": result.url,
"filename": result.filename,
"type": "audio",
"resource_id": result.resource_id,
"file_size_bytes": result.file_size_bytes,
"duration_seconds": result.duration_seconds,
}
@router.post(
"/delete-file",
summary="删除未绑定上传文件",
description=(
"删除当前用户自己的未绑定上传文件,并释放 UploadResource 上传容量。"
"仅允许删除 bind_status=pending、delete_policy=user_deletable、未绑定 source_model/source_id 的资源。"
"删除顺序为主事务先 soft delete 并 commitcommit 成功后再清理真实文件。"
"该接口保留给单文件删除;批量删除请使用 DELETE /upload-resources/history/batch。"
),
responses={400: {"description": "文件路径无效、文件已被模块任务使用或不可单独删除"}, 401: {"description": "未登录或 Token 无效"}, 403: {"description": "无权删除此文件"}},
)
async def delete_upload(
url: str = Query(..., description="文件URL,如 /uploads/images/2024/01/01/video_img_xxx.png"),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""删除未绑定业务记录的上传文件,并释放 UploadResource 容量。"""
user_id = str(current_user.id)
pending_ids: list[str] = []
legacy_paths: list[str] = []
try:
result = await delete_unbound_upload_resource(db, user=current_user, url=url)
pending_ids = list(result.pop("_pending_physical_delete_resource_ids", []) or [])
legacy_paths = list(result.pop("_legacy_pending_delete_paths", []) or [])
await db.commit()
except Exception as exc: # noqa: BLE001
await safe_rollback_with_log(
db,
event_type=UploadResourceEventEnum.DELETE_UPLOAD_ROLLBACK_FAILED.value,
message="删除上传文件主事务回滚失败",
user_id=user_id,
detail={"url": url},
original_exc=exc,
)
log_upload_resource_exception(
event_type=UploadResourceEventEnum.DELETE_UPLOAD_FAILED.value,
message=f"删除上传文件失败: {exc}",
user_id=user_id,
detail={"url": url},
exc=exc,
)
raise
if pending_ids or legacy_paths:
try:
await cleanup_upload_resource_files_after_commit(db, resource_ids=pending_ids, legacy_paths=legacy_paths)
await db.commit()
except Exception as exc: # noqa: BLE001
await safe_rollback_with_log(
db,
event_type=UploadResourceEventEnum.DELETE_UPLOAD_ROLLBACK_FAILED.value,
message="删除上传文件 cleanup 事务回滚失败",
user_id=user_id,
detail={"url": url, "pending_ids": pending_ids, "legacy_paths": legacy_paths},
original_exc=exc,
)
log_upload_resource_exception(
event_type=UploadResourceEventEnum.DELETE_UPLOAD_CLEANUP_FAILED.value,
message=f"删除上传文件后清理真实文件失败: {exc}",
user_id=user_id,
resource_ids=pending_ids,
detail={"url": url, "legacy_paths": legacy_paths},
exc=exc,
)
return result