import asyncio import json import logging import os import time from datetime import datetime, timezone from sqlalchemy.ext.asyncio import AsyncSession from app.models.api.api_key import ApiKey from app.schemas.api_v3.image import ApiImageGenerateRequest, ApiImageGenerateResponse, ApiImageGenerateDataItem from app.schemas.api_v3.video import ApiVideoCreateRequest, ApiVideoCreateResponse from app.services.api_v3 import task_service, engine_service, upscale_service from app.services.api_v3.quota_service import can_start_video_task from app.services.api_v3.pricing_service import calc_api_video_price, calc_api_image_price, PricingNotConfiguredError from app.services.api_v3.logging_service import log_model_request 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( db: AsyncSession, key: ApiKey, callable_models: list[dict], req: ApiVideoCreateRequest, ) -> ApiVideoCreateResponse: """提交视频生成任务(异步)。 流程: 1. 检查并发视频任务数 2. 解析引擎 3. 构建超分快照 4. 创建任务记录 5. 入队 Celery 任务 6. 返回 task_id """ # 1. 检查是否可以立即启动(并发限制) can_start = await can_start_video_task(key, db) # 2. 解析引擎 engine_id, engine = await engine_service.resolve_engine_by_model_name( db, req.model, callable_models, "video" ) engine_snapshot = engine_service.build_engine_snapshot(engine) # 3. 构建超分快照 provider_resolution, upscale_enabled, upscale_snapshot_json = await upscale_service.build_api_upscale_snapshot( db, key.id, req.resolution or "480p", req.ratio ) # 如果超分要求不同的生成分辨率,使用超分的 final_provider_resolution = provider_resolution or req.resolution or "480p" def _get_max_supported_duration(engine) -> int | None: """从引擎 supported_durations 获取最大时长。""" try: durations = json.loads(engine.supported_durations) if engine.supported_durations else [] return max(durations) if durations else engine.max_duration except (json.JSONDecodeError, TypeError, ValueError): return engine.max_duration # 3.5 验证传入的媒体文件是否符合引擎配置要求 from app.services.api_v3.file_service import process_media_url from app.services.video_upscale.media_service import probe_video from fastapi import HTTPException, status # noqa: F401 input_image_count = 0 input_video_count = 0 input_audio_count = 0 input_video_duration = 0.0 local_media_refs = [] # 存储本地路径 # 统计各类媒体数量 for p in req.content: ptype = p.type if ptype == "image_url": input_image_count += 1 elif ptype == "video_url": input_video_count += 1 elif ptype == "audio_url": input_audio_count += 1 # 视频引擎校验(本函数仅处理视频生成) # 校验图片数量限制 if input_image_count > (engine.max_image_count or 0): raise HTTPException( status_code=400, detail=f"该引擎最多支持 {engine.max_image_count} 张参考图片,当前传入 {input_image_count} 张", ) # 校验视频数量限制 if input_video_count > (engine.max_video_count or 0): raise HTTPException( status_code=400, detail=f"该引擎最多支持 {engine.max_video_count} 个参考视频,当前传入 {input_video_count} 个", ) # 校验音频数量限制 if input_audio_count > (engine.max_audio_count or 0): raise HTTPException( status_code=400, detail=f"该引擎最多支持 {engine.max_audio_count} 个参考音频,当前传入 {input_audio_count} 个", ) for p in req.content: ptype = p.type if ptype == "text": local_media_refs.append(p.model_dump(exclude_none=True)) continue original_url = "" if ptype == "image_url" and p.image_url: original_url = p.image_url.get("url", "") elif ptype == "video_url" and p.video_url: original_url = p.video_url.get("url", "") elif ptype == "audio_url" and p.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 # 如果是视频/音频,探测实际时长并校验 if ptype in ("video_url", "audio_url") and local_path: try: media_info = await probe_video(local_path) if media_info and media_info.duration_seconds: duration = media_info.duration_seconds # 校验最低时长(2秒) if duration < 2.0: raise HTTPException( status_code=400, detail=f"上传的{ptype.replace('_url', '')}时长不能低于2秒,当前时长: {duration:.1f}秒", ) # 校验最高时长(根据引擎 supported_durations 最大值) max_duration = _get_max_supported_duration(engine) if max_duration and duration > max_duration: raise HTTPException( status_code=400, detail=f"上传的{ptype.replace('_url', '')}时长不能超过{max_duration}秒,当前时长: {duration:.1f}秒", ) if ptype == "video_url": input_video_duration += duration except HTTPException: raise except Exception as exc: logger.warning("Failed to probe media duration: %s", exc) local_media_refs.append({ "type": ptype, ptype: {"url": local_path}, "role": p.role, }) # 3.6 计算价格(基于实际探测的视频时长) try: estimated_price = await calc_api_video_price( db, duration=req.duration or 5, resolution=req.resolution or "480p", engine_id=engine_id, input_video_duration=input_video_duration, input_image_count=input_image_count, ) except PricingNotConfiguredError as exc: from fastapi import HTTPException, status raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc), ) # 预检配额 if key.quota_limit is not None and key.quota_used + estimated_price > key.quota_limit: from fastapi import HTTPException, status raise HTTPException( status_code=status.HTTP_429_TOO_MANY_REQUESTS, detail=f"配额不足 (需要 {estimated_price:.2f} 元, 剩余 {key.quota_limit - key.quota_used:.2f} 元)", ) # 预扣配额 if estimated_price > 0: key.quota_used = round((key.quota_used or 0.0) + estimated_price, 2) # 4. 创建任务记录(幂等性已在路由层检查) content_dicts = [p.model_dump(exclude_none=True) for p in req.content] task = await task_service.create_video_task( db=db, api_key_id=key.id, model_name=req.model, engine_id=engine_id, engine_snapshot=engine_snapshot, content=content_dicts, ratio=req.ratio, duration=req.duration, resolution=req.resolution, provider_generation_resolution=final_provider_resolution, upscale_enabled=upscale_enabled, upscale_snapshot_json=upscale_snapshot_json, idempotency_key=req.idempotency_key, local_media_refs=local_media_refs, ) task.credits_cost = estimated_price # 记录预扣金额 # 根据并发限制决定立即执行还是排队 if can_start: # 立即执行 task.status = "pending" task.pipeline_stage = "queued" else: # 排队等待 task.status = "queued" task.pipeline_stage = "waiting_concurrency" await db.commit() # 提交时即记录使用日志(配额已预扣) try: from app.services.api_v3.usage_log_service import record_usage quota_before = key.quota_used - estimated_price # 扣减前的余额 quota_after = key.quota_used # 扣减后的余额 price_detail = { "base_price": getattr(locals(), "base_price", 0), "per_second_price": getattr(locals(), "per_second_price", 0), "duration": req.duration, "resolution": req.resolution, "ratio": req.ratio, "total": estimated_price, } await record_usage( db=db, api_key_id=key.id, request_type="video_create", model_name=req.model, gen_type="video", status="success", task_id=task.id, credits_cost=estimated_price, price_action="deduct", resolution=req.resolution, duration=req.duration, quota_before=quota_before, quota_after=quota_after, price_detail_json=json.dumps(price_detail, ensure_ascii=False), ) except Exception as log_exc: logger.error("Failed to record usage on submit: %s", log_exc) # 记录模型调用日志 log_model_request( engine_id=engine_id, model_name=req.model, task_id=task.id, params={ "ratio": req.ratio, "duration": req.duration, "resolution": req.resolution, "generate_audio": req.generate_audio, "watermark": req.watermark, "content_count": len(req.content), "queued": not can_start, }, ) # 5. 只有立即执行的才入队 Celery if can_start: from app.tasks.api_generation_tasks import api_create_generation_task api_create_generation_task.apply_async( args=[task.id], queue="gen_api_create", ) status_str = "queued" if can_start else "pending_queue" logger.info("API video task created: task_id=%s model=%s key=%s price=%.2f status=%s", task.id, req.model, key.id, estimated_price, status_str) return ApiVideoCreateResponse(id=task.id) async def generate_image_sync( db: AsyncSession, key: ApiKey, callable_models: list[dict], req: "ApiImageGenerateRequest", start_time: float, ) -> ApiImageGenerateResponse: """同步生成图片。 流程: 1. 解析引擎 2. 创建任务记录 3. 调用 Volcano Ark SDK(同步) 4. 下载图片 5. 更新任务状态 6. 记录使用日志 7. 返回结果 """ from app.services.api_v3.usage_log_service import record_usage # 1. 解析引擎 engine_id, engine = await engine_service.resolve_engine_by_model_name( db, req.model, callable_models, "image" ) engine_snapshot = engine_service.build_engine_snapshot(engine) # 2. 创建任务记录 task = await task_service.create_image_task( db=db, api_key_id=key.id, model_name=req.model, engine_id=engine_id, engine_snapshot=engine_snapshot, prompt=req.prompt, size=req.size, ) # 2.5 验证传入的媒体文件是否符合引擎配置要求 # 校验参考图片数量限制 input_image_count = len(req.image) if req.image else 0 if input_image_count > (engine.max_reference_image_count or 0): raise HTTPException( status_code=400, detail=f"该引擎最多支持 {engine.max_reference_image_count} 张参考图片,当前传入 {input_image_count} 张", ) # 校验组图数量限制 generation_count = req.generation_count or 1 if generation_count > (engine.multi_image_max_images or 1): raise HTTPException( status_code=400, detail=f"该引擎最多支持生成 {engine.multi_image_max_images} 张图片,当前请求 {generation_count} 张", ) try: estimated_price = await calc_api_image_price( db, image_size=req.size or "2K", engine_id=engine_id, input_image_count=input_image_count, ) except PricingNotConfiguredError as exc: from fastapi import HTTPException, status raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc), ) # 预检配额 if key.quota_limit is not None and key.quota_used + estimated_price > key.quota_limit: from fastapi import HTTPException, status raise HTTPException( status_code=status.HTTP_429_TOO_MANY_REQUESTS, detail=f"配额不足 (需要 {estimated_price:.2f} 元, 剩余 {key.quota_limit - key.quota_used:.2f} 元)", ) # 预扣配额 if estimated_price > 0: key.quota_used = round((key.quota_used or 0.0) + estimated_price, 2) task.credits_cost = estimated_price await db.commit() # 提交时即记录使用日志(配额已预扣) try: quota_before = key.quota_used - estimated_price quota_after = key.quota_used price_detail = { "base_price": getattr(locals(), "base_price", 0), "size": req.size, "generation_count": req.generation_count or 1, "total": estimated_price, } await record_usage( db=db, api_key_id=key.id, request_type="image_generate", model_name=req.model, gen_type="image", status="success", task_id=task.id, credits_cost=estimated_price, price_action="deduct", resolution=req.size, quota_before=quota_before, quota_after=quota_after, price_detail_json=json.dumps(price_detail, ensure_ascii=False), ) except Exception as log_exc: logger.error("Failed to record usage on image submit: %s", log_exc) try: # 设置总体超时(120秒,防止同步请求长时间挂起) _IMAGE_GEN_TIMEOUT = 120 # 3. 调用 Volcano Ark SDK(同步函数,在线程中执行) from app.services.image_gen import submit_image_task, download_image from app.config import settings # 构建 media_references,下载图片到本地 from app.services.api_v3.file_service import process_media_url image_refs = [] if req.image: for url in req.image: try: local_path = await process_media_url(url, "image") except Exception as exc: logger.warning("Failed to download image %s: %s", url[:80], exc) local_path = url image_refs.append({"type": "image", "url": local_path}) # 临时设置 media_references task.media_references = json.dumps(image_refs, ensure_ascii=False) if image_refs else None task.image_size = req.size or "2K" await db.flush() # 在线程中执行同步 SDK 调用(带超时保护) result = await asyncio.wait_for( asyncio.to_thread( submit_image_task, db, engine, task, True, # include_media_references req.generation_count or 1, ), timeout=_IMAGE_GEN_TIMEOUT, ) # 4. 下载图片 items = result.get("items", []) downloaded_items: list[ApiImageGenerateDataItem] = [] for item in items: url = item.get("remote_result_url") if url: # 下载到本地 date_dir = datetime.now().strftime("%Y%m%d") dest_dir = f"./storage/generate/api/images/{date_dir}" os.makedirs(dest_dir, exist_ok=True) dest_path = os.path.join(dest_dir, f"{task.id}_{item.get('generation_index', 1)}.png") try: 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: logger.warning("Image download failed: %s", dl_err) # 将本地路径转为完整 URL image_url = _make_image_url(dest_path) downloaded_items.append(ApiImageGenerateDataItem( url=image_url, size=item.get("size"), output_format=item.get("output_format"), )) elif item.get("error_message"): downloaded_items.append(ApiImageGenerateDataItem( url=None, )) # 5. 使用预扣金额(不再重复扣减) task.credits_cost = estimated_price # 6. 更新任务状态 task.status = "completed" task.pipeline_stage = "done" task.generated_at = datetime.now(timezone.utc) if downloaded_items and downloaded_items[0].url: task.image_url = downloaded_items[0].url await db.commit() # 提交时已记录使用日志,成功时无需重复记录 duration_ms = int((time.perf_counter() - start_time) * 1000) return ApiImageGenerateResponse( created=result.get("created", int(time.time())), data=downloaded_items, 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: # 失败:退回预扣配额 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 = str(exc) task.credits_cost = 0 # 实际消耗为0(已退回) await db.commit() duration_ms = int((time.perf_counter() - start_time) * 1000) quota_after_refund = key.quota_used # 退回后的余额 await record_usage( db=db, api_key_id=key.id, request_type="image_generate", model_name=req.model, gen_type="image", status="failed", task_id=task.id, credits_cost=estimated_price, refund_amount=estimated_price, request_duration_ms=duration_ms, error_message=str(exc), error_code="generation_failed", price_action="refund", resolution=req.size, generation_count=req.generation_count or 1, quota_before=quota_after_refund, quota_after=quota_after_refund + estimated_price, ) await db.commit() raise