Files
video-gen/video-gen-api/app/services/video_prompt_schema_config_service.py

252 lines
12 KiB
Python

from __future__ import annotations
import copy
import json
from typing import Any
from fastapi import HTTPException
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.common import (
VIDEO_SCHEMA_CONFIG_DATABASE_SOURCE,
VIDEO_SCHEMA_CONFIG_DEFAULT_SOURCE,
VIDEO_SCHEMA_CONFIG_DESCRIPTION,
VIDEO_SCHEMA_CONFIG_KEY,
VIDEO_SCHEMA_CONFIG_VERSION,
VIDEO_SCHEMA_CONFIG_KEY_MAX_LEN,
VIDEO_SCHEMA_CONFIG_LABEL_MAX_LEN,
VIDEO_SCHEMA_CONFIG_DESC_MAX_LEN,
VIDEO_SCHEMA_EDITABLE_TEXT_MAX_LEN,
VIDEO_SCHEMA_FLOW_CONTENT_MAX_LEN,
VIDEO_SCHEMA_MAX_FIELD_COUNT_PER_SECTION,
VIDEO_SCHEMA_MAX_SECTION_COUNT,
VIDEO_SCHEMA_MAX_SEGMENT_COUNT_PER_RULE,
VIDEO_SCHEMA_MAX_TIME_RULE_COUNT,
VIDEO_SCHEMA_TIME_DESC_MAX_LEN,
VIDEO_SCHEMA_TIME_STAGE_MAX_LEN,
)
from app.models.system_config import SystemConfig
from app.services.hot_opening_video_prompt_service import (
build_dynamic_schema,
build_schema_config_snapshot,
default_video_prompt_schema_config,
normalize_video_prompt_schema_config,
)
from app.utils.id_gen import generate_id
def _json_dumps(data: Any) -> str:
return json.dumps(data, ensure_ascii=False, separators=(",", ":"))
def _json_loads(value: str | None) -> dict[str, Any] | None:
if not value or not str(value).strip():
return None
try:
data = json.loads(value)
except Exception as exc:
raise HTTPException(status_code=400, detail=f"视频提词 Schema 配置不是合法 JSON: {exc}") from exc
if not isinstance(data, dict):
raise HTTPException(status_code=400, detail="视频提词 Schema 配置必须是 JSON 对象")
return data
async def _get_record(db: AsyncSession) -> SystemConfig | None:
result = await db.execute(select(SystemConfig).where(SystemConfig.key == VIDEO_SCHEMA_CONFIG_KEY).limit(1))
return result.scalar_one_or_none()
def _validate_text_length(path: str, value: Any, max_length: int) -> None:
text = str(value or "")
if len(text) > max_length:
raise HTTPException(status_code=400, detail=f"字段【{path}】长度不能超过 {max_length} 个字符")
def validate_video_prompt_schema_config(data: dict[str, Any]) -> dict[str, Any]:
normalized = normalize_video_prompt_schema_config(data)
sections = normalized.get("sections") if isinstance(normalized.get("sections"), list) else []
if len(sections) > VIDEO_SCHEMA_MAX_SECTION_COUNT:
raise HTTPException(status_code=400, detail=f"最多允许配置 {VIDEO_SCHEMA_MAX_SECTION_COUNT} 个一级分组")
seen_sections: set[str] = set()
for section in sections:
key = str(section.get("key") or "")
if not key:
raise HTTPException(status_code=400, detail="一级分组 key 不能为空")
if key in seen_sections:
raise HTTPException(status_code=400, detail=f"一级分组 key 重复: {key}")
seen_sections.add(key)
_validate_text_length(f"{key}.key", key, VIDEO_SCHEMA_CONFIG_KEY_MAX_LEN)
_validate_text_length(f"{key}.label", section.get("label"), VIDEO_SCHEMA_CONFIG_LABEL_MAX_LEN)
_validate_text_length(f"{key}.description", section.get("description"), VIDEO_SCHEMA_CONFIG_DESC_MAX_LEN)
section_type = str(section.get("type") or "object")
if section_type == "object":
children = section.get("children") if isinstance(section.get("children"), list) else []
if len(children) > VIDEO_SCHEMA_MAX_FIELD_COUNT_PER_SECTION:
raise HTTPException(status_code=400, detail=f"分组【{key}】最多允许 {VIDEO_SCHEMA_MAX_FIELD_COUNT_PER_SECTION} 个字段")
seen_children: set[str] = set()
for child in children:
child_key = str(child.get("key") or "")
if not child_key:
raise HTTPException(status_code=400, detail=f"分组【{key}】存在空字段 key")
if child_key in seen_children:
raise HTTPException(status_code=400, detail=f"分组【{key}】字段 key 重复: {child_key}")
seen_children.add(child_key)
_validate_text_length(f"{key}.{child_key}.key", child_key, VIDEO_SCHEMA_CONFIG_KEY_MAX_LEN)
_validate_text_length(f"{key}.{child_key}.label", child.get("label"), VIDEO_SCHEMA_CONFIG_LABEL_MAX_LEN)
if child.get("editable"):
max_length = int(child.get("max_length") or VIDEO_SCHEMA_EDITABLE_TEXT_MAX_LEN)
_validate_text_length(f"{key}.{child_key}.value", child.get("value"), max_length)
elif section_type == "flow":
item_fields = section.get("item_fields") if isinstance(section.get("item_fields"), list) else []
if len(item_fields) > VIDEO_SCHEMA_MAX_FIELD_COUNT_PER_SECTION:
raise HTTPException(status_code=400, detail=f"流程【{key}】最多允许 {VIDEO_SCHEMA_MAX_FIELD_COUNT_PER_SECTION} 个对象字段")
seen_flow_fields: set[str] = set()
for field in item_fields:
field_key = str(field.get("key") or "")
if not field_key:
raise HTTPException(status_code=400, detail=f"流程【{key}】存在空字段 key")
if field_key == "时间段":
raise HTTPException(status_code=400, detail=f"流程【{key}】字段 key 不能使用锁定字段 时间段")
if field_key in seen_flow_fields:
raise HTTPException(status_code=400, detail=f"流程【{key}】字段 key 重复: {field_key}")
seen_flow_fields.add(field_key)
_validate_text_length(f"{key}.{field_key}.key", field_key, VIDEO_SCHEMA_CONFIG_KEY_MAX_LEN)
_validate_text_length(f"{key}.{field_key}.label", field.get("label"), VIDEO_SCHEMA_CONFIG_LABEL_MAX_LEN)
if field.get("editable"):
max_length = int(field.get("max_length") or VIDEO_SCHEMA_FLOW_CONTENT_MAX_LEN)
_validate_text_length(f"{key}.{field_key}.value", field.get("value"), max_length)
else:
if section.get("editable"):
max_length = int(section.get("max_length") or VIDEO_SCHEMA_EDITABLE_TEXT_MAX_LEN)
_validate_text_length(f"{key}.value", section.get("value"), max_length)
rules = normalized.get("time_plan_rules") if isinstance(normalized.get("time_plan_rules"), list) else []
if len(rules) > VIDEO_SCHEMA_MAX_TIME_RULE_COUNT:
raise HTTPException(status_code=400, detail=f"最多允许配置 {VIDEO_SCHEMA_MAX_TIME_RULE_COUNT} 条秒数切片规则")
for index, rule in enumerate(rules):
min_duration = int(rule.get("min_duration") or 0)
max_duration = int(rule.get("max_duration") or 0)
if min_duration <= 0 or max_duration <= 0 or min_duration > max_duration:
raise HTTPException(status_code=400, detail=f"第 {index + 1} 条秒数切片规则区间不合法")
segments = rule.get("segments") if isinstance(rule.get("segments"), list) else []
ratios = rule.get("ratios") if isinstance(rule.get("ratios"), list) else []
if not segments:
raise HTTPException(status_code=400, detail=f"第 {index + 1} 条秒数切片规则必须至少包含一个片段")
if len(segments) > VIDEO_SCHEMA_MAX_SEGMENT_COUNT_PER_RULE:
raise HTTPException(status_code=400, detail=f"第 {index + 1} 条秒数切片规则最多允许 {VIDEO_SCHEMA_MAX_SEGMENT_COUNT_PER_RULE} 个片段")
if len(ratios) != len(segments):
raise HTTPException(status_code=400, detail=f"第 {index + 1} 条秒数切片规则 ratios 数量必须和 segments 数量一致")
for segment_index, segment in enumerate(segments):
_validate_text_length(f"time_plan_rules[{index}].segments[{segment_index}].stage", segment.get("stage"), VIDEO_SCHEMA_TIME_STAGE_MAX_LEN)
_validate_text_length(f"time_plan_rules[{index}].segments[{segment_index}].description", segment.get("description"), VIDEO_SCHEMA_TIME_DESC_MAX_LEN)
return normalized
async def get_video_prompt_schema_config(db: AsyncSession) -> dict[str, Any]:
default_data = default_video_prompt_schema_config()
record = await _get_record(db)
if not record:
return {
"id": None,
"key": VIDEO_SCHEMA_CONFIG_KEY,
"description": VIDEO_SCHEMA_CONFIG_DESCRIPTION,
"is_enabled": False,
"using_default": True,
"data": copy.deepcopy(default_data),
"default_data": default_data,
"created_at": None,
"updated_at": None,
}
data = _json_loads(record.value)
if not data:
data = copy.deepcopy(default_data)
using_default = True
else:
data = normalize_video_prompt_schema_config(data)
using_default = False
return {
"id": record.id,
"key": record.key,
"description": record.description,
"is_enabled": bool(data.get("enabled", True)) and not using_default,
"using_default": using_default,
"data": data,
"default_data": default_data,
"created_at": record.created_at,
"updated_at": record.updated_at,
}
async def save_video_prompt_schema_config(db: AsyncSession, *, data: dict[str, Any], is_enabled: bool = True) -> dict[str, Any]:
normalized = validate_video_prompt_schema_config(data)
normalized["enabled"] = bool(is_enabled)
record = await _get_record(db)
if not record:
record = SystemConfig(
id=generate_id(),
key=VIDEO_SCHEMA_CONFIG_KEY,
value=_json_dumps(normalized),
description=VIDEO_SCHEMA_CONFIG_DESCRIPTION,
)
db.add(record)
else:
record.value = _json_dumps(normalized)
record.description = VIDEO_SCHEMA_CONFIG_DESCRIPTION
await db.flush()
return await get_video_prompt_schema_config(db)
async def reset_video_prompt_schema_config(db: AsyncSession) -> dict[str, Any]:
return await save_video_prompt_schema_config(db, data=default_video_prompt_schema_config(), is_enabled=True)
async def import_video_prompt_schema_config(db: AsyncSession, *, data: dict[str, Any], is_enabled: bool = True) -> dict[str, Any]:
return await save_video_prompt_schema_config(db, data=data, is_enabled=is_enabled)
async def export_video_prompt_schema_config(db: AsyncSession) -> dict[str, Any]:
current = await get_video_prompt_schema_config(db)
return {
"version": VIDEO_SCHEMA_CONFIG_VERSION,
"key": VIDEO_SCHEMA_CONFIG_KEY,
"is_enabled": current["is_enabled"],
"data": current["data"],
}
async def get_runtime_schema_snapshot(db: AsyncSession) -> dict[str, Any]:
record = await _get_record(db)
if not record:
return build_schema_config_snapshot(default_video_prompt_schema_config(), source=VIDEO_SCHEMA_CONFIG_DEFAULT_SOURCE)
data = _json_loads(record.value)
if not data:
return build_schema_config_snapshot(default_video_prompt_schema_config(), source=VIDEO_SCHEMA_CONFIG_DEFAULT_SOURCE)
normalized = normalize_video_prompt_schema_config(data)
if not normalized.get("enabled", True):
return build_schema_config_snapshot(default_video_prompt_schema_config(), source=VIDEO_SCHEMA_CONFIG_DEFAULT_SOURCE)
return build_schema_config_snapshot(normalized, source=VIDEO_SCHEMA_CONFIG_DATABASE_SOURCE)
def fallback_runtime_schema_snapshot(snapshot: dict[str, Any] | None = None) -> dict[str, Any]:
if isinstance(snapshot, dict) and isinstance(snapshot.get("data"), dict):
return snapshot
return build_schema_config_snapshot(default_video_prompt_schema_config(), source=VIDEO_SCHEMA_CONFIG_DEFAULT_SOURCE)
def preview_runtime_schema(*, data: dict[str, Any], is_enabled: bool, video_config: dict[str, Any]) -> dict[str, Any]:
normalized = validate_video_prompt_schema_config(data)
normalized["enabled"] = bool(is_enabled)
snapshot = build_schema_config_snapshot(
normalized if normalized.get("enabled", True) else default_video_prompt_schema_config(),
source=VIDEO_SCHEMA_CONFIG_DATABASE_SOURCE if normalized.get("enabled", True) else VIDEO_SCHEMA_CONFIG_DEFAULT_SOURCE,
)
runtime_schema = build_dynamic_schema(video_config, snapshot)
return {
"schema_config_snapshot": snapshot,
"runtime_schema": runtime_schema,
"time_plan": runtime_schema.get("动态时间规划") or [],
}