From fd9a259917db61f6989e9491b2b214a06830b90d Mon Sep 17 00:00:00 2001 From: GinHa <15201596918@163.com> Date: Thu, 4 Jun 2026 17:48:37 +0800 Subject: [PATCH] =?UTF-8?q?=E9=A1=B9=E7=9B=AEAI=E7=94=9F=E6=88=90=E6=97=B6?= =?UTF-8?q?=E4=B8=8D=E6=90=BA=E5=B8=A6=E9=99=84=E4=BB=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- video-gen-api/app/api/v1/admin.py | 7 +++- video-gen-api/app/api/v1/generation.py | 14 ++++++- .../services/generation_provider_service.py | 15 ++++++- .../app/services/generation_provider_types.py | 42 +++++++++++++++++++ video-gen-api/app/services/image_gen.py | 21 ++++++---- video-gen-api/app/services/video_gen.py | 29 ++++++++----- video-gen-api/app/services/video_queue.py | 8 +++- 7 files changed, 112 insertions(+), 24 deletions(-) create mode 100644 video-gen-api/app/services/generation_provider_types.py diff --git a/video-gen-api/app/api/v1/admin.py b/video-gen-api/app/api/v1/admin.py index 4439eadb..7dbb9847 100644 --- a/video-gen-api/app/api/v1/admin.py +++ b/video-gen-api/app/api/v1/admin.py @@ -1167,7 +1167,12 @@ async def admin_generate_video( try: from app.services.video_gen import get_active_engine, submit_video_task engine = await get_active_engine(db) - task_id = await submit_video_task(db, engine, record) + task_id = await submit_video_task( + db, + engine, + record, + include_media_references=False, + ) record.seedance_task_id = task_id await db.flush() await task_queue.enqueue(record_id) diff --git a/video-gen-api/app/api/v1/generation.py b/video-gen-api/app/api/v1/generation.py index 83e949f9..beafe0c6 100644 --- a/video-gen-api/app/api/v1/generation.py +++ b/video-gen-api/app/api/v1/generation.py @@ -420,7 +420,12 @@ async def generate( from app.services.video_queue import task_queue engine = await get_active_engine(db) - task_id = await submit_video_task(db, engine, record) + task_id = await submit_video_task( + db, + engine, + record, + include_media_references=False, + ) record.seedance_task_id = task_id await db.flush() await task_queue.enqueue(record_id) @@ -523,7 +528,12 @@ async def retry_generation( if record.gen_type == GenerationType.video: from app.services.video_gen import get_active_engine, submit_video_task, extract_error_message engine = await get_active_engine(db) - task_id = await submit_video_task(db, engine, record) + task_id = await submit_video_task( + db, + engine, + record, + include_media_references=False, + ) record.seedance_task_id = task_id await db.flush() await task_queue.enqueue(record_id) diff --git a/video-gen-api/app/services/generation_provider_service.py b/video-gen-api/app/services/generation_provider_service.py index e221c0ba..f7b2731e 100644 --- a/video-gen-api/app/services/generation_provider_service.py +++ b/video-gen-api/app/services/generation_provider_service.py @@ -68,7 +68,12 @@ async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict started = time.perf_counter() async with provider_limit("ark_video_create", settings.ARK_VIDEO_CREATE_MAX_CONCURRENCY): try: - provider_task_id = await submit_video_task(db, engine, task) + provider_task_id = await submit_video_task( + db, + engine, + task, + include_media_references=True, + ) response = {"task_id": provider_task_id} await log_provider_call( task, @@ -107,7 +112,13 @@ async def _create_image_sync_task(db: AsyncSession, task: ChatGenerationTask) -> started = time.perf_counter() async with provider_limit("ark_image_sync_create", settings.ARK_IMAGE_CREATE_MAX_CONCURRENCY): try: - result = await asyncio.to_thread(submit_image_task, db, engine, task) + result = await asyncio.to_thread( + submit_image_task, + db, + engine, + task, + include_media_references=True, + ) if result.get("error"): raise RuntimeError(result.get("error")) response_data = _try_json(result.get("response_data")) or result diff --git a/video-gen-api/app/services/generation_provider_types.py b/video-gen-api/app/services/generation_provider_types.py new file mode 100644 index 00000000..88a74ab8 --- /dev/null +++ b/video-gen-api/app/services/generation_provider_types.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +from typing import Protocol + + +class ProviderGenerationRecordLike(Protocol): + """图片/视频供应商提交接口需要的任务字段协议。 + + GenerationRecord 与 ChatGenerationTask 都具备这些字段,但二者不是同一个 ORM 模型。 + 使用 Protocol 可以避免把 submit_image_task / submit_video_task 错误限制为某一个具体模型。 + """ + + id: str + original_prompt: str + optimized_prompt: str | None + media_references: str | None + gen_type: str + duration: int | None + aspect_ratio: str | None + resolution: str | None + image_size: str | None + image_proportion: str | None + image_px: str | None + + +class ProviderImageEngineLike(Protocol): + """图片生成提交接口需要的引擎字段协议。""" + + name: str + api_base: str + api_key: str + model_name: str + default_size: str | None + + +class ProviderVideoEngineLike(Protocol): + """视频生成提交接口需要的引擎字段协议。""" + + name: str + api_base: str + api_key: str + model_name: str diff --git a/video-gen-api/app/services/image_gen.py b/video-gen-api/app/services/image_gen.py index 47e3c3c8..f5529af3 100644 --- a/video-gen-api/app/services/image_gen.py +++ b/video-gen-api/app/services/image_gen.py @@ -12,14 +12,16 @@ from volcenginesdkarkruntime import AsyncArk from app.config import settings from app.models.image_engine import ImageEngine -from app.models.generation_record import GenerationRecord from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data -from app.services.error_codes import extract_error_message +from app.services.generation_provider_types import ( + ProviderGenerationRecordLike, + ProviderImageEngineLike, +) logger = logging.getLogger("videogen") -def _log_image_request(engine, record_id: str, request_data: dict): +def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_data: dict): """Log image generation request to log/AiModel/YYYY-MM-DD.log""" if not is_enabled(): return @@ -102,9 +104,11 @@ def _resolve_url(url: str) -> str: def submit_image_task( db, - engine: ImageEngine, - record: GenerationRecord, -) -> str: + engine: ProviderImageEngineLike, + record: ProviderGenerationRecordLike, + *, + include_media_references: bool, +) -> dict: """Submit an image generation task via Ark SDK. Returns image_url.""" from volcenginesdkarkruntime import Ark @@ -114,10 +118,10 @@ def submit_image_task( timeout=300, ) - prompt = record.optimized_prompt + prompt = record.optimized_prompt or record.original_prompt image_urls = [] - if record.media_references: + if include_media_references and record.media_references: try: refs = json.loads(record.media_references) for ref in refs: @@ -137,6 +141,7 @@ def submit_image_task( "output_format": "png", "response_format": "url", "watermark": False, + "include_media_references": include_media_references, } if image_urls: diff --git a/video-gen-api/app/services/video_gen.py b/video-gen-api/app/services/video_gen.py index 43b1e7f5..29857762 100644 --- a/video-gen-api/app/services/video_gen.py +++ b/video-gen-api/app/services/video_gen.py @@ -12,14 +12,16 @@ from volcenginesdkarkruntime import AsyncArk from app.config import settings from app.models.video_engine import VideoEngine -from app.models.generation_record import GenerationRecord from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data -from app.services.error_codes import extract_error_message +from app.services.generation_provider_types import ( + ProviderGenerationRecordLike, + ProviderVideoEngineLike, +) logger = logging.getLogger("videogen") -def _log_video_request(engine, record_id: str, request_data: dict): +def _log_video_request(engine: ProviderVideoEngineLike, record_id: str, request_data: dict): """Log video generation request to log/AiModel/YYYY-MM-DD.log""" if not is_enabled(): return @@ -103,8 +105,10 @@ def _resolve_url(url: str) -> str: async def submit_video_task( db: AsyncSession, - engine: VideoEngine, - record: GenerationRecord, + engine: ProviderVideoEngineLike, + record: ProviderGenerationRecordLike, + *, + include_media_references: bool, ) -> str: """Submit a video generation task via Ark SDK. Returns task_id.""" client = AsyncArk( @@ -112,10 +116,11 @@ async def submit_video_task( api_key=engine.api_key, ) - content = [{"type": "text", "text": record.optimized_prompt}] + content = [{"type": "text", "text": record.optimized_prompt or record.original_prompt}] - # Add reference images/videos from media_references - if record.media_references: + # GenerationRecord 的附件已在 API 提词优化阶段参与过 optimized_prompt 生成; + # ChatGenerationTask 不走 API 提词优化,所以创建供应商任务时仍需携带附件。 + if include_media_references and record.media_references: try: refs = json.loads(record.media_references) for ref in refs: @@ -140,8 +145,12 @@ async def submit_video_task( "watermark": False, } - # Log request to AiModel log - _log_video_request(engine, record.id, request_payload) + # Log request to AiModel log. include_media_references 只用于排查日志,不传给供应商 API。 + _log_video_request( + engine, + record.id, + {**request_payload, "include_media_references": include_media_references}, + ) try: result = await client.content_generation.tasks.create(**request_payload) diff --git a/video-gen-api/app/services/video_queue.py b/video-gen-api/app/services/video_queue.py index f94388db..7129972b 100644 --- a/video-gen-api/app/services/video_queue.py +++ b/video-gen-api/app/services/video_queue.py @@ -207,7 +207,13 @@ class TaskQueue: try: engine = await get_active_image_engine(db) - poll_result = await asyncio.to_thread(submit_image_task, db, engine, record) + poll_result = await asyncio.to_thread( + submit_image_task, + db, + engine, + record, + include_media_references=False, + ) if poll_result["error"] == "": remote_url = poll_result.get("image_url")