项目AI生成时不携带附件
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user