Files
video-gen/video-gen-api/app/api/v1/generation.py
T
2026-07-23 18:37:04 +08:00

1263 lines
50 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 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.models.system_config import SystemConfig
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.credits import deduct_credits, add_credits, calc_text_credits
from app.services.llm import optimize_prompt
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
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 (
CHARGE_TEXT_PROMPT,
OWNER_GENERATION_RECORD,
build_credit_biz_key,
charge_generation_media_for_record,
get_next_credit_attempt_no,
)
from app.services.generation.ai.engine_service import (
get_image_engine,
get_video_engine,
image_supported_sizes,
parse_json_list,
)
from app.services.generation.pipeline.generation_record_config_service import (
ensure_generation_record_config_frozen,
freeze_generation_record_config_with_log,
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.credit_record_meta_service import build_generation_record_prompt_meta
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 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 _canonical_references(value: object) -> str:
return json.dumps(value or [], ensure_ascii=False, sort_keys=True, separators=(",", ":"))
def _idempotency_config_matches(record: GenerationRecord, req: OptimizeParams) -> bool:
try:
existing_references = json.loads(record.media_references) if record.media_references else []
except (TypeError, json.JSONDecodeError):
return False
if (
str(record.project_id) != str(req.project_id)
or record.original_prompt != req.prompt
or record.gen_type != req.gen_type.value
or str(record.engine_id or "") != str(req.engine_id)
or bool(record.include_media_references) != bool(req.include_media_references)
or _canonical_references(existing_references) != _canonical_references(req.references)
):
return False
if req.gen_type == GenerationType.video:
return (
record.duration == req.duration
and record.aspect_ratio == req.aspect_ratio
and record.resolution == req.resolution
)
return (
record.image_size == req.image_size
and record.image_proportion == req.image_proportion
and record.image_px == req.image_px
)
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"}:
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", "generating"},
"client_status": client_status,
"operation_phase": operation_phase,
}
def _frozen_engine_view(record: GenerationRecord) -> SimpleNamespace:
return frozen_generation_record_engine_view(record)
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),
**_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 = {
"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),
):
user_id_snapshot = str(current_user.id)
# Validate project, engine and the complete generation configuration before
# charging prompt credits or invoking the LLM.
proj_result = await db.execute(
select(Project).where(
Project.id == req.project_id,
Project.user_id == user_id_snapshot,
Project.deleted_at.is_(None),
).limit(1)
)
project = proj_result.scalar_one_or_none()
if not project:
raise HTTPException(status_code=404, detail="项目不存在")
log_generation_record_config_event(
event_type=GenerationRecordEventTypeEnum.PROMPT_CONFIG_VALIDATE_START,
event_status=LogEventStatusEnum.STARTED,
source=GenerationRecordConfigSourceEnum.PROMPT_OPTIMIZE,
record=GenerationRecord(
id=req.idempotency_key or "pending",
user_id=user_id_snapshot,
project_id=req.project_id,
original_prompt=req.prompt,
gen_type=req.gen_type.value,
duration=req.duration if req.gen_type == GenerationType.video else None,
aspect_ratio=req.aspect_ratio if req.gen_type == GenerationType.video else None,
resolution=req.resolution 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,
include_media_references=bool(req.include_media_references),
media_references=json.dumps(req.references, ensure_ascii=False) if req.references else None,
engine_id=req.engine_id,
),
detail={
"project_id": req.project_id,
"gen_type": req.gen_type.value,
"engine_id": req.engine_id,
"include_media_references": bool(req.include_media_references),
"reference_count": len(req.references or []),
"source": GenerationRecordConfigSourceEnum.PROMPT_OPTIMIZE.value,
},
)
# Idempotency must be checked before the temporary prompt-credit hold. The
# key is bound to one immutable generation configuration; reusing it with a
# different engine, parameter set or attachment selection is rejected.
if req.idempotency_key:
existing = await db.execute(
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.user_id == user_id_snapshot,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
GenerationRecord.idempotency_key == req.idempotency_key,
)
.order_by(GenerationRecord.created_at.desc())
.limit(1)
)
row = existing.first()
if row:
existing_record, project_name = row
if not _idempotency_config_matches(existing_record, req):
raise HTTPException(
status_code=409,
detail="幂等键已绑定其他生成配置,请重新提交",
)
if not existing_record.optimized_prompt or not _record_config_complete(existing_record):
raise HTTPException(
status_code=409,
detail="幂等记录配置不完整,请使用新的幂等键重新提交",
)
refs = await resolve_private_portrait_reference_display_urls(
db,
json.loads(existing_record.media_references) if existing_record.media_references else None,
user_id=user_id_snapshot,
)
return OptimizeResult(
optimized_prompt=existing_record.optimized_prompt,
text_credits_cost=existing_record.text_credits_cost or 0.0,
text_tokens_used=existing_record.text_tokens_used or 0,
record=_record_to_out(existing_record, project_name, refs_override=refs),
)
if req.gen_type == GenerationType.video:
if req.duration not in DURATIONS:
raise HTTPException(status_code=400, detail=f"视频时长必须为{DURATIONS}秒之一")
if req.aspect_ratio not in ASPECT_RATIOS:
raise HTTPException(status_code=400, detail="不支持的画面比例")
if req.resolution not in RESOLUTIONS:
raise HTTPException(status_code=400, detail="不支持的分辨率")
engine = await get_video_engine(db, req.engine_id)
_validate_video_engine_selection(
engine,
aspect_ratio=req.aspect_ratio,
resolution=req.resolution,
duration=int(req.duration),
)
else:
if req.image_size not in IMAGE_SIZES:
raise HTTPException(status_code=400, detail=f"图片分辨率必须为{IMAGE_SIZES}之一")
if not req.image_proportion or not req.image_px:
raise HTTPException(status_code=400, detail="图片生成需要指定比例和像素尺寸")
engine = await get_image_engine(db, req.engine_id)
_validate_image_engine_selection(engine, image_size=req.image_size)
project_name_snapshot = str(project.name)
project_industry_snapshot = str(project.industry or "")
engine_snapshot_source = SimpleNamespace(
**{
key: value
for key, value in vars(engine).items()
if key != "_sa_instance_state"
}
)
reference_usage = calculate_media_reference_usage(
json.dumps(req.references, ensure_ascii=False) if req.references else None,
include=bool(req.include_media_references),
)
validate_media_reference_usage_for_engine(
reference_usage,
gen_type=req.gen_type.value,
engine=engine,
)
log_generation_record_config_event(
event_type=GenerationRecordEventTypeEnum.PROMPT_CONFIG_VALIDATE_SUCCESS,
event_status=LogEventStatusEnum.SUCCESS,
source=GenerationRecordConfigSourceEnum.PROMPT_OPTIMIZE,
record=GenerationRecord(
id=req.idempotency_key or "pending",
user_id=user_id_snapshot,
project_id=req.project_id,
original_prompt=req.prompt,
gen_type=req.gen_type.value,
duration=req.duration if req.gen_type == GenerationType.video else None,
aspect_ratio=req.aspect_ratio if req.gen_type == GenerationType.video else None,
resolution=req.resolution 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,
include_media_references=bool(req.include_media_references),
media_references=json.dumps(req.references, ensure_ascii=False) if req.references else None,
engine_id=req.engine_id,
),
detail={
"project_id": req.project_id,
"gen_type": req.gen_type.value,
"engine_id": req.engine_id,
"include_media_references": bool(req.include_media_references),
"reference_count": len(req.references or []),
"media_reference_usage": getattr(reference_usage, "__dict__", None) or str(reference_usage),
"source": GenerationRecordConfigSourceEnum.PROMPT_OPTIMIZE.value,
},
)
hold_credits = 5
hold_result = await db.execute(
select(SystemConfig).where(SystemConfig.key == "optimize_hold_credits").limit(1)
)
hold_row = hold_result.scalar_one_or_none()
if hold_row and hold_row.value:
try:
hold_credits = max(0, int(hold_row.value))
except (ValueError, TypeError):
hold_credits = 5
hold_scope = req.idempotency_key or generate_id()
hold_biz_key = f"optimize_hold:{hold_scope}"
hold_refund_biz_key = f"optimize_hold_refund:{hold_scope}"
await deduct_credits(
db,
user_id_snapshot,
hold_credits,
f"AI创作预扣积分 - {project_name_snapshot}",
biz_key=hold_biz_key,
)
await db.commit()
try:
optimized, token_usage = await optimize_prompt(
db,
req.prompt,
user_id=user_id_snapshot,
industry_key=project_industry_snapshot,
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.value,
log_module="generation_record",
log_step="prompt_optimize",
log_project_id=req.project_id,
)
except Exception as exc:
from app.services.error_codes import extract_error_message
await db.rollback()
await add_credits(
db,
user_id_snapshot,
hold_credits,
f"AI创作预扣积分退还 - {project_name_snapshot}",
record_type="refund",
biz_key=hold_refund_biz_key,
refund_for_biz_key=hold_biz_key,
)
await db.commit()
raise HTTPException(
status_code=502,
detail=f"AI模型调用失败: {extract_error_message(exc, '提示词')}",
) from exc
try:
text_credits = await calc_text_credits(
db,
int(token_usage.get("input_tokens", 0) or 0),
int(token_usage.get("output_tokens", 0) or 0),
)
record = GenerationRecord(
id=generate_id(),
user_id=user_id_snapshot,
project_id=req.project_id,
original_prompt=req.prompt,
optimized_prompt=optimized,
gen_type=req.gen_type.value,
duration=req.duration if req.gen_type == GenerationType.video else None,
aspect_ratio=req.aspect_ratio if req.gen_type == GenerationType.video else None,
resolution=req.resolution 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,
status="prompt_optimized",
pipeline_stage=None,
credits_cost=0,
text_credits_cost=round(text_credits, 2),
text_tokens_used=int(token_usage.get("total_tokens", 0) or 0),
media_references=json.dumps(req.references, ensure_ascii=False) if req.references else None,
include_media_references=bool(req.include_media_references),
idempotency_key=req.idempotency_key,
)
if req.gen_type == GenerationType.video:
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=req.resolution,
aspect_ratio=req.aspect_ratio,
supported_provider_resolutions=parse_json_list(
engine_snapshot_source.supported_resolutions,
[],
),
)
record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json
else:
record.provider_generation_resolution = None
record.video_upscale_enabled_snapshot = False
record.video_upscale_snapshot_json = None
freeze_generation_record_config_with_log(
record,
engine=engine_snapshot_source,
source=GenerationRecordConfigSourceEnum.PROMPT_OPTIMIZE,
)
db.add(record)
await db.flush()
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,
)
# Release the hold and charge the exact prompt usage in one transaction.
await add_credits(
db,
user_id_snapshot,
hold_credits,
f"AI创作预扣积分退还 - {project_name_snapshot}",
related_id=record.id,
record_type="refund",
biz_key=hold_refund_biz_key,
refund_for_biz_key=hold_biz_key,
)
await deduct_credits(
db,
user_id_snapshot,
text_credits,
f"提示词优化 - {project_name_snapshot}",
related_id=record.id,
biz_key=prompt_biz_key,
record_meta=prompt_meta,
)
record_id_snapshot = str(record.id)
await db.commit()
except Exception:
await db.rollback()
# Any local pricing/snapshot/persistence failure after the provider call
# must release the committed hold. The refund key is idempotent.
await add_credits(
db,
user_id_snapshot,
hold_credits,
f"AI创作预扣积分退还 - {project_name_snapshot}",
record_type="refund",
biz_key=hold_refund_biz_key,
refund_for_biz_key=hold_biz_key,
)
await db.commit()
raise
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 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)
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)
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