1
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""生成任务领域服务。"""
|
||||
@@ -0,0 +1 @@
|
||||
"""AI 创作生成编排服务。"""
|
||||
@@ -0,0 +1,122 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.common import MAX_GENERATION_COUNT, MIN_GENERATION_COUNT
|
||||
from app.enums.generation_provider import IMAGE_MULTI_OUTPUT_MAX, IMAGE_MULTI_REFERENCE_MAX
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.models.video_engine import VideoEngine
|
||||
|
||||
IMAGE_DEFAULT_SIZE = "2K"
|
||||
IMAGE_DEFAULT_PROPORTION = "1:1"
|
||||
IMAGE_DEFAULT_PX = "2048x2048"
|
||||
VIDEO_DEFAULT_DURATION = 4
|
||||
VIDEO_DEFAULT_RATIO = "16:9"
|
||||
VIDEO_DEFAULT_RESOLUTION = "480p"
|
||||
|
||||
|
||||
def normalize_px(value: str | None) -> str | None:
|
||||
if not value:
|
||||
return value
|
||||
return value.replace("×", "x").replace("X", "x").replace("×x", "x").replace("x×", "x")
|
||||
|
||||
|
||||
def parse_json_list(value: str | None, fallback: list):
|
||||
try:
|
||||
parsed = json.loads(value or "")
|
||||
return parsed if isinstance(parsed, list) else fallback
|
||||
except Exception:
|
||||
return fallback
|
||||
|
||||
|
||||
def image_supported_sizes(engine: ImageEngine) -> dict:
|
||||
try:
|
||||
data = json.loads(engine.supported_sizes or "{}")
|
||||
return data if isinstance(data, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def normalize_generation_count(value: int | None) -> int:
|
||||
try:
|
||||
count = int(value or MIN_GENERATION_COUNT)
|
||||
except (TypeError, ValueError):
|
||||
count = MIN_GENERATION_COUNT
|
||||
return min(MAX_GENERATION_COUNT, max(MIN_GENERATION_COUNT, count))
|
||||
|
||||
|
||||
async def get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngine:
|
||||
query = select(ImageEngine).where(ImageEngine.is_active == True)
|
||||
if engine_id:
|
||||
query = query.where(ImageEngine.id == engine_id)
|
||||
else:
|
||||
query = query.order_by(ImageEngine.priority.desc()).limit(1)
|
||||
result = await db.execute(query)
|
||||
engine = result.scalar_one_or_none()
|
||||
if not engine:
|
||||
raise HTTPException(status_code=400, detail="没有可用的图片引擎")
|
||||
return engine
|
||||
|
||||
|
||||
async def get_video_engine(db: AsyncSession, engine_id: str | None) -> VideoEngine:
|
||||
query = select(VideoEngine).where(VideoEngine.is_active == True)
|
||||
if engine_id:
|
||||
query = query.where(VideoEngine.id == engine_id)
|
||||
else:
|
||||
query = query.order_by(VideoEngine.priority.desc())
|
||||
result = await db.execute(query.limit(1))
|
||||
engine = result.scalar_one_or_none()
|
||||
if not engine:
|
||||
raise HTTPException(status_code=400, detail="没有可用的视频引擎")
|
||||
return engine
|
||||
|
||||
|
||||
def build_image_snapshot(engine: ImageEngine, size: str, proportion: str, px: str) -> dict:
|
||||
return {
|
||||
"engine_type": "image",
|
||||
"id": engine.id,
|
||||
"name": engine.name,
|
||||
"provider": engine.provider,
|
||||
"api_base": engine.api_base,
|
||||
"api_key_masked": "****" if engine.api_key else "",
|
||||
"model_name": engine.model_name,
|
||||
"generate_url": engine.generate_url,
|
||||
"supported_models": parse_json_list(engine.supported_models, []),
|
||||
"default_size": engine.default_size,
|
||||
"multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)),
|
||||
"max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)),
|
||||
"multi_image_max_images": int(getattr(engine, "multi_image_max_images", IMAGE_MULTI_OUTPUT_MAX) or IMAGE_MULTI_OUTPUT_MAX),
|
||||
"max_reference_image_count": int(getattr(engine, "max_reference_image_count", IMAGE_MULTI_REFERENCE_MAX) or 0),
|
||||
"output_format": (getattr(engine, "output_format", "") or "").lower().strip(),
|
||||
"selected_size": size,
|
||||
"selected_proportion": proportion,
|
||||
"selected_px": px,
|
||||
}
|
||||
|
||||
|
||||
def build_video_snapshot(engine: VideoEngine, ratio: str, resolution: str, duration: int) -> dict:
|
||||
return {
|
||||
"engine_type": "video",
|
||||
"id": engine.id,
|
||||
"name": engine.name,
|
||||
"provider": engine.provider,
|
||||
"api_base": engine.api_base,
|
||||
"api_key_masked": "****" if engine.api_key else "",
|
||||
"model_name": engine.model_name,
|
||||
"generate_url": engine.generate_url,
|
||||
"query_url": engine.query_url,
|
||||
"supported_ratios": parse_json_list(engine.supported_ratios, []),
|
||||
"supported_resolutions": parse_json_list(engine.supported_resolutions, []),
|
||||
"supported_durations": parse_json_list(engine.supported_durations, []),
|
||||
"max_duration": engine.max_duration,
|
||||
"max_audio_count": engine.max_audio_count,
|
||||
"multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)),
|
||||
"max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)),
|
||||
"selected_ratio": ratio,
|
||||
"selected_resolution": resolution,
|
||||
"selected_duration": duration,
|
||||
}
|
||||
@@ -0,0 +1,532 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import SimpleNamespace
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.generation_provider import IMAGE_PROVIDER_CLAIM_LEASE_SECONDS
|
||||
from app.enums.generation_task import (
|
||||
ChatGenerationPipelineStage,
|
||||
ChatGenerationTaskEventType,
|
||||
ChatGenerationTaskStatus,
|
||||
GenerationMode,
|
||||
GenerationType,
|
||||
)
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.generation.ai.task_group_service import aggregate_main_task_status, load_children_map
|
||||
from app.services.generation.log_service import log_task_event
|
||||
from app.services.generation.provider_service import (
|
||||
create_image_sync_batch_result_with_engine,
|
||||
get_runtime_engine,
|
||||
)
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.image_gen import ImageProviderError
|
||||
from app.services.operation_log_service import build_exception_detail, log_operation_event
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ImageBatchClaim:
|
||||
acquired: bool
|
||||
main_task_id: str
|
||||
claim_token: str | None = None
|
||||
task_snapshot: SimpleNamespace | None = None
|
||||
runtime_engine: SimpleNamespace | None = None
|
||||
existing_child_ids: list[str] | None = None
|
||||
reason: str | None = None
|
||||
|
||||
|
||||
def _now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
def _json(value) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
return json.dumps(value, ensure_ascii=False, default=str)
|
||||
|
||||
|
||||
def _aware(value: datetime | None) -> datetime | None:
|
||||
if value is None:
|
||||
return None
|
||||
if value.tzinfo is None:
|
||||
return value.replace(tzinfo=timezone.utc)
|
||||
return value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _lease_alive(task: ChatGenerationTask, now: datetime | None = None) -> bool:
|
||||
lease_until = _aware(task.provider_create_lease_until)
|
||||
return bool(task.provider_create_claim_token and lease_until and lease_until > (now or _now()))
|
||||
|
||||
|
||||
def _task_snapshot(main: ChatGenerationTask) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
id=str(main.id),
|
||||
user_id=str(main.user_id),
|
||||
generation_mode=str(main.generation_mode),
|
||||
generation_count=int(main.generation_count or 1),
|
||||
original_prompt=main.original_prompt,
|
||||
optimized_prompt=main.optimized_prompt,
|
||||
media_references=main.media_references,
|
||||
gen_type=main.gen_type,
|
||||
duration=main.duration,
|
||||
aspect_ratio=main.aspect_ratio,
|
||||
resolution=main.resolution,
|
||||
image_size=main.image_size,
|
||||
image_proportion=main.image_proportion,
|
||||
image_px=main.image_px,
|
||||
engine_id=main.engine_id,
|
||||
)
|
||||
|
||||
|
||||
async def _claim_image_main_batch(db: AsyncSession, main_task_id: str) -> ImageBatchClaim:
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == main_task_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
ChatGenerationTask.gen_type == GenerationType.IMAGE.value,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
main = result.scalar_one_or_none()
|
||||
if not main:
|
||||
await db.rollback()
|
||||
return ImageBatchClaim(False, main_task_id, reason="main_missing")
|
||||
|
||||
children_map = await load_children_map(db, [main.id], include_deleted=True)
|
||||
existing_children = children_map.get(main.id, [])
|
||||
if existing_children:
|
||||
child_ids = [str(child.id) for child in existing_children if child.deleted_at is None]
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.commit()
|
||||
return ImageBatchClaim(False, main_task_id, existing_child_ids=child_ids, reason="already_split")
|
||||
|
||||
if main.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
status = str(main.status)
|
||||
await db.rollback()
|
||||
return ImageBatchClaim(False, main_task_id, reason=f"status_{status}")
|
||||
|
||||
now = _now()
|
||||
if _lease_alive(main, now):
|
||||
user_id = str(main.user_id)
|
||||
group_id = str(main.id)
|
||||
lease_until = main.provider_create_lease_until
|
||||
await db.rollback()
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="IMAGE_MAIN_CLAIM_REJECTED",
|
||||
event_status="skipped",
|
||||
source="celery",
|
||||
user_id=user_id,
|
||||
group_id=group_id,
|
||||
task_id=group_id,
|
||||
detail={"reason": "lease_alive", "lease_until": lease_until},
|
||||
)
|
||||
return ImageBatchClaim(False, main_task_id, reason="lease_alive")
|
||||
|
||||
deadline = _aware(main.deadline_at)
|
||||
if deadline and deadline <= now:
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=main,
|
||||
error_message="图片批量生成任务超时",
|
||||
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
await db.commit()
|
||||
return ImageBatchClaim(False, main_task_id, reason="deadline_expired")
|
||||
|
||||
claim_token = uuid4().hex
|
||||
main.provider_create_claim_token = claim_token
|
||||
main.provider_create_started_at = now
|
||||
main.provider_create_lease_until = now + timedelta(seconds=IMAGE_PROVIDER_CLAIM_LEASE_SECONDS)
|
||||
main.pipeline_stage = ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value
|
||||
runtime_engine = await get_runtime_engine(db, main)
|
||||
snapshot = _task_snapshot(main)
|
||||
user_id = str(main.user_id)
|
||||
generation_count = int(main.generation_count or 1)
|
||||
lease_until = main.provider_create_lease_until
|
||||
await db.commit()
|
||||
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="IMAGE_MAIN_CLAIM_ACQUIRED",
|
||||
event_status="success",
|
||||
source="celery",
|
||||
user_id=user_id,
|
||||
group_id=main_task_id,
|
||||
task_id=main_task_id,
|
||||
detail={
|
||||
"generation_count": generation_count,
|
||||
"claim_token_suffix": claim_token[-8:],
|
||||
"lease_until": lease_until,
|
||||
},
|
||||
)
|
||||
return ImageBatchClaim(
|
||||
True,
|
||||
main_task_id,
|
||||
claim_token=claim_token,
|
||||
task_snapshot=snapshot,
|
||||
runtime_engine=runtime_engine,
|
||||
)
|
||||
|
||||
|
||||
def _validate_provider_batch(provider_result: dict, generation_count: int) -> list[dict]:
|
||||
items = provider_result.get("items") or []
|
||||
if not isinstance(items, list):
|
||||
raise RuntimeError("图片供应商返回 items 结构异常")
|
||||
|
||||
success_items: list[dict] = []
|
||||
errors: list[str] = []
|
||||
for position, item in enumerate(items, start=1):
|
||||
if not isinstance(item, dict):
|
||||
errors.append(f"第{position}项返回结构无效")
|
||||
continue
|
||||
if item.get("error_message") or item.get("error_code"):
|
||||
errors.append(
|
||||
f"第{position}项: {item.get('error_message') or item.get('error_code') or '生成失败'}"
|
||||
)
|
||||
continue
|
||||
remote_url = str(item.get("remote_result_url") or "").strip()
|
||||
if not remote_url:
|
||||
errors.append(f"第{position}项: 供应商未返回图片地址")
|
||||
continue
|
||||
normalized = dict(item)
|
||||
normalized["generation_index"] = position
|
||||
success_items.append(normalized)
|
||||
|
||||
generated_images = int(provider_result.get("generated_images") or 0)
|
||||
if generated_images and generated_images != len(success_items):
|
||||
errors.append(
|
||||
f"usage.generated_images={generated_images} 与有效图片数 {len(success_items)} 不一致"
|
||||
)
|
||||
if len(items) != generation_count:
|
||||
errors.append(f"返回条目数应为 {generation_count},实际 {len(items)}")
|
||||
if len(success_items) != generation_count:
|
||||
errors.append(f"成功图片数应为 {generation_count},实际 {len(success_items)}")
|
||||
if errors:
|
||||
raise RuntimeError("图片组图未全部成功;" + ";".join(errors))
|
||||
return success_items
|
||||
|
||||
|
||||
async def _fail_claimed_main(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
main_task_id: str,
|
||||
claim_token: str,
|
||||
error_message: str,
|
||||
event_type: ChatGenerationTaskEventType,
|
||||
exception: Exception | None = None,
|
||||
) -> bool:
|
||||
try:
|
||||
await db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == main_task_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
main = result.scalar_one_or_none()
|
||||
if not main or main.provider_create_claim_token != claim_token:
|
||||
await db.rollback()
|
||||
return False
|
||||
|
||||
existing_map = await load_children_map(db, [main.id], include_deleted=True)
|
||||
if existing_map.get(main.id):
|
||||
# child 已经落库后不再允许图片生成退款。
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.commit()
|
||||
return False
|
||||
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=main,
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
task_id = str(main.id)
|
||||
user_id = str(main.user_id)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task_id=task_id,
|
||||
event_type=event_type.value,
|
||||
to_status=ChatGenerationTaskStatus.FAILED.value,
|
||||
to_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
message=error_message,
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=event_type.value,
|
||||
event_status="failed",
|
||||
source="celery",
|
||||
user_id=user_id,
|
||||
group_id=task_id,
|
||||
task_id=task_id,
|
||||
message=error_message,
|
||||
detail=build_exception_detail(exception) if exception else {"message": error_message},
|
||||
error=error_message,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
async def _split_children(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
main_task_id: str,
|
||||
claim_token: str,
|
||||
provider_result: dict,
|
||||
provider_items: list[dict],
|
||||
) -> list[str]:
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == main_task_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
ChatGenerationTask.gen_type == GenerationType.IMAGE.value,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
main = result.scalar_one_or_none()
|
||||
if not main:
|
||||
raise RuntimeError("图片主任务不存在或已删除")
|
||||
if main.provider_create_claim_token != claim_token:
|
||||
raise RuntimeError("图片主任务执行租约已失效,拒绝拆分子任务")
|
||||
if main.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
raise RuntimeError(f"图片主任务当前状态不允许拆分: {main.status}")
|
||||
|
||||
existing_map = await load_children_map(db, [main.id], include_deleted=True)
|
||||
existing = existing_map.get(main.id, [])
|
||||
if existing:
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.commit()
|
||||
return [str(child.id) for child in existing if child.deleted_at is None]
|
||||
|
||||
expected_count = max(1, int(main.generation_count or 1))
|
||||
if len(provider_items) != expected_count:
|
||||
raise RuntimeError(f"图片批量拆分数量不一致,期望 {expected_count},实际 {len(provider_items)}")
|
||||
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_START.value,
|
||||
event_status="started",
|
||||
source="celery",
|
||||
user_id=main.user_id,
|
||||
group_id=main.id,
|
||||
task_id=main.id,
|
||||
detail={"generation_count": expected_count},
|
||||
)
|
||||
|
||||
children: list[ChatGenerationTask] = []
|
||||
for item in provider_items:
|
||||
index = int(item.get("generation_index") or 0)
|
||||
if index < 1 or index > expected_count:
|
||||
raise RuntimeError(f"无效的图片生成序号: {index}")
|
||||
child = ChatGenerationTask(
|
||||
id=generate_id(),
|
||||
user_id=main.user_id,
|
||||
original_prompt=main.original_prompt,
|
||||
optimized_prompt=main.optimized_prompt,
|
||||
gen_type=main.gen_type,
|
||||
image_size=main.image_size,
|
||||
image_proportion=main.image_proportion,
|
||||
image_px=main.image_px,
|
||||
status=ChatGenerationTaskStatus.GENERATING.value,
|
||||
pipeline_stage=ChatGenerationPipelineStage.RESULT_READY.value,
|
||||
generation_mode=GenerationMode.CHATAPI_CHILD.value,
|
||||
parent_task_id=main.id,
|
||||
generation_count=expected_count,
|
||||
generation_index=index,
|
||||
media_references=main.media_references,
|
||||
remote_result_url=item.get("remote_result_url"),
|
||||
engine_id=main.engine_id,
|
||||
engine_snapshot_json=main.engine_snapshot_json,
|
||||
provider_response_json=_json(item.get("response_data") or {}),
|
||||
# 图片生成计费和 token 都归属于 main;child 只负责下载和资源展示。
|
||||
credits_cost=0,
|
||||
image_tokens_used=0,
|
||||
deadline_at=main.deadline_at,
|
||||
)
|
||||
children.append(child)
|
||||
|
||||
children.sort(key=lambda child: int(child.generation_index or 0))
|
||||
db.add_all(children)
|
||||
main.provider_response_json = _json(provider_result.get("response_data") or provider_result)
|
||||
main.image_tokens_used = int(provider_result.get("image_tokens") or 0)
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.flush()
|
||||
child_ids = [str(child.id) for child in children]
|
||||
main_id = str(main.id)
|
||||
main_user_id = str(main.user_id)
|
||||
await aggregate_main_task_status(db, parent_task_id=main_id)
|
||||
await db.commit()
|
||||
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_SUCCESS.value,
|
||||
event_status="success",
|
||||
source="celery",
|
||||
user_id=main_user_id,
|
||||
group_id=main_id,
|
||||
task_id=main_id,
|
||||
detail={"child_task_ids": child_ids},
|
||||
)
|
||||
return child_ids
|
||||
|
||||
|
||||
async def _enqueue_child_downloads(db: AsyncSession, child_ids: list[str]) -> dict[str, list[str]]:
|
||||
if not child_ids:
|
||||
return {"enqueued": [], "failed": []}
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id.in_(child_ids),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.order_by(ChatGenerationTask.generation_index.asc())
|
||||
)
|
||||
children = list(result.scalars().all())
|
||||
from app.tasks.generation_download_tasks import enqueue_download_task
|
||||
|
||||
enqueued: list[str] = []
|
||||
failed: list[str] = []
|
||||
for child in children:
|
||||
if child.status == ChatGenerationTaskStatus.COMPLETED.value:
|
||||
continue
|
||||
if child.pipeline_stage in {
|
||||
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
|
||||
ChatGenerationPipelineStage.DOWNLOADING.value,
|
||||
ChatGenerationPipelineStage.RETRY_WAITING.value,
|
||||
}:
|
||||
continue
|
||||
celery_task_id = await enqueue_download_task(db, child, reason="image_batch_split")
|
||||
if celery_task_id:
|
||||
enqueued.append(str(child.id))
|
||||
else:
|
||||
failed.append(str(child.id))
|
||||
|
||||
if children and children[0].parent_task_id:
|
||||
await aggregate_main_task_status(db, parent_task_id=str(children[0].parent_task_id))
|
||||
await db.commit()
|
||||
return {"enqueued": enqueued, "failed": failed}
|
||||
|
||||
|
||||
async def run_image_main_batch(db: AsyncSession, main_task: ChatGenerationTask) -> list[str]:
|
||||
"""单次同步组图,全部成功后原子拆分 child。
|
||||
|
||||
绝不在组图 API 失败后退化为 N 次单图请求。
|
||||
"""
|
||||
main_task_id = str(main_task.id)
|
||||
claim = await _claim_image_main_batch(db, main_task_id)
|
||||
if claim.existing_child_ids is not None:
|
||||
await _enqueue_child_downloads(db, claim.existing_child_ids)
|
||||
return claim.existing_child_ids
|
||||
if not claim.acquired or not claim.claim_token or not claim.task_snapshot or not claim.runtime_engine:
|
||||
return []
|
||||
|
||||
generation_count = max(1, int(claim.task_snapshot.generation_count or 1))
|
||||
try:
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_PROVIDER_START.value,
|
||||
event_status="started",
|
||||
source="celery",
|
||||
user_id=claim.task_snapshot.user_id,
|
||||
group_id=main_task_id,
|
||||
task_id=main_task_id,
|
||||
detail={"generation_count": generation_count},
|
||||
)
|
||||
provider_result = await create_image_sync_batch_result_with_engine(
|
||||
claim.task_snapshot,
|
||||
claim.runtime_engine,
|
||||
generation_count=generation_count,
|
||||
)
|
||||
provider_items = _validate_provider_batch(provider_result, generation_count)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_PROVIDER_SUCCESS.value,
|
||||
event_status="success",
|
||||
source="celery",
|
||||
user_id=claim.task_snapshot.user_id,
|
||||
group_id=main_task_id,
|
||||
task_id=main_task_id,
|
||||
detail={
|
||||
"generation_count": generation_count,
|
||||
"result_count": len(provider_items),
|
||||
"image_tokens": int(provider_result.get("image_tokens") or 0),
|
||||
"single_provider_request": True,
|
||||
"fallback_to_single_requests": False,
|
||||
},
|
||||
)
|
||||
except Exception as exc:
|
||||
message = exc.safe_message if isinstance(exc, ImageProviderError) else str(exc)
|
||||
await _fail_claimed_main(
|
||||
db,
|
||||
main_task_id=main_task_id,
|
||||
claim_token=claim.claim_token,
|
||||
error_message=message or "图片批量生成失败",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_PROVIDER_FAILED,
|
||||
exception=exc,
|
||||
)
|
||||
return []
|
||||
|
||||
try:
|
||||
child_ids = await _split_children(
|
||||
db,
|
||||
main_task_id=main_task_id,
|
||||
claim_token=claim.claim_token,
|
||||
provider_result=provider_result,
|
||||
provider_items=provider_items,
|
||||
)
|
||||
except Exception as exc:
|
||||
await _fail_claimed_main(
|
||||
db,
|
||||
main_task_id=main_task_id,
|
||||
claim_token=claim.claim_token,
|
||||
error_message=f"图片批量结果拆分失败: {exc}",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_FAILED,
|
||||
exception=exc,
|
||||
)
|
||||
return []
|
||||
|
||||
# child 已提交后,下载投递失败不属于图片生成失败,不退款、不重新请求供应商。
|
||||
enqueue_result = await _enqueue_child_downloads(db, child_ids)
|
||||
if enqueue_result["failed"]:
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="DOWNLOAD_ENQUEUE_FAILED",
|
||||
event_status="failed",
|
||||
source="celery",
|
||||
user_id=claim.task_snapshot.user_id,
|
||||
group_id=main_task_id,
|
||||
task_id=main_task_id,
|
||||
detail={
|
||||
"failed_child_task_ids": enqueue_result["failed"],
|
||||
"enqueued_child_task_ids": enqueue_result["enqueued"],
|
||||
"provider_regenerated": False,
|
||||
"generation_refunded": False,
|
||||
},
|
||||
)
|
||||
return child_ids
|
||||
+140
-358
@@ -1,77 +1,54 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime, timedelta, timezone, date
|
||||
from datetime import datetime, date
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import and_, func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.config import settings
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.models.project import Project
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.models.user import User
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.enums.audio_reference import (
|
||||
AUDIO_ALLOWED_EXTENSIONS,
|
||||
AUDIO_MAX_COUNT_LIMIT,
|
||||
AUDIO_MAX_DURATION_SECONDS,
|
||||
AUDIO_MAX_TOTAL_DURATION_SECONDS,
|
||||
AUDIO_MIN_DURATION_SECONDS,
|
||||
)
|
||||
from app.enums.generation_task import CHAT_TOP_LEVEL_MODES, GenerationMode
|
||||
from app.enums.generation_history import (
|
||||
GenerationHistorySourceEnum,
|
||||
get_generation_history_source_label,
|
||||
get_generation_history_task_mode,
|
||||
get_generation_history_task_modes,
|
||||
normalize_generation_history_source,
|
||||
HISTORY_DAY_PAGE_SIZE_MAX,
|
||||
HISTORY_GROUP_ITEM_LIMIT,
|
||||
)
|
||||
from app.schemas.generation_ai import (
|
||||
GenerationAIEngineGroupOut,
|
||||
GenerationAIEngineOptionsOut,
|
||||
GenerationAIImageEngineOptionOut,
|
||||
GenerationAIRecordHistoryItemOut,
|
||||
GenerationAITaskCreate,
|
||||
GenerationAITaskOut,
|
||||
GenerationAIVideoEngineOptionOut,
|
||||
)
|
||||
from app.services.generation_billing_service import (
|
||||
OWNER_CHAT_GENERATION_TASK,
|
||||
charge_generation_media_by_params,
|
||||
)
|
||||
from app.services.resource_accounting_service import (
|
||||
SOURCE_MODEL_CHAT_TASK,
|
||||
SOURCE_MODEL_GENERATION_RECORD,
|
||||
batch_get_generated_resource_info_map,
|
||||
soft_delete_chat_task_resources,
|
||||
)
|
||||
from app.services.resource_signed_url_service import build_resource_signed_url
|
||||
from app.services.generation_history_meta_service import (
|
||||
from app.services.generation.history_meta_service import (
|
||||
GenerationHistoryMeta,
|
||||
batch_load_generation_history_meta_map,
|
||||
build_empty_history_meta,
|
||||
)
|
||||
from app.services.resource_capacity_service import assert_user_resource_capacity_available
|
||||
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls, resolve_private_portrait_reference_display_urls, resolve_private_portrait_references
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
IMAGE_DEFAULT_SIZE = "2K"
|
||||
IMAGE_DEFAULT_PROPORTION = "1:1"
|
||||
IMAGE_DEFAULT_PX = "2048x2048"
|
||||
VIDEO_DEFAULT_DURATION = 4
|
||||
VIDEO_DEFAULT_RATIO = "16:9"
|
||||
VIDEO_DEFAULT_RESOLUTION = "480p"
|
||||
|
||||
HISTORY_DAY_PAGE_SIZE_MAX = 10
|
||||
HISTORY_GROUP_ITEM_LIMIT = 10
|
||||
|
||||
def normalize_px(value: str | None) -> str | None:
|
||||
if not value:
|
||||
return value
|
||||
return value.replace("×", "x").replace("X", "x").replace("×x", "x").replace("x×", "x")
|
||||
|
||||
from app.services.generation.ai.task_group_service import get_display_status, load_children_map
|
||||
from app.services.generation.ai.engine_service import (
|
||||
image_supported_sizes,
|
||||
normalize_generation_count,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls
|
||||
|
||||
def _json(data: Any) -> str | None:
|
||||
if data is None:
|
||||
@@ -88,7 +65,12 @@ def _parse_json(text: str | None):
|
||||
return None
|
||||
|
||||
|
||||
async def _resolve_task_reference_display_map(db: AsyncSession, tasks: list[ChatGenerationTask], *, user_id: str | None = None) -> dict[str, list[dict] | None]:
|
||||
async def _resolve_task_reference_display_map(
|
||||
db: AsyncSession,
|
||||
tasks: list[ChatGenerationTask],
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
) -> dict[str, list[dict] | None]:
|
||||
return await batch_resolve_private_portrait_reference_display_urls(
|
||||
db,
|
||||
{task.id: _parse_json(task.media_references) for task in tasks},
|
||||
@@ -96,7 +78,12 @@ async def _resolve_task_reference_display_map(db: AsyncSession, tasks: list[Chat
|
||||
)
|
||||
|
||||
|
||||
async def _resolve_generation_record_reference_display_map(db: AsyncSession, records: list[GenerationRecord], *, user_id: str | None = None) -> dict[str, list[dict] | None]:
|
||||
async def _resolve_generation_record_reference_display_map(
|
||||
db: AsyncSession,
|
||||
records: list[GenerationRecord],
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
) -> dict[str, list[dict] | None]:
|
||||
return await batch_resolve_private_portrait_reference_display_urls(
|
||||
db,
|
||||
{record.id: _parse_json(record.media_references) for record in records},
|
||||
@@ -104,90 +91,6 @@ async def _resolve_generation_record_reference_display_map(db: AsyncSession, rec
|
||||
)
|
||||
|
||||
|
||||
async def _get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngine:
|
||||
query = select(ImageEngine).where(ImageEngine.is_active == True)
|
||||
if engine_id:
|
||||
query = query.where(ImageEngine.id == engine_id)
|
||||
else:
|
||||
query = query.order_by(ImageEngine.priority.desc()).limit(1)
|
||||
result = await db.execute(query)
|
||||
engine = result.scalar_one_or_none()
|
||||
if not engine:
|
||||
raise HTTPException(status_code=400, detail="没有可用的图片引擎")
|
||||
return engine
|
||||
|
||||
|
||||
async def _get_video_engine(db: AsyncSession, engine_id: str | None) -> VideoEngine:
|
||||
query = select(VideoEngine).where(VideoEngine.is_active == True)
|
||||
if engine_id:
|
||||
query = query.where(VideoEngine.id == engine_id)
|
||||
else:
|
||||
query = query.order_by(VideoEngine.priority.desc())
|
||||
|
||||
query = query.limit(1)
|
||||
result = await db.execute(query)
|
||||
engine = result.scalar_one_or_none()
|
||||
if not engine:
|
||||
raise HTTPException(status_code=400, detail="没有可用的视频引擎")
|
||||
return engine
|
||||
|
||||
|
||||
def _image_supported_sizes(engine: ImageEngine) -> dict:
|
||||
try:
|
||||
data = json.loads(engine.supported_sizes or "{}")
|
||||
return data if isinstance(data, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def _parse_list(value: str | None, fallback: list):
|
||||
try:
|
||||
parsed = json.loads(value or "")
|
||||
return parsed if isinstance(parsed, list) else fallback
|
||||
except Exception:
|
||||
return fallback
|
||||
|
||||
|
||||
def _build_image_snapshot(engine: ImageEngine, size: str, proportion: str, px: str) -> dict:
|
||||
return {
|
||||
"engine_type": "image",
|
||||
"id": engine.id,
|
||||
"name": engine.name,
|
||||
"provider": engine.provider,
|
||||
"api_base": engine.api_base,
|
||||
"api_key_masked": "****" if engine.api_key else "",
|
||||
"model_name": engine.model_name,
|
||||
"generate_url": engine.generate_url,
|
||||
"supported_models": _parse_list(engine.supported_models, []),
|
||||
"default_size": engine.default_size,
|
||||
"selected_size": size,
|
||||
"selected_proportion": proportion,
|
||||
"selected_px": px,
|
||||
}
|
||||
|
||||
|
||||
def _build_video_snapshot(engine: VideoEngine, ratio: str, resolution: str, duration: int) -> dict:
|
||||
return {
|
||||
"engine_type": "video",
|
||||
"id": engine.id,
|
||||
"name": engine.name,
|
||||
"provider": engine.provider,
|
||||
"api_base": engine.api_base,
|
||||
"api_key_masked": "****" if engine.api_key else "",
|
||||
"model_name": engine.model_name,
|
||||
"generate_url": engine.generate_url,
|
||||
"query_url": engine.query_url,
|
||||
"supported_ratios": _parse_list(engine.supported_ratios, []),
|
||||
"supported_resolutions": _parse_list(engine.supported_resolutions, []),
|
||||
"supported_durations": _parse_list(engine.supported_durations, []),
|
||||
"max_duration": engine.max_duration,
|
||||
"max_audio_count": engine.max_audio_count,
|
||||
"selected_ratio": ratio,
|
||||
"selected_resolution": resolution,
|
||||
"selected_duration": duration,
|
||||
}
|
||||
|
||||
|
||||
async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEngineOptionsOut:
|
||||
"""获取当前启用的图片/视频生成引擎,供前端创建任务时选择 engine_id。"""
|
||||
image_result = await db.execute(
|
||||
@@ -207,11 +110,15 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
|
||||
name=engine.name,
|
||||
provider=engine.provider,
|
||||
model_name=engine.model_name,
|
||||
supported_models=_parse_list(engine.supported_models, []),
|
||||
supported_sizes=_image_supported_sizes(engine),
|
||||
supported_models=parse_json_list(engine.supported_models, []),
|
||||
supported_sizes=image_supported_sizes(engine),
|
||||
default_size=engine.default_size,
|
||||
priority=engine.priority or 0,
|
||||
max_image_count=engine.max_image_count,
|
||||
multi_generation_enabled=bool(getattr(engine, "multi_generation_enabled", False)),
|
||||
max_generation_count=normalize_generation_count(getattr(engine, "max_generation_count", 1)),
|
||||
multi_image_max_images=int(getattr(engine, "multi_image_max_images", 15) or 15),
|
||||
max_reference_image_count=int(getattr(engine, "max_reference_image_count", 14) or 0),
|
||||
)
|
||||
for engine in image_result.scalars().all()
|
||||
]
|
||||
@@ -221,9 +128,9 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
|
||||
name=engine.name,
|
||||
provider=engine.provider,
|
||||
model_name=engine.model_name,
|
||||
supported_ratios=_parse_list(engine.supported_ratios, []),
|
||||
supported_resolutions=_parse_list(engine.supported_resolutions, []),
|
||||
supported_durations=_parse_list(engine.supported_durations, []),
|
||||
supported_ratios=parse_json_list(engine.supported_ratios, []),
|
||||
supported_resolutions=parse_json_list(engine.supported_resolutions, []),
|
||||
supported_durations=parse_json_list(engine.supported_durations, []),
|
||||
max_duration=engine.max_duration,
|
||||
priority=engine.priority or 0,
|
||||
max_image_count=engine.max_image_count,
|
||||
@@ -231,6 +138,8 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
|
||||
max_audio_count=engine.max_audio_count,
|
||||
supports_first_last_frame=engine.supports_first_last_frame,
|
||||
supports_universal_reference=engine.supports_universal_reference,
|
||||
multi_generation_enabled=bool(getattr(engine, "multi_generation_enabled", False)),
|
||||
max_generation_count=normalize_generation_count(getattr(engine, "max_generation_count", 1)),
|
||||
)
|
||||
for engine in video_result.scalars().all()
|
||||
]
|
||||
@@ -240,193 +149,6 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
|
||||
)
|
||||
|
||||
|
||||
async def create_async_generation_task(db: AsyncSession, current_user: User, req: GenerationAITaskCreate) -> ChatGenerationTask:
|
||||
"""Create a project-independent chat generation task.
|
||||
|
||||
Important: this writes chat_generation_tasks, not generation_records, so chat
|
||||
image/video generation no longer needs or validates a project_id.
|
||||
"""
|
||||
gen_type = req.gen_type.lower().strip()
|
||||
if gen_type not in ("image", "video"):
|
||||
raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video")
|
||||
|
||||
if req.idempotency_key:
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.user_id == current_user.id,
|
||||
ChatGenerationTask.idempotency_key == req.idempotency_key,
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
).order_by(ChatGenerationTask.created_at.desc()).limit(1)
|
||||
)
|
||||
existing = result.scalar_one_or_none()
|
||||
if existing:
|
||||
return existing
|
||||
|
||||
refs = [r.model_dump(exclude_none=True) for r in (req.media_references or [])]
|
||||
refs = await resolve_private_portrait_references(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
media_references=refs,
|
||||
gen_type=gen_type,
|
||||
)
|
||||
now = datetime.now(timezone.utc)
|
||||
task_id = generate_id()
|
||||
|
||||
await assert_user_resource_capacity_available(db, current_user.id)
|
||||
|
||||
if gen_type == "image":
|
||||
if any((r.get("type") or "").lower() == "audio" for r in refs):
|
||||
raise HTTPException(status_code=400, detail="图片生成不支持音频参考素材")
|
||||
engine = await _get_image_engine(db, req.engine_id)
|
||||
sizes = _image_supported_sizes(engine)
|
||||
size = req.image_size or engine.default_size or IMAGE_DEFAULT_SIZE
|
||||
proportion = req.image_proportion or IMAGE_DEFAULT_PROPORTION
|
||||
px = normalize_px(req.image_px)
|
||||
if sizes:
|
||||
if size not in sizes:
|
||||
raise HTTPException(status_code=400, detail=f"图片分辨率档位不支持: {size}")
|
||||
if proportion not in sizes.get(size, {}):
|
||||
raise HTTPException(status_code=400, detail=f"图片比例不支持: {proportion}")
|
||||
px = px or normalize_px((sizes.get(size) or {}).get(proportion))
|
||||
px = px or IMAGE_DEFAULT_PX
|
||||
media_billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=task_id,
|
||||
gen_type="image",
|
||||
image_size=size,
|
||||
engine_id=engine.id,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="AI创作-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
)
|
||||
snapshot = _build_image_snapshot(engine, size, proportion, px)
|
||||
task = ChatGenerationTask(
|
||||
id=task_id,
|
||||
user_id=current_user.id,
|
||||
original_prompt=req.original_prompt,
|
||||
gen_type="image",
|
||||
image_size=size,
|
||||
image_proportion=proportion,
|
||||
image_px=px,
|
||||
status="generating",
|
||||
generation_mode="chatapi_async",
|
||||
pipeline_stage="queued",
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=_json(snapshot),
|
||||
media_references=_json(refs) if refs else None,
|
||||
credits_cost=round(media_billing.total_charged, 2),
|
||||
idempotency_key=req.idempotency_key,
|
||||
deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES),
|
||||
)
|
||||
else:
|
||||
engine = await _get_video_engine(db, req.engine_id)
|
||||
ratio = req.aspect_ratio or VIDEO_DEFAULT_RATIO
|
||||
resolution = req.resolution or VIDEO_DEFAULT_RESOLUTION
|
||||
duration = req.duration or VIDEO_DEFAULT_DURATION
|
||||
ratios = _parse_list(engine.supported_ratios, [])
|
||||
resolutions = _parse_list(engine.supported_resolutions, [])
|
||||
durations = _parse_list(engine.supported_durations, [])
|
||||
if ratios and ratio not in ratios:
|
||||
raise HTTPException(status_code=400, detail=f"视频比例不支持: {ratio}")
|
||||
if resolutions and resolution not in resolutions:
|
||||
raise HTTPException(status_code=400, detail=f"视频分辨率不支持: {resolution}")
|
||||
if durations and duration not in durations:
|
||||
raise HTTPException(status_code=400, detail=f"视频时长不支持: {duration}")
|
||||
if engine.max_duration and duration > engine.max_duration:
|
||||
raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration} 秒")
|
||||
|
||||
input_video_duration = 0.0
|
||||
if refs:
|
||||
video_refs = [r for r in refs if (r.get("type") or "").lower() == "video"]
|
||||
for ref in video_refs:
|
||||
ref_duration = float(ref.get("duration") or 0)
|
||||
if ref_duration < 2:
|
||||
raise HTTPException(status_code=400, detail=f"视频素材最短不能少于 2 秒")
|
||||
input_video_duration += ref_duration
|
||||
if input_video_duration > 15:
|
||||
raise HTTPException(status_code=400, detail=f"所有视频素材总时长不能超过 15 秒,当前 {input_video_duration:.1f} 秒")
|
||||
|
||||
audio_refs = [r for r in refs if (r.get("type") or "").lower() == "audio"]
|
||||
if audio_refs:
|
||||
max_audio_count = int(engine.max_audio_count or 0)
|
||||
if max_audio_count <= 0:
|
||||
raise HTTPException(status_code=400, detail="当前视频引擎不支持音频参考素材")
|
||||
if max_audio_count > AUDIO_MAX_COUNT_LIMIT:
|
||||
max_audio_count = AUDIO_MAX_COUNT_LIMIT
|
||||
if len(audio_refs) > max_audio_count:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"参考音频最多可传 {max_audio_count} 段,当前 {len(audio_refs)} 段",
|
||||
)
|
||||
|
||||
input_audio_duration = 0.0
|
||||
for ref in audio_refs:
|
||||
raw_duration = ref.get("duration")
|
||||
if raw_duration is None:
|
||||
raw_duration = 0.0
|
||||
try:
|
||||
ref_duration = float(raw_duration)
|
||||
except (TypeError, ValueError):
|
||||
ref_duration = 0.0
|
||||
|
||||
if ref_duration < AUDIO_MIN_DURATION_SECONDS or ref_duration > AUDIO_MAX_DURATION_SECONDS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"单段参考音频时长必须在 {AUDIO_MIN_DURATION_SECONDS}-{AUDIO_MAX_DURATION_SECONDS} 秒之间",
|
||||
)
|
||||
input_audio_duration += ref_duration
|
||||
|
||||
if input_audio_duration > AUDIO_MAX_TOTAL_DURATION_SECONDS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"所有参考音频总时长不能超过 {AUDIO_MAX_TOTAL_DURATION_SECONDS} 秒,当前 {input_audio_duration:.1f} 秒",
|
||||
)
|
||||
|
||||
media_billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=task_id,
|
||||
gen_type="video",
|
||||
duration=duration,
|
||||
resolution=resolution,
|
||||
engine_id=engine.id,
|
||||
input_video_duration=input_video_duration if input_video_duration > 0 else None,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="AI创作-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
)
|
||||
snapshot = _build_video_snapshot(engine, ratio, resolution, duration)
|
||||
task = ChatGenerationTask(
|
||||
id=task_id,
|
||||
user_id=current_user.id,
|
||||
original_prompt=req.original_prompt,
|
||||
gen_type="video",
|
||||
duration=duration,
|
||||
aspect_ratio=ratio,
|
||||
resolution=resolution,
|
||||
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
|
||||
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
|
||||
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
|
||||
status="generating",
|
||||
generation_mode="chatapi_async",
|
||||
pipeline_stage="queued",
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=_json(snapshot),
|
||||
media_references=_json(refs) if refs else None,
|
||||
credits_cost=round(media_billing.total_charged, 2),
|
||||
idempotency_key=req.idempotency_key,
|
||||
deadline_at=now + timedelta(hours=settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS),
|
||||
)
|
||||
|
||||
db.add(task)
|
||||
await db.flush()
|
||||
return task
|
||||
|
||||
|
||||
def _resolve_error_message(error_message: str | None) -> str | None:
|
||||
"""匹配 ARK_ERRORS 字典,将原始错误码转换为友好提示。
|
||||
与 app/api/v1/generation.py 的 _record_to_out 保持一致。
|
||||
@@ -479,26 +201,36 @@ def record_to_out(
|
||||
file_name: str | None = None,
|
||||
history_meta: GenerationHistoryMeta | None = None,
|
||||
media_references: list[dict] | None = None,
|
||||
child_items: list[GenerationAITaskOut] | None = None,
|
||||
) -> GenerationAITaskOut:
|
||||
refs = media_references if media_references is not None else _parse_json(task.media_references)
|
||||
snapshot = engine_snapshot_out(_parse_json(task.engine_snapshot_json))
|
||||
|
||||
source = GenerationHistorySourceEnum.CHAT_TASK
|
||||
try:
|
||||
source = GenerationHistorySourceEnum(
|
||||
"chat_task" if task.generation_mode == "chatapi_async" else str(task.generation_mode or "chat_task")
|
||||
)
|
||||
if task.generation_mode in {
|
||||
GenerationMode.CHATAPI_ASYNC.value,
|
||||
GenerationMode.CHATAPI_MAIN.value,
|
||||
GenerationMode.CHATAPI_CHILD.value,
|
||||
}:
|
||||
source = GenerationHistorySourceEnum.CHAT_TASK
|
||||
else:
|
||||
source = GenerationHistorySourceEnum(str(task.generation_mode or "chat_task"))
|
||||
except ValueError:
|
||||
source = GenerationHistorySourceEnum.CHAT_TASK
|
||||
meta = history_meta or build_empty_history_meta(source)
|
||||
|
||||
is_deleted = task.deleted_at is not None
|
||||
is_main = task.generation_mode == GenerationMode.CHATAPI_MAIN.value
|
||||
hide_resource = is_deleted or is_main
|
||||
|
||||
return GenerationAITaskOut(
|
||||
id=task.id,
|
||||
user_id=task.user_id if is_admin else None,
|
||||
user_name=getattr(task, "username", None) if is_admin else None,
|
||||
project_id=None,
|
||||
generated_resource_id=generated_resource_id,
|
||||
file_name=file_name,
|
||||
generated_resource_id=None if hide_resource else generated_resource_id,
|
||||
file_name=None if hide_resource else file_name,
|
||||
history_source=meta.get("history_source"),
|
||||
history_source_label=meta.get("history_source_label"),
|
||||
module_project_id=meta.get("module_project_id"),
|
||||
@@ -515,10 +247,14 @@ def record_to_out(
|
||||
shot_segment_label=meta.get("shot_segment_label"),
|
||||
gen_type=task.gen_type,
|
||||
generation_mode=task.generation_mode,
|
||||
parent_task_id=task.parent_task_id,
|
||||
generation_count=max(1, min(5, int(task.generation_count or 1))),
|
||||
generation_index=task.generation_index,
|
||||
display_status=get_display_status(task),
|
||||
pipeline_stage=task.pipeline_stage,
|
||||
status=task.status,
|
||||
original_prompt=task.original_prompt,
|
||||
# optimized_prompt=task.optimized_prompt,
|
||||
optimized_prompt=task.optimized_prompt,
|
||||
duration=task.duration,
|
||||
aspect_ratio=task.aspect_ratio,
|
||||
resolution=task.resolution,
|
||||
@@ -528,10 +264,9 @@ def record_to_out(
|
||||
media_references=refs,
|
||||
provider_task_id=task.provider_task_id,
|
||||
seedance_task_id=task.seedance_task_id,
|
||||
# remote_result_url=task.remote_result_url,
|
||||
image_url=build_resource_signed_url(task.image_url) if task.image_url else "",
|
||||
video_url=build_resource_signed_url(task.video_url) if task.video_url else "",
|
||||
video_cover_url=build_resource_signed_url(task.video_cover_url) if task.video_cover_url else "",
|
||||
image_url="" if hide_resource else (build_resource_signed_url(task.image_url) if task.image_url else ""),
|
||||
video_url="" if hide_resource else (build_resource_signed_url(task.video_url) if task.video_url else ""),
|
||||
video_cover_url="" if hide_resource else (build_resource_signed_url(task.video_cover_url) if task.video_cover_url else ""),
|
||||
engine_id=task.engine_id,
|
||||
engine_snapshot=snapshot,
|
||||
credits_cost=task.credits_cost or 0.0,
|
||||
@@ -541,33 +276,90 @@ def record_to_out(
|
||||
video_tokens_used=task.video_tokens_used or 0,
|
||||
retry_count=task.retry_count or 0,
|
||||
poll_count=task.poll_count or 0,
|
||||
error_message=_resolve_error_message(task.error_message),
|
||||
error_message=task.error_message if is_main else _resolve_error_message(task.error_message),
|
||||
created_at=task.created_at,
|
||||
generated_at=task.generated_at,
|
||||
child_items=child_items or [],
|
||||
)
|
||||
|
||||
def engine_snapshot_out(snapshot: dict) -> dict:
|
||||
"""
|
||||
从完整的 engine_snapshot 中过滤出需要返回的字段
|
||||
"""
|
||||
"""从完整引擎快照中过滤前端允许展示的字段。"""
|
||||
if not snapshot:
|
||||
return {}
|
||||
keys = (
|
||||
"engine_type", "id", "name", "provider", "model_name",
|
||||
"supported_models", "default_size", "selected_size",
|
||||
"selected_proportion", "selected_px", "supported_ratios",
|
||||
"supported_resolutions", "supported_durations", "max_duration",
|
||||
"max_audio_count", "selected_ratio", "selected_resolution",
|
||||
"selected_duration", "generation_count", "multi_generation_enabled",
|
||||
"max_generation_count", "multi_image_max_images", "max_reference_image_count", "output_format",
|
||||
)
|
||||
result = {key: snapshot.get(key) for key in keys if key in snapshot}
|
||||
result.setdefault("generation_count", 1)
|
||||
return result
|
||||
|
||||
return {
|
||||
"engine_type": snapshot.get("engine_type"),
|
||||
"id": snapshot.get("id"),
|
||||
"name": snapshot.get("name"),
|
||||
"provider": snapshot.get("provider"),
|
||||
# "api_base": snapshot.get("api_base"),
|
||||
# "api_key_masked": snapshot.get("api_key_masked"),
|
||||
"model_name": snapshot.get("model_name"),
|
||||
# "generate_url": snapshot.get("generate_url"),
|
||||
"supported_models": snapshot.get("supported_models", []),
|
||||
"default_size": snapshot.get("default_size"),
|
||||
"selected_size": snapshot.get("selected_size"),
|
||||
"selected_proportion": snapshot.get("selected_proportion"),
|
||||
"selected_px": snapshot.get("selected_px")
|
||||
}
|
||||
|
||||
async def build_task_out_list(
|
||||
db: AsyncSession,
|
||||
tasks: list[ChatGenerationTask],
|
||||
*,
|
||||
is_admin: bool = False,
|
||||
viewer_user_id: str | None = None,
|
||||
) -> list[GenerationAITaskOut]:
|
||||
"""批量回填主任务子项、资源账本和参考素材,避免列表 N+1。"""
|
||||
if not tasks:
|
||||
return []
|
||||
parent_ids = [
|
||||
task.id for task in tasks
|
||||
if task.generation_mode == GenerationMode.CHATAPI_MAIN.value
|
||||
]
|
||||
children_map = await load_children_map(db, parent_ids, include_deleted=True)
|
||||
children = [child for items in children_map.values() for child in items]
|
||||
resource_task_ids = [
|
||||
task.id for task in [*tasks, *children]
|
||||
if task.generation_mode != GenerationMode.CHATAPI_MAIN.value and task.deleted_at is None
|
||||
]
|
||||
resource_info_map = await batch_get_generated_resource_info_map(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=resource_task_ids,
|
||||
)
|
||||
reference_display_map = await _resolve_task_reference_display_map(
|
||||
db,
|
||||
tasks,
|
||||
user_id=viewer_user_id,
|
||||
)
|
||||
|
||||
output: list[GenerationAITaskOut] = []
|
||||
for task in tasks:
|
||||
refs = reference_display_map.get(task.id)
|
||||
child_out: list[GenerationAITaskOut] = []
|
||||
for child in children_map.get(task.id, []):
|
||||
if is_admin:
|
||||
child.username = getattr(task, "username", None)
|
||||
resource = resource_info_map.get(child.id, {})
|
||||
child_out.append(
|
||||
record_to_out(
|
||||
child,
|
||||
is_admin=is_admin,
|
||||
generated_resource_id=resource.get("resource_id"),
|
||||
file_name=resource.get("file_name"),
|
||||
media_references=refs,
|
||||
)
|
||||
)
|
||||
resource = resource_info_map.get(task.id, {})
|
||||
output.append(
|
||||
record_to_out(
|
||||
task,
|
||||
is_admin=is_admin,
|
||||
generated_resource_id=resource.get("resource_id"),
|
||||
file_name=resource.get("file_name"),
|
||||
media_references=refs,
|
||||
child_items=child_out,
|
||||
)
|
||||
)
|
||||
return output
|
||||
|
||||
async def list_async_generation_tasks(
|
||||
db: AsyncSession,
|
||||
@@ -594,7 +386,7 @@ async def list_async_generation_tasks(
|
||||
query = select(ChatGenerationTask)
|
||||
|
||||
query = query.where(
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.generation_mode.in_(list(CHAT_TOP_LEVEL_MODES)),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
|
||||
@@ -623,7 +415,7 @@ async def list_async_generation_tasks(
|
||||
total = (await db.execute(count_query)).scalar_one()
|
||||
|
||||
result = await db.execute(
|
||||
query.order_by(ChatGenerationTask.created_at.desc())
|
||||
query.order_by(ChatGenerationTask.created_at.desc(), ChatGenerationTask.id.desc())
|
||||
.offset((page - 1) * page_size)
|
||||
.limit(page_size)
|
||||
)
|
||||
@@ -677,12 +469,12 @@ def _parse_history_date(value: str) -> date:
|
||||
|
||||
|
||||
def _history_base_filters(user_id: str, gen_type: str, source: GenerationHistorySourceEnum):
|
||||
task_mode = get_generation_history_task_mode(source)
|
||||
if not task_mode:
|
||||
task_modes = get_generation_history_task_modes(source)
|
||||
if not task_modes:
|
||||
raise HTTPException(status_code=400, detail="history_source 不支持查询 ChatGenerationTask 历史")
|
||||
return [
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode == task_mode.value,
|
||||
ChatGenerationTask.generation_mode.in_([mode.value for mode in task_modes]),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.status == "completed",
|
||||
ChatGenerationTask.gen_type == gen_type,
|
||||
@@ -1193,13 +985,3 @@ async def list_generation_history_day_items(
|
||||
],
|
||||
}
|
||||
|
||||
async def soft_delete_chat_generation_task(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task: ChatGenerationTask,
|
||||
deleted_at: datetime | None = None,
|
||||
) -> int:
|
||||
"""软删 ChatGenerationTask 并联动软删资源账本,返回释放的 active 空间字节数。"""
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
task.deleted_at = deleted_at
|
||||
return await soft_delete_chat_task_resources(db, task.id, deleted_at=deleted_at)
|
||||
@@ -0,0 +1,587 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.audio_reference import (
|
||||
AUDIO_MAX_COUNT_LIMIT,
|
||||
AUDIO_MAX_DURATION_SECONDS,
|
||||
AUDIO_MAX_TOTAL_DURATION_SECONDS,
|
||||
AUDIO_MIN_DURATION_SECONDS,
|
||||
)
|
||||
from app.enums.generation_task import CHAT_TOP_LEVEL_MODES, GenerationMode, GenerationType
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.user import User
|
||||
from app.schemas.generation_ai import GenerationAITaskCreate
|
||||
from app.services.generation.ai.engine_service import (
|
||||
IMAGE_DEFAULT_PROPORTION,
|
||||
IMAGE_DEFAULT_PX,
|
||||
IMAGE_DEFAULT_SIZE,
|
||||
VIDEO_DEFAULT_DURATION,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
build_image_snapshot,
|
||||
build_video_snapshot,
|
||||
get_image_engine,
|
||||
get_video_engine,
|
||||
image_supported_sizes,
|
||||
normalize_generation_count,
|
||||
normalize_px,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.generation.billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params
|
||||
from app.services.operation_log_service import log_operation_event
|
||||
from app.services.private_portrait.reference_resolver import resolve_private_portrait_references
|
||||
from app.services.resource_capacity_service import assert_user_resource_capacity_available
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class GenerationTaskCreateResult:
|
||||
top_level_task_id: str
|
||||
enqueue_task_ids: list[str] = field(default_factory=list)
|
||||
child_task_ids: list[str] = field(default_factory=list)
|
||||
generation_count: int = 1
|
||||
gen_type: str = GenerationType.IMAGE.value
|
||||
created: bool = True
|
||||
|
||||
|
||||
def _json(data: Any) -> str | None:
|
||||
if data is None:
|
||||
return None
|
||||
return json.dumps(data, ensure_ascii=False, default=str)
|
||||
|
||||
|
||||
async def find_existing_top_level_task(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
user_id: str,
|
||||
idempotency_key: str | None,
|
||||
) -> ChatGenerationTask | None:
|
||||
if not idempotency_key:
|
||||
return None
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.idempotency_key == idempotency_key,
|
||||
ChatGenerationTask.generation_mode.in_(list(CHAT_TOP_LEVEL_MODES)),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.order_by(ChatGenerationTask.created_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
def _validate_video_references(refs: list[dict], *, max_audio_count: int) -> float:
|
||||
input_video_duration = 0.0
|
||||
for ref in refs:
|
||||
if (ref.get("type") or "").lower() != GenerationType.VIDEO.value:
|
||||
continue
|
||||
try:
|
||||
ref_duration = float(ref.get("duration") or 0)
|
||||
except (TypeError, ValueError):
|
||||
ref_duration = 0.0
|
||||
if ref_duration < 2:
|
||||
raise HTTPException(status_code=400, detail="视频素材最短不能少于 2 秒")
|
||||
input_video_duration += ref_duration
|
||||
if input_video_duration > 15:
|
||||
raise HTTPException(status_code=400, detail=f"所有视频素材总时长不能超过 15 秒,当前 {input_video_duration:.1f} 秒")
|
||||
|
||||
audio_refs = [ref for ref in refs if (ref.get("type") or "").lower() == "audio"]
|
||||
if audio_refs:
|
||||
allowed_count = min(AUDIO_MAX_COUNT_LIMIT, max(0, int(max_audio_count or 0)))
|
||||
if allowed_count <= 0:
|
||||
raise HTTPException(status_code=400, detail="当前视频引擎不支持音频参考素材")
|
||||
if len(audio_refs) > allowed_count:
|
||||
raise HTTPException(status_code=400, detail=f"参考音频最多可传 {allowed_count} 段,当前 {len(audio_refs)} 段")
|
||||
|
||||
input_audio_duration = 0.0
|
||||
for ref in audio_refs:
|
||||
try:
|
||||
ref_duration = float(ref.get("duration") or 0)
|
||||
except (TypeError, ValueError):
|
||||
ref_duration = 0.0
|
||||
if ref_duration < AUDIO_MIN_DURATION_SECONDS or ref_duration > AUDIO_MAX_DURATION_SECONDS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"单段参考音频时长必须在 {AUDIO_MIN_DURATION_SECONDS}-{AUDIO_MAX_DURATION_SECONDS} 秒之间",
|
||||
)
|
||||
input_audio_duration += ref_duration
|
||||
if input_audio_duration > AUDIO_MAX_TOTAL_DURATION_SECONDS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"所有参考音频总时长不能超过 {AUDIO_MAX_TOTAL_DURATION_SECONDS} 秒,当前 {input_audio_duration:.1f} 秒",
|
||||
)
|
||||
return input_video_duration
|
||||
|
||||
|
||||
def _base_task_kwargs(
|
||||
*,
|
||||
task_id: str,
|
||||
user_id: str,
|
||||
req: GenerationAITaskCreate,
|
||||
gen_type: str,
|
||||
generation_mode: str,
|
||||
generation_count: int,
|
||||
engine_id: str,
|
||||
engine_snapshot_json: str,
|
||||
media_references_json: str | None,
|
||||
deadline_at: datetime,
|
||||
parent_task_id: str | None = None,
|
||||
generation_index: int | None = None,
|
||||
credits_cost: float = 0.0,
|
||||
idempotency_key: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"id": task_id,
|
||||
"user_id": user_id,
|
||||
"original_prompt": req.original_prompt,
|
||||
"gen_type": gen_type,
|
||||
"status": "generating",
|
||||
"generation_mode": generation_mode,
|
||||
"pipeline_stage": "queued",
|
||||
"parent_task_id": parent_task_id,
|
||||
"generation_count": generation_count,
|
||||
"generation_index": generation_index,
|
||||
"engine_id": engine_id,
|
||||
"engine_snapshot_json": engine_snapshot_json,
|
||||
"media_references": media_references_json,
|
||||
"credits_cost": round(float(credits_cost or 0), 2),
|
||||
"idempotency_key": idempotency_key,
|
||||
"deadline_at": deadline_at,
|
||||
}
|
||||
|
||||
|
||||
async def create_generation_task_group(
|
||||
db: AsyncSession,
|
||||
current_user: User,
|
||||
req: GenerationAITaskCreate,
|
||||
) -> GenerationTaskCreateResult:
|
||||
"""创建单份 chatapi_async 或多份 chatapi_main/chatapi_child 任务组。
|
||||
|
||||
本函数只 flush,不主动 commit。调用方提交成功后才能投递 Celery。
|
||||
"""
|
||||
gen_type = (req.gen_type or "").lower().strip()
|
||||
if gen_type not in (GenerationType.IMAGE.value, GenerationType.VIDEO.value):
|
||||
raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video")
|
||||
|
||||
existing = await find_existing_top_level_task(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
idempotency_key=req.idempotency_key,
|
||||
)
|
||||
if existing:
|
||||
return GenerationTaskCreateResult(
|
||||
top_level_task_id=existing.id,
|
||||
generation_count=int(existing.generation_count or 1),
|
||||
gen_type=existing.gen_type,
|
||||
created=False,
|
||||
)
|
||||
|
||||
refs = [item.model_dump(exclude_none=True) for item in (req.media_references or [])]
|
||||
refs = await resolve_private_portrait_references(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
media_references=refs,
|
||||
gen_type=gen_type,
|
||||
)
|
||||
media_references_json = _json(refs) if refs else None
|
||||
await assert_user_resource_capacity_available(db, current_user.id)
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
main_id = generate_id()
|
||||
child_ids: list[str] = []
|
||||
enqueue_ids: list[str] = []
|
||||
total_billed_credits = 0.0
|
||||
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_CREATE_START",
|
||||
event_status="started",
|
||||
source="service",
|
||||
user_id=current_user.id,
|
||||
group_id=main_id,
|
||||
detail={
|
||||
"gen_type": gen_type,
|
||||
"requested_generation_count": normalize_generation_count(req.generation_count),
|
||||
"idempotency_key_present": bool(req.idempotency_key),
|
||||
},
|
||||
)
|
||||
|
||||
if gen_type == GenerationType.IMAGE.value:
|
||||
if any((ref.get("type") or "").lower() == "audio" for ref in refs):
|
||||
raise HTTPException(status_code=400, detail="图片生成不支持音频参考素材")
|
||||
|
||||
engine = await get_image_engine(db, req.engine_id)
|
||||
generation_count = normalize_generation_count(req.generation_count)
|
||||
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
|
||||
max_generation_count = normalize_generation_count(getattr(engine, "max_generation_count", 1))
|
||||
if generation_count > 1 and not multi_generation_enabled:
|
||||
raise HTTPException(status_code=400, detail="当前图片引擎未开启多份生成,本次生成数量只能为 1")
|
||||
if generation_count > max_generation_count:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"当前图片引擎本次最多允许生成 {max_generation_count} 份",
|
||||
)
|
||||
|
||||
reference_image_count = sum(
|
||||
1 for ref in refs if (ref.get("type") or "").lower() == GenerationType.IMAGE.value
|
||||
)
|
||||
max_reference_count = max(0, int(getattr(engine, "max_reference_image_count", 14) or 0))
|
||||
multi_image_max_images = max(1, int(getattr(engine, "multi_image_max_images", 15) or 15))
|
||||
if reference_image_count > max_reference_count:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"当前图片引擎最多支持 {max_reference_count} 张参考图,当前 {reference_image_count} 张",
|
||||
)
|
||||
if generation_count > 1 and reference_image_count + generation_count > multi_image_max_images:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
f"参考图数量与生成数量合计不能超过 {multi_image_max_images} 张,"
|
||||
f"当前参考图 {reference_image_count} 张、生成 {generation_count} 张"
|
||||
),
|
||||
)
|
||||
|
||||
sizes = image_supported_sizes(engine)
|
||||
size = req.image_size or engine.default_size or IMAGE_DEFAULT_SIZE
|
||||
proportion = req.image_proportion or IMAGE_DEFAULT_PROPORTION
|
||||
px = normalize_px(req.image_px)
|
||||
if sizes:
|
||||
if size not in sizes:
|
||||
raise HTTPException(status_code=400, detail=f"图片分辨率档位不支持: {size}")
|
||||
if proportion not in sizes.get(size, {}):
|
||||
raise HTTPException(status_code=400, detail=f"图片比例不支持: {proportion}")
|
||||
px = px or normalize_px((sizes.get(size) or {}).get(proportion))
|
||||
px = px or IMAGE_DEFAULT_PX
|
||||
|
||||
mode = GenerationMode.CHATAPI_ASYNC.value if generation_count == 1 else GenerationMode.CHATAPI_MAIN.value
|
||||
billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=main_id,
|
||||
gen_type=GenerationType.IMAGE.value,
|
||||
image_size=size,
|
||||
engine_id=engine.id,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="AI创作-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
quantity=generation_count,
|
||||
)
|
||||
image_snapshot = build_image_snapshot(engine, size, proportion, px)
|
||||
image_snapshot["generation_count"] = generation_count
|
||||
snapshot_json = _json(image_snapshot) or "{}"
|
||||
total_billed_credits = round(float(billing.total_charged or 0), 2)
|
||||
task = ChatGenerationTask(
|
||||
**_base_task_kwargs(
|
||||
task_id=main_id,
|
||||
user_id=current_user.id,
|
||||
req=req,
|
||||
gen_type=GenerationType.IMAGE.value,
|
||||
generation_mode=mode,
|
||||
generation_count=generation_count,
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=snapshot_json,
|
||||
media_references_json=media_references_json,
|
||||
deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES),
|
||||
credits_cost=billing.total_charged,
|
||||
idempotency_key=req.idempotency_key,
|
||||
),
|
||||
image_size=size,
|
||||
image_proportion=proportion,
|
||||
image_px=px,
|
||||
)
|
||||
db.add(task)
|
||||
enqueue_ids.append(task.id)
|
||||
else:
|
||||
engine = await get_video_engine(db, req.engine_id)
|
||||
generation_count = normalize_generation_count(req.generation_count)
|
||||
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
|
||||
max_generation_count = normalize_generation_count(getattr(engine, "max_generation_count", 1))
|
||||
if generation_count > 1 and not multi_generation_enabled:
|
||||
raise HTTPException(status_code=400, detail="当前视频引擎未开启多份生成,本次生成数量只能为 1")
|
||||
if generation_count > max_generation_count:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"当前视频引擎本次最多允许生成 {max_generation_count} 份",
|
||||
)
|
||||
ratio = req.aspect_ratio or VIDEO_DEFAULT_RATIO
|
||||
resolution = req.resolution or VIDEO_DEFAULT_RESOLUTION
|
||||
duration = req.duration or VIDEO_DEFAULT_DURATION
|
||||
ratios = parse_json_list(engine.supported_ratios, [])
|
||||
resolutions = parse_json_list(engine.supported_resolutions, [])
|
||||
durations = parse_json_list(engine.supported_durations, [])
|
||||
if ratios and ratio not in ratios:
|
||||
raise HTTPException(status_code=400, detail=f"视频比例不支持: {ratio}")
|
||||
if resolutions and resolution not in resolutions:
|
||||
raise HTTPException(status_code=400, detail=f"视频分辨率不支持: {resolution}")
|
||||
if durations and duration not in durations:
|
||||
raise HTTPException(status_code=400, detail=f"视频时长不支持: {duration}")
|
||||
if engine.max_duration and duration > engine.max_duration:
|
||||
raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration} 秒")
|
||||
|
||||
input_video_duration = _validate_video_references(refs, max_audio_count=engine.max_audio_count)
|
||||
video_snapshot = build_video_snapshot(engine, ratio, resolution, duration)
|
||||
video_snapshot["generation_count"] = generation_count
|
||||
snapshot_json = _json(video_snapshot) or "{}"
|
||||
deadline_at = now + timedelta(hours=settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS)
|
||||
|
||||
if generation_count == 1:
|
||||
billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=main_id,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
duration=duration,
|
||||
resolution=resolution,
|
||||
engine_id=engine.id,
|
||||
input_video_duration=input_video_duration if input_video_duration > 0 else None,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="AI创作-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
)
|
||||
total_billed_credits = round(float(billing.total_charged or 0), 2)
|
||||
task = ChatGenerationTask(
|
||||
**_base_task_kwargs(
|
||||
task_id=main_id,
|
||||
user_id=current_user.id,
|
||||
req=req,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
generation_mode=GenerationMode.CHATAPI_ASYNC.value,
|
||||
generation_count=1,
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=snapshot_json,
|
||||
media_references_json=media_references_json,
|
||||
deadline_at=deadline_at,
|
||||
credits_cost=billing.total_charged,
|
||||
idempotency_key=req.idempotency_key,
|
||||
),
|
||||
duration=duration,
|
||||
aspect_ratio=ratio,
|
||||
resolution=resolution,
|
||||
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
|
||||
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
|
||||
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
|
||||
)
|
||||
db.add(task)
|
||||
enqueue_ids.append(task.id)
|
||||
else:
|
||||
main_task = ChatGenerationTask(
|
||||
**_base_task_kwargs(
|
||||
task_id=main_id,
|
||||
user_id=current_user.id,
|
||||
req=req,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
generation_mode=GenerationMode.CHATAPI_MAIN.value,
|
||||
generation_count=generation_count,
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=snapshot_json,
|
||||
media_references_json=media_references_json,
|
||||
deadline_at=deadline_at,
|
||||
idempotency_key=req.idempotency_key,
|
||||
),
|
||||
duration=duration,
|
||||
aspect_ratio=ratio,
|
||||
resolution=resolution,
|
||||
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
|
||||
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
|
||||
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
|
||||
)
|
||||
db.add(main_task)
|
||||
await db.flush()
|
||||
|
||||
total_credits = 0.0
|
||||
children: list[ChatGenerationTask] = []
|
||||
for generation_index in range(1, generation_count + 1):
|
||||
child_id = generate_id()
|
||||
billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=child_id,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
duration=duration,
|
||||
resolution=resolution,
|
||||
engine_id=engine.id,
|
||||
input_video_duration=input_video_duration if input_video_duration > 0 else None,
|
||||
project_name="AI生成任务",
|
||||
description_prefix=f"AI创作-第{generation_index}份-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
)
|
||||
child = ChatGenerationTask(
|
||||
**_base_task_kwargs(
|
||||
task_id=child_id,
|
||||
user_id=current_user.id,
|
||||
req=req,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
generation_mode=GenerationMode.CHATAPI_CHILD.value,
|
||||
generation_count=generation_count,
|
||||
generation_index=generation_index,
|
||||
parent_task_id=main_id,
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=snapshot_json,
|
||||
media_references_json=media_references_json,
|
||||
deadline_at=deadline_at,
|
||||
credits_cost=billing.total_charged,
|
||||
),
|
||||
duration=duration,
|
||||
aspect_ratio=ratio,
|
||||
resolution=resolution,
|
||||
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
|
||||
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
|
||||
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
|
||||
)
|
||||
children.append(child)
|
||||
child_ids.append(child_id)
|
||||
enqueue_ids.append(child_id)
|
||||
total_credits = round(total_credits + billing.total_charged, 2)
|
||||
db.add_all(children)
|
||||
main_task.credits_cost = total_credits
|
||||
total_billed_credits = total_credits
|
||||
|
||||
await db.flush()
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_BILLING_SUCCESS",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=current_user.id,
|
||||
group_id=main_id,
|
||||
detail={
|
||||
"gen_type": gen_type,
|
||||
"generation_count": generation_count,
|
||||
"total_billed_credits": total_billed_credits,
|
||||
},
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_CHILDREN_CREATED" if child_ids else "BATCH_MAIN_CREATED",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=current_user.id,
|
||||
group_id=main_id,
|
||||
detail={
|
||||
"gen_type": gen_type,
|
||||
"generation_count": generation_count,
|
||||
"child_task_ids": child_ids,
|
||||
"enqueue_task_ids": enqueue_ids,
|
||||
},
|
||||
)
|
||||
return GenerationTaskCreateResult(
|
||||
top_level_task_id=main_id,
|
||||
enqueue_task_ids=enqueue_ids,
|
||||
child_task_ids=child_ids,
|
||||
generation_count=generation_count,
|
||||
gen_type=gen_type,
|
||||
created=True,
|
||||
)
|
||||
|
||||
|
||||
async def enqueue_created_generation_tasks(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task_ids: list[str],
|
||||
) -> list[str]:
|
||||
"""在业务事务提交后投递任务;返回投递失败的任务ID。
|
||||
|
||||
投递失败会在补偿事务中将对应任务置为失败并幂等退款,视频子任务
|
||||
同时触发主任务状态汇总。调用方不应在初始事务提交前调用本函数。
|
||||
"""
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
from app.services.generation.log_service import log_task_event
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.tasks.generation_create_tasks import chatapi_create_generation_task
|
||||
|
||||
normalized_ids = list(dict.fromkeys(str(item) for item in task_ids if item))
|
||||
meta_result = await db.execute(
|
||||
select(
|
||||
ChatGenerationTask.id,
|
||||
ChatGenerationTask.user_id,
|
||||
ChatGenerationTask.parent_task_id,
|
||||
ChatGenerationTask.generation_index,
|
||||
).where(ChatGenerationTask.id.in_(normalized_ids))
|
||||
) if normalized_ids else None
|
||||
task_meta = {
|
||||
str(row.id): {
|
||||
"user_id": str(row.user_id),
|
||||
"parent_task_id": str(row.parent_task_id) if row.parent_task_id else None,
|
||||
"generation_index": row.generation_index,
|
||||
}
|
||||
for row in (meta_result.all() if meta_result is not None else [])
|
||||
}
|
||||
|
||||
failed_ids: list[str] = []
|
||||
for task_id in normalized_ids:
|
||||
meta = task_meta.get(task_id, {})
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_ENQUEUE_START",
|
||||
event_status="started",
|
||||
source="api",
|
||||
user_id=meta.get("user_id"),
|
||||
group_id=meta.get("parent_task_id") or task_id,
|
||||
task_id=task_id,
|
||||
detail={"generation_index": meta.get("generation_index")},
|
||||
)
|
||||
try:
|
||||
chatapi_create_generation_task.delay(task_id)
|
||||
await log_task_event(
|
||||
task_id=task_id,
|
||||
event_type="CHILD_ENQUEUE_SUCCESS",
|
||||
to_status="generating",
|
||||
to_stage="queued",
|
||||
detail={"task_id": task_id},
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_ENQUEUE_SUCCESS",
|
||||
event_status="success",
|
||||
source="api",
|
||||
user_id=meta.get("user_id"),
|
||||
group_id=meta.get("parent_task_id") or task_id,
|
||||
task_id=task_id,
|
||||
detail={"generation_index": meta.get("generation_index")},
|
||||
)
|
||||
except Exception as exc:
|
||||
failed_ids.append(task_id)
|
||||
await db.rollback()
|
||||
failed_task = await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task_id=task_id,
|
||||
error_message=f"任务队列投递失败: {exc}",
|
||||
pipeline_stage="failed",
|
||||
)
|
||||
await aggregate_parent_for_child(db, failed_task)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task_id=task_id,
|
||||
event_type="CHILD_ENQUEUE_FAILED",
|
||||
to_status="failed",
|
||||
to_stage="failed",
|
||||
message=str(exc),
|
||||
detail={"task_id": task_id},
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_ENQUEUE_FAILED",
|
||||
event_status="failed",
|
||||
source="api",
|
||||
user_id=getattr(failed_task, "user_id", None),
|
||||
group_id=getattr(failed_task, "parent_task_id", None) or task_id,
|
||||
task_id=task_id,
|
||||
message=str(exc),
|
||||
detail={"physical_files_deleted": False},
|
||||
error=str(exc),
|
||||
)
|
||||
return failed_ids
|
||||
@@ -0,0 +1,426 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import Counter, defaultdict
|
||||
from datetime import datetime, timezone
|
||||
from typing import Iterable, Sequence
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.generation_task import (
|
||||
ChatGenerationPipelineStage,
|
||||
ChatGenerationTaskStatus,
|
||||
GenerationMode,
|
||||
)
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.operation_log_service import log_operation_event
|
||||
from app.services.resource_accounting_service import (
|
||||
SOURCE_MODEL_CHAT_TASK,
|
||||
soft_delete_resources_by_source,
|
||||
)
|
||||
|
||||
|
||||
ACTIVE_STAGES = {
|
||||
ChatGenerationPipelineStage.QUEUED.value,
|
||||
ChatGenerationPipelineStage.PREPARING.value,
|
||||
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
ChatGenerationPipelineStage.WAITING_REMOTE.value,
|
||||
ChatGenerationPipelineStage.POLLING.value,
|
||||
ChatGenerationPipelineStage.RESULT_READY.value,
|
||||
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
|
||||
ChatGenerationPipelineStage.DOWNLOADING.value,
|
||||
ChatGenerationPipelineStage.RETRY_WAITING.value,
|
||||
}
|
||||
|
||||
|
||||
def is_task_active(task: ChatGenerationTask) -> bool:
|
||||
return task.deleted_at is None and (
|
||||
task.status == ChatGenerationTaskStatus.GENERATING.value
|
||||
or (task.pipeline_stage or "") in ACTIVE_STAGES
|
||||
)
|
||||
|
||||
|
||||
def get_display_status(task: ChatGenerationTask) -> str:
|
||||
if task.deleted_at is not None:
|
||||
return "deleted"
|
||||
if task.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value:
|
||||
return "download_failed"
|
||||
return task.status or ChatGenerationTaskStatus.PENDING.value
|
||||
|
||||
|
||||
async def load_children_map(
|
||||
db: AsyncSession,
|
||||
parent_ids: Sequence[str] | Iterable[str],
|
||||
*,
|
||||
include_deleted: bool = True,
|
||||
) -> dict[str, list[ChatGenerationTask]]:
|
||||
ids = list(dict.fromkeys(str(item) for item in parent_ids if item))
|
||||
if not ids:
|
||||
return {}
|
||||
query = select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.parent_task_id.in_(ids),
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
)
|
||||
if not include_deleted:
|
||||
query = query.where(ChatGenerationTask.deleted_at.is_(None))
|
||||
result = await db.execute(
|
||||
query.order_by(
|
||||
ChatGenerationTask.parent_task_id.asc(),
|
||||
ChatGenerationTask.generation_index.asc(),
|
||||
ChatGenerationTask.created_at.asc(),
|
||||
)
|
||||
)
|
||||
grouped: dict[str, list[ChatGenerationTask]] = defaultdict(list)
|
||||
for task in result.scalars().all():
|
||||
if task.parent_task_id:
|
||||
grouped[task.parent_task_id].append(task)
|
||||
return dict(grouped)
|
||||
|
||||
|
||||
async def load_task_and_children(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task_id: str,
|
||||
user_id: str | None = None,
|
||||
include_deleted_children: bool = True,
|
||||
) -> tuple[ChatGenerationTask | None, list[ChatGenerationTask]]:
|
||||
query = select(ChatGenerationTask).where(ChatGenerationTask.id == task_id)
|
||||
if user_id:
|
||||
query = query.where(ChatGenerationTask.user_id == user_id)
|
||||
result = await db.execute(query.limit(1))
|
||||
task = result.scalar_one_or_none()
|
||||
if not task:
|
||||
return None, []
|
||||
if task.generation_mode == GenerationMode.CHATAPI_CHILD.value and task.parent_task_id:
|
||||
parent_result = await db.execute(
|
||||
select(ChatGenerationTask).where(ChatGenerationTask.id == task.parent_task_id).limit(1)
|
||||
)
|
||||
parent = parent_result.scalar_one_or_none()
|
||||
return parent or task, [task]
|
||||
if task.generation_mode != GenerationMode.CHATAPI_MAIN.value:
|
||||
return task, []
|
||||
children_map = await load_children_map(
|
||||
db,
|
||||
[task.id],
|
||||
include_deleted=include_deleted_children,
|
||||
)
|
||||
return task, children_map.get(task.id, [])
|
||||
|
||||
|
||||
def _generation_result_status(task: ChatGenerationTask) -> str:
|
||||
"""返回任务真实生成结果,不受资源软删除影响。"""
|
||||
if task.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value:
|
||||
return "download_failed"
|
||||
if task.status == ChatGenerationTaskStatus.FAILED.value or (task.pipeline_stage or "") in {
|
||||
ChatGenerationPipelineStage.FAILED.value,
|
||||
ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
}:
|
||||
return "failed"
|
||||
if is_task_active(task):
|
||||
return "generating"
|
||||
if task.status == ChatGenerationTaskStatus.COMPLETED.value:
|
||||
return "completed"
|
||||
return task.status or "pending"
|
||||
|
||||
|
||||
def _build_summary(children: list[ChatGenerationTask]) -> str | None:
|
||||
if not children:
|
||||
return None
|
||||
result_counters: Counter[str] = Counter(_generation_result_status(child) for child in children)
|
||||
labels = {
|
||||
"completed": "完成",
|
||||
"failed": "生成失败",
|
||||
"download_failed": "下载失败",
|
||||
"generating": "生成中",
|
||||
"pending": "待处理",
|
||||
}
|
||||
parts = [f"{count}项{labels.get(status, status)}" for status, count in result_counters.items() if count]
|
||||
deleted_count = sum(1 for child in children if child.deleted_at is not None)
|
||||
if deleted_count:
|
||||
parts.append(f"{deleted_count}项资源已删除")
|
||||
return f"{len(children)}项中" + ",".join(parts)
|
||||
|
||||
|
||||
async def aggregate_main_task_status(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
parent_task_id: str,
|
||||
) -> ChatGenerationTask | None:
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == parent_task_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
main = result.scalar_one_or_none()
|
||||
if not main or main.deleted_at is not None:
|
||||
return main
|
||||
|
||||
children_map = await load_children_map(db, [parent_task_id], include_deleted=True)
|
||||
children = children_map.get(parent_task_id, [])
|
||||
if not children:
|
||||
return main
|
||||
|
||||
previous_status = main.status
|
||||
previous_stage = main.pipeline_stage
|
||||
active_children = [child for child in children if is_task_active(child)]
|
||||
failed_children = [
|
||||
child
|
||||
for child in children
|
||||
if (
|
||||
child.status == ChatGenerationTaskStatus.FAILED.value
|
||||
or (child.pipeline_stage or "") in {
|
||||
ChatGenerationPipelineStage.FAILED.value,
|
||||
ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value,
|
||||
}
|
||||
)
|
||||
]
|
||||
completed_children = [
|
||||
child for child in children if child.status == ChatGenerationTaskStatus.COMPLETED.value
|
||||
]
|
||||
|
||||
if active_children:
|
||||
main.status = ChatGenerationTaskStatus.GENERATING.value
|
||||
main.pipeline_stage = active_children[0].pipeline_stage or ChatGenerationPipelineStage.QUEUED.value
|
||||
main.generated_at = None
|
||||
main.error_message = _build_summary(children)
|
||||
elif failed_children:
|
||||
main.status = ChatGenerationTaskStatus.FAILED.value
|
||||
main.pipeline_stage = (
|
||||
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value
|
||||
if any(child.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value for child in failed_children)
|
||||
else ChatGenerationPipelineStage.FAILED.value
|
||||
)
|
||||
main.generated_at = max(
|
||||
(child.generated_at for child in completed_children if child.generated_at),
|
||||
default=datetime.now(timezone.utc),
|
||||
)
|
||||
main.error_message = _build_summary(children)
|
||||
else:
|
||||
# 所有子任务真实生成结果均成功;资源是否软删除不改变生成历史终态。
|
||||
main.status = ChatGenerationTaskStatus.COMPLETED.value
|
||||
main.pipeline_stage = ChatGenerationPipelineStage.DONE.value
|
||||
main.generated_at = max(
|
||||
(child.generated_at for child in children if child.generated_at),
|
||||
default=main.generated_at or datetime.now(timezone.utc),
|
||||
)
|
||||
main.error_message = None
|
||||
|
||||
if main.gen_type == "video":
|
||||
main.credits_cost = round(sum(float(child.credits_cost or 0) for child in children), 2)
|
||||
main.text_credits_cost = round(sum(float(child.text_credits_cost or 0) for child in children), 2)
|
||||
main.text_tokens_used = sum(int(child.text_tokens_used or 0) for child in children)
|
||||
main.image_tokens_used = sum(int(child.image_tokens_used or 0) for child in children)
|
||||
main.video_tokens_used = sum(int(child.video_tokens_used or 0) for child in children)
|
||||
main.retry_count = sum(int(child.retry_count or 0) for child in children)
|
||||
main.poll_count = sum(int(child.poll_count or 0) for child in children)
|
||||
|
||||
await db.flush()
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="MAIN_STATUS_AGGREGATED",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=main.user_id,
|
||||
group_id=main.id,
|
||||
task_id=main.id,
|
||||
detail={
|
||||
"before_status": previous_status,
|
||||
"before_stage": previous_stage,
|
||||
"after_status": main.status,
|
||||
"after_stage": main.pipeline_stage,
|
||||
"summary": _build_summary(children),
|
||||
},
|
||||
)
|
||||
return main
|
||||
|
||||
|
||||
async def aggregate_parent_for_child(db: AsyncSession, child: ChatGenerationTask | None) -> ChatGenerationTask | None:
|
||||
if not child or child.generation_mode != GenerationMode.CHATAPI_CHILD.value or not child.parent_task_id:
|
||||
return None
|
||||
# 项目关闭了 autoflush,先显式 flush 子任务的终态,确保聚合查询读取到本事务最新状态。
|
||||
await db.flush()
|
||||
return await aggregate_main_task_status(db, parent_task_id=str(child.parent_task_id))
|
||||
|
||||
|
||||
async def soft_delete_child_tasks_batch(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
child_task_ids: Sequence[str] | Iterable[str],
|
||||
user_id: str,
|
||||
deleted_at: datetime | None = None,
|
||||
require_completed: bool = False,
|
||||
) -> int:
|
||||
ids = list(dict.fromkeys(str(item) for item in child_task_ids if item))
|
||||
if not ids:
|
||||
return 0
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id.in_(ids),
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
)
|
||||
.with_for_update()
|
||||
)
|
||||
children = list(result.scalars().all())
|
||||
found_ids = {str(child.id) for child in children}
|
||||
missing_ids = [item for item in ids if item not in found_ids]
|
||||
if missing_ids:
|
||||
raise HTTPException(status_code=404, detail=f"子任务不存在: {','.join(missing_ids)}")
|
||||
|
||||
active_children = [child for child in children if child.deleted_at is None]
|
||||
running_ids = [child.id for child in active_children if is_task_active(child)]
|
||||
if running_ids:
|
||||
raise HTTPException(status_code=400, detail=f"仍有 {len(running_ids)} 个子任务生成中,暂不能删除")
|
||||
if require_completed:
|
||||
invalid_ids = [
|
||||
child.id for child in active_children
|
||||
if child.status != ChatGenerationTaskStatus.COMPLETED.value or child.generated_at is None
|
||||
]
|
||||
if invalid_ids:
|
||||
raise HTTPException(status_code=409, detail=f"只有生成完成的资源才能从素材云删除: {','.join(invalid_ids)}")
|
||||
|
||||
source_ids = [str(child.id) for child in active_children]
|
||||
freed_size = await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=source_ids,
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
parent_ids = list(dict.fromkeys(str(child.parent_task_id) for child in active_children if child.parent_task_id))
|
||||
for child in active_children:
|
||||
child.deleted_at = deleted_at
|
||||
await db.flush()
|
||||
for parent_id in parent_ids:
|
||||
await aggregate_main_task_status(db, parent_task_id=parent_id)
|
||||
return int(freed_size or 0)
|
||||
|
||||
|
||||
async def soft_delete_child_task(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
child_task_id: str,
|
||||
user_id: str,
|
||||
deleted_at: datetime | None = None,
|
||||
) -> int:
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
detail_result = await db.execute(
|
||||
select(
|
||||
ChatGenerationTask.parent_task_id,
|
||||
ChatGenerationTask.generation_index,
|
||||
).where(
|
||||
ChatGenerationTask.id == child_task_id,
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
).limit(1)
|
||||
)
|
||||
detail = detail_result.one_or_none()
|
||||
if not detail:
|
||||
raise HTTPException(status_code=404, detail="子任务不存在")
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_RESOURCE_DELETE_START",
|
||||
event_status="started",
|
||||
source="service",
|
||||
user_id=user_id,
|
||||
group_id=detail.parent_task_id,
|
||||
task_id=child_task_id,
|
||||
detail={"generation_index": detail.generation_index},
|
||||
)
|
||||
freed_size = await soft_delete_child_tasks_batch(
|
||||
db,
|
||||
child_task_ids=[child_task_id],
|
||||
user_id=user_id,
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_RESOURCE_DELETE_SUCCESS",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=user_id,
|
||||
group_id=detail.parent_task_id,
|
||||
task_id=child_task_id,
|
||||
detail={"generation_index": detail.generation_index, "freed_size_bytes": freed_size},
|
||||
)
|
||||
return freed_size
|
||||
|
||||
|
||||
async def soft_delete_top_level_task_group(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task_id: str,
|
||||
user_id: str,
|
||||
deleted_at: datetime | None = None,
|
||||
) -> int:
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode.in_(
|
||||
[GenerationMode.CHATAPI_ASYNC.value, GenerationMode.CHATAPI_MAIN.value]
|
||||
),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
task = result.scalar_one_or_none()
|
||||
if not task:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
|
||||
if task.generation_mode == GenerationMode.CHATAPI_ASYNC.value:
|
||||
if is_task_active(task):
|
||||
raise HTTPException(status_code=400, detail="当前任务正在生成中,暂不能删除")
|
||||
freed_size = await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=[task.id],
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
task.deleted_at = deleted_at
|
||||
await db.flush()
|
||||
return int(freed_size or 0)
|
||||
|
||||
children_map = await load_children_map(db, [task.id], include_deleted=True)
|
||||
children = children_map.get(task.id, [])
|
||||
active_ids = [child.id for child in children if is_task_active(child)]
|
||||
if active_ids:
|
||||
raise HTTPException(status_code=400, detail=f"任务组仍有 {len(active_ids)} 个子任务生成中,暂不能删除")
|
||||
|
||||
active_children = [child for child in children if child.deleted_at is None]
|
||||
child_ids = [child.id for child in active_children]
|
||||
freed_size = await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=child_ids,
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
for child in active_children:
|
||||
child.deleted_at = deleted_at
|
||||
task.deleted_at = deleted_at
|
||||
await db.flush()
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_GROUP_DELETE_SUCCESS",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=user_id,
|
||||
group_id=task.id,
|
||||
task_id=task.id,
|
||||
detail={
|
||||
"child_task_ids": child_ids,
|
||||
"freed_size_bytes": int(freed_size or 0),
|
||||
"physical_files_deleted": False,
|
||||
},
|
||||
)
|
||||
return int(freed_size or 0)
|
||||
+8
-4
@@ -470,6 +470,7 @@ async def charge_generation_media_by_params(
|
||||
source_step_id: str | None = None,
|
||||
source_step_code: str | None = None,
|
||||
billing_scene: str | None = None,
|
||||
quantity: int = 1,
|
||||
) -> BillingSummary:
|
||||
"""图片/视频媒体生成扣费。
|
||||
|
||||
@@ -478,6 +479,7 @@ async def charge_generation_media_by_params(
|
||||
"""
|
||||
project_name = project_name or "AI生成任务"
|
||||
gen_type = (gen_type or "").lower().strip()
|
||||
quantity = max(1, int(quantity or 1))
|
||||
attempt_no = attempt_no or await get_next_credit_attempt_no(
|
||||
db,
|
||||
owner_type=owner_type,
|
||||
@@ -508,13 +510,14 @@ async def charge_generation_media_by_params(
|
||||
|
||||
if gen_type == "image":
|
||||
size = image_size or "2K"
|
||||
amount = await calc_image_credits(db, size, engine_id=engine_id)
|
||||
unit_amount = await calc_image_credits(db, size, engine_id=engine_id)
|
||||
amount = round(unit_amount * quantity, 2)
|
||||
items.append(
|
||||
await deduct_credits_locked_once(
|
||||
db,
|
||||
user_id=user_id,
|
||||
amount=amount,
|
||||
description=f"{description_prefix}图片生成",
|
||||
description=f"{description_prefix}图片生成" + (f"×{quantity}" if quantity > 1 else ""),
|
||||
related_id=record_id,
|
||||
charge_key=CHARGE_MEDIA,
|
||||
biz_key=biz_key,
|
||||
@@ -523,17 +526,18 @@ async def charge_generation_media_by_params(
|
||||
)
|
||||
)
|
||||
elif gen_type == "video":
|
||||
amount = await calc_video_credits(
|
||||
unit_amount = await calc_video_credits(
|
||||
db, duration or 5, resolution or "720p",
|
||||
engine_id=engine_id,
|
||||
input_video_duration=input_video_duration,
|
||||
)
|
||||
amount = round(unit_amount * quantity, 2)
|
||||
items.append(
|
||||
await deduct_credits_locked_once(
|
||||
db,
|
||||
user_id=user_id,
|
||||
amount=amount,
|
||||
description=f"{description_prefix}视频生成",
|
||||
description=f"{description_prefix}视频生成" + (f"×{quantity}" if quantity > 1 else ""),
|
||||
related_id=record_id,
|
||||
charge_key=CHARGE_MEDIA,
|
||||
biz_key=biz_key,
|
||||
+51
-11
@@ -1,6 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
from typing import Iterable, Sequence
|
||||
|
||||
@@ -24,12 +23,13 @@ from app.models.module_generation_step import ModuleGenerationStep
|
||||
from app.models.shot_replicate_segment import ShotReplicateSegment
|
||||
from app.models.user import User
|
||||
from app.schemas.generation_ai import GenerationAIHistoryBatchDeleteOut
|
||||
from app.services.generation.ai.task_group_service import soft_delete_child_tasks_batch
|
||||
from app.services.module_generation_flow_base_service import is_active_chat_generation_task
|
||||
from app.services.module_generation_log_service import log_module_event_file
|
||||
from app.services.operation_log_service import log_operation_event
|
||||
# from app.services.operation_log import log_operation
|
||||
from app.services.resource_accounting_service import (
|
||||
SOURCE_MODEL_CHAT_TASK,
|
||||
SOURCE_MODEL_GENERATION_RECORD,
|
||||
SOURCE_MODEL_SHOT_SEGMENT,
|
||||
soft_delete_generation_record_resources,
|
||||
soft_delete_resources_by_source,
|
||||
@@ -238,14 +238,19 @@ async def _delete_chat_tasks(
|
||||
.where(
|
||||
ChatGenerationTask.id.in_(ids),
|
||||
ChatGenerationTask.user_id == current_user.id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_ASYNC.value,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.generation_mode.in_([
|
||||
GenerationMode.CHATAPI_ASYNC.value,
|
||||
GenerationMode.CHATAPI_CHILD.value,
|
||||
]),
|
||||
)
|
||||
.with_for_update()
|
||||
)
|
||||
tasks = list(result.scalars().all())
|
||||
_raise_missing_if_any(ids=ids, found_ids=[task.id for task in tasks], message="AI 创作记录不存在或已删除")
|
||||
|
||||
already_deleted_ids = [str(task.id) for task in tasks if task.deleted_at is not None]
|
||||
_raise_invalid_if_any(invalid_ids=already_deleted_ids, message="AI 创作记录不存在或已删除", status_code=404)
|
||||
|
||||
invalid_ids = [
|
||||
task.id
|
||||
for task in tasks
|
||||
@@ -253,14 +258,49 @@ async def _delete_chat_tasks(
|
||||
]
|
||||
_raise_invalid_if_any(invalid_ids=invalid_ids, message="AI 创作记录只有生成完成后才能删除")
|
||||
|
||||
freed_size = await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=[task.id for task in tasks],
|
||||
deleted_at=deleted_at,
|
||||
async_ids = [str(task.id) for task in tasks if task.generation_mode == GenerationMode.CHATAPI_ASYNC.value]
|
||||
child_ids = [str(task.id) for task in tasks if task.generation_mode == GenerationMode.CHATAPI_CHILD.value]
|
||||
|
||||
freed_size = 0
|
||||
if async_ids:
|
||||
freed_size += int(await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=async_ids,
|
||||
deleted_at=deleted_at,
|
||||
) or 0)
|
||||
async_id_set = set(async_ids)
|
||||
for task in tasks:
|
||||
if str(task.id) in async_id_set:
|
||||
task.deleted_at = deleted_at
|
||||
|
||||
if child_ids:
|
||||
freed_size += await soft_delete_child_tasks_batch(
|
||||
db,
|
||||
child_task_ids=child_ids,
|
||||
user_id=current_user.id,
|
||||
deleted_at=deleted_at,
|
||||
require_completed=True,
|
||||
)
|
||||
|
||||
parent_task_ids = list(dict.fromkeys(
|
||||
str(task.parent_task_id) for task in tasks if task.parent_task_id
|
||||
))
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_RESOURCE_DELETE_SUCCESS",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=current_user.id,
|
||||
group_id=parent_task_ids[0] if len(parent_task_ids) == 1 else None,
|
||||
detail={
|
||||
"batch": True,
|
||||
"task_ids": [str(task.id) for task in tasks],
|
||||
"parent_task_ids": parent_task_ids,
|
||||
"freed_size_bytes": int(freed_size or 0),
|
||||
"physical_files_deleted": False,
|
||||
},
|
||||
)
|
||||
for task in tasks:
|
||||
task.deleted_at = deleted_at
|
||||
|
||||
return _build_out(
|
||||
source=source,
|
||||
+1
-1
@@ -14,7 +14,7 @@ from app.config import settings
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.model_config import ModelConfig
|
||||
from app.models.token_usage import TokenUsage
|
||||
from app.services.generation_log_service import log_provider_call
|
||||
from app.services.generation.log_service import log_provider_call
|
||||
from app.services.provider_limit import provider_limit
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
+92
-38
@@ -13,10 +13,11 @@ from app.config import settings
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.services.generation_log_service import log_provider_call
|
||||
from app.services.image_gen import poll_image_task_status, submit_image_task
|
||||
from app.services.generation.log_service import log_provider_call
|
||||
from app.services.image_gen import ImageProviderError, poll_image_task_status, submit_image_task
|
||||
from app.services.provider_limit import provider_limit
|
||||
from app.services.video_gen import poll_task_status, submit_video_task
|
||||
from app.types.generation.provider import ImageProviderBatchResult
|
||||
|
||||
|
||||
def _loads(data: str | None) -> dict:
|
||||
@@ -29,8 +30,17 @@ def _loads(data: str | None) -> dict:
|
||||
return {}
|
||||
|
||||
|
||||
def _try_json(value: Any) -> Any:
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
try:
|
||||
return json.loads(value)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
|
||||
"""Use frozen snapshot for historical params, current DB row only for secret api_key."""
|
||||
"""使用任务快照冻结历史参数,只从当前引擎记录读取密钥。"""
|
||||
snapshot = _loads(task.engine_snapshot_json)
|
||||
if not task.engine_id:
|
||||
raise ValueError("缺少 engine_id")
|
||||
@@ -51,6 +61,31 @@ async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
|
||||
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"),
|
||||
multi_generation_enabled=bool(
|
||||
snapshot.get("multi_generation_enabled")
|
||||
if snapshot.get("multi_generation_enabled") is not None
|
||||
else getattr(engine, "multi_generation_enabled", False)
|
||||
),
|
||||
max_generation_count=int(
|
||||
snapshot.get("max_generation_count")
|
||||
or getattr(engine, "max_generation_count", 1)
|
||||
or 1
|
||||
),
|
||||
multi_image_max_images=int(
|
||||
snapshot.get("multi_image_max_images")
|
||||
or getattr(engine, "multi_image_max_images", 15)
|
||||
or 15
|
||||
),
|
||||
max_reference_image_count=int(
|
||||
snapshot.get("max_reference_image_count")
|
||||
if snapshot.get("max_reference_image_count") is not None
|
||||
else getattr(engine, "max_reference_image_count", 14)
|
||||
),
|
||||
output_format=(
|
||||
snapshot.get("output_format")
|
||||
if snapshot.get("output_format") is not None
|
||||
else getattr(engine, "output_format", "")
|
||||
) or "",
|
||||
)
|
||||
|
||||
|
||||
@@ -58,22 +93,16 @@ async def create_provider_task(db: AsyncSession, task: ChatGenerationTask) -> di
|
||||
if task.gen_type == "video":
|
||||
return await _create_video_task(db, task)
|
||||
if task.gen_type == "image":
|
||||
return await _create_image_sync_task(db, task)
|
||||
return await create_image_sync_result(db, task)
|
||||
raise ValueError(f"不支持的生成类型: {task.gen_type}")
|
||||
|
||||
|
||||
async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
|
||||
"""Create video provider task through the original Ark SDK async task API."""
|
||||
engine = await get_runtime_engine(db, task)
|
||||
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,
|
||||
include_media_references=True,
|
||||
)
|
||||
provider_task_id = await submit_video_task(None, engine, task, include_media_references=True)
|
||||
response = {"task_id": provider_task_id}
|
||||
await log_provider_call(
|
||||
task,
|
||||
@@ -101,65 +130,90 @@ async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict
|
||||
raise
|
||||
|
||||
|
||||
async def _create_image_sync_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
|
||||
"""Run the original synchronous image generation SDK under Celery control.
|
||||
|
||||
The legacy image SDK returns a final remote image URL immediately. We do
|
||||
NOT use image_generation.tasks.create here, so image generation stays aligned
|
||||
with the old working flow while no longer blocking the FastAPI request.
|
||||
"""
|
||||
async def create_image_sync_batch_result(
|
||||
db: AsyncSession,
|
||||
task: ChatGenerationTask,
|
||||
*,
|
||||
generation_count: int,
|
||||
) -> ImageProviderBatchResult:
|
||||
engine = await get_runtime_engine(db, task)
|
||||
return await create_image_sync_batch_result_with_engine(
|
||||
task,
|
||||
engine,
|
||||
generation_count=generation_count,
|
||||
)
|
||||
|
||||
|
||||
async def create_image_sync_batch_result_with_engine(
|
||||
task: ChatGenerationTask,
|
||||
engine: Any,
|
||||
*,
|
||||
generation_count: int,
|
||||
) -> ImageProviderBatchResult:
|
||||
"""执行一次同步图片请求。
|
||||
|
||||
generation_count > 1 时是一次组图 API 调用;失败后绝不退化为多次单图调用。
|
||||
"""
|
||||
count = max(1, int(generation_count or 1))
|
||||
started = time.perf_counter()
|
||||
api_type = "image_sync_batch_create" if count > 1 else "image_sync_create"
|
||||
async with provider_limit("ark_image_sync_create", settings.ARK_IMAGE_CREATE_MAX_CONCURRENCY):
|
||||
try:
|
||||
result = await asyncio.to_thread(
|
||||
submit_image_task,
|
||||
db,
|
||||
None,
|
||||
engine,
|
||||
task,
|
||||
include_media_references=True,
|
||||
generation_count=count,
|
||||
)
|
||||
if result.get("error"):
|
||||
raise RuntimeError(result.get("error"))
|
||||
response_data = _try_json(result.get("response_data")) or result
|
||||
response_data = result.get("response_data") or result
|
||||
await log_provider_call(
|
||||
task,
|
||||
provider=engine.provider,
|
||||
api_type="image_sync_create",
|
||||
api_type=api_type,
|
||||
model=engine.model_name,
|
||||
engine_id=task.engine_id,
|
||||
status="success",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
provider_task_id=None,
|
||||
response_data=response_data,
|
||||
total_tokens=int(result.get("image_tokens", 0) or 0),
|
||||
)
|
||||
return {
|
||||
"task_id": None,
|
||||
"remote_result_url": result.get("image_url"),
|
||||
"image_tokens": result.get("image_tokens", 0) or 0,
|
||||
"response_data": response_data,
|
||||
}
|
||||
return result
|
||||
except Exception as exc:
|
||||
error_message = exc.safe_message if isinstance(exc, ImageProviderError) else str(exc)
|
||||
await log_provider_call(
|
||||
task,
|
||||
provider=engine.provider,
|
||||
api_type="image_sync_create",
|
||||
api_type=api_type,
|
||||
model=engine.model_name,
|
||||
engine_id=task.engine_id,
|
||||
status="failed",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
error_message=str(exc),
|
||||
error_message=error_message,
|
||||
response_data=exc.as_dict() if isinstance(exc, ImageProviderError) else None,
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
def _try_json(text: Any) -> Any:
|
||||
if not isinstance(text, str):
|
||||
return text
|
||||
try:
|
||||
return json.loads(text)
|
||||
except Exception:
|
||||
return None
|
||||
async def create_image_sync_result(db: AsyncSession, task: ChatGenerationTask) -> dict:
|
||||
result = await create_image_sync_batch_result(db, task, generation_count=1)
|
||||
items = result.get("items") or []
|
||||
if len(items) != 1:
|
||||
raise RuntimeError(f"图片供应商单图返回数量异常,期望 1,实际 {len(items)}")
|
||||
item = items[0]
|
||||
if item.get("error_message"):
|
||||
raise RuntimeError(item.get("error_message") or "图片生成失败")
|
||||
image_url = item.get("remote_result_url")
|
||||
if not image_url:
|
||||
raise RuntimeError("图片供应商未返回有效图片地址")
|
||||
return {
|
||||
"task_id": None,
|
||||
"remote_result_url": image_url,
|
||||
"image_tokens": int(result.get("image_tokens", 0) or 0),
|
||||
"response_data": result.get("response_data") or {},
|
||||
}
|
||||
|
||||
|
||||
async def poll_provider_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
|
||||
+127
-4
@@ -15,6 +15,7 @@ from app.enums.generation_task import (
|
||||
ChatGenerationPipelineStage,
|
||||
ChatGenerationTaskEventType,
|
||||
ChatGenerationTaskStatus,
|
||||
GenerationMode,
|
||||
GenerationType,
|
||||
)
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
@@ -25,10 +26,10 @@ from app.services.celery_download_recovery_service import (
|
||||
postpone_download_active_check,
|
||||
remove_download_active,
|
||||
)
|
||||
from app.services.generation_log_service import log_task_event
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation_poll_schedule_service import ensure_video_poll_fields, is_poll_not_due, is_video_generation_task
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation.log_service import log_task_event
|
||||
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.poll_schedule_service import ensure_video_poll_fields, is_poll_not_due, is_video_generation_task
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.redis_registry_service import (
|
||||
redis_get_due_registry_ids,
|
||||
redis_get_registry_payloads,
|
||||
@@ -341,6 +342,8 @@ async def _mark_timeout(
|
||||
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
await aggregate_parent_for_child(db, task)
|
||||
await db.commit()
|
||||
await _remove_poll_active(task.id)
|
||||
await log_task_event(
|
||||
@@ -367,6 +370,8 @@ async def _mark_failed(
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
await aggregate_parent_for_child(db, task)
|
||||
await db.commit()
|
||||
await _remove_poll_active(task.id)
|
||||
await log_task_event(task, event_type=event_type, message=task.error_message, detail=detail)
|
||||
@@ -590,6 +595,99 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
|
||||
checked_ids: set[str] = set()
|
||||
results: dict[str, int] = {}
|
||||
|
||||
# 图片多份主任务只补投递,不在恢复服务内直接调用供应商。
|
||||
# 有效 claim 未过期时必须跳过,防止与正在运行的 Worker 重复调用组图 API。
|
||||
from app.tasks.generation_create_tasks import chatapi_create_generation_task
|
||||
image_main_cursor: str | None = None
|
||||
image_main_batch_size = max(1, int(settings.GENERATION_RECOVERY_BATCH_SIZE or 100))
|
||||
while True:
|
||||
image_main_query = select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
ChatGenerationTask.gen_type == GenerationType.IMAGE.value,
|
||||
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
|
||||
ChatGenerationTask.pipeline_stage.in_([
|
||||
ChatGenerationPipelineStage.QUEUED.value,
|
||||
ChatGenerationPipelineStage.PREPARING.value,
|
||||
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
]),
|
||||
)
|
||||
if image_main_cursor:
|
||||
image_main_query = image_main_query.where(ChatGenerationTask.id > image_main_cursor)
|
||||
image_main_result = await db.execute(
|
||||
image_main_query.order_by(ChatGenerationTask.id.asc())
|
||||
.limit(image_main_batch_size)
|
||||
.with_for_update()
|
||||
)
|
||||
image_mains = list(image_main_result.scalars().all())
|
||||
if not image_mains:
|
||||
break
|
||||
|
||||
for main in image_mains:
|
||||
main_id = str(main.id)
|
||||
image_main_cursor = main_id
|
||||
checked_ids.add(main_id)
|
||||
|
||||
child_result = await db.execute(
|
||||
select(ChatGenerationTask.id)
|
||||
.where(
|
||||
ChatGenerationTask.parent_task_id == main_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
)
|
||||
.limit(1)
|
||||
)
|
||||
if child_result.scalar_one_or_none() is not None:
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.commit()
|
||||
results["image_main_already_split"] = results.get("image_main_already_split", 0) + 1
|
||||
continue
|
||||
|
||||
now = _now()
|
||||
lease_until = ensure_aware_utc(main.provider_create_lease_until)
|
||||
lease_alive = bool(main.provider_create_claim_token and lease_until and lease_until > now)
|
||||
if lease_alive:
|
||||
await db.commit()
|
||||
results["image_main_claim_alive"] = results.get("image_main_claim_alive", 0) + 1
|
||||
continue
|
||||
|
||||
if _is_expired(main.deadline_at, now):
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=main,
|
||||
error_message="图片批量生成任务超时",
|
||||
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
await db.commit()
|
||||
results["image_main_timeout"] = results.get("image_main_timeout", 0) + 1
|
||||
continue
|
||||
|
||||
if main.provider_create_claim_token or main.provider_create_lease_until:
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
main.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
|
||||
await log_task_event(
|
||||
main,
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_MAIN_CLAIM_EXPIRED.value,
|
||||
message="图片主任务供应商执行租约已过期,恢复重新投递",
|
||||
)
|
||||
await db.commit()
|
||||
try:
|
||||
chatapi_create_generation_task.apply_async(
|
||||
args=[main_id],
|
||||
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
|
||||
countdown=0,
|
||||
)
|
||||
results["recover_image_main_create"] = results.get("recover_image_main_create", 0) + 1
|
||||
except Exception as exc:
|
||||
logger.exception("恢复投递图片主任务失败 task_id=%s: %s", main_id, exc)
|
||||
results["recover_image_main_enqueue_failed"] = results.get("recover_image_main_enqueue_failed", 0) + 1
|
||||
|
||||
if len(image_mains) < image_main_batch_size:
|
||||
break
|
||||
|
||||
due_poll_ids = await redis_get_due_registry_ids(
|
||||
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
|
||||
limit=int(settings.POLL_RECOVERY_BATCH_SIZE or settings.GENERATION_RECOVERY_BATCH_SIZE or 100),
|
||||
@@ -673,6 +771,31 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
|
||||
if len(tasks) < batch_size or progressed_this_round <= 0:
|
||||
break
|
||||
|
||||
# 子任务可能在 worker 中断前已进入终态但主任务尚未汇总,按稳定游标完整重算全部主任务。
|
||||
from app.services.generation.ai.task_group_service import aggregate_main_task_status
|
||||
reconciled = 0
|
||||
main_cursor: str | None = None
|
||||
while True:
|
||||
main_query = select(ChatGenerationTask.id).where(
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
)
|
||||
if main_cursor:
|
||||
main_query = main_query.where(ChatGenerationTask.id > main_cursor)
|
||||
main_result = await db.execute(main_query.order_by(ChatGenerationTask.id.asc()).limit(batch_size))
|
||||
parent_ids = list(main_result.scalars().all())
|
||||
if not parent_ids:
|
||||
break
|
||||
for parent_task_id in parent_ids:
|
||||
main_cursor = str(parent_task_id)
|
||||
await aggregate_main_task_status(db, parent_task_id=str(parent_task_id))
|
||||
await db.commit()
|
||||
reconciled += 1
|
||||
if len(parent_ids) < batch_size:
|
||||
break
|
||||
if reconciled:
|
||||
results["reconcile_main"] = reconciled
|
||||
|
||||
return {
|
||||
"checked": len(checked_ids),
|
||||
"db_checked": total_db_checked,
|
||||
+2
-4
@@ -1,7 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Iterable
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
@@ -11,7 +9,7 @@ from app.models.credit_record import CreditRecord
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.services.credits import refund_credits
|
||||
from app.services.credit_record_meta_service import build_refund_meta_from_charge
|
||||
from app.services.generation_billing_service import (
|
||||
from app.services.generation.billing_service import (
|
||||
CHARGE_MEDIA,
|
||||
OWNER_CHAT_GENERATION_TASK,
|
||||
OWNER_GENERATION_RECORD,
|
||||
@@ -182,7 +180,7 @@ async def mark_chat_generation_task_failed_and_refund_once(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate", "shot_replicate"]),
|
||||
ChatGenerationTask.generation_mode.in_(["chatapi_async", "chatapi_main", "chatapi_child", "hot_opening_replicate", "shot_replicate"]),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
+10
-8
@@ -11,21 +11,21 @@ from app.config import settings
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.user import User
|
||||
from app.schemas.generation_ai import GenerationAIReference, GenerationAITaskCreate
|
||||
from app.services.generation_ai_service import (
|
||||
from app.services.generation.ai.engine_service import (
|
||||
IMAGE_DEFAULT_PROPORTION,
|
||||
IMAGE_DEFAULT_PX,
|
||||
IMAGE_DEFAULT_SIZE,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
_build_image_snapshot,
|
||||
_build_video_snapshot,
|
||||
_get_image_engine,
|
||||
_get_video_engine,
|
||||
_image_supported_sizes,
|
||||
_parse_list,
|
||||
build_image_snapshot as _build_image_snapshot,
|
||||
build_video_snapshot as _build_video_snapshot,
|
||||
get_image_engine as _get_image_engine,
|
||||
get_video_engine as _get_video_engine,
|
||||
image_supported_sizes as _image_supported_sizes,
|
||||
normalize_px,
|
||||
parse_json_list as _parse_list,
|
||||
)
|
||||
from app.services.generation_billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params
|
||||
from app.services.generation.billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params
|
||||
from app.services.resource_capacity_service import assert_user_resource_capacity_available
|
||||
from app.services.private_portrait.reference_resolver import resolve_private_portrait_references
|
||||
from app.utils.id_gen import generate_id
|
||||
@@ -128,6 +128,7 @@ async def create_chat_generation_task_for_module(
|
||||
billing_scene=billing_scene,
|
||||
)
|
||||
snapshot = _build_image_snapshot(engine, size, proportion, px)
|
||||
snapshot["generation_count"] = 1
|
||||
task = ChatGenerationTask(
|
||||
id=task_id,
|
||||
user_id=current_user.id,
|
||||
@@ -182,6 +183,7 @@ async def create_chat_generation_task_for_module(
|
||||
billing_scene=billing_scene,
|
||||
)
|
||||
snapshot = _build_video_snapshot(engine, ratio, selected_resolution, selected_duration)
|
||||
snapshot["generation_count"] = 1
|
||||
task = ChatGenerationTask(
|
||||
id=task_id,
|
||||
user_id=current_user.id,
|
||||
@@ -1,42 +0,0 @@
|
||||
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
|
||||
@@ -1,14 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from copy import deepcopy
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import String, cast, func, or_, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum
|
||||
@@ -35,20 +33,19 @@ from app.schemas.hot_opening_replicate import (
|
||||
HotOpeningVideoGenerationOut,
|
||||
HotOpeningVideoPromptSchemaUpdateRequest,
|
||||
)
|
||||
from app.services.generation_ai_service import (
|
||||
from app.services.generation.ai.engine_service import (
|
||||
VIDEO_DEFAULT_DURATION,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
_get_video_engine,
|
||||
_parse_list,
|
||||
get_video_engine,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.generation_billing_service import charge_module_prompt_usage
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation_task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.generation.billing_service import charge_module_prompt_usage
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation.task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.hot_opening_video_prompt_service import build_final_video_prompt, optimize_hot_opening_video_prompt, patch_video_prompt_schema_from_client
|
||||
from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event
|
||||
from app.services.llm import optimize_prompt
|
||||
from app.services.resource_accounting_service import soft_delete_chat_task_resources
|
||||
from app.services.module_generation_flow_base_service import (
|
||||
chat_tasks_by_id as _base_chat_tasks_by_id,
|
||||
create_module_step as _base_create_step,
|
||||
@@ -1021,10 +1018,10 @@ async def generate_image_from_prompt(
|
||||
|
||||
|
||||
async def _resolve_video_prompt_config(db: AsyncSession, req: HotOpeningGenerateVideoPromptRequest) -> dict[str, Any]:
|
||||
engine = await _get_video_engine(db, req.engine_id)
|
||||
supported_ratios = _parse_list(engine.supported_ratios, [])
|
||||
supported_resolutions = _parse_list(engine.supported_resolutions, [])
|
||||
supported_durations = _parse_list(engine.supported_durations, [])
|
||||
engine = await get_video_engine(db, req.engine_id)
|
||||
supported_ratios = parse_json_list(engine.supported_ratios, [])
|
||||
supported_resolutions = parse_json_list(engine.supported_resolutions, [])
|
||||
supported_durations = parse_json_list(engine.supported_durations, [])
|
||||
|
||||
default_ratio = getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_RATIO", None) or VIDEO_DEFAULT_RATIO
|
||||
default_resolution = getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_RESOLUTION", None) or VIDEO_DEFAULT_RESOLUTION
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
import mimetypes
|
||||
import os
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy import select
|
||||
@@ -11,10 +10,16 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from volcenginesdkarkruntime import AsyncArk
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.generation_provider import (
|
||||
MULTI_IMAGE_PROMPT_TEMPLATE,
|
||||
ImageProviderErrorType,
|
||||
)
|
||||
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data
|
||||
from app.services.generation_provider_types import (
|
||||
from app.services.log_config import LOG_DATE_FORMAT, LOG_DIR, encrypt_data, is_enabled
|
||||
from app.types.generation.provider import (
|
||||
ImageProviderBatchResult,
|
||||
ImageProviderItem,
|
||||
ProviderGenerationRecordLike,
|
||||
ProviderImageEngineLike,
|
||||
)
|
||||
@@ -22,8 +27,39 @@ from app.services.generation_provider_types import (
|
||||
logger = logging.getLogger("videogen")
|
||||
|
||||
|
||||
class ImageProviderError(RuntimeError):
|
||||
"""可被生成任务状态机安全收敛的图片供应商异常。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
*,
|
||||
error_type: ImageProviderErrorType = ImageProviderErrorType.UNKNOWN,
|
||||
error_code: str | None = None,
|
||||
retryable: bool = False,
|
||||
http_status: int | None = None,
|
||||
provider_request_id: str | None = None,
|
||||
):
|
||||
super().__init__(message)
|
||||
self.safe_message = message
|
||||
self.error_type = error_type
|
||||
self.error_code = error_code
|
||||
self.retryable = retryable
|
||||
self.http_status = http_status
|
||||
self.provider_request_id = provider_request_id
|
||||
|
||||
def as_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"error_type": self.error_type.value,
|
||||
"error_code": self.error_code,
|
||||
"message": self.safe_message,
|
||||
"retryable": self.retryable,
|
||||
"http_status": self.http_status,
|
||||
"provider_request_id": self.provider_request_id,
|
||||
}
|
||||
|
||||
|
||||
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
|
||||
try:
|
||||
@@ -41,14 +77,13 @@ def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_
|
||||
"request": request_encrypted,
|
||||
"request_length": len(request_str),
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
with open(log_file, "a", encoding="utf-8") as file:
|
||||
file.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _log_image_response(record_id: str, response_data: dict, error: str | None = None):
|
||||
"""Log image generation response to log/AiModel/YYYY-MM-DD.log"""
|
||||
if not is_enabled():
|
||||
return
|
||||
try:
|
||||
@@ -63,17 +98,13 @@ def _log_image_response(record_id: str, response_data: dict, error: str | None =
|
||||
"response": response_encrypted,
|
||||
"error": error,
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
with open(log_file, "a", encoding="utf-8") as file:
|
||||
file.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
async def get_active_image_engine(db: AsyncSession) -> ImageEngine:
|
||||
"""Get the active image engine with highest priority."""
|
||||
result = await db.execute(
|
||||
select(ImageEngine)
|
||||
.where(ImageEngine.is_active == True)
|
||||
@@ -103,105 +134,291 @@ def _resolve_url(url: str) -> str:
|
||||
return f"{settings.BASE_URL.rstrip('/')}/{url.lstrip('/')}"
|
||||
|
||||
|
||||
def _value(obj: Any, name: str, default: Any = None) -> Any:
|
||||
if obj is None:
|
||||
return default
|
||||
if isinstance(obj, dict):
|
||||
return obj.get(name, default)
|
||||
return getattr(obj, name, default)
|
||||
|
||||
|
||||
def _jsonable(value: Any) -> Any:
|
||||
if value is None or isinstance(value, (str, int, float, bool)):
|
||||
return value
|
||||
if isinstance(value, dict):
|
||||
return {str(key): _jsonable(item) for key, item in value.items()}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [_jsonable(item) for item in value]
|
||||
if hasattr(value, "model_dump"):
|
||||
try:
|
||||
return _jsonable(value.model_dump())
|
||||
except Exception:
|
||||
pass
|
||||
if hasattr(value, "to_dict"):
|
||||
try:
|
||||
return _jsonable(value.to_dict())
|
||||
except Exception:
|
||||
pass
|
||||
result: dict[str, Any] = {}
|
||||
for key in ("url", "b64_json", "size", "output_format", "error", "code", "message"):
|
||||
item = getattr(value, key, None)
|
||||
if item is not None:
|
||||
result[key] = _jsonable(item)
|
||||
return result or str(value)
|
||||
|
||||
|
||||
def _safe_text(value: Any, *, limit: int = 1000) -> str:
|
||||
text = str(value or "").strip()
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def _classify_provider_exception(exc: Exception) -> ImageProviderError:
|
||||
if isinstance(exc, ImageProviderError):
|
||||
return exc
|
||||
if isinstance(exc, (httpx.TimeoutException, TimeoutError)):
|
||||
return ImageProviderError(
|
||||
"图片生成请求超时,请稍后重试",
|
||||
error_type=ImageProviderErrorType.TIMEOUT,
|
||||
retryable=True,
|
||||
)
|
||||
|
||||
status_code = getattr(exc, "status_code", None)
|
||||
request_id = getattr(exc, "request_id", None) or getattr(exc, "x_request_id", None)
|
||||
code = getattr(exc, "code", None)
|
||||
raw_message = _safe_text(getattr(exc, "message", None) or exc)
|
||||
lowered = raw_message.lower()
|
||||
|
||||
if status_code == 429 or "rate limit" in lowered or "限流" in raw_message:
|
||||
error_type = ImageProviderErrorType.RATE_LIMIT
|
||||
retryable = True
|
||||
message = "图片生成请求过于频繁,请稍后重试"
|
||||
elif status_code in {401, 403} or "api key" in lowered or "unauthorized" in lowered:
|
||||
error_type = ImageProviderErrorType.AUTH
|
||||
retryable = False
|
||||
message = "图片引擎鉴权失败,请联系管理员检查配置"
|
||||
elif status_code and int(status_code) >= 500:
|
||||
error_type = ImageProviderErrorType.PROVIDER_INTERNAL
|
||||
retryable = True
|
||||
message = "图片供应商服务异常,请稍后重试"
|
||||
elif "sequential_image_generation" in lowered or "not support" in lowered or "unsupported" in lowered:
|
||||
error_type = ImageProviderErrorType.CAPABILITY_MISMATCH
|
||||
retryable = False
|
||||
message = "图片引擎组图能力配置与供应商实际能力不匹配,请联系管理员"
|
||||
elif "content" in lowered and ("risk" in lowered or "moderation" in lowered or "policy" in lowered):
|
||||
error_type = ImageProviderErrorType.CONTENT_REJECTED
|
||||
retryable = False
|
||||
message = "图片内容未通过供应商审核,请调整提示词后重试"
|
||||
elif status_code and 400 <= int(status_code) < 500:
|
||||
error_type = ImageProviderErrorType.INVALID_REQUEST
|
||||
retryable = False
|
||||
message = "图片生成参数不被供应商支持,请联系管理员检查引擎配置"
|
||||
elif isinstance(exc, httpx.HTTPError):
|
||||
error_type = ImageProviderErrorType.NETWORK
|
||||
retryable = True
|
||||
message = "图片供应商网络连接异常,请稍后重试"
|
||||
else:
|
||||
error_type = ImageProviderErrorType.UNKNOWN
|
||||
retryable = False
|
||||
message = raw_message or "图片生成失败"
|
||||
|
||||
return ImageProviderError(
|
||||
message,
|
||||
error_type=error_type,
|
||||
error_code=_safe_text(code, limit=128) or None,
|
||||
retryable=retryable,
|
||||
http_status=int(status_code) if status_code is not None else None,
|
||||
provider_request_id=_safe_text(request_id, limit=128) or None,
|
||||
)
|
||||
|
||||
|
||||
def build_multi_image_provider_prompt(prompt: str, generation_count: int) -> str:
|
||||
base_prompt = (prompt or "").strip()
|
||||
if generation_count <= 1:
|
||||
return base_prompt
|
||||
suffix = MULTI_IMAGE_PROMPT_TEMPLATE.format(count=generation_count)
|
||||
return f"{base_prompt}\n\n{suffix}" if base_prompt else suffix
|
||||
|
||||
|
||||
def submit_image_task(
|
||||
db,
|
||||
engine: ProviderImageEngineLike,
|
||||
record: ProviderGenerationRecordLike,
|
||||
*,
|
||||
include_media_references: bool,
|
||||
) -> dict:
|
||||
"""Submit an image generation task via Ark SDK. Returns image_url."""
|
||||
from volcenginesdkarkruntime import Ark
|
||||
|
||||
client = Ark(
|
||||
base_url=engine.api_base,
|
||||
api_key=engine.api_key,
|
||||
timeout=300,
|
||||
)
|
||||
generation_count: int = 1,
|
||||
) -> ImageProviderBatchResult:
|
||||
"""通过 Ark 同步图片接口生成单图或单次组图。
|
||||
|
||||
prompt = record.optimized_prompt or record.original_prompt
|
||||
image_urls = []
|
||||
generation_count > 1 时只执行一次 sequential_auto 请求;任何失败都直接抛出,
|
||||
绝不退化为多次单图请求。
|
||||
"""
|
||||
from volcenginesdkarkruntime import Ark
|
||||
|
||||
count = max(1, int(generation_count or 1))
|
||||
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
|
||||
max_generation_count = max(1, min(5, int(getattr(engine, "max_generation_count", 1) or 1)))
|
||||
if count > 1 and not multi_generation_enabled:
|
||||
raise ImageProviderError(
|
||||
"当前图片引擎未开启多份生成",
|
||||
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
|
||||
)
|
||||
if count > max_generation_count:
|
||||
raise ImageProviderError(
|
||||
f"当前图片引擎最多允许生成 {max_generation_count} 份",
|
||||
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
|
||||
)
|
||||
|
||||
client = Ark(base_url=engine.api_base, api_key=engine.api_key, timeout=300)
|
||||
original_prompt = record.optimized_prompt or record.original_prompt
|
||||
provider_prompt = build_multi_image_provider_prompt(original_prompt, count)
|
||||
image_urls: list[str] = []
|
||||
|
||||
if include_media_references and record.media_references:
|
||||
try:
|
||||
refs = json.loads(record.media_references)
|
||||
for ref in refs:
|
||||
ref_type = ref.get("type")
|
||||
ref_url = ref.get("url", "")
|
||||
if ref_type == "image" and ref_url:
|
||||
resolved = _resolve_url(ref_url)
|
||||
image_urls.append(resolved)
|
||||
for ref in refs if isinstance(refs, list) else []:
|
||||
if (ref.get("type") or "").lower() == "image" and ref.get("url"):
|
||||
image_urls.append(_resolve_url(ref["url"]))
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
image_urls = []
|
||||
|
||||
request_payload = {
|
||||
request_log_payload: dict[str, Any] = {
|
||||
"model": engine.model_name,
|
||||
"prompt": prompt,
|
||||
"prompt": provider_prompt,
|
||||
"size": record.image_size or engine.default_size,
|
||||
"sequential_image_generation": "disabled",
|
||||
"output_format": "png",
|
||||
"response_format": "url",
|
||||
"watermark": False,
|
||||
"include_media_references": include_media_references,
|
||||
}
|
||||
|
||||
request_sdk_payload: dict[str, Any] = dict(request_log_payload)
|
||||
if image_urls:
|
||||
request_payload["image"] = image_urls
|
||||
request_log_payload["image"] = image_urls
|
||||
request_sdk_payload["image"] = image_urls
|
||||
output_format = (getattr(engine, "output_format", "") or "").lower().strip()
|
||||
if output_format:
|
||||
request_log_payload["output_format"] = output_format
|
||||
request_sdk_payload["output_format"] = output_format
|
||||
if count > 1:
|
||||
try:
|
||||
from volcenginesdkarkruntime.types.images import SequentialImageGenerationOptions
|
||||
except Exception:
|
||||
try:
|
||||
from volcenginesdkarkruntime.types.images.image_generate_params import (
|
||||
SequentialImageGenerationOptions,
|
||||
)
|
||||
except Exception as import_exc:
|
||||
raise ImageProviderError(
|
||||
"当前图片引擎运行依赖缺少组图参数对象,请升级火山 Ark SDK 后重试",
|
||||
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
|
||||
) from import_exc
|
||||
|
||||
_log_image_request(engine, record.id, request_payload)
|
||||
request_log_payload["sequential_image_generation"] = "auto"
|
||||
request_log_payload["sequential_image_generation_options"] = {"max_images": count}
|
||||
request_log_payload["stream"] = False
|
||||
|
||||
request_sdk_payload["sequential_image_generation"] = "auto"
|
||||
request_sdk_payload["sequential_image_generation_options"] = SequentialImageGenerationOptions(
|
||||
max_images=count,
|
||||
)
|
||||
request_sdk_payload["stream"] = False
|
||||
|
||||
_log_image_request(engine, record.id, request_log_payload)
|
||||
|
||||
try:
|
||||
result = client.images.generate(
|
||||
model=engine.model_name,
|
||||
prompt=prompt,
|
||||
size=record.image_size or engine.default_size,
|
||||
output_format="png",
|
||||
response_format="url",
|
||||
watermark=False,
|
||||
image=image_urls if image_urls else None,
|
||||
)
|
||||
image_url = result.data[0].url
|
||||
|
||||
response_data = {
|
||||
"model": result.model,
|
||||
"created": result.created,
|
||||
"data": [{"url": item.url, "size": item.size} for item in result.data] if result.data else [],
|
||||
"usage": {
|
||||
"generated_images": result.usage.generated_images if hasattr(result.usage, 'generated_images') else 0,
|
||||
"output_tokens": result.usage.output_tokens if hasattr(result.usage, 'output_tokens') else 0,
|
||||
"total_tokens": result.usage.total_tokens if hasattr(result.usage, 'total_tokens') else 0,
|
||||
result = client.images.generate(**request_sdk_payload)
|
||||
top_error = _value(result, "error")
|
||||
if top_error:
|
||||
error_code = _value(top_error, "code")
|
||||
error_message = _value(top_error, "message") or str(top_error)
|
||||
raise ImageProviderError(
|
||||
_safe_text(error_message) or "图片供应商返回失败",
|
||||
error_type=ImageProviderErrorType.INVALID_REQUEST,
|
||||
error_code=_safe_text(error_code, limit=128) or None,
|
||||
)
|
||||
|
||||
raw_data = _value(result, "data", []) or []
|
||||
if not isinstance(raw_data, (list, tuple)):
|
||||
raise ImageProviderError(
|
||||
"图片供应商返回 data 结构异常",
|
||||
error_type=ImageProviderErrorType.INVALID_RESPONSE,
|
||||
)
|
||||
|
||||
items: list[ImageProviderItem] = []
|
||||
response_items: list[dict[str, Any]] = []
|
||||
for index, raw_item in enumerate(raw_data, start=1):
|
||||
item_error = _value(raw_item, "error")
|
||||
if item_error:
|
||||
error_code = _safe_text(_value(item_error, "code"), limit=128)
|
||||
error_message = _safe_text(_value(item_error, "message") or item_error)
|
||||
items.append({
|
||||
"generation_index": index,
|
||||
"error_code": error_code,
|
||||
"error_message": error_message or "单张图片生成失败",
|
||||
"response_data": _jsonable(raw_item),
|
||||
})
|
||||
response_items.append(_jsonable(raw_item))
|
||||
continue
|
||||
|
||||
url = _safe_text(_value(raw_item, "url"), limit=4000)
|
||||
b64_json = _safe_text(_value(raw_item, "b64_json"), limit=100) if not url else ""
|
||||
item: ImageProviderItem = {
|
||||
"generation_index": index,
|
||||
"remote_result_url": url,
|
||||
"size": _safe_text(_value(raw_item, "size"), limit=64),
|
||||
"output_format": _safe_text(_value(raw_item, "output_format"), limit=32),
|
||||
"response_data": _jsonable(raw_item),
|
||||
}
|
||||
if b64_json:
|
||||
item["b64_json"] = b64_json
|
||||
items.append(item)
|
||||
response_items.append(_jsonable(raw_item))
|
||||
|
||||
usage = _value(result, "usage")
|
||||
generated_images = int(_value(usage, "generated_images", 0) or 0)
|
||||
total_tokens = int(_value(usage, "total_tokens", 0) or 0)
|
||||
response_data = {
|
||||
"model": _value(result, "model", engine.model_name),
|
||||
"created": _value(result, "created"),
|
||||
"data": response_items,
|
||||
"usage": {
|
||||
"generated_images": generated_images,
|
||||
"input_images": int(_value(usage, "input_images", 0) or 0),
|
||||
"output_tokens": int(_value(usage, "output_tokens", 0) or 0),
|
||||
"total_tokens": total_tokens,
|
||||
},
|
||||
}
|
||||
except httpx.TimeoutException:
|
||||
error_msg = "图片生成超时,请稍后重试"
|
||||
logger.error(f"Image generation timeout for record {record.id}")
|
||||
_log_image_response(record.id, {}, error_msg)
|
||||
raise TimeoutError(error_msg)
|
||||
except Exception as e:
|
||||
error_msg = str(e)
|
||||
logger.error(f"Image generation failed for record {record.id}: {error_msg}")
|
||||
_log_image_response(record.id, {}, error_msg)
|
||||
raise
|
||||
_log_image_response(record.id, response_data)
|
||||
return {
|
||||
"items": items,
|
||||
"model": str(response_data["model"] or ""),
|
||||
"created": int(response_data["created"] or 0),
|
||||
"generated_images": generated_images,
|
||||
"image_tokens": total_tokens,
|
||||
"response_data": response_data,
|
||||
}
|
||||
except Exception as exc:
|
||||
provider_error = _classify_provider_exception(exc)
|
||||
logger.error(
|
||||
"Image generation failed for record %s: type=%s code=%s message=%s",
|
||||
record.id,
|
||||
provider_error.error_type.value,
|
||||
provider_error.error_code,
|
||||
provider_error.safe_message,
|
||||
)
|
||||
_log_image_response(record.id, provider_error.as_dict(), provider_error.safe_message)
|
||||
raise provider_error from exc
|
||||
finally:
|
||||
client.close()
|
||||
|
||||
return {
|
||||
"image_url": image_url,
|
||||
"image_tokens": getattr(result.usage, "total_tokens", 0),
|
||||
"response_data": json.dumps(response_data, ensure_ascii=False, default=str),
|
||||
"error": str(result.error) if result.error else "",
|
||||
}
|
||||
try:
|
||||
client.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict:
|
||||
"""Query image task status via Ark SDK. Returns {status, image_url, response_data}."""
|
||||
client = AsyncArk(
|
||||
base_url=engine.api_base,
|
||||
api_key=engine.api_key,
|
||||
)
|
||||
|
||||
result = await client.image_generation.tasks.get(task_id=task_id)
|
||||
await client.close()
|
||||
client = AsyncArk(base_url=engine.api_base, api_key=engine.api_key)
|
||||
try:
|
||||
result = await client.image_generation.tasks.get(task_id=task_id)
|
||||
finally:
|
||||
await client.close()
|
||||
|
||||
response_dict = {
|
||||
"id": result.id,
|
||||
@@ -239,13 +456,11 @@ async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict:
|
||||
|
||||
|
||||
async def download_image(image_url: str, dest_path: str) -> str:
|
||||
"""Download image to local storage."""
|
||||
os.makedirs(os.path.dirname(dest_path), exist_ok=True)
|
||||
|
||||
async with httpx.AsyncClient(timeout=300) as client:
|
||||
async with client.stream("GET", image_url) as response:
|
||||
response.raise_for_status()
|
||||
with open(dest_path, "wb") as f:
|
||||
with open(dest_path, "wb") as file:
|
||||
async for chunk in response.aiter_bytes(chunk_size=8192):
|
||||
f.write(chunk)
|
||||
return dest_path
|
||||
file.write(chunk)
|
||||
return dest_path
|
||||
|
||||
@@ -4,7 +4,7 @@ from collections.abc import Iterable
|
||||
from datetime import datetime
|
||||
from typing import Any, TypedDict
|
||||
|
||||
from sqlalchemy import and_, func, or_, select
|
||||
from sqlalchemy import and_, case, func, or_, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.generation_task import GenerationType
|
||||
@@ -12,7 +12,7 @@ from app.enums.recent_generation import (
|
||||
RECENT_GENERATION_ALL_MODULES,
|
||||
RECENT_GENERATION_CHAT_TASK_MODULES,
|
||||
RECENT_GENERATION_COMPLETED_STATUS,
|
||||
RECENT_GENERATION_MODULE_TO_TASK_MODE,
|
||||
RECENT_GENERATION_MODULE_TO_TASK_MODES,
|
||||
RECENT_GENERATION_TASK_MODE_VALUE_TO_MODULE,
|
||||
RecentGenerationModuleEnum,
|
||||
RecentGenerationResourceTypeEnum,
|
||||
@@ -175,11 +175,12 @@ async def _list_chat_task_recent_rows(
|
||||
modules: list[RecentGenerationModuleEnum],
|
||||
limit: int,
|
||||
) -> list[dict[str, Any]]:
|
||||
task_mode_values = [
|
||||
RECENT_GENERATION_MODULE_TO_TASK_MODE[module].value
|
||||
task_mode_values = list(dict.fromkeys(
|
||||
task_mode.value
|
||||
for module in modules
|
||||
if module in RECENT_GENERATION_CHAT_TASK_MODULES
|
||||
]
|
||||
for task_mode in RECENT_GENERATION_MODULE_TO_TASK_MODES[module]
|
||||
))
|
||||
if not task_mode_values:
|
||||
return []
|
||||
|
||||
@@ -189,10 +190,19 @@ async def _list_chat_task_recent_rows(
|
||||
ChatGenerationTask.created_at,
|
||||
)
|
||||
|
||||
module_partition_expr = case(
|
||||
(
|
||||
ChatGenerationTask.generation_mode.in_(["chatapi_async", "chatapi_child"]),
|
||||
RecentGenerationModuleEnum.CHAT_AI.value,
|
||||
),
|
||||
else_=ChatGenerationTask.generation_mode,
|
||||
)
|
||||
|
||||
ranked_subquery = (
|
||||
select(
|
||||
ChatGenerationTask.id.label("generation_id"),
|
||||
ChatGenerationTask.generation_mode.label("generation_mode"),
|
||||
module_partition_expr.label("module_key"),
|
||||
ChatGenerationTask.gen_type.label("gen_type"),
|
||||
ChatGenerationTask.image_url.label("image_url"),
|
||||
ChatGenerationTask.video_url.label("video_url"),
|
||||
@@ -200,7 +210,7 @@ async def _list_chat_task_recent_rows(
|
||||
generated_time_expr.label("generated_time"),
|
||||
func.row_number()
|
||||
.over(
|
||||
partition_by=ChatGenerationTask.generation_mode,
|
||||
partition_by=module_partition_expr,
|
||||
order_by=(generated_time_expr.desc(), ChatGenerationTask.created_at.desc()),
|
||||
)
|
||||
.label("row_num"),
|
||||
@@ -218,7 +228,7 @@ async def _list_chat_task_recent_rows(
|
||||
stmt = (
|
||||
select(ranked_subquery)
|
||||
.where(ranked_subquery.c.row_num <= limit)
|
||||
.order_by(ranked_subquery.c.generation_mode.asc(), ranked_subquery.c.generated_time.desc())
|
||||
.order_by(ranked_subquery.c.module_key.asc(), ranked_subquery.c.generated_time.desc())
|
||||
)
|
||||
|
||||
return [dict(row) for row in (await db.execute(stmt)).mappings().all()]
|
||||
|
||||
@@ -1,14 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from copy import deepcopy
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum
|
||||
@@ -35,16 +33,16 @@ from app.schemas.shot_replicate import (
|
||||
ShotReplicateVideoGenerationOut,
|
||||
ShotReplicateVideoPromptSchemaUpdateRequest,
|
||||
)
|
||||
from app.services.generation_ai_service import (
|
||||
from app.services.generation.ai.engine_service import (
|
||||
VIDEO_DEFAULT_DURATION,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
_get_video_engine,
|
||||
_parse_list,
|
||||
get_video_engine,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.generation_billing_service import charge_module_prompt_usage
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation_task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.generation.billing_service import charge_module_prompt_usage
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation.task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.hot_opening_video_prompt_service import (
|
||||
build_final_video_prompt,
|
||||
optimize_hot_opening_video_prompt as optimize_shot_replicate_video_prompt,
|
||||
@@ -52,7 +50,6 @@ from app.services.hot_opening_video_prompt_service import (
|
||||
)
|
||||
from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event
|
||||
from app.services.llm import optimize_prompt
|
||||
from app.services.resource_accounting_service import soft_delete_chat_task_resources
|
||||
from app.services.module_generation_flow_base_service import (
|
||||
assert_project_has_no_active_chat_tasks as _base_assert_project_has_no_active_chat_tasks,
|
||||
chat_tasks_by_id as _base_chat_tasks_by_id,
|
||||
@@ -976,10 +973,10 @@ async def generate_image_from_prompt(
|
||||
|
||||
|
||||
async def _resolve_video_prompt_config(db: AsyncSession, req: ShotReplicateGenerateVideoPromptRequest) -> dict[str, Any]:
|
||||
engine = await _get_video_engine(db, req.engine_id)
|
||||
supported_ratios = _parse_list(engine.supported_ratios, [])
|
||||
supported_resolutions = _parse_list(engine.supported_resolutions, [])
|
||||
supported_durations = _parse_list(engine.supported_durations, [])
|
||||
engine = await get_video_engine(db, req.engine_id)
|
||||
supported_ratios = parse_json_list(engine.supported_ratios, [])
|
||||
supported_resolutions = parse_json_list(engine.supported_resolutions, [])
|
||||
supported_durations = parse_json_list(engine.supported_durations, [])
|
||||
|
||||
default_ratio = getattr(settings, "SHOT_REPLICATE_DEFAULT_VIDEO_RATIO", None) or VIDEO_DEFAULT_RATIO
|
||||
default_resolution = getattr(settings, "SHOT_REPLICATE_DEFAULT_VIDEO_RESOLUTION", None) or VIDEO_DEFAULT_RESOLUTION
|
||||
|
||||
@@ -14,7 +14,7 @@ from app.config import settings
|
||||
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data
|
||||
from app.services.generation_provider_types import (
|
||||
from app.types.generation.provider import (
|
||||
ProviderGenerationRecordLike,
|
||||
ProviderVideoEngineLike,
|
||||
)
|
||||
|
||||
@@ -17,7 +17,7 @@ from app.services.resource_accounting_service import (
|
||||
)
|
||||
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
|
||||
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
|
||||
|
||||
logger = logging.getLogger("videogen")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user