feat(#1970): 素材原子化切片 P1 - 数据层/切片逻辑/原子片段级选片 #1974

Merged
auto-approve-bot merged 4 commits from feat/1970-atom-clips into develop 2026-09-18 03:57:07 +08:00
27 changed files with 2099 additions and 31 deletions
+58
View File
@@ -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)
+14 -4
View File
@@ -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)
+1 -1
View File
@@ -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:
+64 -11
View File
@@ -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,
*,
+133 -13
View File
@@ -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]] = {}
+1
View File
@@ -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",
+5
View File
@@ -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()
+15
View File
@@ -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
+42 -1
View File
@@ -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)"""
+22
View File
@@ -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",
+81
View File
@@ -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,
)
+104
View File
@@ -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
+264
View File
@@ -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
+215
View File
@@ -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
+13 -1
View File
@@ -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
+110
View File
@@ -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"
+205
View File
@@ -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}
+151
View File
@@ -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"}
+43
View File
@@ -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"]