Files
video-gen/video-gen-api/app/api/v1/generation.py
T
2026-07-21 19:28:02 +08:00

969 lines
38 KiB
Python
Raw 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 datetime, timezone, timedelta
CST = timezone(timedelta(hours=8))
from fastapi import APIRouter, Depends, HTTPException, Query, Request, UploadFile, File, status
from fastapi.responses import RedirectResponse
from sqlalchemy import select, func
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
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,
GenerateParams,
GenerationRecordOut,
GenerationRecordPageListOut,
OptimizeResult,
UpdatePromptRequest,
GenerationType,
DURATIONS,
ASPECT_RATIOS,
RESOLUTIONS,
IMAGE_SIZES,
)
from app.services.generation.pipeline.db_lock_service import (
DatabaseRowLockBusy,
execute_with_lock_timeout,
)
from app.services.credits import deduct_credits, calc_text_credits
from app.services.llm import optimize_prompt
from app.services.video_url import generate_temp_url, validate_and_get_record_id, get_video_stream_url
from app.services.resource_accounting_service import (
record_generation_record_generated_resource,
safe_file_size,
)
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls, resolve_private_portrait_reference_display_urls
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 GenerationRecordPipelineStage
from app.services.generation.billing_service import (
CHARGE_TEXT_PROMPT,
OWNER_GENERATION_RECORD,
build_credit_biz_key,
charge_generation_media_by_params,
charge_generation_media_for_record,
get_next_credit_attempt_no,
)
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
from app.services.generation.ai.engine_service import (
get_image_engine,
get_video_engine,
image_supported_sizes,
parse_json_list,
)
from app.services.generation.media_reference_service import (
calculate_media_reference_usage,
validate_media_reference_usage_for_engine,
)
from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot
from app.services.credit_record_meta_service import build_generation_record_prompt_meta
from app.services.video_cover_service import async_create_video_cover_for_local_video
from app.enums.audio_reference import (
AUDIO_ALLOWED_EXTENSIONS,
AUDIO_ALLOWED_MIME_TYPES,
AUDIO_MAX_FILE_SIZE_MB,
)
from app.utils.id_gen import generate_id
from app.utils.exceptions import InsufficientCreditsError, RecordNotFoundError, InvalidStatusError
router = APIRouter(prefix="/generation-records", tags=["generation"])
logger = logging.getLogger("videogen")
def _engine_snapshot(record: GenerationRecord) -> dict | None:
if not record.engine_snapshot_json:
return None
try:
value = json.loads(record.engine_snapshot_json)
except (TypeError, json.JSONDecodeError):
return None
return value if isinstance(value, dict) else None
def _validate_video_engine_selection(engine, *, aspect_ratio: str, resolution: str, duration: int) -> None:
ratios = [str(item) for item in parse_json_list(engine.supported_ratios, [])]
resolutions = [str(item) for item in parse_json_list(engine.supported_resolutions, [])]
durations = [int(item) for item in parse_json_list(engine.supported_durations, []) if str(item).isdigit()]
if ratios and aspect_ratio not in ratios:
raise HTTPException(status_code=400, detail="当前视频引擎不支持所选画面比例")
if resolutions and resolution not in resolutions:
raise HTTPException(status_code=400, detail="当前视频引擎不支持所选分辨率")
if durations and duration not in durations:
raise HTTPException(status_code=400, detail="当前视频引擎不支持所选时长")
if int(engine.max_duration or 0) > 0 and duration > int(engine.max_duration):
raise HTTPException(status_code=400, detail="生成时长超过当前视频引擎上限")
def _validate_image_engine_selection(engine, *, image_size: str) -> None:
sizes = image_supported_sizes(engine)
if sizes and image_size not in sizes:
raise HTTPException(status_code=400, detail="当前图片引擎不支持所选画面分辨率")
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),
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 = {
"prompt_optimized",
"generating",
"failed",
"completed",
}
if status and status not in allowed_statuses:
raise HTTPException(
status_code=400,
detail="状态参数错误,仅支持: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),
):
record = None
# 积分不足直接返回
if (current_user.credits or 0) < 5:
raise HTTPException(status_code=402, detail="积分不足,请充值")
# Validate parameters based on generation type
if req.gen_type == GenerationType.video:
if req.duration not in DURATIONS:
raise HTTPException(status_code=400, detail=f"视频时长必须为{DURATIONS}秒之一")
if not req.duration:
raise HTTPException(status_code=400, detail="视频生成需要指定时长")
elif req.gen_type == GenerationType.image:
if req.image_size not in IMAGE_SIZES:
raise HTTPException(status_code=400, detail=f"图片分辨率必须为{IMAGE_SIZES}之一")
if not req.image_size:
raise HTTPException(status_code=400, detail="图片生成需要指定画面分辨率")
# Idempotency check: if key provided, return existing record if found
if req.idempotency_key:
existing = await db.execute(
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.user_id == current_user.id,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
GenerationRecord.idempotency_key == req.idempotency_key,
GenerationRecord.gen_type == req.gen_type,
GenerationRecord.status == "prompt_optimized",
)
.order_by(GenerationRecord.created_at.desc())
.limit(1)
)
row = existing.first()
if row:
record, project_name = row
refs = await resolve_private_portrait_reference_display_urls(db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id)
return OptimizeResult(
optimized_prompt=record.optimized_prompt or "",
text_credits_cost=record.text_credits_cost or 0.00,
text_tokens_used=record.text_tokens_used or 0,
record=_record_to_out(record, project_name, refs_override=refs),
)
# Check project exists and belongs to user
proj_result = await db.execute(
select(Project).where(
Project.id == req.project_id,
Project.user_id == current_user.id,
Project.deleted_at.is_(None),
)
.limit(1)
)
project = proj_result.scalar_one_or_none()
if not project:
raise HTTPException(status_code=404, detail="项目不存在")
# Optimize prompt via LLM with type-specific context
try:
optimized, token_usage = await optimize_prompt(
db, req.prompt,
user_id=current_user.id,
industry_key=project.industry,
duration=req.duration if req.gen_type == GenerationType.video else None,
image_size=req.image_size if req.gen_type == GenerationType.image else None,
image_proportion=req.image_proportion if req.gen_type == GenerationType.image else None,
image_px=req.image_px if req.gen_type == GenerationType.image else None,
references=req.references,
gen_type=req.gen_type,
)
# LLM 成功后再创建记录;LLM 失败不写 GenerationRecord。
record = GenerationRecord(
id=generate_id(),
user_id=current_user.id,
project_id=req.project_id,
original_prompt=req.prompt,
gen_type=req.gen_type,
duration=req.duration,
image_size=req.image_size,
image_proportion=req.image_proportion,
image_px=req.image_px,
status="optimizing",
credits_cost=0,
text_credits_cost=0,
text_tokens_used=0,
media_references=json.dumps(req.references) if req.references else None,
idempotency_key=req.idempotency_key,
)
db.add(record)
await db.flush()
await db.commit()
except Exception as e:
from app.services.error_codes import extract_error_message
if record:
record.status = "failed"
record.error_message = extract_error_message(e, "提示词")
await db.flush()
await db.commit()
error_message = extract_error_message(e, "提示词")
raise HTTPException(
status_code=502,
detail=f"AI模型调用失败: {error_message}"
)
text_credits = await calc_text_credits(
db, token_usage["input_tokens"], token_usage["output_tokens"],
)
failed_record_id = record.id
failed_user_id = current_user.id
try:
prompt_attempt_no = 1
prompt_biz_key = build_credit_biz_key(
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
attempt_no=prompt_attempt_no,
charge_kind=CHARGE_TEXT_PROMPT,
action="charge",
)
prompt_meta = await build_generation_record_prompt_meta(
db,
record_id=record.id,
attempt_no=prompt_attempt_no,
charge_kind=CHARGE_TEXT_PROMPT,
usage=token_usage,
)
await deduct_credits(
db, current_user.id, text_credits,
f"提示词优化 - {project.name}",
related_id=record.id,
biz_key=prompt_biz_key,
record_meta=prompt_meta,
)
except InsufficientCreditsError as e:
# /optimize 阶段只处理提示词优化扣费。
# 提示词积分不足时,之前已落库的 optimizing 记录必须改为 failed,避免前端长期显示生成中。
# 此阶段没有媒体生成扣费,不调用生成失败退款逻辑。
await db.rollback()
try:
result = await execute_with_lock_timeout(
db,
select(GenerationRecord)
.where(
GenerationRecord.id == failed_record_id,
GenerationRecord.user_id == failed_user_id,
GenerationRecord.deleted_at.is_(None),
)
.with_for_update()
.limit(1),
)
except DatabaseRowLockBusy:
# Preserve the original 402 response; a later admin/manual check can
# reconcile the rare record-state update lock conflict.
raise e
failed_record = result.scalar_one_or_none()
if failed_record:
failed_record.status = "failed"
failed_record.error_message = e.detail
failed_record.optimized_prompt = None
failed_record.text_credits_cost = 0
failed_record.credits_cost = 0
failed_record.text_tokens_used = token_usage.get("total_tokens", 0)
await db.flush()
# 这里必须主动提交,否则后续抛出 402 后 get_db 会 rollbackfailed 状态会被回滚。
await db.commit()
raise e
record.optimized_prompt = optimized
record.status = "prompt_optimized"
record.text_credits_cost = round(text_credits, 2)
record.text_tokens_used = token_usage["total_tokens"]
await db.flush()
refs = await resolve_private_portrait_reference_display_urls(db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id)
return OptimizeResult(
optimized_prompt=optimized,
text_credits_cost=round(text_credits, 2),
# text_tokens_used=token_usage["total_tokens"],
record=_record_to_out(record, project.name, refs_override=refs),
)
@router.post("/{record_id}/generate")
async def generate_record_resource(
record_id: str,
req: GenerateParams,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
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 == current_user.id,
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 not in ("prompt_optimized", "failed"):
raise InvalidStatusError("当前状态不允许生成")
if record.pipeline_stage == GenerationRecordPipelineStage.UPSCALE_FAILED.value:
raise InvalidStatusError("该任务生成失败,请联系客服进行修复")
await assert_user_resource_capacity_available(db, current_user.id)
attempt_no = await get_next_credit_attempt_no(
db, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id
)
selected_engine_id = req.engine_id or record.engine_id
record.include_media_references = bool(req.include_media_references)
from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record,
prepare_generation_record_execution,
)
if record.gen_type == GenerationType.video:
aspect_ratio = req.aspect_ratio or record.aspect_ratio
resolution = req.resolution or record.resolution
if aspect_ratio not in ASPECT_RATIOS:
raise HTTPException(status_code=400, detail="不支持的画面比例")
if resolution not in RESOLUTIONS:
raise HTTPException(status_code=400, detail="不支持的分辨率")
engine = await get_video_engine(db, selected_engine_id)
_validate_video_engine_selection(
engine,
aspect_ratio=aspect_ratio,
resolution=resolution,
duration=int(record.duration or 5),
)
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
supported_provider_resolutions = parse_json_list(engine.supported_resolutions, [])
provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot(
db,
target_resolution=resolution,
aspect_ratio=aspect_ratio,
supported_provider_resolutions=supported_provider_resolutions,
)
record.aspect_ratio = aspect_ratio
record.resolution = resolution
record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json
else:
engine = await get_image_engine(db, selected_engine_id)
image_size = req.image_size or record.image_size or engine.default_size or "2K"
_validate_image_engine_selection(engine, image_size=image_size)
record.image_size = image_size
record.provider_generation_resolution = None
record.video_upscale_enabled_snapshot = False
record.video_upscale_snapshot_json = None
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=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=engine.id,
)
record.credits_cost = round(
float(record.credits_cost or 0) + float(billing.total_charged or 0), 2
)
prepare_generation_record_execution(record, engine=engine, 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_generate"
)
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=current_user.id,
)
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),
):
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 == current_user.id,
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.pipeline_stage == GenerationRecordPipelineStage.UPSCALE_FAILED.value:
raise InvalidStatusError("该任务生成失败,请联系客服进行修复")
await assert_user_resource_capacity_available(db, current_user.id)
attempt_no = await get_next_credit_attempt_no(
db, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id
)
from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record,
prepare_generation_record_execution,
)
if record.gen_type == GenerationType.video:
engine = await get_video_engine(db, record.engine_id)
_validate_video_engine_selection(
engine,
aspect_ratio=record.aspect_ratio or "16:9",
resolution=record.resolution or "480p",
duration=int(record.duration or 5),
)
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot(
db,
target_resolution=record.resolution or "480p",
aspect_ratio=record.aspect_ratio or "16:9",
supported_provider_resolutions=parse_json_list(engine.supported_resolutions, []),
)
record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json
else:
engine = await get_image_engine(db, record.engine_id)
_validate_image_engine_selection(
engine, image_size=record.image_size or engine.default_size or "2K"
)
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=engine,
)
billing = await charge_generation_media_for_record(
db,
record=record,
project_name=project_name,
description_prefix="资源生成重试-",
attempt_no=attempt_no,
engine_id=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)
prepare_generation_record_execution(record, engine=engine, 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=current_user.id,
)
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