import asyncio import json import logging import os from datetime import datetime, timezone from urllib.parse import urlparse from sqlalchemy import or_, select from app.config import settings from app.enums.generation_status import GenerationRecordPipelineStage from app.models.base import async_session from app.models.generation_record import GenerationRecord from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once from app.services.image_gen import download_image, get_active_image_engine from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot from app.services.provider_limit import provider_limit from app.services.resource_accounting_service import ( record_generation_record_generated_resource, safe_file_size, ) from app.services.video_cover_service import create_video_cover_for_local_video from app.services.video_gen import _log_video_response, download_video, get_active_engine, poll_task_status from app.services.video_upscale.media_service import build_part_mp4_path, probe_video, safe_remove logger = logging.getLogger("videogen") POLL_INTERVAL = 30 MAX_POLLS = 60 _PROVIDER_RECOVERABLE_STAGES = { None, "", GenerationRecordPipelineStage.CREATING_PROVIDER_TASK.value, GenerationRecordPipelineStage.WAITING_REMOTE.value, GenerationRecordPipelineStage.POLLING.value, GenerationRecordPipelineStage.RESULT_READY.value, GenerationRecordPipelineStage.DOWNLOADING.value, } def _is_provider_stage(record: GenerationRecord) -> bool: return (record.pipeline_stage or "") in _PROVIDER_RECOVERABLE_STAGES def _source_date_dir(record: GenerationRecord) -> str: created = record.created_at if created is None: created = datetime.now(timezone.utc) 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, ) -> tuple[str, int]: date_dir = _source_date_dir(record) dest_dir = os.path.join(settings.STORAGE_LOCAL_PATH, "_upscale_source", date_dir) os.makedirs(dest_dir, exist_ok=True) final_path = os.path.join(dest_dir, f"{record.id}.source.mp4") if os.path.isfile(final_path) and os.path.getsize(final_path) > 0: await probe_video(final_path) return final_path, safe_file_size(final_path) part_path = build_part_mp4_path(final_path) try: async with provider_limit("result_download", settings.RESULT_DOWNLOAD_MAX_CONCURRENCY): await download_video(remote_url, part_path) if not os.path.isfile(part_path) or os.path.getsize(part_path) <= 0: raise RuntimeError("超分源视频下载完成但临时文件为空") await probe_video(part_path) os.replace(part_path, final_path) return final_path, safe_file_size(final_path) except Exception: safe_remove(part_path) raise async def handle_generation_record_video_succeeded( db, record: GenerationRecord, *, remote_url: str, provider_response: dict | None, video_tokens: int = 0, ) -> bool: """处理 GenerationRecord 原视频生成成功。 返回 True 表示已进入超分队列;False 表示按原流程直接完成。 """ record.video_tokens_used = int(video_tokens or 0) await sync_generation_record_media_token_snapshot(db, record, provider_response=provider_response or {}) if bool(record.video_upscale_enabled_snapshot) and record.video_upscale_snapshot_json: record.pipeline_stage = GenerationRecordPipelineStage.DOWNLOADING.value await db.flush() source_path, source_size = await _download_generation_record_upscale_source(record, remote_url) from app.services.video_upscale.task_service import enqueue_upscale_task, prepare_video_upscale_task upscale = await prepare_video_upscale_task( db, generation_record=record, source_local_path=source_path, source_file_size_bytes=source_size, source_remote_url=remote_url, ) # enqueue_upscale_task 会先提交数据库,再投递 Celery;投递失败由 gen_recovery 补投。 await enqueue_upscale_task(db, upscale=upscale, reason="generation_record_source_ready") logger.info("GenerationRecord 已进入视频超分队列: record_id=%s upscale_task_id=%s", record.id, upscale.id) return True storage_path = None file_size_bytes = 0 if settings.STORAGE_TYPE == "local" and remote_url: try: date_dir = _source_date_dir(record) dest_dir = os.path.join(settings.STORAGE_LOCAL_PATH, date_dir) os.makedirs(dest_dir, exist_ok=True) dest = os.path.join(dest_dir, f"{record.id}.mp4") await download_video(remote_url, dest) record.video_url = f"/generate/videos/{date_dir}/{record.id}.mp4" cover_url, _cover_storage_path = create_video_cover_for_local_video( record_id=record.id, video_path=dest, date_dir=date_dir, log_prefix=f"GenerationRecord视频封面生成 record_id={record.id}", ) record.video_cover_url = cover_url 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.video_url = remote_url else: record.video_url = remote_url record.status = "completed" record.pipeline_stage = GenerationRecordPipelineStage.DONE.value record.generated_at = datetime.now(timezone.utc) record.error_message = None if record.video_url: await record_generation_record_generated_resource( db, record, resource_url=record.video_url, storage_path=storage_path, file_size_bytes=file_size_bytes, remote_url=remote_url, generated_at=record.generated_at, ) await db.commit() return False class TaskQueue: def __init__(self): self.queue: asyncio.Queue[str] = asyncio.Queue() self.running = False self._active: dict[str, int] = {} async def enqueue(self, record_id: str): await self.queue.put(record_id) async def recover(self): async with async_session() as db: result = await db.execute( select(GenerationRecord).where( GenerationRecord.status == "generating", GenerationRecord.seedance_task_id.isnot(None), GenerationRecord.deleted_at.is_(None), or_( GenerationRecord.pipeline_stage.is_(None), GenerationRecord.pipeline_stage.in_( [stage for stage in _PROVIDER_RECOVERABLE_STAGES if stage] ), ), ) ) records = result.scalars().all() for record in records: await self.queue.put(record.id) logger.info("Recovered task: %s (seedance: %s stage=%s)", record.id, record.seedance_task_id, record.pipeline_stage) async def run(self): self.running = True logger.info("Video queue started") while self.running: try: record_id = await asyncio.wait_for(self.queue.get(), timeout=5.0) except asyncio.TimeoutError: continue try: await self._process(record_id) except Exception as exc: logger.exception("Error processing %s: %s", record_id, exc) finally: self.queue.task_done() logger.info("Video queue stopped") async def _process(self, record_id: str): async with async_session() as db: result = await db.execute( select(GenerationRecord).where( GenerationRecord.id == record_id, GenerationRecord.deleted_at.is_(None), ).with_for_update().limit(1) ) record = result.scalar_one_or_none() if not record or record.status != "generating": return if record.gen_type == "video": if not _is_provider_stage(record): return await self._process_video(db, record) else: await self._process_image(db, record) async def _process_video(self, db, record: GenerationRecord): record_id = record.id if not record.seedance_task_id: record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value await mark_generation_record_failed_and_refund_once(db, record=record, error_message="缺少外部任务ID") await db.commit() return record.pipeline_stage = GenerationRecordPipelineStage.POLLING.value try: engine = await get_active_engine(db) poll_result = await poll_task_status(engine, record.seedance_task_id) except Exception as exc: logger.error("Poll error for %s: %s", record_id, exc) count = self._active.get(record_id, 0) + 1 self._active[record_id] = count if count >= MAX_POLLS: record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value await mark_generation_record_failed_and_refund_once(db, record=record, error_message=f"轮询超时: {exc}") await db.commit() self._active.pop(record_id, None) else: await db.commit() await asyncio.sleep(POLL_INTERVAL) await self.queue.put(record_id) return status = poll_result["status"] try: resp_data = json.loads(poll_result.get("response_data", "{}")) except (json.JSONDecodeError, TypeError): resp_data = {} _log_video_response(record_id, resp_data, poll_result.get("error")) if status == "succeeded": file_url = str(poll_result.get("video_url") or "").strip() if not file_url: record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value await mark_generation_record_failed_and_refund_once(db, record=record, error_message="供应商成功但未返回视频地址") await db.commit() return record.pipeline_stage = GenerationRecordPipelineStage.RESULT_READY.value await db.flush() try: await handle_generation_record_video_succeeded( db, record, remote_url=file_url, provider_response=resp_data, video_tokens=poll_result.get("video_tokens", 0), ) except Exception as exc: await db.rollback() result = await db.execute( select(GenerationRecord).where(GenerationRecord.id == record_id).with_for_update().limit(1) ) failed_record = result.scalar_one_or_none() if failed_record: failed_record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value await mark_generation_record_failed_and_refund_once( db, record=failed_record, error_message=f"视频结果下载失败: {exc}", ) await db.commit() logger.exception("GenerationRecord 视频成功结果处理失败: %s", record_id) self._active.pop(record_id, None) return if status == "failed": record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value await mark_generation_record_failed_and_refund_once( db, record=record, error_message=poll_result.get("error", "视频生成失败"), ) self._active.pop(record_id, None) await db.commit() logger.info("Video task failed: %s", record_id) return count = self._active.get(record_id, 0) + 1 self._active[record_id] = count if count >= MAX_POLLS: record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value await mark_generation_record_failed_and_refund_once(db, record=record, error_message="视频生成超时") self._active.pop(record_id, None) await db.commit() logger.info("Video task timed out: %s", record_id) else: await db.commit() await asyncio.sleep(POLL_INTERVAL) await self.queue.put(record_id) 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 try: engine = await get_active_image_engine(db) poll_result = await asyncio.to_thread( submit_image_task, db, engine, record, include_media_references=False, ) if not isinstance(poll_result, dict): raise RuntimeError("图片供应商返回结构异常") 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=(getattr(exc, "safe_message", None) or str(exc) or "图片生成失败"), ) await db.commit() logger.error("Image task failed: %s, error: %s", record_id, exc, exc_info=True) def stop(self): """Signal the queue to stop.""" self.running = False task_queue = TaskQueue()