feat(#1970): 素材原子化切片 P1 - 数据层/切片逻辑/原子片段级选片 #1974
@@ -0,0 +1,58 @@
|
||||
"""add asset_atom_clips table
|
||||
|
||||
Revision ID: 079_asset_atom_clips
|
||||
Revises: 078_drop_script_title_fields
|
||||
Create Date: 2026-09-17
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "079_asset_atom_clips"
|
||||
down_revision = "078_drop_script_title_fields"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"asset_atom_clips",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column(
|
||||
"asset_id",
|
||||
sa.String(36),
|
||||
sa.ForeignKey("assets.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("start_time", sa.Float(), nullable=False),
|
||||
sa.Column("end_time", sa.Float(), nullable=False),
|
||||
sa.Column("duration", sa.Float(), nullable=False),
|
||||
sa.Column("clip_index", sa.Integer(), nullable=False),
|
||||
sa.Column("tags", sa.JSON(), nullable=False, server_default=sa.text("'[]'")),
|
||||
sa.Column("scene_change_at", sa.Float(), nullable=True),
|
||||
sa.Column(
|
||||
"is_fallback",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.text("false"),
|
||||
),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("NOW()"),
|
||||
),
|
||||
)
|
||||
# 按素材查片段并按索引排序(复合索引前缀可独立用于 asset_id 过滤)
|
||||
op.create_index(
|
||||
"ix_asset_atom_clips_asset_index",
|
||||
"asset_atom_clips",
|
||||
["asset_id", "clip_index"],
|
||||
unique=True,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_asset_atom_clips_asset_index", table_name="asset_atom_clips")
|
||||
op.drop_table("asset_atom_clips")
|
||||
@@ -0,0 +1,37 @@
|
||||
"""add edit_plan_clips.atom_clip_id for #1970
|
||||
|
||||
Revision ID: 080_edit_plan_clips_atom_clip_id
|
||||
Revises: 079_asset_atom_clips
|
||||
Create Date: 2026-09-17
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "080_edit_plan_clips_atom_clip_id"
|
||||
down_revision = "079_asset_atom_clips"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"edit_plan_clips",
|
||||
sa.Column(
|
||||
"atom_clip_id",
|
||||
sa.String(36),
|
||||
nullable=False,
|
||||
server_default=sa.text("''"),
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_edit_plan_clips_atom_clip_id",
|
||||
"edit_plan_clips",
|
||||
["atom_clip_id"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_edit_plan_clips_atom_clip_id", table_name="edit_plan_clips")
|
||||
op.drop_column("edit_plan_clips", "atom_clip_id")
|
||||
@@ -476,9 +476,14 @@ def create_generation_task(
|
||||
variant_plan_ids.append(_plan0.id)
|
||||
|
||||
# #1855 P0:批次区间避让表,从变体0实际clips构建初始值(公共函数)
|
||||
from app.services.generation_common import collect_plan_atom_clip_ids as _collect_atom_ids
|
||||
from app.services.generation_common import collect_plan_segments as _collect_segments
|
||||
|
||||
_batch_segments = _collect_segments(_plan0.id, _plan_svc._clip_repo)
|
||||
# #1970:批次内原子片段硬避让集合
|
||||
_batch_atom_ids: list[str] = _collect_atom_ids(
|
||||
_plan0.id, _plan_svc._clip_repo
|
||||
)
|
||||
|
||||
# 变体 1..N-1 独立选片(传入累积batch_segments做素材区间避让)
|
||||
for task_index in range(1, count):
|
||||
@@ -493,6 +498,7 @@ def create_generation_task(
|
||||
name_suffix=f"批量{task_index + 1}",
|
||||
voice_duration=voice_durations[task_index] if task_index < len(voice_durations) else 0.0,
|
||||
batch_segments=_batch_segments,
|
||||
batch_used_atom_ids=_batch_atom_ids,
|
||||
)
|
||||
break
|
||||
except ValueError as ve:
|
||||
@@ -529,6 +535,10 @@ def create_generation_task(
|
||||
_new_segs = _collect_segments(variant.id, _plan_svc._clip_repo)
|
||||
for _aid, _ivs in _new_segs.items():
|
||||
_batch_segments.setdefault(_aid, []).extend(_ivs)
|
||||
# #1970:同步累积原子片段ID
|
||||
_batch_atom_ids.extend(
|
||||
_collect_atom_ids(variant.id, _plan_svc._clip_repo)
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("[生成任务] 变体%d 区间收集失败(不阻断)", task_index)
|
||||
|
||||
|
||||
@@ -55,7 +55,9 @@ _DOUYIN_DEBUG_ERRORS = os.environ.get("DOUYIN_DEBUG_ERRORS", "").lower() in (
|
||||
"1",
|
||||
"true",
|
||||
"yes",
|
||||
) or os.environ.get("APP_ENV", "").lower() in ("staging", "dev", "development", "test")
|
||||
) or os.environ.get(
|
||||
"APP_ENV", ""
|
||||
).lower() in ("staging", "dev", "development", "test")
|
||||
|
||||
_TAIL_PUNCT = ".,;:!?,。;:!?))]》" + chr(34) + chr(39) + "<>"
|
||||
_URL_EXTRACT_RE = re.compile(r"https?://\S+", re.IGNORECASE)
|
||||
@@ -140,6 +142,7 @@ def _extract_and_validate_douyin_url(raw_input):
|
||||
|
||||
def _mk_post_json(self, path, payload):
|
||||
import httpx
|
||||
|
||||
if not self.is_available:
|
||||
raise MediaKitError("MediaKit API Key 未配置", code="NotConfigured")
|
||||
url = self._base_url + path
|
||||
@@ -168,6 +171,7 @@ def _mk_post_json(self, path, payload):
|
||||
|
||||
def _mk_get_json(self, path):
|
||||
import httpx
|
||||
|
||||
if not self.is_available:
|
||||
raise MediaKitError("MediaKit API Key 未配置", code="NotConfigured")
|
||||
url = self._base_url + path
|
||||
@@ -283,7 +287,7 @@ def _direct_url_download_and_local_asr(direct_url, page_url, temp_dir):
|
||||
raise
|
||||
except httpx.TimeoutException:
|
||||
logger.warning("直链下载超时: %s", page_url)
|
||||
raise HTTPException(status_code=status.HTTP_504_GATEWAY_TIMEOUT, detail="视频下载超时,请稍后重试")
|
||||
raise HTTPException(status_code=status.HTTP_504_GATEWAY_TIMEOUT, detail="视频下载超时,请稍后重试") from None
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.exception("直链下载失败: url=%s err=%s", page_url, exc)
|
||||
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail="视频下载失败: " + str(exc)[:200]) from exc
|
||||
@@ -420,7 +424,10 @@ def extract_from_douyin(
|
||||
if text:
|
||||
logger.info(
|
||||
"抖音 MediaKit ASR 成功: source=%s text_len=%d duration=%.1f total_time=%.1fs",
|
||||
result.source, len(text), duration, time.time() - t0,
|
||||
result.source,
|
||||
len(text),
|
||||
duration,
|
||||
time.time() - t0,
|
||||
)
|
||||
else:
|
||||
logger.info("抖音 MediaKit ASR 返回空文本(无旁白/BGM视频)")
|
||||
@@ -440,7 +447,9 @@ def extract_from_douyin(
|
||||
if text:
|
||||
logger.info(
|
||||
"抖音本地 ASR 成功: source=%s text_len=%d total_time=%.1fs",
|
||||
result.source, len(text), time.time() - t0,
|
||||
result.source,
|
||||
len(text),
|
||||
time.time() - t0,
|
||||
)
|
||||
last_err_stage = "asr"
|
||||
except HTTPException as exc:
|
||||
@@ -539,6 +548,7 @@ def ai_generate_titles(
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="文案内容不能为空")
|
||||
count = max(1, min(5, request.count))
|
||||
from app.services.ai_service import generate_smart_titles
|
||||
|
||||
result = generate_smart_titles(description=content, style="viral", count=count)
|
||||
titles = result.get("titles", [])[:count]
|
||||
return AiGenerateTitlesResponse(titles=titles)
|
||||
|
||||
@@ -52,7 +52,7 @@ def _extract_url_from_text(text: str) -> str:
|
||||
if not text:
|
||||
return ""
|
||||
m = re.search(r"https?://\S+", text)
|
||||
return m.group(0).rstrip("。,!?!?,,;;\"'))】") if m else ""
|
||||
return m.group(0).rstrip("。,!?!?,,;;\"'))】") if m else "" # noqa: B005
|
||||
|
||||
|
||||
def _canonicalize_url(url: str, timeout: int = 8) -> str:
|
||||
|
||||
@@ -423,6 +423,7 @@ class EditPlanService:
|
||||
clip_type=clip.clip_type,
|
||||
order=clip.order,
|
||||
asset_id=clip.asset_id,
|
||||
atom_clip_id=clip_item.get("atom_clip_id", ""),
|
||||
text_content=clip.text_content,
|
||||
start_time=clip.start_time,
|
||||
duration=clip.duration,
|
||||
@@ -474,6 +475,7 @@ class EditPlanService:
|
||||
voice_duration: float = 0.0,
|
||||
rng=None,
|
||||
batch_segments: dict[str, list[tuple[float, float]]] | None = None,
|
||||
batch_used_atom_ids: set[str] | list[str] | None = None,
|
||||
) -> EditPlan:
|
||||
"""为批量变体生成独立 plan:完整重跑单视频选片流程(#1743)。
|
||||
|
||||
@@ -608,18 +610,69 @@ class EditPlanService:
|
||||
st = float(c.start_time or 0.0)
|
||||
batch_segments_resolved.setdefault(c.asset_id, []).append((st, st + float(c.duration)))
|
||||
|
||||
clips_data = reselect_clips_for_variant(
|
||||
source_clips_data,
|
||||
pool_ids,
|
||||
asset_durations=durations,
|
||||
asset_scene_points=scene_points,
|
||||
historical_used_segments=historical,
|
||||
batch_segments=batch_segments_resolved,
|
||||
target_durations=target_durations,
|
||||
rng=rng,
|
||||
)
|
||||
clips_data = None
|
||||
# #1970 原子片段级变体重选:候选素材已切片时优先按原子片段选片
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import (
|
||||
SQLAlchemyAssetAtomClipRepository,
|
||||
)
|
||||
from packages.domain.atom_clip_resolver import flatten_candidates, load_atom_clips_for_assets
|
||||
from packages.domain.atom_clip_selector import reselect_clips_from_atoms
|
||||
|
||||
# 片段区间写回素材 metadata(与落库同事务;replace_all_clips_transactional 内 commit)
|
||||
atom_repo = SQLAlchemyAssetAtomClipRepository(db)
|
||||
|
||||
# 兜底切片只需要时长;本方法已查出 durations,封装一个只读假素材仓储
|
||||
class _DurationOnlyAssetRepo:
|
||||
def __init__(self, durations_map: dict[str, float]) -> None:
|
||||
self._durations = durations_map
|
||||
|
||||
def get(self, asset_id: str):
|
||||
if asset_id not in self._durations:
|
||||
return None
|
||||
|
||||
class _A:
|
||||
pass
|
||||
|
||||
a = _A()
|
||||
a.duration = self._durations[asset_id]
|
||||
return a
|
||||
|
||||
clips_by_asset = load_atom_clips_for_assets(
|
||||
pool_ids,
|
||||
atom_clip_repo=atom_repo,
|
||||
asset_repo=_DurationOnlyAssetRepo(durations),
|
||||
)
|
||||
atom_candidates = flatten_candidates(clips_by_asset)
|
||||
if atom_candidates:
|
||||
# 历史成片已用原子片段(降权);批次内前序变体已用(硬避让)
|
||||
historical_atom_ids = set(
|
||||
self._clip_repo.list_recent_atom_clip_ids_by_user(
|
||||
created_by_user_id or source.created_by_user_id or "",
|
||||
limit=200,
|
||||
)
|
||||
)
|
||||
clips_data = reselect_clips_from_atoms(
|
||||
source_clips_data,
|
||||
atom_candidates,
|
||||
historical_atom_ids=historical_atom_ids,
|
||||
batch_used_atom_ids=(set(batch_used_atom_ids) if batch_used_atom_ids else None),
|
||||
rng=rng,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning("原子片段变体重选失败,回退整条素材选片", exc_info=True)
|
||||
clips_data = None
|
||||
|
||||
if clips_data is None:
|
||||
clips_data = reselect_clips_for_variant(
|
||||
source_clips_data,
|
||||
pool_ids,
|
||||
asset_durations=durations,
|
||||
asset_scene_points=scene_points,
|
||||
historical_used_segments=historical,
|
||||
batch_segments=batch_segments_resolved,
|
||||
target_durations=target_durations,
|
||||
rng=rng,
|
||||
) # 片段区间写回素材 metadata(与落库同事务;replace_all_clips_transactional 内 commit)
|
||||
for item in clips_data:
|
||||
aid = item.get("asset_id", "")
|
||||
if aid:
|
||||
|
||||
@@ -157,6 +157,33 @@ def collect_plan_segments(
|
||||
return segs
|
||||
|
||||
|
||||
def collect_plan_atom_clip_ids(
|
||||
plan_id: str,
|
||||
clip_repo: Any,
|
||||
*,
|
||||
page_size: int = 500,
|
||||
) -> list[str]:
|
||||
"""分页读取 plan 所有 clips,收集已选用的原子片段 ID(#1970)。
|
||||
|
||||
用于批量变体间原子片段级硬避让:同一原子片段在同批次内只用一次。
|
||||
旧路径 clips 的 atom_clip_id 为空串,自动忽略。
|
||||
"""
|
||||
ids: list[str] = []
|
||||
sk, pg = 0, page_size
|
||||
while True:
|
||||
batch = clip_repo.list_by_plan(plan_id, skip=sk, limit=pg)
|
||||
if not batch:
|
||||
break
|
||||
for c in batch:
|
||||
acid = getattr(c, "atom_clip_id", "") or ""
|
||||
if acid:
|
||||
ids.append(acid)
|
||||
if len(batch) < pg:
|
||||
break
|
||||
sk += pg
|
||||
return ids
|
||||
|
||||
|
||||
def resolve_latest_plan_by_template(
|
||||
db: Session,
|
||||
*,
|
||||
|
||||
@@ -22,6 +22,11 @@ from packages.adapters.sqlalchemy_impl import (
|
||||
SQLAlchemyEditPlanClipRepository,
|
||||
SQLAlchemyEditPlanRepository,
|
||||
)
|
||||
from packages.domain.atom_clip_resolver import load_atom_clips_for_assets
|
||||
from packages.domain.atom_clip_selector import (
|
||||
estimate_required_clip_count,
|
||||
select_atom_clips,
|
||||
)
|
||||
from packages.domain.config_schemas import normalize_plan_config
|
||||
from packages.domain.edit_plan import EditPlan
|
||||
from packages.domain.edit_plan_clip import EditPlanClip
|
||||
@@ -52,10 +57,12 @@ class PlanGeneratorService:
|
||||
基于模板 + 素材,自动生成 EditPlan 及 EditPlanClip 列表。
|
||||
"""
|
||||
|
||||
def __init__(self, db: Session, asset_repo=None) -> None:
|
||||
def __init__(self, db: Session, asset_repo=None, atom_clip_repo=None) -> None:
|
||||
self._plan_repo = SQLAlchemyEditPlanRepository(db)
|
||||
self._clip_repo = SQLAlchemyEditPlanClipRepository(db)
|
||||
self._asset_repo = asset_repo
|
||||
# #1970 原子化切片:可选注入;未注入时走旧的整条素材选片路径(向后兼容)
|
||||
self._atom_clip_repo = atom_clip_repo
|
||||
|
||||
# ── 公开接口 ─────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -121,18 +128,34 @@ class PlanGeneratorService:
|
||||
|
||||
# 4. 按 editing_mode 分配素材
|
||||
if asset_ids:
|
||||
# 获取素材时长信息,用于随机起始时间
|
||||
asset_durations = None
|
||||
if self._asset_repo:
|
||||
asset_durations = self._fetch_asset_durations(asset_ids)
|
||||
self._distribute_assets(
|
||||
clips,
|
||||
asset_ids,
|
||||
editing_mode,
|
||||
random_selection=random_preview,
|
||||
asset_durations=asset_durations,
|
||||
user_id=created_by_user_id,
|
||||
)
|
||||
# #1970 原子化切片:素材 clip 从 atom_clips 表选取(未就绪自动内存兜底)。
|
||||
# 预览随机模式保持旧路径(整条素材 + 随机起点),与现有预览契约一致。
|
||||
atom_applied = False
|
||||
if not random_preview and self._atom_clip_repo is not None:
|
||||
try:
|
||||
atom_applied = self._distribute_atom_clips(
|
||||
clips,
|
||||
asset_ids,
|
||||
editing_mode,
|
||||
user_id=created_by_user_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning("原子片段选片失败,回退整条素材选片", exc_info=True)
|
||||
atom_applied = False
|
||||
|
||||
if not atom_applied:
|
||||
# 获取素材时长信息,用于随机起始时间
|
||||
asset_durations = None
|
||||
if self._asset_repo:
|
||||
asset_durations = self._fetch_asset_durations(asset_ids)
|
||||
self._distribute_assets(
|
||||
clips,
|
||||
asset_ids,
|
||||
editing_mode,
|
||||
random_selection=random_preview,
|
||||
asset_durations=asset_durations,
|
||||
user_id=created_by_user_id,
|
||||
)
|
||||
|
||||
# 5. 持久化所有 clips 并计算总时长
|
||||
created_clips: list[EditPlanClip] = []
|
||||
@@ -259,6 +282,103 @@ class PlanGeneratorService:
|
||||
external_used_segments=external_used_segments,
|
||||
)
|
||||
|
||||
def _distribute_atom_clips(
|
||||
self,
|
||||
clips: list[EditPlanClip],
|
||||
asset_ids: list[str],
|
||||
editing_mode: str,
|
||||
*,
|
||||
user_id: str = "",
|
||||
) -> bool:
|
||||
"""#1970 原子化切片选片(就地修改 clips,未持久化).
|
||||
|
||||
从 ``asset_atom_clips`` 表按原子片段选取;老素材/切片未就绪的素材
|
||||
内存兜底切片。同一原子片段在一次方案中只用一次;跨视频避让走
|
||||
edit_plan_clips.atom_clip_id 最近使用记录。
|
||||
|
||||
Returns:
|
||||
True 表示原子片段选片成功;False 表示无可用片段,调用方应回退
|
||||
到旧的整条素材 distribute_assets。
|
||||
"""
|
||||
# 1. 加载候选原子片段(DB + 兜底)
|
||||
clips_by_asset = load_atom_clips_for_assets(
|
||||
asset_ids,
|
||||
atom_clip_repo=self._atom_clip_repo,
|
||||
asset_repo=self._asset_repo,
|
||||
)
|
||||
if not clips_by_asset:
|
||||
return False
|
||||
|
||||
# 2. 最近使用片段(跨视频原子片段级避让)
|
||||
recently_used: set[str] = set()
|
||||
if user_id and hasattr(self._clip_repo, "list_recent_atom_clip_ids_by_user"):
|
||||
try:
|
||||
recently_used = set(self._clip_repo.list_recent_atom_clip_ids_by_user(user_id, limit=200))
|
||||
except Exception:
|
||||
logger.warning("跨视频原子片段避让查询失败", exc_info=True)
|
||||
|
||||
# 3. 片段需求估算:无配音时按 clips 数量;voice_over 的配音总时长存于
|
||||
# clip.config["voice_duration"],按 平均片段时长≈需要片段数 估算
|
||||
voice_total = 0.0
|
||||
for c in clips:
|
||||
cfg_vd = c.config.get("voice_duration") if c.config else None
|
||||
if cfg_vd:
|
||||
voice_total += float(cfg_vd)
|
||||
avg_clip_target = sum(float(c.duration or 0.0) for c in clips) / max(len(clips), 1)
|
||||
required_count = estimate_required_clip_count(
|
||||
voice_total or sum(float(c.duration or 0.0) for c in clips),
|
||||
avg_clip_target or 3.5,
|
||||
)
|
||||
required_count = max(required_count, len(clips))
|
||||
|
||||
rng = random.Random()
|
||||
|
||||
# 4. 正式生成:先按素材 smart_score 对素材池排序,再展开为片段池
|
||||
# (同素材的片段保持连续,高分素材的片段排在前面优先入选)
|
||||
if self._asset_repo:
|
||||
asset_order = self._sort_assets_by_smart_score(list(clips_by_asset.keys()))
|
||||
ordered: dict[str, list] = {}
|
||||
for aid in asset_order:
|
||||
if aid in clips_by_asset:
|
||||
ordered[aid] = clips_by_asset[aid]
|
||||
clips_by_asset = ordered
|
||||
|
||||
candidates: list = []
|
||||
for asset_clips in clips_by_asset.values():
|
||||
candidates.extend(asset_clips)
|
||||
|
||||
# 5. 逐虚拟片段选片:评分排序,同片段不重复使用
|
||||
used_atom_ids: set[str] = set()
|
||||
asset_usage: dict[str, int] = {}
|
||||
assigned = 0
|
||||
for clip in clips:
|
||||
# 对每个虚拟片段重新评分(usage_count 随选择动态变化)
|
||||
scored = select_atom_clips(
|
||||
candidates,
|
||||
target_duration=float(clip.duration or 0.0),
|
||||
used_atom_clip_ids=used_atom_ids,
|
||||
asset_usage_counts=asset_usage,
|
||||
recently_used_atom_ids=recently_used,
|
||||
required_count=required_count,
|
||||
limit=1,
|
||||
rng=rng,
|
||||
)
|
||||
if not scored:
|
||||
# 候选耗尽(同片段不可重复),交由调用方回退或留白
|
||||
continue
|
||||
picked = scored[0]
|
||||
clip.asset_id = picked.asset_id
|
||||
clip.atom_clip_id = picked.atom_clip_id
|
||||
clip.start_time = round(picked.start_time, 3)
|
||||
clip.duration = round(picked.duration, 3)
|
||||
used_atom_ids.add(picked.atom_clip_id)
|
||||
asset_usage[picked.asset_id] = asset_usage.get(picked.asset_id, 0) + 1
|
||||
assigned += 1
|
||||
|
||||
if assigned == 0:
|
||||
return False
|
||||
return True
|
||||
|
||||
def _fetch_asset_scene_points(self, asset_ids: list[str]) -> dict[str, list[float]]:
|
||||
"""从素材 metadata 读取场景切换点缓存(无缓存的素材不包含在结果中)。"""
|
||||
points_map: dict[str, list[float]] = {}
|
||||
|
||||
@@ -27,6 +27,7 @@ celery_app.conf.broker_transport_options = {"visibility_timeout": 4 * 60 * 60}
|
||||
celery_app.conf.imports = (
|
||||
"worker_app.tasks.health",
|
||||
"worker_app.tasks.ingest",
|
||||
"worker_app.tasks.atom_clips",
|
||||
"worker_app.tasks.classification",
|
||||
"worker_app.tasks.generation",
|
||||
"worker_app.tasks.voice_extraction",
|
||||
|
||||
@@ -53,12 +53,17 @@ def __getattr__(name: str):
|
||||
from .batch_thumbnail import batch_generate_thumbnails
|
||||
|
||||
return batch_generate_thumbnails
|
||||
elif name == "generate_atom_clips":
|
||||
from .atom_clips import generate_atom_clips
|
||||
|
||||
return generate_atom_clips
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
|
||||
__all__ = [
|
||||
"batch_generate_thumbnails",
|
||||
"classify_asset",
|
||||
"generate_atom_clips",
|
||||
"generate_video",
|
||||
"healthcheck",
|
||||
"ingest_asset",
|
||||
|
||||
@@ -0,0 +1,81 @@
|
||||
"""素材原子切片 Celery 任务 — #1970 智能剪辑流程重构 P1.
|
||||
|
||||
素材入库预处理完成(ingest 置 READY)后异步触发:
|
||||
根据素材时长和已缓存的 scdet 切换点计算原子片段并落库。
|
||||
失败不阻断素材入库主流程(atom_clips 未就绪时选片有内存兜底)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from celery.utils.log import get_task_logger
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import (
|
||||
SQLAlchemyAssetAtomClipRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
|
||||
from packages.domain.atom_clip_service import compute_atom_clips
|
||||
from packages.domain.plan_generator_utils import extract_scene_points_from_metadata
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
|
||||
|
||||
@celery_app.task(name="worker.generate_atom_clips")
|
||||
def generate_atom_clips(asset_id: str) -> dict:
|
||||
"""为单条视频素材生成原子片段。
|
||||
|
||||
Returns:
|
||||
任务结果 dict:status / asset_id / clips_count。
|
||||
"""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
asset_repo = SQLAlchemyAssetRepository(db)
|
||||
atom_repo = SQLAlchemyAssetAtomClipRepository(db)
|
||||
|
||||
asset = asset_repo.find_by_id(asset_id)
|
||||
if asset is None:
|
||||
return {"status": "skipped", "reason": "asset not found", "asset_id": asset_id}
|
||||
|
||||
# 仅视频素材切片
|
||||
if asset.mime_type and not asset.mime_type.startswith("video/"):
|
||||
return {"status": "skipped", "reason": "not a video", "asset_id": asset_id}
|
||||
if not asset.duration or asset.duration <= 0:
|
||||
return {"status": "skipped", "reason": "invalid duration", "asset_id": asset_id}
|
||||
|
||||
# 已生成过则幂等跳过(重新切片需先显式删除)
|
||||
existing = atom_repo.count_by_asset(asset_id)
|
||||
if existing > 0:
|
||||
return {
|
||||
"status": "skipped",
|
||||
"reason": "already generated",
|
||||
"asset_id": asset_id,
|
||||
"clips_count": existing,
|
||||
}
|
||||
|
||||
scene_points = extract_scene_points_from_metadata(asset.metadata)
|
||||
# P1 阶段继承素材的标签 ID;片段级语义标签是 P2 功能
|
||||
tags = list(getattr(asset, "tag_ids", []) or [])
|
||||
|
||||
clips = compute_atom_clips(
|
||||
asset_id=asset_id,
|
||||
duration=float(asset.duration),
|
||||
scene_change_points=scene_points,
|
||||
tags=tags,
|
||||
)
|
||||
if not clips:
|
||||
return {"status": "skipped", "reason": "no clips computed", "asset_id": asset_id}
|
||||
|
||||
atom_repo.batch_create(clips)
|
||||
logger.info(
|
||||
"[atom_clips] asset_id=%s 生成 %d 个原子片段",
|
||||
asset_id,
|
||||
len(clips),
|
||||
)
|
||||
return {"status": "completed", "asset_id": asset_id, "clips_count": len(clips)}
|
||||
except Exception as exc: # noqa: BLE001 - 后台任务兜底,失败不阻断主流程
|
||||
db.rollback()
|
||||
logger.exception("[atom_clips] asset_id=%s 生成失败: %s", asset_id, exc)
|
||||
return {"status": "failed", "asset_id": asset_id, "error": str(exc)}
|
||||
finally:
|
||||
db.close()
|
||||
@@ -808,6 +808,21 @@ def ingest_asset(job_id: str) -> dict:
|
||||
|
||||
db.commit()
|
||||
|
||||
# ── #1970 素材原子切片:视频 READY 后异步触发,失败不阻断入库 ──
|
||||
# atom_clips 未就绪时选片逻辑有内存兜底(compute_fallback_clips)。
|
||||
try:
|
||||
if media_type == "video" and float(asset.duration or 0) > 0:
|
||||
celery_app.send_task(
|
||||
"worker.generate_atom_clips",
|
||||
args=[asset.id],
|
||||
)
|
||||
except Exception as atom_err: # noqa: BLE001
|
||||
logger.warning(
|
||||
"触发原子切片任务失败(不影响入库): asset_id=%s err=%s",
|
||||
asset.id,
|
||||
atom_err,
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "completed",
|
||||
"job_id": job.id,
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
"""素材原子片段仓储 SQLAlchemy 实现。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetAtomClipModel
|
||||
from packages.domain.asset_atom_clip import AssetAtomClip
|
||||
|
||||
|
||||
class SQLAlchemyAssetAtomClipRepository:
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
def create(self, clip: AssetAtomClip) -> AssetAtomClip:
|
||||
model = self._to_model(clip)
|
||||
self.session.add(model)
|
||||
self.session.flush()
|
||||
self.session.commit()
|
||||
return clip
|
||||
|
||||
def batch_create(self, clips: list[AssetAtomClip]) -> list[AssetAtomClip]:
|
||||
if not clips:
|
||||
return []
|
||||
models = [self._to_model(c) for c in clips]
|
||||
self.session.add_all(models)
|
||||
self.session.flush()
|
||||
self.session.commit()
|
||||
return clips
|
||||
|
||||
def find_by_asset(self, asset_id: str) -> list[AssetAtomClip]:
|
||||
models = (
|
||||
self.session.query(AssetAtomClipModel)
|
||||
.filter(AssetAtomClipModel.asset_id == asset_id)
|
||||
.order_by(AssetAtomClipModel.clip_index.asc())
|
||||
.all()
|
||||
)
|
||||
return [self._to_domain(m) for m in models]
|
||||
|
||||
def find_by_id(self, clip_id: str) -> AssetAtomClip | None:
|
||||
model = self.session.query(AssetAtomClipModel).filter(AssetAtomClipModel.id == clip_id).first()
|
||||
if model is None:
|
||||
return None
|
||||
return self._to_domain(model)
|
||||
|
||||
def find_by_ids(self, clip_ids: list[str]) -> list[AssetAtomClip]:
|
||||
if not clip_ids:
|
||||
return []
|
||||
models = self.session.query(AssetAtomClipModel).filter(AssetAtomClipModel.id.in_(clip_ids)).all()
|
||||
return [self._to_domain(m) for m in models]
|
||||
|
||||
def delete_by_asset(self, asset_id: str) -> int:
|
||||
count = (
|
||||
self.session.query(AssetAtomClipModel)
|
||||
.filter(AssetAtomClipModel.asset_id == asset_id)
|
||||
.delete(synchronize_session=False)
|
||||
)
|
||||
self.session.commit()
|
||||
return count
|
||||
|
||||
def count_by_asset(self, asset_id: str) -> int:
|
||||
return self.session.query(AssetAtomClipModel).filter(AssetAtomClipModel.asset_id == asset_id).count()
|
||||
|
||||
def find_candidates_for_selection(
|
||||
self,
|
||||
asset_ids: list[str],
|
||||
*,
|
||||
min_duration: float | None = None,
|
||||
max_duration: float | None = None,
|
||||
limit: int = 100,
|
||||
) -> list[AssetAtomClip]:
|
||||
"""按筛选条件查找候选原子片段,按时长排序。用于选片逻辑。"""
|
||||
query = self.session.query(AssetAtomClipModel).filter(AssetAtomClipModel.asset_id.in_(asset_ids))
|
||||
if min_duration is not None:
|
||||
query = query.filter(AssetAtomClipModel.duration >= min_duration)
|
||||
if max_duration is not None:
|
||||
query = query.filter(AssetAtomClipModel.duration <= max_duration)
|
||||
query = query.order_by(AssetAtomClipModel.clip_index.asc())
|
||||
if limit > 0:
|
||||
query = query.limit(limit)
|
||||
models = query.all()
|
||||
return [self._to_domain(m) for m in models]
|
||||
|
||||
def _to_model(self, clip: AssetAtomClip) -> AssetAtomClipModel:
|
||||
return AssetAtomClipModel(
|
||||
id=clip.id,
|
||||
asset_id=clip.asset_id,
|
||||
start_time=clip.start_time,
|
||||
end_time=clip.end_time,
|
||||
duration=clip.duration,
|
||||
clip_index=clip.clip_index,
|
||||
tags=clip.tags,
|
||||
scene_change_at=clip.scene_change_at,
|
||||
is_fallback=clip.is_fallback,
|
||||
created_at=clip.created_at or datetime.now(UTC),
|
||||
)
|
||||
|
||||
def _to_domain(self, model: AssetAtomClipModel) -> AssetAtomClip:
|
||||
return AssetAtomClip(
|
||||
id=model.id,
|
||||
asset_id=model.asset_id,
|
||||
start_time=model.start_time,
|
||||
end_time=model.end_time,
|
||||
duration=model.duration,
|
||||
clip_index=model.clip_index,
|
||||
tags=model.tags or [],
|
||||
scene_change_at=model.scene_change_at,
|
||||
is_fallback=model.is_fallback,
|
||||
created_at=model.created_at,
|
||||
)
|
||||
@@ -50,6 +50,7 @@ class SQLAlchemyEditPlanClipRepository:
|
||||
order=clip.order,
|
||||
template_clip_config_id=clip.template_clip_config_id,
|
||||
asset_id=clip.asset_id,
|
||||
atom_clip_id=getattr(clip, "atom_clip_id", "") or "",
|
||||
text_content=clip.text_content,
|
||||
start_time=clip.start_time,
|
||||
duration=clip.duration,
|
||||
@@ -74,6 +75,7 @@ class SQLAlchemyEditPlanClipRepository:
|
||||
model.order = clip.order
|
||||
model.template_clip_config_id = clip.template_clip_config_id
|
||||
model.asset_id = clip.asset_id
|
||||
model.atom_clip_id = getattr(clip, "atom_clip_id", "") or ""
|
||||
model.text_content = clip.text_content
|
||||
model.start_time = clip.start_time
|
||||
model.duration = clip.duration
|
||||
@@ -120,6 +122,7 @@ class SQLAlchemyEditPlanClipRepository:
|
||||
order=model.order,
|
||||
template_clip_config_id=model.template_clip_config_id or "",
|
||||
asset_id=model.asset_id or "",
|
||||
atom_clip_id=getattr(model, "atom_clip_id", "") or "",
|
||||
text_content=model.text_content or "",
|
||||
start_time=model.start_time or 0.0,
|
||||
duration=model.duration or 0.0,
|
||||
@@ -193,3 +196,53 @@ class SQLAlchemyEditPlanClipRepository:
|
||||
result[asset_id].append((start_time or 0.0, (start_time or 0.0) + (duration or 0.0)))
|
||||
|
||||
return result
|
||||
|
||||
def list_recent_atom_clip_ids_by_user(
|
||||
self,
|
||||
user_id: str,
|
||||
*,
|
||||
limit: int = 200,
|
||||
) -> list[str]:
|
||||
"""#1970 跨视频原子片段级避让:查询用户最近成片用过的 atom_clip_id.
|
||||
|
||||
只统计已完成 plan 下已渲染且 atom_clip_id 非空的 clips,按 plan
|
||||
创建时间倒序,返回去重后的 ID 列表。
|
||||
"""
|
||||
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
|
||||
|
||||
if not user_id:
|
||||
return []
|
||||
|
||||
recent_plan_ids = [
|
||||
row[0]
|
||||
for row in self.session.query(EditPlanModel.id)
|
||||
.filter(
|
||||
EditPlanModel.created_by_user_id == user_id,
|
||||
EditPlanModel.status == "completed",
|
||||
)
|
||||
.order_by(EditPlanModel.created_at.desc())
|
||||
.limit(50)
|
||||
.all()
|
||||
]
|
||||
if not recent_plan_ids:
|
||||
return []
|
||||
|
||||
rows = (
|
||||
self.session.query(EditPlanClipModel.atom_clip_id)
|
||||
.filter(
|
||||
EditPlanClipModel.plan_id.in_(recent_plan_ids),
|
||||
EditPlanClipModel.status == "rendered",
|
||||
EditPlanClipModel.atom_clip_id.isnot(None),
|
||||
EditPlanClipModel.atom_clip_id != "",
|
||||
)
|
||||
.all()
|
||||
)
|
||||
seen: set[str] = set()
|
||||
ordered: list[str] = []
|
||||
for (atom_clip_id,) in rows:
|
||||
if atom_clip_id and atom_clip_id not in seen:
|
||||
seen.add(atom_clip_id)
|
||||
ordered.append(atom_clip_id)
|
||||
if len(ordered) >= limit:
|
||||
break
|
||||
return ordered
|
||||
|
||||
@@ -1,7 +1,20 @@
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import JSON, Boolean, Column, DateTime, Float, Index, Integer, String, Text, UniqueConstraint, text
|
||||
from sqlalchemy import (
|
||||
JSON,
|
||||
Boolean,
|
||||
Column,
|
||||
DateTime,
|
||||
Float,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
String,
|
||||
Text,
|
||||
UniqueConstraint,
|
||||
text,
|
||||
)
|
||||
from sqlalchemy.orm import declarative_base
|
||||
|
||||
Base: Any = declarative_base()
|
||||
@@ -234,6 +247,8 @@ class EditPlanClipModel(Base):
|
||||
order = Column(Integer, nullable=False)
|
||||
template_clip_config_id = Column(String(36), nullable=False, default="", index=True)
|
||||
asset_id = Column(String(36), nullable=False, default="", index=True)
|
||||
# #1970 原子化切片:片段选中的原子片段 ID(空串表示旧的整条素材选取路径)
|
||||
atom_clip_id = Column(String(36), nullable=False, default="", index=True)
|
||||
text_content = Column(Text, nullable=False, default="")
|
||||
start_time = Column(Float, nullable=False, default=0.0)
|
||||
duration = Column(Float, nullable=False, default=0.0)
|
||||
@@ -801,6 +816,32 @@ class PointsOrderModel(Base):
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class AssetAtomClipModel(Base):
|
||||
"""素材原子片段 ORM 模型 (#1970 智能剪辑流程重构)。
|
||||
|
||||
逻辑切分单元,不物理切割视频文件。
|
||||
"""
|
||||
|
||||
__tablename__ = "asset_atom_clips"
|
||||
__table_args__ = (UniqueConstraint("asset_id", "clip_index", name="uq_asset_atom_clips_asset_index"),)
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
asset_id = Column(
|
||||
String(36),
|
||||
ForeignKey("assets.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
index=True,
|
||||
)
|
||||
start_time = Column(Float, nullable=False)
|
||||
end_time = Column(Float, nullable=False)
|
||||
duration = Column(Float, nullable=False)
|
||||
clip_index = Column(Integer, nullable=False)
|
||||
tags = Column(JSON, nullable=False, default=list)
|
||||
scene_change_at = Column(Float, nullable=True)
|
||||
is_fallback = Column(Boolean, nullable=False, default=False)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class DailyUsageRecordModel(Base):
|
||||
"""每日使用记录 ORM 模型 (#1895)"""
|
||||
|
||||
|
||||
@@ -1,5 +1,18 @@
|
||||
"""Domain package for core business entities and rules."""
|
||||
|
||||
from . import atom_clip_resolver
|
||||
from .asset_atom_clip import AssetAtomClip
|
||||
from .atom_clip_selector import (
|
||||
ScoredAtomClip,
|
||||
clips_to_segments,
|
||||
estimate_required_clip_count,
|
||||
score_atom_clip,
|
||||
select_atom_clips,
|
||||
)
|
||||
from .atom_clip_service import (
|
||||
compute_atom_clips,
|
||||
compute_fallback_clips,
|
||||
)
|
||||
from .classification import (
|
||||
AssetClassification,
|
||||
ClassificationJob,
|
||||
@@ -37,6 +50,15 @@ from .voice_library import VoiceLibraryItem
|
||||
|
||||
__all__ = [
|
||||
"Asset",
|
||||
"AssetAtomClip",
|
||||
"ScoredAtomClip",
|
||||
"clips_to_segments",
|
||||
"compute_atom_clips",
|
||||
"compute_fallback_clips",
|
||||
"estimate_required_clip_count",
|
||||
"score_atom_clip",
|
||||
"select_atom_clips",
|
||||
"atom_clip_resolver",
|
||||
"AssetClassification",
|
||||
"DailyUsageRecord",
|
||||
"PointsAccount",
|
||||
|
||||
@@ -0,0 +1,81 @@
|
||||
"""素材原子片段(Atom Clip)领域实体 — #1970 智能剪辑流程重构。
|
||||
|
||||
原子片段是素材的逻辑切分单元,不物理切割视频文件。
|
||||
每条记录指向某条素材的一段 [start_time, end_time] 区间。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
|
||||
|
||||
@dataclass
|
||||
class AssetAtomClip:
|
||||
"""素材原子片段。
|
||||
|
||||
Attributes:
|
||||
id: 唯一标识。
|
||||
asset_id: 所属素材 ID。
|
||||
start_time: 片段起始时间(秒,浮点)。
|
||||
end_time: 片段结束时间(秒,浮点)。
|
||||
duration: 片段时长 = end_time - start_time(秒)。
|
||||
clip_index: 在同一素材内的顺序编号(从 0 开始)。
|
||||
tags: 继承自素材的标签,JSONB 存储,可为空列表。
|
||||
scene_change_at: 片段尾部是否对齐了 scdet 镜头切换点(存储该切点的精确时间),
|
||||
未对齐时为 None。
|
||||
is_fallback: 是否为兜底逻辑在内存中生成的临时片段(不入库)。
|
||||
created_at: 创建时间。
|
||||
"""
|
||||
|
||||
id: str
|
||||
asset_id: str
|
||||
start_time: float
|
||||
end_time: float
|
||||
duration: float
|
||||
clip_index: int
|
||||
tags: list[str] = field(default_factory=list)
|
||||
scene_change_at: float | None = None
|
||||
is_fallback: bool = False
|
||||
created_at: datetime | None = None
|
||||
|
||||
def __post_init__(self):
|
||||
if not self.id:
|
||||
self.id = str(uuid.uuid4())
|
||||
if self.duration <= 0:
|
||||
self.duration = round(self.end_time - self.start_time, 3)
|
||||
if self.duration < 0:
|
||||
raise ValueError(f"duration must be >= 0, got start={self.start_time}, end={self.end_time}")
|
||||
if self.start_time < 0:
|
||||
raise ValueError(f"start_time must be >= 0, got {self.start_time}")
|
||||
if self.end_time <= self.start_time:
|
||||
raise ValueError(f"end_time must be > start_time, got start={self.start_time}, end={self.end_time}")
|
||||
if self.clip_index < 0:
|
||||
raise ValueError(f"clip_index must be >= 0, got {self.clip_index}")
|
||||
if self.created_at is None:
|
||||
self.created_at = datetime.now(UTC)
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
asset_id: str,
|
||||
start_time: float,
|
||||
end_time: float,
|
||||
clip_index: int,
|
||||
tags: list[str] | None = None,
|
||||
scene_change_at: float | None = None,
|
||||
is_fallback: bool = False,
|
||||
) -> AssetAtomClip:
|
||||
"""工厂方法:创建一个新的原子片段。"""
|
||||
return cls(
|
||||
id="", # __post_init__ 会自动生成
|
||||
asset_id=asset_id,
|
||||
start_time=round(start_time, 3),
|
||||
end_time=round(end_time, 3),
|
||||
duration=round(end_time - start_time, 3),
|
||||
clip_index=clip_index,
|
||||
tags=tags or [],
|
||||
scene_change_at=scene_change_at,
|
||||
is_fallback=is_fallback,
|
||||
)
|
||||
@@ -0,0 +1,104 @@
|
||||
"""原子片段加载与兜底 — #1970 智能剪辑流程重构 P1.
|
||||
|
||||
选片前从 ``asset_atom_clips`` 表加载素材池的原子片段;老素材/切片任务尚未
|
||||
完成/切片失败导致某些素材没有片段时,按需求兜底:内存中按 3-6 秒临时均匀
|
||||
切片(不存库,片段标记 is_fallback=True)。
|
||||
|
||||
本模块对 repository 做鸭子类型约束(只需 find_by_asset / find_candidates_for_selection
|
||||
和 asset_repo.get),方便 API 侧(SQLAlchemy)与 worker 侧复用,也便于单测注入内存假实现。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from packages.domain.asset_atom_clip import AssetAtomClip
|
||||
from packages.domain.atom_clip_service import compute_fallback_clips
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 兜底均匀切片步长(秒),落在 3~6s 区间中段
|
||||
FALLBACK_CLIP_SECONDS = 4.5
|
||||
|
||||
|
||||
def load_atom_clips_for_assets(
|
||||
asset_ids: list[str],
|
||||
*,
|
||||
atom_clip_repo,
|
||||
asset_repo=None,
|
||||
) -> dict[str, list[AssetAtomClip]]:
|
||||
"""加载素材池的原子片段(缺失素材走内存兜底).
|
||||
|
||||
Args:
|
||||
asset_ids: 候选素材 ID(去重保序)。
|
||||
atom_clip_repo: AssetAtomClipRepository 实现(需有
|
||||
``find_candidates_for_selection`` 或 ``find_by_asset``)。
|
||||
asset_repo: 可选,素材仓储(需有 ``get``),用于读取时长兜底切片。
|
||||
为 None 时,没有原子片段的素材直接跳过(不兜底)。
|
||||
|
||||
Returns:
|
||||
{asset_id: [AssetAtomClip, ...]},仅包含至少有一个片段的素材,
|
||||
片段按 clip_index 排序。
|
||||
"""
|
||||
result: dict[str, list[AssetAtomClip]] = {}
|
||||
unique_ids = list(dict.fromkeys(asset_ids))
|
||||
if not unique_ids:
|
||||
return result
|
||||
|
||||
# 1. 批量查询已生成的原子片段
|
||||
persisted: dict[str, list[AssetAtomClip]] = {}
|
||||
try:
|
||||
if hasattr(atom_clip_repo, "find_candidates_for_selection"):
|
||||
clips = atom_clip_repo.find_candidates_for_selection(unique_ids, limit=0)
|
||||
else:
|
||||
clips = []
|
||||
for asset_id in unique_ids:
|
||||
clips.extend(atom_clip_repo.find_by_asset(asset_id))
|
||||
for clip in clips:
|
||||
persisted.setdefault(clip.asset_id, []).append(clip)
|
||||
except Exception:
|
||||
logger.warning("加载 atom_clips 失败,全部走内存兜底", exc_info=True)
|
||||
persisted = {}
|
||||
|
||||
for asset_id in unique_ids:
|
||||
clips = persisted.get(asset_id)
|
||||
if clips:
|
||||
clips.sort(key=lambda c: c.clip_index)
|
||||
result[asset_id] = clips
|
||||
continue
|
||||
|
||||
# 2. 兜底:内存均匀切片(不存库)
|
||||
if asset_repo is None:
|
||||
continue
|
||||
duration = _safe_asset_duration(asset_repo, asset_id)
|
||||
if duration <= 0:
|
||||
continue
|
||||
result[asset_id] = compute_fallback_clips(
|
||||
asset_id,
|
||||
duration,
|
||||
clip_seconds=FALLBACK_CLIP_SECONDS,
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def flatten_candidates(
|
||||
clips_by_asset: dict[str, list[AssetAtomClip]],
|
||||
) -> list[AssetAtomClip]:
|
||||
"""把 {asset_id: [clips]} 摊平为候选片段列表(素材顺序内片段有序)。"""
|
||||
flat: list[AssetAtomClip] = []
|
||||
for clips in clips_by_asset.values():
|
||||
flat.extend(clips)
|
||||
return flat
|
||||
|
||||
|
||||
def _safe_asset_duration(asset_repo, asset_id: str) -> float:
|
||||
"""安全读取素材时长,任何异常返回 0。"""
|
||||
try:
|
||||
asset = asset_repo.get(asset_id)
|
||||
if asset is None:
|
||||
return 0.0
|
||||
return float(getattr(asset, "duration", 0.0) or 0.0)
|
||||
except Exception:
|
||||
logger.warning("读取素材时长失败: asset_id=%s", asset_id, exc_info=True)
|
||||
return 0.0
|
||||
@@ -0,0 +1,264 @@
|
||||
"""原子片段级选片核心 — #1970 智能剪辑流程重构 P1.
|
||||
|
||||
选片单元从"整条素材 + 随机起点"升级为"原子片段(atom clip)":
|
||||
|
||||
- 每个 EditPlanClip 指向一个 atom_clip_id(含 asset_id + start/end);
|
||||
- 同一素材的不同原子片段可被同一视频多次选用;
|
||||
- 同一原子片段在一个视频内只用一次;
|
||||
- 跨变体/跨任务的避让升级为原子片段级(同 asset 的不同片段天然不重叠);
|
||||
- atom_clips 未就绪(老素材/切片失败)时由调用方走内存兜底切片,
|
||||
再不行回退到现有的整条素材随机起点逻辑。
|
||||
|
||||
本模块是纯函数:原子片段数据由调用方从 repository 读取后注入,不直接碰 DB,
|
||||
便于单元测试。评分维度与 smart_match 保持一致(质量分、时长适配、新鲜度、
|
||||
未使用加分),只是评分对象从素材变为原子片段。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import random
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from packages.domain.asset_atom_clip import AssetAtomClip
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ScoredAtomClip:
|
||||
"""带评分的候选原子片段。"""
|
||||
|
||||
clip: AssetAtomClip
|
||||
score: float
|
||||
|
||||
@property
|
||||
def atom_clip_id(self) -> str:
|
||||
return self.clip.id
|
||||
|
||||
@property
|
||||
def asset_id(self) -> str:
|
||||
return self.clip.asset_id
|
||||
|
||||
@property
|
||||
def start_time(self) -> float:
|
||||
return self.clip.start_time
|
||||
|
||||
@property
|
||||
def end_time(self) -> float:
|
||||
return self.clip.end_time
|
||||
|
||||
@property
|
||||
def duration(self) -> float:
|
||||
return self.clip.duration
|
||||
|
||||
|
||||
# 评分权重(与 smart_match.score_asset 的维度对齐)
|
||||
W_QUALITY = 0.35
|
||||
W_DURATION_FIT = 0.30
|
||||
W_FRESHNESS = 0.15
|
||||
W_UNUSED_BONUS = 0.10
|
||||
W_ASSET_BALANCE = 0.10
|
||||
|
||||
# 评分随机噪声上限(与 SCORE_RANDOM_NOISE_MAX 同量级,避免反复选同一组合)
|
||||
SCORE_NOISE_MAX = 0.05
|
||||
|
||||
|
||||
def score_atom_clip(
|
||||
clip: AssetAtomClip,
|
||||
*,
|
||||
target_duration: float,
|
||||
asset_quality: dict[str, float] | None = None,
|
||||
asset_freshness: dict[str, float] | None = None,
|
||||
used_in_video: set[str] | None = None,
|
||||
asset_usage_counts: dict[str, int] | None = None,
|
||||
recently_used: set[str] | None = None,
|
||||
required_count: int = 1,
|
||||
total_candidates: int = 1,
|
||||
) -> float:
|
||||
"""评估单个原子片段对某个目标槽位的适配分(越高越优先).
|
||||
|
||||
评分维度:
|
||||
- 质量分(继承素材质量,缺省中性 0.6);
|
||||
- 时长适配(片段时长越接近目标越好,覆盖不满显著扣分);
|
||||
- 新鲜度(缺省中性 0.5);
|
||||
- 未使用加分(本视频内未用过 +1,已用 0);
|
||||
- 素材均衡(同一素材在本视频用得越多,其剩余片段扣分越多,鼓励分散到多素材);
|
||||
- 跨视频/历史使用降权(recently_used 中的片段扣分,不硬禁)。
|
||||
"""
|
||||
asset_quality = asset_quality or {}
|
||||
asset_freshness = asset_freshness or {}
|
||||
used_in_video = used_in_video or set()
|
||||
asset_usage_counts = asset_usage_counts or {}
|
||||
recently_used = recently_used or set()
|
||||
|
||||
quality = asset_quality.get(clip.asset_id, 0.6)
|
||||
|
||||
if target_duration > 0:
|
||||
coverage = min(1.0, clip.duration / target_duration)
|
||||
overshoot = max(0.0, (clip.duration - target_duration) / target_duration)
|
||||
duration_fit = max(0.0, coverage - 0.15 * overshoot)
|
||||
else:
|
||||
duration_fit = 0.5
|
||||
|
||||
freshness = asset_freshness.get(clip.asset_id, 0.5)
|
||||
unused_bonus = 0.0 if clip.id in used_in_video else 1.0
|
||||
|
||||
# 素材均衡:该素材已被本视频选用 k 次,其片段逐次扣分
|
||||
times_used = asset_usage_counts.get(clip.asset_id, 0)
|
||||
balance = 1.0 / (1.0 + times_used)
|
||||
|
||||
# 跨视频/历史使用降权(不硬禁)
|
||||
history_penalty = 0.35 if clip.id in recently_used else 0.0
|
||||
|
||||
score = (
|
||||
W_QUALITY * quality
|
||||
+ W_DURATION_FIT * duration_fit
|
||||
+ W_FRESHNESS * freshness
|
||||
+ W_UNUSED_BONUS * unused_bonus
|
||||
+ W_ASSET_BALANCE * balance
|
||||
- history_penalty
|
||||
)
|
||||
return score
|
||||
|
||||
|
||||
def select_atom_clips(
|
||||
candidates: list[AssetAtomClip],
|
||||
*,
|
||||
target_duration: float = 0.0,
|
||||
used_atom_clip_ids: set[str] | None = None,
|
||||
asset_usage_counts: dict[str, int] | None = None,
|
||||
recently_used_atom_ids: set[str] | None = None,
|
||||
required_count: int = 1,
|
||||
limit: int = 0,
|
||||
asset_quality: dict[str, float] | None = None,
|
||||
asset_freshness: dict[str, float] | None = None,
|
||||
rng: random.Random | None = None,
|
||||
) -> list[ScoredAtomClip]:
|
||||
"""为一个目标槽位从候选原子片段中评分选片(纯函数).
|
||||
|
||||
Args:
|
||||
candidates: 候选原子片段(可跨多素材)。
|
||||
target_duration: 槽位目标时长(秒)。
|
||||
used_atom_clip_ids: 本视频已用过的原子片段 ID(硬排除,同片段不重复)。
|
||||
asset_usage_counts: 本视频各素材已选片段数(均衡评分用)。
|
||||
recently_used_atom_ids: 跨视频/历史成片用过的片段 ID(降权,不硬禁)。
|
||||
required_count: 整个视频需要的片段总数(预留,供覆盖策略判断)。
|
||||
limit: 最多返回条数;<=0 表示返回全部排序结果。
|
||||
asset_quality / asset_freshness: 评分注入。
|
||||
rng: 可选随机源(测试注入)。
|
||||
|
||||
Returns:
|
||||
评分降序的 ScoredAtomClip 列表(已排除本视频用过的片段)。
|
||||
"""
|
||||
rng = rng or random.Random()
|
||||
used = used_atom_clip_ids or set()
|
||||
asset_usage_counts = asset_usage_counts or {}
|
||||
recently_used = recently_used_atom_ids or set()
|
||||
|
||||
available = [c for c in candidates if c.id not in used]
|
||||
scored: list[ScoredAtomClip] = []
|
||||
for clip in available:
|
||||
base = score_atom_clip(
|
||||
clip,
|
||||
target_duration=target_duration,
|
||||
asset_quality=asset_quality,
|
||||
asset_freshness=asset_freshness,
|
||||
used_in_video=used,
|
||||
asset_usage_counts=asset_usage_counts,
|
||||
recently_used=recently_used,
|
||||
required_count=required_count,
|
||||
total_candidates=len(candidates),
|
||||
)
|
||||
noise = rng.uniform(0.0, SCORE_NOISE_MAX)
|
||||
scored.append(ScoredAtomClip(clip=clip, score=base + noise))
|
||||
|
||||
scored.sort(key=lambda s: s.score, reverse=True)
|
||||
if limit and limit > 0:
|
||||
return scored[:limit]
|
||||
return scored
|
||||
|
||||
|
||||
def clips_to_segments(clips: list[AssetAtomClip]) -> dict[str, list[tuple[float, float]]]:
|
||||
"""把选中的原子片段转换为旧的 {asset_id: [(start, end), ...]} 区间结构.
|
||||
|
||||
用于与现有跨变体区间避让(variant_plan_selector / metadata.used_segments)对接。
|
||||
原子片段级天然不重叠,同素材多片段直接形成多段不重叠区间。
|
||||
"""
|
||||
segments: dict[str, list[tuple[float, float]]] = {}
|
||||
for clip in clips:
|
||||
segments.setdefault(clip.asset_id, []).append((clip.start_time, clip.end_time))
|
||||
for asset_id in segments:
|
||||
segments[asset_id].sort()
|
||||
return segments
|
||||
|
||||
|
||||
def estimate_required_clip_count(
|
||||
voice_total_duration: float,
|
||||
average_clip_duration: float = 4.5,
|
||||
) -> int:
|
||||
"""配音总时长 / 平均片段时长 ≈ 需要的片段数(至少 1)。"""
|
||||
if voice_total_duration <= 0 or average_clip_duration <= 0:
|
||||
return 1
|
||||
return max(1, round(voice_total_duration / average_clip_duration))
|
||||
|
||||
|
||||
def reselect_clips_from_atoms(
|
||||
source_clips: list[dict[str, Any]],
|
||||
candidates: list[AssetAtomClip],
|
||||
*,
|
||||
historical_atom_ids: set[str] | None = None,
|
||||
batch_used_atom_ids: set[str] | None = None,
|
||||
rng: random.Random | None = None,
|
||||
) -> list[dict[str, Any]] | None:
|
||||
"""#1970 变体重选的原子片段级实现.
|
||||
|
||||
与 variant_plan_selector.reselect_clips_for_variant 对应:保留源 plan 的
|
||||
片段骨架(order/clip_type/文案/转场),从候选原子片段中为每个 main 片段
|
||||
选取一个原子片段;同变体/批次内同一片段不可重复,历史成片用过的片段降权。
|
||||
|
||||
Returns:
|
||||
新 clips_data(dict 列表,含 asset_id/atom_clip_id/start_time/duration),
|
||||
候选不足(main 片段多于去重后片段数)时返回 None,由调用方回退整条素材路径。
|
||||
非 main 片段(intro/outro 等)原样保留不分配素材。
|
||||
"""
|
||||
if not source_clips or not candidates:
|
||||
return None
|
||||
|
||||
rng = rng or random.Random()
|
||||
main_indexes = [i for i, c in enumerate(source_clips) if c.get("clip_type", "main") == "main"]
|
||||
if len(main_indexes) > len({c.id for c in candidates}):
|
||||
return None
|
||||
|
||||
used: set[str] = set(batch_used_atom_ids or ())
|
||||
result: list[dict[str, Any]] = [dict(c) for c in source_clips]
|
||||
asset_usage: dict[str, int] = {}
|
||||
|
||||
for idx in main_indexes:
|
||||
skeleton = source_clips[idx]
|
||||
target_duration = float(skeleton.get("duration") or 0.0)
|
||||
ranked = select_atom_clips(
|
||||
candidates,
|
||||
target_duration=target_duration,
|
||||
used_atom_clip_ids=used,
|
||||
asset_usage_counts=asset_usage,
|
||||
recently_used_atom_ids=historical_atom_ids or set(),
|
||||
required_count=len(main_indexes),
|
||||
limit=1,
|
||||
rng=rng,
|
||||
)
|
||||
if not ranked:
|
||||
return None
|
||||
picked = ranked[0]
|
||||
# 段长:片段短于槽位时取片段全长(渲染末帧冻结铺满),长于槽位时按槽位时长 trim
|
||||
new_duration = picked.duration if target_duration <= 0 else min(target_duration, picked.duration)
|
||||
result[idx].update(
|
||||
{
|
||||
"asset_id": picked.asset_id,
|
||||
"atom_clip_id": picked.atom_clip_id,
|
||||
"start_time": round(picked.start_time, 3),
|
||||
"duration": round(new_duration, 3),
|
||||
}
|
||||
)
|
||||
used.add(picked.atom_clip_id)
|
||||
asset_usage[picked.asset_id] = asset_usage.get(picked.asset_id, 0) + 1
|
||||
|
||||
return result
|
||||
@@ -0,0 +1,215 @@
|
||||
"""素材原子切片服务 — #1970 智能剪辑流程重构 P1.
|
||||
|
||||
切片规则(见 docs/smart-edit-flow-redesign-20260916.md §1):
|
||||
- 3~6 秒一个片段,具体时长在此范围内随机(避免固定节奏)
|
||||
- 切点附近 0.5 秒内有 scdet 镜头切换点时,切点偏移到切换处
|
||||
(复用素材 metadata 中已缓存的 scene_change_points,不重新计算)
|
||||
- <6 秒素材整条作为一个片段,不切
|
||||
- 最后一个片段不足 3 秒的合并到前一个;超过 3 秒独立成段
|
||||
- 片段是逻辑索引,不物理切割视频文件
|
||||
|
||||
片段在内存中计算;持久化由上层调用 repository 完成,保证本模块可单测、无 IO 依赖。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import random
|
||||
|
||||
from packages.domain.asset_atom_clip import AssetAtomClip
|
||||
|
||||
# 切片参数(集中常量,便于后续抽配置)
|
||||
MIN_CLIP_SECONDS = 3.0
|
||||
MAX_CLIP_SECONDS = 6.0
|
||||
# 切点与 scdet 切换点的对齐窗口
|
||||
SCENE_SNAP_WINDOW = 0.5
|
||||
# 末段最小独立时长:不足则并入前一段
|
||||
MIN_TAIL_SECONDS = 3.0
|
||||
# 浮点比较容差
|
||||
_EPS = 0.05
|
||||
|
||||
|
||||
def _round3(value: float) -> float:
|
||||
return round(float(value), 3)
|
||||
|
||||
|
||||
def _snap_to_scene(
|
||||
cut: float,
|
||||
scene_points: list[float] | None,
|
||||
lower: float,
|
||||
upper: float,
|
||||
) -> tuple[float, float | None]:
|
||||
"""将切点 ``cut`` 对齐到窗口内最近的 scdet 切换点.
|
||||
|
||||
Args:
|
||||
cut: 原始切点(秒)。
|
||||
scene_points: 候选切换点(秒,已排序),可为空。
|
||||
lower: 允许偏移的下界(不早于当前片段起点)。
|
||||
upper: 允许偏移的上界(不晚于素材总时长)。
|
||||
|
||||
Returns:
|
||||
(对齐后的切点, 命中的切换点);未命中返回 (cut, None)。
|
||||
"""
|
||||
if not scene_points:
|
||||
return cut, None
|
||||
|
||||
best: float | None = None
|
||||
best_dist = SCENE_SNAP_WINDOW
|
||||
for point in scene_points:
|
||||
# 切换点必须严格落在片段内部(不能与边界重合),且在窗口内
|
||||
if point <= lower + _EPS or point >= upper - _EPS:
|
||||
continue
|
||||
dist = abs(point - cut)
|
||||
if dist <= best_dist:
|
||||
best_dist = dist
|
||||
best = point
|
||||
if best is None:
|
||||
return cut, None
|
||||
return _round3(best), _round3(best)
|
||||
|
||||
|
||||
def compute_atom_clips(
|
||||
asset_id: str,
|
||||
duration: float,
|
||||
*,
|
||||
scene_change_points: list[float] | None = None,
|
||||
tags: list[str] | None = None,
|
||||
rng: random.Random | None = None,
|
||||
) -> list[AssetAtomClip]:
|
||||
"""根据素材时长计算原子片段(纯函数,不落库).
|
||||
|
||||
Args:
|
||||
asset_id: 素材 ID。
|
||||
duration: 素材总时长(秒)。
|
||||
scene_change_points: metadata 中缓存的 scdet 切换点(秒)。
|
||||
tags: 继承自素材的标签。
|
||||
rng: 可选随机源(测试可注入固定种子)。
|
||||
|
||||
Returns:
|
||||
有序的原子片段列表(clip_index 从 0 开始)。
|
||||
"""
|
||||
if duration <= 0:
|
||||
return []
|
||||
|
||||
r = rng or random.Random()
|
||||
points = _normalize_scene_points(scene_change_points, duration)
|
||||
|
||||
# <6 秒素材整条作为一个片段,不切
|
||||
if duration < MAX_CLIP_SECONDS:
|
||||
return [
|
||||
AssetAtomClip.create(
|
||||
asset_id=asset_id,
|
||||
start_time=0.0,
|
||||
end_time=_round3(duration),
|
||||
clip_index=0,
|
||||
tags=list(tags or []),
|
||||
)
|
||||
]
|
||||
|
||||
boundaries: list[float] = [0.0]
|
||||
scene_hits: dict[int, float] = {}
|
||||
|
||||
cursor = 0.0
|
||||
while duration - cursor > MAX_CLIP_SECONDS + _EPS:
|
||||
# 在 [3, 6] 内随机决定本段目标时长
|
||||
target_len = r.uniform(MIN_CLIP_SECONDS, MAX_CLIP_SECONDS)
|
||||
raw_cut = cursor + target_len
|
||||
if raw_cut >= duration - _EPS:
|
||||
break
|
||||
cut, hit = _snap_to_scene(raw_cut, points, lower=cursor, upper=duration)
|
||||
|
||||
# 对齐后若导致本段短于 3 秒(切换点太靠近段首),放弃对齐
|
||||
if cut - cursor < MIN_CLIP_SECONDS - _EPS:
|
||||
cut = _round3(raw_cut)
|
||||
hit = None
|
||||
|
||||
boundaries.append(_round3(cut))
|
||||
if hit is not None:
|
||||
scene_hits[len(boundaries) - 1] = hit
|
||||
cursor = cut
|
||||
|
||||
boundaries.append(_round3(duration))
|
||||
|
||||
# 末段处理:最后一个片段不足 3 秒则合并到前一个
|
||||
if len(boundaries) >= 3:
|
||||
tail_start = boundaries[-2]
|
||||
tail_len = duration - tail_start
|
||||
if tail_len < MIN_TAIL_SECONDS - _EPS:
|
||||
boundaries.pop(-2)
|
||||
|
||||
clips: list[AssetAtomClip] = []
|
||||
for index in range(len(boundaries) - 1):
|
||||
start = boundaries[index]
|
||||
end = boundaries[index + 1]
|
||||
if end - start < _EPS:
|
||||
continue
|
||||
# 片段尾部对齐的切换点 = 该片段右边界(若它来自 snap)
|
||||
scene_at = scene_hits.get(index + 1)
|
||||
clips.append(
|
||||
AssetAtomClip.create(
|
||||
asset_id=asset_id,
|
||||
start_time=start,
|
||||
end_time=end,
|
||||
clip_index=index,
|
||||
tags=list(tags or []),
|
||||
scene_change_at=scene_at,
|
||||
)
|
||||
)
|
||||
return clips
|
||||
|
||||
|
||||
def compute_fallback_clips(
|
||||
asset_id: str,
|
||||
duration: float,
|
||||
*,
|
||||
tags: list[str] | None = None,
|
||||
clip_seconds: float = 4.5,
|
||||
) -> list[AssetAtomClip]:
|
||||
"""兜底切片:atom_clips 未就绪时,内存中按固定步长临时均匀切片(不存库).
|
||||
|
||||
与 :func:`compute_atom_clips` 的区别:不随机、不对齐切点,
|
||||
产出的片段标记 ``is_fallback=True``。
|
||||
"""
|
||||
if duration <= 0:
|
||||
return []
|
||||
|
||||
step = min(max(clip_seconds, MIN_CLIP_SECONDS), MAX_CLIP_SECONDS)
|
||||
clips: list[AssetAtomClip] = []
|
||||
cursor = 0.0
|
||||
index = 0
|
||||
while cursor < duration - _EPS:
|
||||
end = min(cursor + step, duration)
|
||||
clips.append(
|
||||
AssetAtomClip.create(
|
||||
asset_id=asset_id,
|
||||
start_time=_round3(cursor),
|
||||
end_time=_round3(end),
|
||||
clip_index=index,
|
||||
tags=list(tags or []),
|
||||
is_fallback=True,
|
||||
)
|
||||
)
|
||||
cursor = end
|
||||
index += 1
|
||||
|
||||
# 末段不足 3 秒合并
|
||||
if len(clips) >= 2 and clips[-1].duration < MIN_TAIL_SECONDS - _EPS:
|
||||
last = clips.pop()
|
||||
prev = clips[-1]
|
||||
merged = AssetAtomClip.create(
|
||||
asset_id=asset_id,
|
||||
start_time=prev.start_time,
|
||||
end_time=last.end_time,
|
||||
clip_index=prev.clip_index,
|
||||
tags=list(tags or []),
|
||||
is_fallback=True,
|
||||
)
|
||||
clips[-1] = merged
|
||||
return clips
|
||||
|
||||
|
||||
def _normalize_scene_points(points: list[float] | None, duration: float) -> list[float]:
|
||||
"""清洗切换点:去重、排序、限定在 (0, duration) 内。"""
|
||||
if not points:
|
||||
return []
|
||||
cleaned = sorted({round(float(p), 3) for p in points if 0 < float(p) < duration})
|
||||
return cleaned
|
||||
@@ -65,6 +65,7 @@ class EditPlanClip:
|
||||
order: int
|
||||
template_clip_config_id: str = ""
|
||||
asset_id: str = ""
|
||||
atom_clip_id: str = ""
|
||||
text_content: str = ""
|
||||
start_time: float = 0.0
|
||||
duration: float = 0.0
|
||||
@@ -85,6 +86,7 @@ class EditPlanClip:
|
||||
*,
|
||||
template_clip_config_id: str = "",
|
||||
asset_id: str = "",
|
||||
atom_clip_id: str = "",
|
||||
text_content: str = "",
|
||||
start_time: float = 0.0,
|
||||
duration: float = 0.0,
|
||||
@@ -117,6 +119,7 @@ class EditPlanClip:
|
||||
order=order,
|
||||
template_clip_config_id=template_clip_config_id.strip() if template_clip_config_id else "",
|
||||
asset_id=asset_id.strip() if asset_id else "",
|
||||
atom_clip_id=atom_clip_id.strip() if atom_clip_id else "",
|
||||
text_content=text_content.strip(),
|
||||
start_time=start_time,
|
||||
duration=duration,
|
||||
@@ -127,16 +130,25 @@ class EditPlanClip:
|
||||
config=config or {},
|
||||
)
|
||||
|
||||
def assign_asset(self, asset_id: str, *, start_time: float | None = None) -> None:
|
||||
def assign_asset(
|
||||
self,
|
||||
asset_id: str,
|
||||
*,
|
||||
start_time: float | None = None,
|
||||
atom_clip_id: str | None = None,
|
||||
) -> None:
|
||||
"""分配素材
|
||||
|
||||
Args:
|
||||
asset_id: 素材 ID
|
||||
start_time: 可选,素材播放起始时间(秒)。如果提供且在有效范围内,则设置;否则保持默认 0.0
|
||||
atom_clip_id: 可选,选中的原子片段 ID(#1970 原子化切片)。
|
||||
"""
|
||||
if not asset_id.strip():
|
||||
raise ValueError("asset_id 不能为空")
|
||||
self.asset_id = asset_id.strip()
|
||||
if atom_clip_id is not None:
|
||||
self.atom_clip_id = atom_clip_id.strip() if atom_clip_id else ""
|
||||
if start_time is not None and start_time >= 0:
|
||||
self.start_time = start_time
|
||||
self.updated_at = datetime.now(UTC)
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
"""素材原子片段仓储接口定义。"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from packages.domain.asset_atom_clip import AssetAtomClip
|
||||
|
||||
|
||||
class AssetAtomClipRepository(ABC):
|
||||
@abstractmethod
|
||||
def create(self, clip: AssetAtomClip) -> AssetAtomClip:
|
||||
"""创建一条原子片段记录。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def batch_create(self, clips: list[AssetAtomClip]) -> list[AssetAtomClip]:
|
||||
"""批量创建原子片段记录。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_by_asset(self, asset_id: str) -> list[AssetAtomClip]:
|
||||
"""查找某个素材的所有原子片段,按 clip_index 排序。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_by_id(self, clip_id: str) -> AssetAtomClip | None:
|
||||
"""按 ID 查找单个原子片段。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_by_ids(self, clip_ids: list[str]) -> list[AssetAtomClip]:
|
||||
"""批量查找原子片段。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def delete_by_asset(self, asset_id: str) -> int:
|
||||
"""删除某素材的所有原子片段(级联删除),返回删除数量。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def count_by_asset(self, asset_id: str) -> int:
|
||||
"""统计某素材的原子片段数量。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_candidates_for_selection(
|
||||
self,
|
||||
asset_ids: list[str],
|
||||
*,
|
||||
min_duration: float | None = None,
|
||||
max_duration: float | None = None,
|
||||
limit: int = 100,
|
||||
) -> list[AssetAtomClip]:
|
||||
"""按素材集合和时长条件查找候选原子片段,按 clip_index 排序。
|
||||
|
||||
选片逻辑一次拉取多条素材的候选片段时使用,避免 N+1 查询。
|
||||
"""
|
||||
pass
|
||||
@@ -0,0 +1,110 @@
|
||||
"""#1970 原子片段 resolver 单元测试:DB 加载 + 内存兜底."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from packages.domain.asset_atom_clip import AssetAtomClip
|
||||
from packages.domain.atom_clip_resolver import (
|
||||
flatten_candidates,
|
||||
load_atom_clips_for_assets,
|
||||
)
|
||||
|
||||
|
||||
def _atom(asset_id: str, idx: int, start: float, end: float) -> AssetAtomClip:
|
||||
return AssetAtomClip(
|
||||
id=f"{asset_id}-clip-{idx}",
|
||||
asset_id=asset_id,
|
||||
start_time=start,
|
||||
end_time=end,
|
||||
duration=round(end - start, 3),
|
||||
clip_index=idx,
|
||||
)
|
||||
|
||||
|
||||
class FakeAtomRepo:
|
||||
def __init__(self, by_asset):
|
||||
self._by_asset = by_asset
|
||||
|
||||
def find_candidates_for_selection(self, asset_ids, *, limit=0):
|
||||
out = []
|
||||
for aid in asset_ids:
|
||||
out.extend(self._by_asset.get(aid, []))
|
||||
return out
|
||||
|
||||
def find_by_asset(self, asset_id):
|
||||
return list(self._by_asset.get(asset_id, []))
|
||||
|
||||
|
||||
class _Asset:
|
||||
def __init__(self, duration):
|
||||
self.duration = duration
|
||||
|
||||
|
||||
class FakeAssetRepo:
|
||||
def __init__(self, durations):
|
||||
self._durations = durations
|
||||
|
||||
def get(self, asset_id):
|
||||
d = self._durations.get(asset_id)
|
||||
return _Asset(d) if d is not None else None
|
||||
|
||||
|
||||
class TestLoadAtomClips:
|
||||
def test_persisted_clips_loaded_sorted(self):
|
||||
clips = [_atom("a", 1, 4.5, 9.0), _atom("a", 0, 0.0, 4.5)]
|
||||
repo = FakeAtomRepo({"a": clips})
|
||||
result = load_atom_clips_for_assets(["a"], atom_clip_repo=repo)
|
||||
assert [c.clip_index for c in result["a"]] == [0, 1]
|
||||
|
||||
def test_dedup_asset_ids_preserves_order(self):
|
||||
repo = FakeAtomRepo({"a": [_atom("a", 0, 0, 4)], "b": [_atom("b", 0, 0, 4)]})
|
||||
result = load_atom_clips_for_assets(["a", "b", "a"], atom_clip_repo=repo)
|
||||
assert list(result.keys()) == ["a", "b"]
|
||||
|
||||
def test_fallback_when_no_persisted_clips(self):
|
||||
"""老素材没有 atom_clips 时,内存按 3-6 秒均匀切片,标记 is_fallback。"""
|
||||
atom_repo = FakeAtomRepo({})
|
||||
asset_repo = FakeAssetRepo({"old": 20.0})
|
||||
result = load_atom_clips_for_assets(["old"], atom_clip_repo=atom_repo, asset_repo=asset_repo)
|
||||
assert "old" in result
|
||||
clips = result["old"]
|
||||
assert clips
|
||||
assert all(c.is_fallback for c in clips)
|
||||
assert abs(clips[-1].end_time - 20.0) < 0.01
|
||||
|
||||
def test_missing_duration_skipped(self):
|
||||
atom_repo = FakeAtomRepo({})
|
||||
asset_repo = FakeAssetRepo({})
|
||||
result = load_atom_clips_for_assets(["ghost"], atom_clip_repo=atom_repo, asset_repo=asset_repo)
|
||||
assert result == {}
|
||||
|
||||
def test_no_asset_repo_skips_empty_assets(self):
|
||||
atom_repo = FakeAtomRepo({})
|
||||
result = load_atom_clips_for_assets(["a"], atom_clip_repo=atom_repo, asset_repo=None)
|
||||
assert result == {}
|
||||
|
||||
def test_mixed_persisted_and_fallback(self):
|
||||
atom_repo = FakeAtomRepo({"new": [_atom("new", 0, 0, 5)]})
|
||||
asset_repo = FakeAssetRepo({"new": 5.0, "old": 10.0})
|
||||
result = load_atom_clips_for_assets(["new", "old"], atom_clip_repo=atom_repo, asset_repo=asset_repo)
|
||||
assert not result["new"][0].is_fallback
|
||||
assert all(c.is_fallback for c in result["old"])
|
||||
|
||||
def test_repo_exception_falls_back(self):
|
||||
class BrokenRepo(FakeAtomRepo):
|
||||
def find_candidates_for_selection(self, asset_ids, *, limit=0):
|
||||
raise RuntimeError("db down")
|
||||
|
||||
asset_repo = FakeAssetRepo({"a": 9.0})
|
||||
result = load_atom_clips_for_assets(["a"], atom_clip_repo=BrokenRepo({}), asset_repo=asset_repo)
|
||||
assert result["a"]
|
||||
assert all(c.is_fallback for c in result["a"])
|
||||
|
||||
def test_empty_input(self):
|
||||
assert load_atom_clips_for_assets([], atom_clip_repo=FakeAtomRepo({})) == {}
|
||||
|
||||
|
||||
class TestFlatten:
|
||||
def test_flatten_order(self):
|
||||
clips = flatten_candidates({"a": [_atom("a", 0, 0, 4)], "b": [_atom("b", 0, 0, 4), _atom("b", 1, 4, 8)]})
|
||||
assert len(clips) == 3
|
||||
assert clips[0].asset_id == "a"
|
||||
@@ -0,0 +1,205 @@
|
||||
"""#1970 原子片段级选片核心单元测试(纯函数,不依赖 DB)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import random
|
||||
|
||||
from packages.domain.asset_atom_clip import AssetAtomClip
|
||||
from packages.domain.atom_clip_selector import (
|
||||
clips_to_segments,
|
||||
estimate_required_clip_count,
|
||||
reselect_clips_from_atoms,
|
||||
score_atom_clip,
|
||||
select_atom_clips,
|
||||
)
|
||||
from packages.domain.atom_clip_service import compute_atom_clips
|
||||
|
||||
|
||||
def _clip(asset_id: str, start: float, end: float, clip_id: str = "") -> AssetAtomClip:
|
||||
return (
|
||||
AssetAtomClip.create(
|
||||
asset_id=asset_id,
|
||||
start_time=start,
|
||||
end_time=end,
|
||||
clip_index=int(start),
|
||||
)
|
||||
if not clip_id
|
||||
else AssetAtomClip(
|
||||
id=clip_id,
|
||||
asset_id=asset_id,
|
||||
start_time=start,
|
||||
end_time=end,
|
||||
duration=round(end - start, 3),
|
||||
clip_index=0,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class TestEstimateCount:
|
||||
def test_basic(self):
|
||||
assert estimate_required_clip_count(30.0, 4.5) == 7
|
||||
assert estimate_required_clip_count(18.0, 4.0) == round(18 / 4)
|
||||
|
||||
def test_invalid_inputs_returns_one(self):
|
||||
assert estimate_required_clip_count(0) == 1
|
||||
assert estimate_required_clip_count(10, 0) == 1
|
||||
assert estimate_required_clip_count(-1) == 1
|
||||
|
||||
|
||||
class TestScore:
|
||||
def test_unused_beats_used(self):
|
||||
c = _clip("a1", 0, 4)
|
||||
s_unused = score_atom_clip(c, target_duration=4.0, used_in_video=set())
|
||||
s_used = score_atom_clip(c, target_duration=4.0, used_in_video={c.id})
|
||||
assert s_unused > s_used
|
||||
|
||||
def test_duration_fit_better_when_closer(self):
|
||||
target = 4.0
|
||||
exact = score_atom_clip(_clip("a", 0, 4.0), target_duration=target)
|
||||
short = score_atom_clip(_clip("b", 0, 1.5), target_duration=target)
|
||||
assert exact > short
|
||||
|
||||
def test_history_penalty(self):
|
||||
c = _clip("a1", 0, 4)
|
||||
normal = score_atom_clip(c, target_duration=4.0)
|
||||
penalized = score_atom_clip(c, target_duration=4.0, recently_used={c.id})
|
||||
assert normal > penalized
|
||||
|
||||
def test_asset_balance_penalizes_repeated_asset(self):
|
||||
c1 = _clip("a", 0, 4)
|
||||
first = score_atom_clip(c1, target_duration=4.0, asset_usage_counts={})
|
||||
third = score_atom_clip(c1, target_duration=4.0, asset_usage_counts={"a": 2})
|
||||
assert first > third
|
||||
|
||||
|
||||
class TestSelect:
|
||||
def test_no_duplicate_atom_within_video(self):
|
||||
pool = compute_atom_clips("a", 30.0, rng=random.Random(1))
|
||||
used: set[str] = set()
|
||||
usage: dict[str, int] = {}
|
||||
chosen = []
|
||||
rng = random.Random(5)
|
||||
for _ in range(4):
|
||||
ranked = select_atom_clips(
|
||||
pool,
|
||||
target_duration=4.0,
|
||||
used_atom_clip_ids=used,
|
||||
asset_usage_counts=usage,
|
||||
required_count=4,
|
||||
limit=1,
|
||||
rng=rng,
|
||||
)
|
||||
assert ranked
|
||||
pick = ranked[0]
|
||||
assert pick.atom_clip_id not in used
|
||||
chosen.append(pick)
|
||||
used.add(pick.atom_clip_id)
|
||||
usage[pick.asset_id] = usage.get(pick.asset_id, 0) + 1
|
||||
assert len(used) == 4
|
||||
|
||||
def test_same_asset_different_clips_allowed(self):
|
||||
pool = compute_atom_clips("a", 30.0, rng=random.Random(2))
|
||||
used: set[str] = set()
|
||||
usage: dict[str, int] = {}
|
||||
rng = random.Random(7)
|
||||
picked_assets = set()
|
||||
for _ in range(3):
|
||||
pick = select_atom_clips(
|
||||
pool,
|
||||
target_duration=4.0,
|
||||
used_atom_clip_ids=used,
|
||||
asset_usage_counts=usage,
|
||||
limit=1,
|
||||
rng=rng,
|
||||
)[0]
|
||||
used.add(pick.atom_clip_id)
|
||||
usage[pick.asset_id] = usage.get(pick.asset_id, 0) + 1
|
||||
picked_assets.add(pick.asset_id)
|
||||
# 单素材池允许同素材多片段
|
||||
assert picked_assets == {"a"}
|
||||
assert len(used) == 3
|
||||
|
||||
def test_exhausted_pool_returns_empty(self):
|
||||
pool = [_clip("a", 0, 4)]
|
||||
ranked = select_atom_clips(pool, used_atom_clip_ids={pool[0].id}, target_duration=4.0)
|
||||
assert ranked == []
|
||||
|
||||
def test_recently_used_deprioritized_not_hard_blocked(self):
|
||||
# 两个片段,recent 中包含更合适的那个;它应被降权但不会从候选中消失
|
||||
fresh = _clip("a", 0, 2.0, clip_id="fresh")
|
||||
recent = _clip("b", 0, 4.0, clip_id="recent")
|
||||
ranked = select_atom_clips(
|
||||
[fresh, recent],
|
||||
target_duration=4.0,
|
||||
recently_used_atom_ids={"recent"},
|
||||
limit=2,
|
||||
rng=random.Random(0), # 噪声 0 不影响
|
||||
)
|
||||
ids = [r.atom_clip_id for r in ranked]
|
||||
assert set(ids) == {"fresh", "recent"}
|
||||
# 降权 + 噪声可能导致排序不稳定,只验证 recent 仍在候选中(不硬禁)
|
||||
|
||||
def test_limit(self):
|
||||
pool = compute_atom_clips("a", 40.0, rng=random.Random(4))
|
||||
ranked = select_atom_clips(pool, target_duration=4.0, limit=3)
|
||||
assert len(ranked) == 3
|
||||
scores = [r.score for r in ranked]
|
||||
assert scores == sorted(scores, reverse=True)
|
||||
|
||||
|
||||
class TestClipsToSegments:
|
||||
def test_grouped_by_asset_sorted(self):
|
||||
clips = [
|
||||
_clip("a", 10, 14),
|
||||
_clip("a", 0, 4),
|
||||
_clip("b", 2, 6),
|
||||
]
|
||||
segs = clips_to_segments(clips)
|
||||
assert segs["a"] == [(0, 4), (10, 14)]
|
||||
assert segs["b"] == [(2, 6)]
|
||||
|
||||
|
||||
class TestReselectFromAtoms:
|
||||
def _src(self, n):
|
||||
return [{"order": i, "clip_type": "main", "duration": 4.0, "start_time": 0.0} for i in range(n)]
|
||||
|
||||
def test_skeleton_preserved_and_unique(self):
|
||||
pool = compute_atom_clips("a", 30.0, rng=random.Random(11)) + compute_atom_clips(
|
||||
"b", 30.0, rng=random.Random(12)
|
||||
)
|
||||
out = reselect_clips_from_atoms(self._src(5), pool, rng=random.Random(13))
|
||||
assert out is not None
|
||||
assert len(out) == 5
|
||||
ids = [c["atom_clip_id"] for c in out]
|
||||
assert len(set(ids)) == 5
|
||||
for c in out:
|
||||
assert c["asset_id"]
|
||||
assert c["start_time"] >= 0
|
||||
assert c["duration"] > 0
|
||||
|
||||
def test_insufficient_candidates_returns_none(self):
|
||||
pool = compute_atom_clips("a", 10.0, rng=random.Random(1))
|
||||
assert reselect_clips_from_atoms(self._src(20), pool) is None
|
||||
|
||||
def test_non_main_clips_left_untouched(self):
|
||||
pool = compute_atom_clips("a", 30.0, rng=random.Random(8))
|
||||
src = [
|
||||
{"order": 0, "clip_type": "intro", "duration": 2.0, "asset_id": "fixed"},
|
||||
{"order": 1, "clip_type": "main", "duration": 4.0},
|
||||
]
|
||||
out = reselect_clips_from_atoms(src, pool, rng=random.Random(3))
|
||||
assert out is not None
|
||||
assert out[0]["asset_id"] == "fixed"
|
||||
assert "atom_clip_id" not in out[0]
|
||||
assert out[1].get("atom_clip_id")
|
||||
|
||||
def test_empty_inputs(self):
|
||||
assert reselect_clips_from_atoms([], [_clip("a", 0, 4)]) is None
|
||||
assert reselect_clips_from_atoms(self._src(2), []) is None
|
||||
|
||||
def test_batch_used_excluded(self):
|
||||
pool = compute_atom_clips("a", 30.0, rng=random.Random(21))
|
||||
batch_used = {pool[0].id}
|
||||
out = reselect_clips_from_atoms(self._src(3), pool, batch_used_atom_ids=batch_used, rng=random.Random(22))
|
||||
assert out is not None
|
||||
assert pool[0].id not in {c["atom_clip_id"] for c in out}
|
||||
@@ -0,0 +1,151 @@
|
||||
"""#1970 素材原子化切片逻辑单元测试(纯函数,不依赖 DB)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import random
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.asset_atom_clip import AssetAtomClip
|
||||
from packages.domain.atom_clip_service import (
|
||||
MAX_CLIP_SECONDS,
|
||||
MIN_CLIP_SECONDS,
|
||||
compute_atom_clips,
|
||||
compute_fallback_clips,
|
||||
)
|
||||
|
||||
|
||||
class TestComputeAtomClips:
|
||||
def test_short_asset_under_6s_single_clip(self):
|
||||
"""<6 秒素材整条作为一个片段,不切。"""
|
||||
for dur in (0.1, 3.0, 5.99):
|
||||
clips = compute_atom_clips("a1", dur, rng=random.Random(1))
|
||||
assert len(clips) == 1
|
||||
assert clips[0].start_time == 0.0
|
||||
assert abs(clips[0].end_time - dur) < 0.01
|
||||
assert clips[0].clip_index == 0
|
||||
|
||||
def test_exactly_6s_single_clip(self):
|
||||
clips = compute_atom_clips("a1", 6.0, rng=random.Random(1))
|
||||
assert len(clips) == 1
|
||||
assert clips[0].start_time == 0.0
|
||||
|
||||
def test_zero_and_negative_duration_returns_empty(self):
|
||||
assert compute_atom_clips("a1", 0) == []
|
||||
assert compute_atom_clips("a1", -1.0) == []
|
||||
|
||||
@pytest.mark.parametrize("seed", range(30))
|
||||
def test_clips_in_3_to_6_range(self, seed):
|
||||
"""除末段外,每段时长在 3~6 秒;末段 >=3 秒。"""
|
||||
clips = compute_atom_clips("a1", 60.0, rng=random.Random(seed))
|
||||
assert len(clips) >= 2
|
||||
for clip in clips[:-1]:
|
||||
assert MIN_CLIP_SECONDS - 0.06 <= clip.duration <= MAX_CLIP_SECONDS + 0.06
|
||||
# 末段 >=3(不足 3 应已合并)
|
||||
assert clips[-1].duration >= MIN_CLIP_SECONDS - 0.06
|
||||
|
||||
@pytest.mark.parametrize("dur", [6.01, 7.0, 9.0, 12.3, 30.0, 45.3, 100.0])
|
||||
def test_full_coverage_no_gaps_no_overlap(self, dur):
|
||||
clips = compute_atom_clips("a1", dur, rng=random.Random(int(dur * 100) % 10000))
|
||||
assert abs(clips[0].start_time) < 0.001
|
||||
assert abs(clips[-1].end_time - dur) < 0.01
|
||||
for prev, nxt in zip(clips, clips[1:], strict=False):
|
||||
assert abs(prev.end_time - nxt.start_time) < 0.001
|
||||
|
||||
def test_clip_index_sequential(self):
|
||||
clips = compute_atom_clips("a1", 40.0, rng=random.Random(5))
|
||||
assert [c.clip_index for c in clips] == list(range(len(clips)))
|
||||
|
||||
def test_tail_shorter_than_3s_merges_into_previous(self):
|
||||
"""末段不足 3 秒必须合并到前一段。"""
|
||||
# 多跑种子,保证任何随机结果都不存在 <3s 的末段
|
||||
for seed in range(100):
|
||||
clips = compute_atom_clips("a1", 7.5, rng=random.Random(seed))
|
||||
assert clips[-1].duration >= MIN_CLIP_SECONDS - 0.06
|
||||
assert abs(clips[-1].end_time - 7.5) < 0.01
|
||||
|
||||
def test_tail_between_3_and_6_stands_alone(self):
|
||||
"""末段 >=3 秒独立成段。"""
|
||||
found_standalone = False
|
||||
for seed in range(100):
|
||||
clips = compute_atom_clips("a1", 9.5, rng=random.Random(seed))
|
||||
if len(clips) == 2:
|
||||
found_standalone = True
|
||||
assert clips[-1].duration >= MIN_CLIP_SECONDS - 0.06
|
||||
assert found_standalone, "9.5s 至少在某些种子下应切为两段"
|
||||
|
||||
def test_scene_change_snap_within_window(self):
|
||||
"""切点 0.5s 窗口内有切换点时,切点对齐到切换处。"""
|
||||
aligned = 0
|
||||
for seed in range(500):
|
||||
clips = compute_atom_clips("a1", 20.0, scene_change_points=[4.52], rng=random.Random(seed))
|
||||
if any(c.scene_change_at == 4.52 for c in clips):
|
||||
aligned += 1
|
||||
hit = next(c for c in clips if c.scene_change_at == 4.52)
|
||||
# 命中片段的右边界即切换点
|
||||
assert abs(hit.end_time - 4.52) < 0.001
|
||||
assert aligned > 0
|
||||
|
||||
def test_scene_change_outside_window_not_force_aligned(self):
|
||||
"""窗口外的切换点不应强行对齐。"""
|
||||
clips = compute_atom_clips("a1", 30.0, scene_change_points=[15.0], rng=random.Random(1))
|
||||
for c in clips:
|
||||
if c.scene_change_at is not None:
|
||||
assert abs(c.end_time - c.scene_change_at) < 0.001
|
||||
|
||||
def test_scene_snap_never_creates_sub_3s_clip(self):
|
||||
"""对齐不能导致片段短于 3 秒。"""
|
||||
for seed in range(100):
|
||||
clips = compute_atom_clips("a1", 40.0, scene_change_points=[3.2, 6.3, 9.4], rng=random.Random(seed))
|
||||
for c in clips:
|
||||
assert c.duration >= MIN_CLIP_SECONDS - 0.06
|
||||
|
||||
def test_scene_points_out_of_duration_ignored(self):
|
||||
clips = compute_atom_clips("a1", 20.0, scene_change_points=[-1.0, 25.0, 4.0], rng=random.Random(3))
|
||||
assert all(c.scene_change_at != -1.0 and c.scene_change_at != 25.0 for c in clips)
|
||||
|
||||
def test_tags_inherited(self):
|
||||
clips = compute_atom_clips("a1", 30.0, tags=["t1", "t2"], rng=random.Random(2))
|
||||
assert all(c.tags == ["t1", "t2"] for c in clips)
|
||||
|
||||
def test_random_not_fixed_rhythm(self):
|
||||
"""随机切片:不同种子产出的切点集合应不同(避免固定节奏)。"""
|
||||
cuts1 = [c.end_time for c in compute_atom_clips("a1", 60.0, rng=random.Random(1))]
|
||||
cuts2 = [c.end_time for c in compute_atom_clips("a1", 60.0, rng=random.Random(2))]
|
||||
assert cuts1 != cuts2
|
||||
|
||||
def test_seed_reproducible(self):
|
||||
"""相同种子结果可复现。"""
|
||||
a = [(c.start_time, c.end_time) for c in compute_atom_clips("a1", 60.0, rng=random.Random(42))]
|
||||
b = [(c.start_time, c.end_time) for c in compute_atom_clips("a1", 60.0, rng=random.Random(42))]
|
||||
assert a == b
|
||||
|
||||
|
||||
class TestComputeFallbackClips:
|
||||
def test_fallback_marked_and_uniform(self):
|
||||
clips = compute_fallback_clips("a1", 20.0, clip_seconds=4.5)
|
||||
assert clips
|
||||
assert all(c.is_fallback for c in clips)
|
||||
for prev, nxt in zip(clips, clips[1:], strict=False):
|
||||
assert abs(prev.end_time - nxt.start_time) < 0.001
|
||||
assert abs(clips[-1].end_time - 20.0) < 0.01
|
||||
|
||||
def test_fallback_tail_merge(self):
|
||||
"""11.5s = 4.5+4.5+2.5 → 末段 2.5<3 合并 → 4.5+7.0。"""
|
||||
clips = compute_fallback_clips("a1", 11.5, clip_seconds=4.5)
|
||||
assert len(clips) == 2
|
||||
assert abs(clips[-1].duration - 7.0) < 0.01
|
||||
|
||||
def test_fallback_short_asset(self):
|
||||
clips = compute_fallback_clips("a1", 2.0)
|
||||
assert len(clips) == 1
|
||||
assert clips[0].is_fallback
|
||||
|
||||
def test_fallback_invalid_duration(self):
|
||||
assert compute_fallback_clips("a1", 0) == []
|
||||
assert compute_fallback_clips("a1", -5) == []
|
||||
|
||||
def test_fallback_clip_has_no_persisted_id(self):
|
||||
clips = compute_fallback_clips("a1", 10.0)
|
||||
# 兜底片段仍有运行时 id(dataclass 生成),但 is_fallback 是判别标记
|
||||
assert all(isinstance(c, AssetAtomClip) for c in clips)
|
||||
@@ -0,0 +1,181 @@
|
||||
"""#1970 PlanGeneratorService 原子片段选片端到端单元测试.
|
||||
|
||||
用 SQLite 内存库 + 真实仓储验证:注入 atom_clip_repo 后,正式生成(非预览)
|
||||
从原子片段选片,EditPlanClip.atom_clip_id 落库;预览模式保持旧路径。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
import pytest
|
||||
from app.services.plan_generator_service import PlanGeneratorService
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import (
|
||||
SQLAlchemyAssetAtomClipRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.models import Base
|
||||
from packages.domain.asset_atom_clip import AssetAtomClip
|
||||
from packages.domain.edit_template import EditTemplate, EditTemplateStatus
|
||||
from packages.domain.editing_mode import EditingMode
|
||||
from packages.domain.template_clip_config import ClipType, TemplateClipConfig
|
||||
|
||||
|
||||
class _FakeAsset:
|
||||
def __init__(self, aid, duration):
|
||||
self.id = aid
|
||||
self.duration = duration
|
||||
self.quality_score = 60.0
|
||||
self.metadata = {}
|
||||
self.created_at = None
|
||||
|
||||
|
||||
class FakeAssetRepo:
|
||||
def __init__(self, durations):
|
||||
self._durations = durations
|
||||
|
||||
def get(self, aid):
|
||||
return _FakeAsset(aid, self._durations[aid]) if aid in self._durations else None
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def db_session():
|
||||
engine = create_engine("sqlite://")
|
||||
# 只建相关表,避免全模型依赖
|
||||
Base.metadata.create_all(
|
||||
engine,
|
||||
tables=[
|
||||
Base.metadata.tables["edit_plans"],
|
||||
Base.metadata.tables["edit_plan_clips"],
|
||||
Base.metadata.tables["asset_atom_clips"],
|
||||
],
|
||||
)
|
||||
connection = engine.connect()
|
||||
Session = sessionmaker(bind=connection)
|
||||
session = Session()
|
||||
yield session
|
||||
session.close()
|
||||
connection.close()
|
||||
|
||||
|
||||
def _template(mode=EditingMode.ONE_TAKE.value):
|
||||
return EditTemplate(
|
||||
id="tpl-1",
|
||||
name="测试模板",
|
||||
editing_mode=mode,
|
||||
status=EditTemplateStatus.ACTIVE,
|
||||
)
|
||||
|
||||
|
||||
def _clip_configs(n=3):
|
||||
return [
|
||||
TemplateClipConfig(
|
||||
id=f"cfg-{i}",
|
||||
template_id="tpl-1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=i,
|
||||
min_duration=3.0,
|
||||
max_duration=6.0,
|
||||
)
|
||||
for i in range(n)
|
||||
]
|
||||
|
||||
|
||||
class TestAtomClipPlanGeneration:
|
||||
def test_generation_uses_atom_clips(self, db_session):
|
||||
atom_repo = SQLAlchemyAssetAtomClipRepository(db_session)
|
||||
# 两个素材各 30s,各切若干片段
|
||||
clips_a = [AssetAtomClip.create("asset-a", i * 5.0, i * 5.0 + 5.0, i) for i in range(6)]
|
||||
clips_b = [AssetAtomClip.create("asset-b", i * 5.0, i * 5.0 + 5.0, i) for i in range(6)]
|
||||
atom_repo.batch_create(clips_a + clips_b)
|
||||
db_session.commit()
|
||||
|
||||
svc = PlanGeneratorService(
|
||||
db_session,
|
||||
asset_repo=FakeAssetRepo({"asset-a": 30.0, "asset-b": 30.0}),
|
||||
atom_clip_repo=atom_repo,
|
||||
)
|
||||
result = svc.generate_from_template(
|
||||
template=_template(),
|
||||
clip_configs=_clip_configs(3),
|
||||
asset_ids=["asset-a", "asset-b"],
|
||||
created_by_user_id="user-1",
|
||||
)
|
||||
clips = result["clips"]
|
||||
assert len(clips) == 3
|
||||
# 每个 clip 都绑定了原子片段
|
||||
atom_ids = [c.atom_clip_id for c in clips]
|
||||
assert all(atom_ids)
|
||||
# 同一原子片段一个视频只用一次
|
||||
assert len(set(atom_ids)) == 3
|
||||
# start_time/duration 与选中片段一致
|
||||
for c in clips:
|
||||
assert c.start_time >= 0
|
||||
assert 0 < c.duration <= 6.0 + 0.01
|
||||
# asset_id 与 atom_clip 归属一致
|
||||
for c in clips:
|
||||
assert c.asset_id.startswith("asset-")
|
||||
|
||||
def test_fallback_when_atom_clips_not_ready(self, db_session):
|
||||
"""素材没有 atom_clips 时内存兜底切片,仍能选出片段。"""
|
||||
atom_repo = SQLAlchemyAssetAtomClipRepository(db_session)
|
||||
svc = PlanGeneratorService(
|
||||
db_session,
|
||||
asset_repo=FakeAssetRepo({"old-asset": 20.0}),
|
||||
atom_clip_repo=atom_repo,
|
||||
)
|
||||
result = svc.generate_from_template(
|
||||
template=_template(),
|
||||
clip_configs=_clip_configs(3),
|
||||
asset_ids=["old-asset"],
|
||||
created_by_user_id="user-1",
|
||||
)
|
||||
clips = result["clips"]
|
||||
# 兜底片段不落库、无持久 ID,clip 不绑定 atom_clip_id(回退旧路径)或绑定运行时 ID
|
||||
# 关键:必须成功选出素材,不报错
|
||||
assert all(c.asset_id == "old-asset" for c in clips)
|
||||
|
||||
def test_preview_mode_keeps_legacy_path(self, db_session):
|
||||
"""随机预览模式走旧路径,不要求 atom clips。"""
|
||||
atom_repo = SQLAlchemyAssetAtomClipRepository(db_session)
|
||||
svc = PlanGeneratorService(
|
||||
db_session,
|
||||
asset_repo=FakeAssetRepo({"asset-a": 30.0, "asset-b": 30.0, "asset-c": 30.0}),
|
||||
atom_clip_repo=atom_repo,
|
||||
)
|
||||
result = svc.generate_from_template(
|
||||
template=_template(),
|
||||
clip_configs=_clip_configs(3),
|
||||
asset_ids=["asset-a", "asset-b", "asset-c"],
|
||||
created_by_user_id="user-1",
|
||||
random_preview=True,
|
||||
)
|
||||
clips = result["clips"]
|
||||
assert len(clips) == 3
|
||||
assert {c.asset_id for c in clips} == {"asset-a", "asset-b", "asset-c"}
|
||||
# 预览路径不绑定 atom_clip_id
|
||||
assert all(not c.atom_clip_id for c in clips)
|
||||
|
||||
def test_no_atom_repo_uses_legacy_path(self, db_session):
|
||||
"""未注入 atom_clip_repo(旧调用方)时行为不变。"""
|
||||
svc = PlanGeneratorService(
|
||||
db_session,
|
||||
asset_repo=FakeAssetRepo({"asset-a": 30.0, "asset-b": 30.0, "asset-c": 30.0}),
|
||||
)
|
||||
result = svc.generate_from_template(
|
||||
template=_template(),
|
||||
clip_configs=_clip_configs(3),
|
||||
asset_ids=["asset-a", "asset-b", "asset-c"],
|
||||
created_by_user_id="user-1",
|
||||
)
|
||||
clips = result["clips"]
|
||||
assert len(clips) == 3
|
||||
assert {c.asset_id for c in clips} == {"asset-a", "asset-b", "asset-c"}
|
||||
@@ -297,3 +297,46 @@ class TestResolveLatestPlanByTemplate:
|
||||
with caplog.at_level("WARNING"):
|
||||
assert resolve_latest_plan_by_template(db, template_id="tpl", user_id="u") is None
|
||||
assert any("查找最新plan失败" in rec.message for rec in caplog.records)
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# collect_plan_atom_clip_ids (#1970)
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def _make_atom_clip(atom_clip_id):
|
||||
c = MagicMock()
|
||||
c.atom_clip_id = atom_clip_id
|
||||
return c
|
||||
|
||||
|
||||
class TestCollectPlanAtomClipIds:
|
||||
def test_empty_plan_returns_empty(self):
|
||||
from app.services.generation_common import collect_plan_atom_clip_ids
|
||||
|
||||
repo = MagicMock()
|
||||
repo.list_by_plan.return_value = []
|
||||
assert collect_plan_atom_clip_ids("p1", repo) == []
|
||||
|
||||
def test_collects_non_empty_ids_and_ignores_blank(self):
|
||||
from app.services.generation_common import collect_plan_atom_clip_ids
|
||||
|
||||
repo = MagicMock()
|
||||
repo.list_by_plan.side_effect = [
|
||||
[
|
||||
_make_atom_clip("atom-1"),
|
||||
_make_atom_clip(""),
|
||||
_make_atom_clip("atom-2"),
|
||||
],
|
||||
[],
|
||||
]
|
||||
assert collect_plan_atom_clip_ids("p1", repo) == ["atom-1", "atom-2"]
|
||||
|
||||
def test_missing_attribute_treated_as_blank(self):
|
||||
from app.services.generation_common import collect_plan_atom_clip_ids
|
||||
|
||||
legacy = MagicMock()
|
||||
del legacy.atom_clip_id # 旧对象无该属性
|
||||
repo = MagicMock()
|
||||
repo.list_by_plan.side_effect = [[legacy, _make_atom_clip("atom-9")], []]
|
||||
assert collect_plan_atom_clip_ids("p1", repo) == ["atom-9"]
|
||||
|
||||
Reference in New Issue
Block a user