Files
video-gen/video-gen-api/app/services/api_v3/task_service.py
T
root 0c511f3451 1、增加调用 AI 视频生成能力和虚拟素材库管理的对外api
2、增加后台apikkey管理
3、增加apikey单独的模型定价
4、增加apikey调用情况
5、完善所有数据的注释增加
2026-08-06 13:13:28 +08:00

216 lines
7.1 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
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:
"""创建视频生成任务记录。"""
# 提取文本提示词
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 兼容)
from app.services.api_v3.file_service import process_media_url
media_refs = [] # 扁平格式: {"type": "image", "url": "...", "role": "..."}
local_media_refs = [] # 本地下载路径(嵌套格式)
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"),
})
# 下载文件到本地
try:
local_path = await process_media_url(original_url, media_type)
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)