项目AI生成时不携带附件

This commit is contained in:
2026-06-04 17:48:37 +08:00
parent ca560e977b
commit fd9a259917
7 changed files with 112 additions and 24 deletions
+6 -1
View File
@@ -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)
+12 -2
View File
@@ -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)
@@ -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
@@ -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
+13 -8
View File
@@ -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:
+19 -10
View File
@@ -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)
+7 -1
View File
@@ -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")