Files
video-gen/video-gen-api/app/services/generation_ai_service.py
T

724 lines
24 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from __future__ import annotations
import json
from datetime import datetime, timedelta, timezone, date
from typing import Any
from fastapi import HTTPException
from sqlalchemy import 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.schemas.generation_ai import (
GenerationAIRecordHistoryItemOut,
GenerationAITaskCreate,
GenerationAITaskOut,
)
from app.services.generation_billing_service import charge_generation_media_by_params
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")
def _json(data: Any) -> str | None:
if data is None:
return None
return json.dumps(data, ensure_ascii=False, default=str)
def _parse_json(text: str | None):
if not text:
return None
try:
return json.loads(text)
except Exception:
return None
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()).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,
"selected_ratio": ratio,
"selected_resolution": resolution,
"selected_duration": duration,
}
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",
).order_by(ChatGenerationTask.created_at.desc()).limit(1)
)
existing = result.scalar_one_or_none()
if existing:
return existing
refs = [r.model_dump() for r in (req.media_references or [])]
now = datetime.now(timezone.utc)
task_id = generate_id()
if gen_type == "image":
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,
project_name="AI生成任务",
description_prefix="Chat任务",
)
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} 秒")
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,
project_name="AI生成任务",
description_prefix="Chat任务",
)
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(minutes=settings.CHATAPI_ASYNC_VIDEO_DEADLINE_MINUTES),
)
db.add(task)
await db.flush()
return task
def record_to_out(task: ChatGenerationTask) -> GenerationAITaskOut:
refs = _parse_json(task.media_references)
snapshot = engine_snapshot_out(_parse_json(task.engine_snapshot_json))
return GenerationAITaskOut(
id=task.id,
# project_id=None,
gen_type=task.gen_type,
generation_mode=task.generation_mode,
pipeline_stage=task.pipeline_stage,
status=task.status,
original_prompt=task.original_prompt,
# optimized_prompt=task.optimized_prompt,
duration=task.duration,
aspect_ratio=task.aspect_ratio,
resolution=task.resolution,
image_size=task.image_size,
image_proportion=task.image_proportion,
image_px=task.image_px,
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=task.image_url,
video_url=task.video_url,
engine_id=task.engine_id,
engine_snapshot=snapshot,
credits_cost=task.credits_cost or 0.0,
text_credits_cost=task.text_credits_cost or 0.0,
text_tokens_used=task.text_tokens_used or 0,
image_tokens_used=task.image_tokens_used or 0,
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=task.error_message,
created_at=task.created_at,
generated_at=task.generated_at,
)
def engine_snapshot_out(snapshot: dict) -> dict:
"""
从完整的 engine_snapshot 中过滤出需要返回的字段
"""
if not snapshot:
return {}
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 list_async_generation_tasks(
db: AsyncSession,
user_id: str,
gen_type: str | None,
status: str | None,
page: int,
page_size: int,
):
query = select(ChatGenerationTask).where(
ChatGenerationTask.user_id == user_id,
ChatGenerationTask.generation_mode == "chatapi_async",
)
if gen_type:
query = query.where(ChatGenerationTask.gen_type == gen_type)
if status:
query = query.where(ChatGenerationTask.status == status)
count_query = select(func.count()).select_from(query.subquery())
total = (await db.execute(count_query)).scalar_one()
result = await db.execute(
query.order_by(ChatGenerationTask.created_at.desc())
.offset((page - 1) * page_size)
.limit(page_size)
)
return total, list(result.scalars().all())
def _normalize_history_gen_type(gen_type: str | None) -> str:
value = (gen_type or "").lower().strip()
if value not in ("image", "video"):
raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video")
return value
def _normalize_history_source(history_source: str | None) -> str:
"""Normalize history source query param.
默认保持原来的 chat_generation_tasks 历史;只有显式传 generation_record
才切换旧 generation_records 历史,避免影响现有前端。
"""
value = (history_source or "chat_task").lower().strip()
if value in ("", "chat", "chat_task", "chat_generation_task", "chat_generation_tasks"):
return "chat_task"
if value in ("record", "records", "generation_record", "generation_records"):
return "generation_record"
raise HTTPException(status_code=400, detail="history_source 仅支持 chat_task 或 generation_record")
def _history_day_to_str(value) -> str:
if isinstance(value, datetime):
return value.date().strftime("%Y-%m-%d")
if isinstance(value, date):
return value.strftime("%Y-%m-%d")
return str(value)[:10]
def _parse_history_date(value: str) -> date:
try:
return datetime.strptime(value, "%Y-%m-%d").date()
except ValueError:
raise HTTPException(status_code=400, detail="generated_date 格式必须是 YYYY-MM-DD")
def _history_base_filters(user_id: str, gen_type: str):
return [
ChatGenerationTask.user_id == user_id,
ChatGenerationTask.generation_mode == "chatapi_async",
ChatGenerationTask.status == "completed",
ChatGenerationTask.gen_type == gen_type,
ChatGenerationTask.generated_at.is_not(None),
]
def _generation_record_history_base_filters(user_id: str, gen_type: str):
return [
GenerationRecord.user_id == user_id,
GenerationRecord.status == "completed",
GenerationRecord.gen_type == gen_type,
GenerationRecord.generated_at.is_not(None),
]
def generation_record_to_history_out(
record: GenerationRecord,
project_name: str | None = None,
) -> GenerationAIRecordHistoryItemOut:
refs = _parse_json(record.media_references)
return GenerationAIRecordHistoryItemOut(
id=record.id,
source_type="generation_record",
project_id=record.project_id,
project_name=project_name,
gen_type=record.gen_type,
generation_mode="generation_record",
pipeline_stage=None,
status=record.status,
original_prompt=record.original_prompt,
# optimized_prompt=None,
duration=record.duration,
aspect_ratio=record.aspect_ratio,
resolution=record.resolution,
image_size=record.image_size,
image_proportion=record.image_proportion,
image_px=record.image_px,
references=refs,
media_references=refs,
provider_task_id=record.seedance_task_id,
seedance_task_id=record.seedance_task_id,
remote_result_url=None,
image_url=record.image_url,
video_url=record.video_url,
engine_id=None,
engine_snapshot=None,
credits_cost=record.credits_cost or 0.0,
text_credits_cost=record.text_credits_cost or 0.0,
text_tokens_used=record.text_tokens_used or 0,
image_tokens_used=record.image_tokens_used or 0,
video_tokens_used=record.video_tokens_used or 0,
retry_count=0,
poll_count=0,
error_message=record.error_message,
created_at=record.created_at,
generated_at=record.generated_at,
)
async def list_generation_record_history_grouped_days(
db: AsyncSession,
user_id: str,
gen_type: str,
page: int,
page_size: int,
):
"""
按生成日期倒序返回旧 generation_records 历史记录分组。
- 每页最多返回 10 个生成日期
- 每个日期分组内最多返回倒序前 10 条旧记录
- 只返回 completed 成功记录
"""
gen_type = _normalize_history_gen_type(gen_type)
page = max(page, 1)
page_size = min(max(page_size, 1), HISTORY_DAY_PAGE_SIZE_MAX)
filters = _generation_record_history_base_filters(user_id, gen_type)
day_expr = func.date(GenerationRecord.generated_at).label("generated_date")
days_subquery = (
select(day_expr)
.where(*filters)
.group_by(day_expr)
.subquery()
)
total_days = (
await db.execute(select(func.count()).select_from(days_subquery))
).scalar_one()
day_rows_result = await db.execute(
select(
day_expr,
func.count(GenerationRecord.id).label("total"),
)
.where(*filters)
.group_by(day_expr)
.order_by(day_expr.desc())
.offset((page - 1) * page_size)
.limit(page_size)
)
day_rows = day_rows_result.all()
groups = []
for generated_day, day_total in day_rows:
item_result = await db.execute(
select(GenerationRecord, Project.name.label("project_name"))
.outerjoin(Project, GenerationRecord.project_id == Project.id)
.where(
*filters,
func.date(GenerationRecord.generated_at) == generated_day,
)
.order_by(GenerationRecord.generated_at.desc(), GenerationRecord.created_at.desc())
.limit(HISTORY_GROUP_ITEM_LIMIT)
)
rows = item_result.all()
groups.append(
{
"generated_date": _history_day_to_str(generated_day),
"total": int(day_total or 0),
"items": [
generation_record_to_history_out(record, project_name)
for record, project_name in rows
],
}
)
return {
"total_days": int(total_days or 0),
"page": page,
"page_size": page_size,
"groups": groups,
}
async def list_generation_record_history_day_items(
db: AsyncSession,
user_id: str,
gen_type: str,
generated_date: str,
page: int,
page_size: int,
):
"""
获取旧 generation_records 指定生成日期下的历史记录分页。
"""
gen_type = _normalize_history_gen_type(gen_type)
target_day = _parse_history_date(generated_date)
page = max(page, 1)
page_size = min(max(page_size, 1), 100)
filters = _generation_record_history_base_filters(user_id, gen_type)
day_expr = func.date(GenerationRecord.generated_at)
total = (
await db.execute(
select(func.count(GenerationRecord.id)).where(
*filters,
day_expr == target_day,
)
)
).scalar_one()
result = await db.execute(
select(GenerationRecord, Project.name.label("project_name"))
.outerjoin(Project, GenerationRecord.project_id == Project.id)
.where(
*filters,
day_expr == target_day,
)
.order_by(GenerationRecord.generated_at.desc(), GenerationRecord.created_at.desc())
.offset((page - 1) * page_size)
.limit(page_size)
)
rows = result.all()
return {
"generated_date": target_day.strftime("%Y-%m-%d"),
"total": int(total or 0),
"page": page,
"page_size": page_size,
"items": [
generation_record_to_history_out(record, project_name)
for record, project_name in rows
],
}
async def list_generation_history_grouped_days(
db: AsyncSession,
user_id: str,
gen_type: str,
page: int,
page_size: int,
history_source: str | None = None,
):
"""
按生成日期倒序返回历史记录分组。
- 每页最多返回 10 个生成日期
- 每个日期分组内最多返回倒序前 10 条任务
- 只返回 completed 成功任务
"""
source = _normalize_history_source(history_source)
if source == "generation_record":
return await list_generation_record_history_grouped_days(
db=db,
user_id=user_id,
gen_type=gen_type,
page=page,
page_size=page_size,
)
gen_type = _normalize_history_gen_type(gen_type)
page = max(page, 1)
page_size = min(max(page_size, 1), HISTORY_DAY_PAGE_SIZE_MAX)
filters = _history_base_filters(user_id, gen_type)
day_expr = func.date(ChatGenerationTask.generated_at).label("generated_date")
days_subquery = (
select(day_expr)
.where(*filters)
.group_by(day_expr)
.subquery()
)
total_days = (
await db.execute(select(func.count()).select_from(days_subquery))
).scalar_one()
day_rows_result = await db.execute(
select(
day_expr,
func.count(ChatGenerationTask.id).label("total"),
)
.where(*filters)
.group_by(day_expr)
.order_by(day_expr.desc())
.offset((page - 1) * page_size)
.limit(page_size)
)
day_rows = day_rows_result.all()
groups = []
for generated_day, day_total in day_rows:
item_result = await db.execute(
select(ChatGenerationTask)
.where(
*filters,
func.date(ChatGenerationTask.generated_at) == generated_day,
)
.order_by(ChatGenerationTask.generated_at.desc(), ChatGenerationTask.created_at.desc())
.limit(HISTORY_GROUP_ITEM_LIMIT)
)
tasks = list(item_result.scalars().all())
groups.append(
{
"generated_date": _history_day_to_str(generated_day),
"total": int(day_total or 0),
"items": [record_to_out(task) for task in tasks],
}
)
return {
"total_days": int(total_days or 0),
"page": page,
"page_size": page_size,
"groups": groups,
}
async def list_generation_history_day_items(
db: AsyncSession,
user_id: str,
gen_type: str,
generated_date: str,
page: int,
page_size: int,
history_source: str | None = None,
):
"""
获取指定生成日期下的历史记录分页。
用于前端点击某一天后,继续加载该日期下的第 2 页、第 3 页数据。
"""
source = _normalize_history_source(history_source)
if source == "generation_record":
return await list_generation_record_history_day_items(
db=db,
user_id=user_id,
gen_type=gen_type,
generated_date=generated_date,
page=page,
page_size=page_size,
)
gen_type = _normalize_history_gen_type(gen_type)
target_day = _parse_history_date(generated_date)
page = max(page, 1)
page_size = min(max(page_size, 1), 100)
filters = _history_base_filters(user_id, gen_type)
day_expr = func.date(ChatGenerationTask.generated_at)
total = (
await db.execute(
select(func.count(ChatGenerationTask.id)).where(
*filters,
day_expr == target_day,
)
)
).scalar_one()
result = await db.execute(
select(ChatGenerationTask)
.where(
*filters,
day_expr == target_day,
)
.order_by(ChatGenerationTask.generated_at.desc(), ChatGenerationTask.created_at.desc())
.offset((page - 1) * page_size)
.limit(page_size)
)
tasks = list(result.scalars().all())
return {
"generated_date": target_day.strftime("%Y-%m-%d"),
"total": int(total or 0),
"page": page,
"page_size": page_size,
"items": [record_to_out(task) for task in tasks],
}