2、图片返回的 URL 是本地路径 3、视频生成中媒体文件重复下载 4、幂等性检查无数据库唯一约束 5、幂等键冲突返回 409 改为返回已有任务信息 6、虚拟素材库配额校验 TOCTOU 7、项目级联删除与独立素材删除任务并发冲突
234 lines
8.0 KiB
Python
234 lines
8.0 KiB
Python
import json
|
||
import logging
|
||
from datetime import datetime, timedelta, timezone
|
||
|
||
from sqlalchemy import select
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
|
||
from app.models.api.api_generation_task import ApiGenerationTask
|
||
from app.schemas.api_v3.video import ApiVideoStatusResponse
|
||
from app.utils.id_gen import generate_id
|
||
|
||
logger = logging.getLogger("videogen")
|
||
|
||
|
||
async def create_video_task(
|
||
db: AsyncSession,
|
||
api_key_id: str,
|
||
model_name: str,
|
||
engine_id: str,
|
||
engine_snapshot: dict,
|
||
content: list[dict],
|
||
ratio: str | None,
|
||
duration: int | None,
|
||
resolution: str | None,
|
||
provider_generation_resolution: str | None,
|
||
upscale_enabled: bool,
|
||
upscale_snapshot_json: str | None,
|
||
idempotency_key: str | None = None,
|
||
local_media_refs: list[dict] | None = None,
|
||
) -> ApiGenerationTask:
|
||
"""创建视频生成任务记录。
|
||
|
||
如果调用方已下载好媒体文件(local_media_refs),则直接复用,避免重复下载。
|
||
"""
|
||
# 提取文本提示词
|
||
text_parts = [p.get("text", "") for p in content if p.get("type") == "text"]
|
||
original_prompt = " ".join(text_parts) if text_parts else content[0].get("text", "") if content else ""
|
||
|
||
# 构建 media_references(扁平格式,便于外部读取)
|
||
# 构建 local_media_json(嵌套格式,与 Volcano SDK 兼容)
|
||
media_refs = [] # 扁平格式: {"type": "image", "url": "...", "role": "..."}
|
||
|
||
for p in content:
|
||
ptype = p.get("type", "")
|
||
if ptype == "text":
|
||
continue
|
||
|
||
# 提取原始 URL(从嵌套格式中提取)
|
||
original_url = ""
|
||
media_type = ptype.replace("_url", "") # image_url -> image
|
||
if ptype == "image_url" and p.get("image_url"):
|
||
original_url = p["image_url"].get("url", "")
|
||
elif ptype == "video_url" and p.get("video_url"):
|
||
original_url = p["video_url"].get("url", "")
|
||
elif ptype == "audio_url" and p.get("audio_url"):
|
||
original_url = p["audio_url"].get("url", "")
|
||
|
||
# 存储扁平格式到 media_references
|
||
media_refs.append({
|
||
"type": media_type,
|
||
"url": original_url,
|
||
"role": p.get("role"),
|
||
})
|
||
|
||
# 如果调用方已传入 local_media_refs(已下载),直接使用,不再重复下载
|
||
if local_media_refs is None:
|
||
from app.services.api_v3.file_service import process_media_url
|
||
|
||
local_media_refs = [] # 本地下载路径(嵌套格式)
|
||
for p in content:
|
||
ptype = p.get("type", "")
|
||
if ptype == "text":
|
||
continue
|
||
|
||
original_url = ""
|
||
if ptype == "image_url" and p.get("image_url"):
|
||
original_url = p["image_url"].get("url", "")
|
||
elif ptype == "video_url" and p.get("video_url"):
|
||
original_url = p["video_url"].get("url", "")
|
||
elif ptype == "audio_url" and p.get("audio_url"):
|
||
original_url = p["audio_url"].get("url", "")
|
||
|
||
# 下载文件到本地
|
||
try:
|
||
local_path = await process_media_url(original_url, ptype.replace("_url", ""))
|
||
except Exception as exc:
|
||
logger.warning("Failed to download media %s: %s", original_url[:80], exc)
|
||
local_path = original_url
|
||
|
||
# 本地路径使用嵌套格式(与 Volcano SDK 兼容)
|
||
local_media_refs.append({
|
||
"type": ptype,
|
||
ptype: {"url": local_path},
|
||
"role": p.get("role"),
|
||
})
|
||
|
||
media_references_json = json.dumps(media_refs, ensure_ascii=False) if media_refs else None
|
||
local_media_json = json.dumps(local_media_refs, ensure_ascii=False) if local_media_refs else None
|
||
|
||
now = datetime.now(timezone.utc)
|
||
deadline = now + timedelta(hours=24)
|
||
|
||
task = ApiGenerationTask(
|
||
id=generate_id(),
|
||
api_key_id=api_key_id,
|
||
external_idempotency_key=idempotency_key,
|
||
original_prompt=original_prompt,
|
||
gen_type="video",
|
||
model_name=model_name,
|
||
duration=duration,
|
||
aspect_ratio=ratio,
|
||
resolution=resolution,
|
||
provider_generation_resolution=provider_generation_resolution,
|
||
generation_count=1,
|
||
engine_id=engine_id,
|
||
media_references=media_references_json,
|
||
local_media_json=local_media_json,
|
||
engine_snapshot_json=json.dumps(engine_snapshot, ensure_ascii=False),
|
||
status="pending",
|
||
pipeline_stage="queued",
|
||
deadline_at=deadline,
|
||
video_upscale_enabled_snapshot=upscale_enabled,
|
||
video_upscale_snapshot_json=upscale_snapshot_json,
|
||
)
|
||
db.add(task)
|
||
await db.flush()
|
||
return task
|
||
|
||
|
||
async def create_image_task(
|
||
db: AsyncSession,
|
||
api_key_id: str,
|
||
model_name: str,
|
||
engine_id: str,
|
||
engine_snapshot: dict,
|
||
prompt: str,
|
||
size: str | None,
|
||
idempotency_key: str | None = None,
|
||
) -> ApiGenerationTask:
|
||
"""创建图片生成任务记录。"""
|
||
task = ApiGenerationTask(
|
||
id=generate_id(),
|
||
api_key_id=api_key_id,
|
||
external_idempotency_key=idempotency_key,
|
||
original_prompt=prompt,
|
||
gen_type="image",
|
||
image_size=size,
|
||
generation_count=1,
|
||
engine_id=engine_id,
|
||
engine_snapshot_json=json.dumps(engine_snapshot, ensure_ascii=False),
|
||
status="processing",
|
||
pipeline_stage="creating_provider_task",
|
||
)
|
||
db.add(task)
|
||
await db.flush()
|
||
return task
|
||
|
||
|
||
async def get_task(db: AsyncSession, task_id: str, api_key_id: str) -> ApiGenerationTask | None:
|
||
"""获取任务(带所有权验证)。"""
|
||
result = await db.execute(
|
||
select(ApiGenerationTask).where(
|
||
ApiGenerationTask.id == task_id,
|
||
ApiGenerationTask.api_key_id == api_key_id,
|
||
ApiGenerationTask.deleted_at.is_(None),
|
||
).limit(1)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
|
||
async def find_by_idempotency_key(db: AsyncSession, api_key_id: str, idempotency_key: str) -> ApiGenerationTask | None:
|
||
"""根据幂等键查找已存在的任务。"""
|
||
result = await db.execute(
|
||
select(ApiGenerationTask).where(
|
||
ApiGenerationTask.api_key_id == api_key_id,
|
||
ApiGenerationTask.external_idempotency_key == idempotency_key,
|
||
ApiGenerationTask.deleted_at.is_(None),
|
||
).limit(1)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
|
||
def map_task_to_status_response(task: ApiGenerationTask) -> ApiVideoStatusResponse:
|
||
"""将任务对象映射为状态查询响应。"""
|
||
from app.config import settings
|
||
# 返回完整 URL(包含 BASE_URL)
|
||
video_url = _make_full_url(task.video_url)
|
||
video_cover_url = _make_full_url(task.video_cover_url)
|
||
return ApiVideoStatusResponse(
|
||
task_id=task.id,
|
||
status=_map_status(task.status),
|
||
video_url=video_url,
|
||
video_cover_url=video_cover_url,
|
||
duration=task.duration,
|
||
ratio=task.aspect_ratio,
|
||
resolution=task.resolution,
|
||
error=task.error_message,
|
||
created_at=task.created_at,
|
||
completed_at=task.generated_at,
|
||
)
|
||
|
||
|
||
def _make_full_url(path: str | None) -> str | None:
|
||
"""将本地路径转换为完整 URL。"""
|
||
if not path:
|
||
return None
|
||
from app.config import settings
|
||
# 如果已经是完整 URL,直接返回
|
||
if path.startswith(("http://", "https://")):
|
||
return path
|
||
# 处理 ./storage/generate/... 格式 → /generate/...
|
||
if path.startswith("./storage"):
|
||
url_path = path[len("./storage"):]
|
||
elif path.startswith("/"):
|
||
url_path = path
|
||
else:
|
||
url_path = f"/{path}"
|
||
# 拼接 BASE_URL
|
||
base = settings.BASE_URL.rstrip("/")
|
||
return f"{base}{url_path}"
|
||
|
||
|
||
def _map_status(status: str) -> str:
|
||
"""将内部状态映射为 API 状态。"""
|
||
status_map = {
|
||
"pending": "queued",
|
||
"queued": "pending_queue",
|
||
"generating": "generating",
|
||
"processing": "generating",
|
||
"completed": "completed",
|
||
"failed": "failed",
|
||
}
|
||
return status_map.get(status, status)
|