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)