924 lines
35 KiB
Python
924 lines
35 KiB
Python
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 并 commit,commit 成功后再清理真实文件。"
|
||
"该接口保留给单文件删除;批量删除请使用 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
|