1、图片生成同步接口超时风险
2、图片返回的 URL 是本地路径 3、视频生成中媒体文件重复下载 4、幂等性检查无数据库唯一约束 5、幂等键冲突返回 409 改为返回已有任务信息 6、虚拟素材库配额校验 TOCTOU 7、项目级联删除与独立素材删除任务并发冲突
This commit is contained in:
@@ -83,15 +83,19 @@ async def create_video(
|
|||||||
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
|
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
|
||||||
db: AsyncSession = Depends(get_db),
|
db: AsyncSession = Depends(get_db),
|
||||||
) -> ApiVideoCreateResponse:
|
) -> ApiVideoCreateResponse:
|
||||||
"""创建视频生成任务(异步)。"""
|
"""创建视频生成任务(异步)。
|
||||||
|
|
||||||
|
幂等性说明:如果 idempotency_key 已存在,直接返回已有任务 ID(不会重复创建)。
|
||||||
|
"""
|
||||||
try:
|
try:
|
||||||
# 路由层校验:权限、幂等性
|
# 路由层校验:权限、幂等性
|
||||||
existing_task = await _validate_request(db, key_context, req)
|
existing_task = await _validate_request(db, key_context, req)
|
||||||
if existing_task:
|
if existing_task:
|
||||||
raise HTTPException(
|
logger.info(
|
||||||
status_code=status.HTTP_409_CONFLICT,
|
"Idempotent request: returning existing task %s for key=%s",
|
||||||
detail=f"幂等键已存在: 任务 {req.idempotency_key} 已创建",
|
existing_task.id, req.idempotency_key,
|
||||||
)
|
)
|
||||||
|
return ApiVideoCreateResponse(id=f"zc-{existing_task.id}")
|
||||||
|
|
||||||
# 调用服务层创建任务
|
# 调用服务层创建任务
|
||||||
result = await generation_service.submit_video_generation(
|
result = await generation_service.submit_video_generation(
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
import time
|
import time
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
@@ -16,6 +17,26 @@ from app.services.api_v3.logging_service import log_model_request
|
|||||||
logger = logging.getLogger("videogen")
|
logger = logging.getLogger("videogen")
|
||||||
|
|
||||||
|
|
||||||
|
def _make_image_url(local_path: str) -> str:
|
||||||
|
"""将本地图片路径转为完整可访问 URL。"""
|
||||||
|
if not local_path:
|
||||||
|
return local_path
|
||||||
|
# 如果已经是完整 URL,直接返回
|
||||||
|
if local_path.startswith(("http://", "https://")):
|
||||||
|
return local_path
|
||||||
|
from app.config import settings
|
||||||
|
from app.services.resource_signed_url_service import build_resource_signed_url
|
||||||
|
# 生成签名 URL
|
||||||
|
signed = build_resource_signed_url(local_path)
|
||||||
|
if signed and not signed.startswith(("http://", "https://")):
|
||||||
|
base = settings.BASE_URL.rstrip("/")
|
||||||
|
if signed.startswith("/"):
|
||||||
|
signed = f"{base}{signed}"
|
||||||
|
else:
|
||||||
|
signed = f"{base}/{signed}"
|
||||||
|
return signed or local_path
|
||||||
|
|
||||||
|
|
||||||
async def submit_video_generation(
|
async def submit_video_generation(
|
||||||
db: AsyncSession,
|
db: AsyncSession,
|
||||||
key: ApiKey,
|
key: ApiKey,
|
||||||
@@ -379,10 +400,12 @@ async def generate_image_sync(
|
|||||||
logger.error("Failed to record usage on image submit: %s", log_exc)
|
logger.error("Failed to record usage on image submit: %s", log_exc)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
# 设置总体超时(120秒,防止同步请求长时间挂起)
|
||||||
|
_IMAGE_GEN_TIMEOUT = 120
|
||||||
|
|
||||||
# 3. 调用 Volcano Ark SDK(同步函数,在线程中执行)
|
# 3. 调用 Volcano Ark SDK(同步函数,在线程中执行)
|
||||||
from app.services.image_gen import submit_image_task, download_image
|
from app.services.image_gen import submit_image_task, download_image
|
||||||
from app.config import settings
|
from app.config import settings
|
||||||
import os
|
|
||||||
|
|
||||||
# 构建 media_references,下载图片到本地
|
# 构建 media_references,下载图片到本地
|
||||||
from app.services.api_v3.file_service import process_media_url
|
from app.services.api_v3.file_service import process_media_url
|
||||||
@@ -402,14 +425,17 @@ async def generate_image_sync(
|
|||||||
task.image_size = req.size or "2K"
|
task.image_size = req.size or "2K"
|
||||||
await db.flush()
|
await db.flush()
|
||||||
|
|
||||||
# 在线程中执行同步 SDK 调用
|
# 在线程中执行同步 SDK 调用(带超时保护)
|
||||||
result = await asyncio.to_thread(
|
result = await asyncio.wait_for(
|
||||||
|
asyncio.to_thread(
|
||||||
submit_image_task,
|
submit_image_task,
|
||||||
db,
|
db,
|
||||||
engine,
|
engine,
|
||||||
task,
|
task,
|
||||||
True, # include_media_references
|
True, # include_media_references
|
||||||
req.generation_count or 1,
|
req.generation_count or 1,
|
||||||
|
),
|
||||||
|
timeout=_IMAGE_GEN_TIMEOUT,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 4. 下载图片
|
# 4. 下载图片
|
||||||
@@ -425,12 +451,19 @@ async def generate_image_sync(
|
|||||||
os.makedirs(dest_dir, exist_ok=True)
|
os.makedirs(dest_dir, exist_ok=True)
|
||||||
dest_path = os.path.join(dest_dir, f"{task.id}_{item.get('generation_index', 1)}.png")
|
dest_path = os.path.join(dest_dir, f"{task.id}_{item.get('generation_index', 1)}.png")
|
||||||
try:
|
try:
|
||||||
await download_image(url, dest_path)
|
await asyncio.wait_for(
|
||||||
|
download_image(url, dest_path),
|
||||||
|
timeout=30,
|
||||||
|
)
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
logger.warning("Image download timeout: %s", url[:80])
|
||||||
except Exception as dl_err:
|
except Exception as dl_err:
|
||||||
logger.warning("Image download failed: %s", dl_err)
|
logger.warning("Image download failed: %s", dl_err)
|
||||||
|
|
||||||
|
# 将本地路径转为完整 URL
|
||||||
|
image_url = _make_image_url(dest_path)
|
||||||
downloaded_items.append(ApiImageGenerateDataItem(
|
downloaded_items.append(ApiImageGenerateDataItem(
|
||||||
url=dest_path, # 使用本地路径
|
url=image_url,
|
||||||
size=item.get("size"),
|
size=item.get("size"),
|
||||||
output_format=item.get("output_format"),
|
output_format=item.get("output_format"),
|
||||||
))
|
))
|
||||||
@@ -459,6 +492,21 @@ async def generate_image_sync(
|
|||||||
model=result.get("model", req.model),
|
model=result.get("model", req.model),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
from fastapi import HTTPException, status
|
||||||
|
logger.exception("API image generation timed out (task_id=%s)", task.id)
|
||||||
|
# 超时:退回预扣配额
|
||||||
|
if estimated_price > 0:
|
||||||
|
key.quota_used = round(max(0, (key.quota_used or 0.0) - estimated_price), 2)
|
||||||
|
task.status = "failed"
|
||||||
|
task.error_message = "图片生成超时(超过120秒)"
|
||||||
|
task.credits_cost = 0
|
||||||
|
await db.commit()
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=504,
|
||||||
|
detail="图片生成超时,请稍后重试",
|
||||||
|
)
|
||||||
|
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
# 失败:退回预扣配额
|
# 失败:退回预扣配额
|
||||||
if estimated_price > 0:
|
if estimated_price > 0:
|
||||||
|
|||||||
@@ -28,17 +28,17 @@ async def create_video_task(
|
|||||||
idempotency_key: str | None = None,
|
idempotency_key: str | None = None,
|
||||||
local_media_refs: list[dict] | None = None,
|
local_media_refs: list[dict] | None = None,
|
||||||
) -> ApiGenerationTask:
|
) -> ApiGenerationTask:
|
||||||
"""创建视频生成任务记录。"""
|
"""创建视频生成任务记录。
|
||||||
|
|
||||||
|
如果调用方已下载好媒体文件(local_media_refs),则直接复用,避免重复下载。
|
||||||
|
"""
|
||||||
# 提取文本提示词
|
# 提取文本提示词
|
||||||
text_parts = [p.get("text", "") for p in content if p.get("type") == "text"]
|
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 ""
|
original_prompt = " ".join(text_parts) if text_parts else content[0].get("text", "") if content else ""
|
||||||
|
|
||||||
# 构建 media_references(扁平格式,便于外部读取)
|
# 构建 media_references(扁平格式,便于外部读取)
|
||||||
# 构建 local_media_json(嵌套格式,与 Volcano SDK 兼容)
|
# 构建 local_media_json(嵌套格式,与 Volcano SDK 兼容)
|
||||||
from app.services.api_v3.file_service import process_media_url
|
|
||||||
|
|
||||||
media_refs = [] # 扁平格式: {"type": "image", "url": "...", "role": "..."}
|
media_refs = [] # 扁平格式: {"type": "image", "url": "...", "role": "..."}
|
||||||
local_media_refs = [] # 本地下载路径(嵌套格式)
|
|
||||||
|
|
||||||
for p in content:
|
for p in content:
|
||||||
ptype = p.get("type", "")
|
ptype = p.get("type", "")
|
||||||
@@ -62,9 +62,27 @@ async def create_video_task(
|
|||||||
"role": p.get("role"),
|
"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:
|
try:
|
||||||
local_path = await process_media_url(original_url, media_type)
|
local_path = await process_media_url(original_url, ptype.replace("_url", ""))
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.warning("Failed to download media %s: %s", original_url[:80], exc)
|
logger.warning("Failed to download media %s: %s", original_url[:80], exc)
|
||||||
local_path = original_url
|
local_path = original_url
|
||||||
|
|||||||
@@ -386,6 +386,7 @@ async def delete_v3_project_remote(
|
|||||||
会先级联删除项目下所有素材的远端资源,再删除项目的远端 Group。
|
会先级联删除项目下所有素材的远端资源,再删除项目的远端 Group。
|
||||||
"""
|
"""
|
||||||
# 先删除项目下所有素材的远端资源
|
# 先删除项目下所有素材的远端资源
|
||||||
|
# 使用 with_for_update(skip_locked=True) 避免与独立素材删除任务冲突
|
||||||
assets = (
|
assets = (
|
||||||
await db.execute(
|
await db.execute(
|
||||||
select(VpV3Asset).where(
|
select(VpV3Asset).where(
|
||||||
@@ -393,10 +394,17 @@ async def delete_v3_project_remote(
|
|||||||
VpV3Asset.deleted_at.is_not(None),
|
VpV3Asset.deleted_at.is_not(None),
|
||||||
VpV3Asset.remote_delete_status == PrivatePortraitRemoteDeleteStatus.PENDING.value,
|
VpV3Asset.remote_delete_status == PrivatePortraitRemoteDeleteStatus.PENDING.value,
|
||||||
)
|
)
|
||||||
|
.with_for_update(skip_locked=True)
|
||||||
)
|
)
|
||||||
).scalars().all()
|
).scalars().all()
|
||||||
|
|
||||||
for asset in assets:
|
for asset in assets:
|
||||||
|
# 二次确认:如果独立素材删除任务已处理完该素材,跳过
|
||||||
|
if asset.remote_delete_status not in (
|
||||||
|
PrivatePortraitRemoteDeleteStatus.PENDING.value,
|
||||||
|
PrivatePortraitRemoteDeleteStatus.FAILED.value,
|
||||||
|
):
|
||||||
|
continue
|
||||||
if asset.remote_asset_id:
|
if asset.remote_asset_id:
|
||||||
try:
|
try:
|
||||||
await ArkPrivateAssetClient(for_celery=True).delete_asset(
|
await ArkPrivateAssetClient(for_celery=True).delete_asset(
|
||||||
|
|||||||
@@ -128,8 +128,18 @@ def _check(limit: int | None, used: int | float | None, delta: int | float, fiel
|
|||||||
|
|
||||||
|
|
||||||
async def check_project_quota(db: AsyncSession, *, api_key_id: str, delta: int = 1) -> VpV3ApiKeyQuota:
|
async def check_project_quota(db: AsyncSession, *, api_key_id: str, delta: int = 1) -> VpV3ApiKeyQuota:
|
||||||
"""创建项目前校验配额。"""
|
"""创建项目前校验配额(带行锁,防止并发超配)。"""
|
||||||
quota = await ensure_quota_enabled(db, api_key_id=api_key_id)
|
quota = await ensure_quota_enabled(db, api_key_id=api_key_id)
|
||||||
|
# 用行锁重新读取,保证并发安全
|
||||||
|
quota = (
|
||||||
|
await db.execute(
|
||||||
|
select(VpV3ApiKeyQuota)
|
||||||
|
.where(VpV3ApiKeyQuota.api_key_id == api_key_id)
|
||||||
|
.with_for_update()
|
||||||
|
.limit(1)
|
||||||
|
)
|
||||||
|
).scalar_one()
|
||||||
|
await _refresh_quota_used(db, quota)
|
||||||
_check(quota.project_limit, quota.project_used, delta, "项目数")
|
_check(quota.project_limit, quota.project_used, delta, "项目数")
|
||||||
return quota
|
return quota
|
||||||
|
|
||||||
@@ -141,13 +151,23 @@ async def check_asset_quota(
|
|||||||
asset_count_delta: int = 1,
|
asset_count_delta: int = 1,
|
||||||
file_size_bytes: int | None = None,
|
file_size_bytes: int | None = None,
|
||||||
) -> VpV3ApiKeyQuota:
|
) -> VpV3ApiKeyQuota:
|
||||||
"""上传素材前校验配额。
|
"""上传素材前校验配额(带行锁,防止并发超配)。
|
||||||
|
|
||||||
注:「存储空间上限」已从业务约束中移除(不再做硬性配额限制),仅保留素材数量上限
|
注:「存储空间上限」已从业务约束中移除(不再做硬性配额限制),仅保留素材数量上限
|
||||||
与项目数量上限的校验;storage_mb_used 字段仍会在 get_quota 中刷新用于统计展示。
|
与项目数量上限的校验;storage_mb_used 字段仍会在 get_quota 中刷新用于统计展示。
|
||||||
"""
|
"""
|
||||||
del file_size_bytes # 不再用于配额校验(仅保留形参兼容现有调用点)
|
del file_size_bytes # 不再用于配额校验(仅保留形参兼容现有调用点)
|
||||||
quota = await ensure_quota_enabled(db, api_key_id=api_key_id)
|
quota = await ensure_quota_enabled(db, api_key_id=api_key_id)
|
||||||
|
# 用行锁重新读取,保证并发安全
|
||||||
|
quota = (
|
||||||
|
await db.execute(
|
||||||
|
select(VpV3ApiKeyQuota)
|
||||||
|
.where(VpV3ApiKeyQuota.api_key_id == api_key_id)
|
||||||
|
.with_for_update()
|
||||||
|
.limit(1)
|
||||||
|
)
|
||||||
|
).scalar_one()
|
||||||
|
await _refresh_quota_used(db, quota)
|
||||||
_check(quota.asset_limit, quota.asset_used, asset_count_delta, "素材总数")
|
_check(quota.asset_limit, quota.asset_used, asset_count_delta, "素材总数")
|
||||||
return quota
|
return quota
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user