import asyncio import json import logging import os from datetime import datetime from types import SimpleNamespace from sqlalchemy import select from app.models.base import async_session from app.models.generation_record import GenerationRecord from app.models.image_engine import ImageEngine from app.models.video_engine import VideoEngine from app.enums.model_pricing import PricingSnapshotStage, ProviderCostStatus from app.services.video_gen import poll_task_status, download_video, _log_video_response from app.services.image_gen import download_image, is_sync_image_provider_result_uncertain from app.services.media_token_usage_snapshot_service import ( mark_media_provider_cost_status, sync_generation_record_media_token_snapshot, ) 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.config import settings from app.services.generation_refund_service import mark_generation_record_failed_and_refund_once logger = logging.getLogger("videogen") POLL_INTERVAL = 30 # seconds between polls MAX_POLLS = 60 # max 30 minutes total def _loads(data: str | None) -> dict: if not data: return {} try: value = json.loads(data) return value if isinstance(value, dict) else {} except Exception: return {} async def _get_runtime_engine(db, record: GenerationRecord): """使用生成开始时冻结的引擎快照,数据库行只读取当前密钥。""" if not record.engine_id: raise ValueError("生成记录缺少锁定的 engine_id") snapshot = _loads(record.engine_snapshot_json) model = ImageEngine if record.gen_type == "image" else VideoEngine engine = (await db.execute(select(model).where(model.id == record.engine_id).limit(1))).scalar_one_or_none() if not engine: raise ValueError("锁定的生成引擎不存在") return SimpleNamespace( id=record.engine_id, name=snapshot.get("name") or engine.name, provider=snapshot.get("provider") or engine.provider, api_base=snapshot.get("api_base") or engine.api_base, api_key=engine.api_key, model_name=snapshot.get("model_name") or engine.model_name, generate_url=snapshot.get("generate_url") or getattr(engine, "generate_url", ""), query_url=snapshot.get("query_url") or getattr(engine, "query_url", ""), default_size=snapshot.get("default_size") or getattr(engine, "default_size", "2K"), ) class TaskQueue: def __init__(self): self.queue: asyncio.Queue[str] = asyncio.Queue() self.running = False self._active: dict[str, int] = {} # record_id -> poll count async def enqueue(self, record_id: str): """Add a record to the polling queue.""" await self.queue.put(record_id) async def recover(self): """Recover in-progress tasks from DB on startup.""" 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), ) ) records = result.scalars().all() for record in records: await self.queue.put(record.id) logger.info(f"Recovered task: {record.id} (seedance: {record.seedance_task_id})") async def run(self): """Main polling loop.""" 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 e: logger.error(f"Error processing {record_id}: {e}") finally: self.queue.task_done() logger.info("Video queue stopped") async def _process(self, record_id: str): """Process a single record: poll status and update DB.""" 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": await self._process_video(db, record) else: await self._process_image(db, record) async def _process_video(self, db, record): """Process video generation task.""" record_id = record.id if not record.seedance_task_id: await mark_generation_record_failed_and_refund_once( db, record=record, error_message="缺少外部任务ID", ) await db.commit() return try: engine = await _get_runtime_engine(db, record) poll_result = await poll_task_status(engine, record.seedance_task_id) except Exception as e: logger.error(f"Poll error for {record_id}: {e}") count = self._active.get(record_id, 0) + 1 self._active[record_id] = count if count >= MAX_POLLS: await mark_generation_record_failed_and_refund_once( db, record=record, error_message=f"轮询超时: {e}", ) await db.commit() del self._active[record_id] else: 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 = poll_result.get("video_url", "") resp_data["video_url"] = file_url resp_data["task_id"] = record.seedance_task_id record.provider_response_json = json.dumps(resp_data, ensure_ascii=False, default=str) record.video_tokens_used = poll_result.get("video_tokens", 0) # 视频 Provider 已完成时核算供应商成本;下载失败不影响已发生的供应商费用。 await sync_generation_record_media_token_snapshot( db, record, provider_response=resp_data, stage=PricingSnapshotStage.PROVIDER_ASYNC_COMPLETED.value, ) storage_path = None file_size_bytes = 0 if settings.STORAGE_TYPE == "local" and file_url: try: date_dir = datetime.now().strftime("%Y/%m/%d") 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(file_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 e: logger.warning(f"Download failed, using remote URL: {e}") record.video_url = file_url else: record.video_url = file_url record.status = "completed" record.generated_at = datetime.now() 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=file_url, generated_at=record.generated_at, ) await sync_generation_record_media_token_snapshot( db, record, provider_response=resp_data, stage=PricingSnapshotStage.RESOURCE_DOWNLOAD_COMPLETED.value, ) self._active.pop(record_id, None) await db.commit() logger.info(f"Video task completed: {record_id}") elif status == "failed": 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(f"Video task failed: {record_id}") else: count = self._active.get(record_id, 0) + 1 self._active[record_id] = count if count >= MAX_POLLS: await mark_generation_record_failed_and_refund_once( db, record=record, error_message="视频生成超时", ) self._active.pop(record_id, None) await db.commit() logger.info(f"Video task timed out: {record_id}") else: await db.commit() await asyncio.sleep(POLL_INTERVAL) await self.queue.put(record_id) async def _process_image(self, db, record): """Process image generation task - calls API directly.""" record_id = record.id from app.services.image_gen import submit_image_task, _log_image_response provider_call_started = False provider_call_completed = False try: engine = await _get_runtime_engine(db, record) provider_call_started = True poll_result = await asyncio.to_thread( submit_image_task, db, engine, record, include_media_references=False, ) provider_call_completed = True if poll_result["error"] == "": remote_url = poll_result.get("image_url") try: provider_response = json.loads(poll_result.get("response_data") or "{}") except (json.JSONDecodeError, TypeError): provider_response = {} record.provider_response_json = json.dumps(provider_response, ensure_ascii=False, default=str) record.image_tokens_used = poll_result.get("image_tokens", 0) # 火山图片接口为同步生成:最终响应返回后立即完成供应商成本核算。 await sync_generation_record_media_token_snapshot( db, record, provider_response=provider_response, stage=PricingSnapshotStage.PROVIDER_SYNC_COMPLETED.value, ) 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.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 sync_generation_record_media_token_snapshot( db, record, provider_response=provider_response, stage=PricingSnapshotStage.RESOURCE_DOWNLOAD_COMPLETED.value, ) await db.commit() logger.info(f"Image task completed: {record_id}") else: error_message = poll_result.get("error", "图片生成失败") await mark_media_provider_cost_status( db, owner=record, status=ProviderCostStatus.NOT_INCURRED.value, reason=error_message, usage_stage="provider_sync_failed", ) await mark_generation_record_failed_and_refund_once( db, record=record, error_message=error_message, ) await db.commit() logger.info(f"Image task failed: {record_id}") _log_image_response(record_id, poll_result) except Exception as e: if provider_call_started: uncertain = provider_call_completed or is_sync_image_provider_result_uncertain(e) await mark_media_provider_cost_status( db, owner=record, status=( ProviderCostStatus.PROVIDER_RESULT_UNCERTAIN.value if uncertain else ProviderCostStatus.NOT_INCURRED.value ), reason=str(e), usage_stage="provider_sync_exception", ) await mark_generation_record_failed_and_refund_once( db, record=record, error_message=str(e), ) _log_image_response(record_id, {}, str(e)) await db.commit() logger.error(f"Image task failed: {record_id}, error: {e}") def stop(self): """Signal the queue to stop.""" self.running = False task_queue = TaskQueue()