This commit is contained in:
2026-07-15 13:51:21 +08:00
parent a9190ba4e1
commit 6db989dc42
65 changed files with 4499 additions and 1515 deletions
@@ -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
@@ -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)
@@ -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,
@@ -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,
@@ -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
@@ -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:
@@ -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,
@@ -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()
@@ -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
+305 -90
View File
@@ -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
+1 -1
View File
@@ -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,
)
+1 -1
View File
@@ -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")