1
This commit is contained in:
@@ -0,0 +1,148 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Iterable
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MediaReferenceUsage:
|
||||
image_count: int = 0
|
||||
video_count: int = 0
|
||||
audio_count: int = 0
|
||||
input_video_duration: float = 0.0
|
||||
input_audio_duration: float = 0.0
|
||||
|
||||
|
||||
def parse_media_references(value: str | list[dict] | None) -> list[dict]:
|
||||
if not value:
|
||||
return []
|
||||
data: Any = value
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
data = json.loads(value)
|
||||
except (TypeError, json.JSONDecodeError):
|
||||
return []
|
||||
if not isinstance(data, list):
|
||||
return []
|
||||
return [item for item in data if isinstance(item, dict)]
|
||||
|
||||
|
||||
def _reference_type(item: dict) -> str:
|
||||
value = str(item.get("type") or item.get("media_type") or "").strip().lower()
|
||||
if value in {"image", "video", "audio"}:
|
||||
return value
|
||||
mime = str(item.get("mime_type") or item.get("content_type") or "").lower()
|
||||
if mime.startswith("image/"):
|
||||
return "image"
|
||||
if mime.startswith("video/"):
|
||||
return "video"
|
||||
if mime.startswith("audio/"):
|
||||
return "audio"
|
||||
url = str(item.get("url") or item.get("file_url") or "").lower().split("?", 1)[0]
|
||||
if url.endswith((".png", ".jpg", ".jpeg", ".webp", ".gif", ".bmp")):
|
||||
return "image"
|
||||
if url.endswith((".mp4", ".mov", ".webm", ".mkv", ".avi")):
|
||||
return "video"
|
||||
if url.endswith((".mp3", ".wav", ".m4a", ".aac", ".ogg", ".flac")):
|
||||
return "audio"
|
||||
return ""
|
||||
|
||||
|
||||
def _duration(item: dict) -> float:
|
||||
for key in ("duration", "duration_seconds", "video_duration", "audio_duration"):
|
||||
try:
|
||||
value = float(item.get(key) or 0)
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
if value > 0:
|
||||
return value
|
||||
return 0.0
|
||||
|
||||
|
||||
def calculate_media_reference_usage(
|
||||
references: str | list[dict] | None,
|
||||
*,
|
||||
include: bool,
|
||||
) -> MediaReferenceUsage:
|
||||
if not include:
|
||||
return MediaReferenceUsage()
|
||||
image_count = video_count = audio_count = 0
|
||||
video_duration = audio_duration = 0.0
|
||||
for item in parse_media_references(references):
|
||||
media_type = _reference_type(item)
|
||||
if media_type == "image":
|
||||
image_count += 1
|
||||
elif media_type == "video":
|
||||
video_count += 1
|
||||
video_duration += _duration(item)
|
||||
elif media_type == "audio":
|
||||
audio_count += 1
|
||||
audio_duration += _duration(item)
|
||||
return MediaReferenceUsage(
|
||||
image_count=image_count,
|
||||
video_count=video_count,
|
||||
audio_count=audio_count,
|
||||
input_video_duration=round(video_duration, 3),
|
||||
input_audio_duration=round(audio_duration, 3),
|
||||
)
|
||||
|
||||
|
||||
def filter_references_by_type(
|
||||
references: str | list[dict] | None,
|
||||
*,
|
||||
allowed_types: Iterable[str],
|
||||
max_count: int | None = None,
|
||||
) -> list[dict]:
|
||||
allowed = {str(item).lower() for item in allowed_types}
|
||||
result = [item for item in parse_media_references(references) if _reference_type(item) in allowed]
|
||||
if max_count is not None:
|
||||
return result[: max(0, int(max_count))]
|
||||
return result
|
||||
|
||||
|
||||
def validate_media_reference_usage_for_engine(
|
||||
usage: MediaReferenceUsage,
|
||||
*,
|
||||
gen_type: str,
|
||||
engine: Any,
|
||||
) -> None:
|
||||
"""按所选引擎能力校验最终实际会发送的附件。"""
|
||||
normalized = str(gen_type or "").strip().lower()
|
||||
if normalized == "video":
|
||||
if not bool(getattr(engine, "supports_universal_reference", True)) and (
|
||||
usage.image_count or usage.video_count or usage.audio_count
|
||||
):
|
||||
raise HTTPException(status_code=400, detail="当前视频引擎不支持参考附件")
|
||||
limits = {
|
||||
"图片": int(getattr(engine, "max_image_count", 0) or 0),
|
||||
"视频": int(getattr(engine, "max_video_count", 0) or 0),
|
||||
"音频": int(getattr(engine, "max_audio_count", 0) or 0),
|
||||
}
|
||||
counts = {"图片": usage.image_count, "视频": usage.video_count, "音频": usage.audio_count}
|
||||
for label, count in counts.items():
|
||||
limit = limits[label]
|
||||
if count > limit:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"当前视频引擎最多支持 {limit} 个{label}附件,当前为 {count} 个",
|
||||
)
|
||||
return
|
||||
if normalized == "image":
|
||||
if usage.video_count or usage.audio_count:
|
||||
raise HTTPException(status_code=400, detail="图片生成只能携带图片附件")
|
||||
limit = int(
|
||||
getattr(engine, "max_reference_image_count", None)
|
||||
if getattr(engine, "max_reference_image_count", None) is not None
|
||||
else getattr(engine, "max_image_count", 0)
|
||||
or 0
|
||||
)
|
||||
if usage.image_count > limit:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"当前图片引擎最多支持 {limit} 张参考图,当前为 {usage.image_count} 张",
|
||||
)
|
||||
return
|
||||
raise HTTPException(status_code=400, detail="不支持的生成类型")
|
||||
Reference in New Issue
Block a user