diff --git a/video-gen-api/app/services/video_queue.py b/video-gen-api/app/services/video_queue.py index da756939..b4e35f73 100644 --- a/video-gen-api/app/services/video_queue.py +++ b/video-gen-api/app/services/video_queue.py @@ -3,6 +3,7 @@ import json import logging import os from datetime import datetime, timezone +from urllib.parse import urlparse from sqlalchemy import or_, select @@ -49,6 +50,31 @@ def _source_date_dir(record: GenerationRecord) -> str: return created.strftime("%Y/%m/%d") +def _normalize_image_extension(output_format: str | None, remote_url: str | None = None) -> str: + value = str(output_format or "").strip().lower() + if value in {"jpg", "jpeg"}: + return "jpg" + if value == "png": + return "png" + if value == "webp": + return "webp" + + if remote_url: + try: + path = urlparse(remote_url).path or "" + except Exception: + path = str(remote_url) + suffix = os.path.splitext(path)[1].lower().lstrip(".") + if suffix in {"jpg", "jpeg"}: + return "jpg" + if suffix == "png": + return "png" + if suffix == "webp": + return "webp" + + return "jpg" + + async def _download_generation_record_upscale_source( record: GenerationRecord, remote_url: str, @@ -308,10 +334,10 @@ class TaskQueue: await asyncio.sleep(POLL_INTERVAL) await self.queue.put(record_id) - async def _process_image(self, db, record): + async def _process_image(self, db, record: GenerationRecord): """Process image generation task - calls API directly.""" record_id = record.id - from app.services.image_gen import submit_image_task, _log_image_response + from app.services.image_gen import submit_image_task try: engine = await get_active_image_engine(db) @@ -323,60 +349,81 @@ class TaskQueue: include_media_references=False, ) - if poll_result["error"] == "": - remote_url = poll_result.get("image_url") - storage_path = None - file_size_bytes = 0 - if settings.STORAGE_TYPE == "local" and remote_url: - try: - date_dir = datetime.now().strftime("%Y/%m/%d") - dest_dir = os.path.join(settings.STORAGE_IMAGE_LOCAL_PATH, date_dir) - os.makedirs(dest_dir, exist_ok=True) - dest = os.path.join(dest_dir, f"{record_id}.png") - await download_image(remote_url, dest) - record.image_url = f"/generate/images/{date_dir}/{record_id}.png" - storage_path = dest - file_size_bytes = safe_file_size(dest) - except Exception as e: - logger.warning(f"Download failed, using remote URL: {e}") - record.image_url = remote_url - else: - record.image_url = remote_url - record.image_tokens_used = poll_result.get("image_tokens", 0) - await sync_generation_record_media_token_snapshot(db, record, provider_response=poll_result) - record.status = "completed" - record.generated_at = datetime.now() - if record.image_url: - await record_generation_record_generated_resource( - db, - record, - resource_url=record.image_url, - storage_path=storage_path, - file_size_bytes=file_size_bytes, - remote_url=remote_url, - generated_at=record.generated_at, - ) - await db.commit() - logger.info(f"Image task completed: {record_id}") - else: - await mark_generation_record_failed_and_refund_once( - db, - record=record, - error_message=poll_result.get("error", "图片生成失败"), - ) - await db.commit() - logger.info(f"Image task failed: {record_id}") - _log_image_response(record_id, poll_result) + if not isinstance(poll_result, dict): + raise RuntimeError("图片供应商返回结构异常") - except Exception as e: + items = poll_result.get("items") or [] + if not isinstance(items, list): + raise RuntimeError("图片供应商返回结果列表异常") + if not items: + raise RuntimeError("图片供应商未返回图片结果") + if len(items) != 1: + raise RuntimeError(f"图片供应商单图返回数量异常,期望 1,实际 {len(items)}") + + item = items[0] or {} + if not isinstance(item, dict): + raise RuntimeError("图片供应商返回单项结果结构异常") + + item_error = item.get("error_message") or item.get("error_code") + if item_error: + raise RuntimeError(str(item_error)) + + remote_url = str(item.get("remote_result_url") or "").strip() + if not remote_url: + raise RuntimeError("图片供应商成功响应但没有图片地址") + + storage_path = None + file_size_bytes = 0 + if settings.STORAGE_TYPE == "local": + try: + date_dir = _source_date_dir(record) + dest_dir = os.path.join(settings.STORAGE_IMAGE_LOCAL_PATH, date_dir) + os.makedirs(dest_dir, exist_ok=True) + extension = _normalize_image_extension(item.get("output_format"), remote_url) + dest = os.path.join(dest_dir, f"{record_id}.{extension}") + await download_image(remote_url, dest) + record.image_url = f"/generate/images/{date_dir}/{record_id}.{extension}" + storage_path = dest + file_size_bytes = safe_file_size(dest) + except Exception as exc: + logger.warning("GenerationRecord 图片本地保存失败,回退远程地址: record_id=%s error=%s", record_id, exc) + record.image_url = remote_url + else: + record.image_url = remote_url + + record.image_tokens_used = int(poll_result.get("image_tokens", 0) or 0) + provider_response = poll_result.get("response_data") or {} + await sync_generation_record_media_token_snapshot( + db, + record, + provider_response=provider_response if isinstance(provider_response, dict) else {}, + ) + record.status = "completed" + record.pipeline_stage = GenerationRecordPipelineStage.DONE.value + record.generated_at = datetime.now(timezone.utc) + record.error_message = None + if record.image_url: + await record_generation_record_generated_resource( + db, + record, + resource_url=record.image_url, + storage_path=storage_path, + file_size_bytes=file_size_bytes, + remote_url=remote_url, + generated_at=record.generated_at, + ) + await db.commit() + logger.info("Image task completed: %s", record_id) + + except Exception as exc: + record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value await mark_generation_record_failed_and_refund_once( db, record=record, - error_message=str(e), + error_message=(getattr(exc, "safe_message", None) or str(exc) or "图片生成失败"), ) - _log_image_response(record_id, {}, str(e)) await db.commit() - logger.error(f"Image task failed: {record_id}, error: {e}") + logger.error("Image task failed: %s, error: %s", record_id, exc, exc_info=True) def stop(self): """Signal the queue to stop."""