From 0ac1bdf3b086b2262095556a4409fa2e25cc223c Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Fri, 18 Sep 2026 01:54:27 +0800 Subject: [PATCH 1/4] =?UTF-8?q?feat(#1970):=20=E7=B4=A0=E6=9D=90=E5=8E=9F?= =?UTF-8?q?=E5=AD=90=E5=8C=96=E5=88=87=E7=89=87=20P1=20-=20=E6=95=B0?= =?UTF-8?q?=E6=8D=AE=E5=B1=82/=E5=88=87=E7=89=87=E9=80=BB=E8=BE=91/?= =?UTF-8?q?=E5=8E=9F=E5=AD=90=E7=89=87=E6=AE=B5=E7=BA=A7=E9=80=89=E7=89=87?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 asset_atom_clips 表(FK→assets 级联删除)+ Alembic 079 - edit_plan_clips 新增 atom_clip_id 列 + Alembic 080 - AssetAtomClip 领域实体 + Repository 接口/SQLAlchemy 实现 - 切片逻辑:3~6s 随机、scdet 切点对齐、<6s 不切、末段<3s 合并 - Celery worker.generate_atom_clips:素材 READY 后异步切片,失败不阻断入库 - 选片改造:PlanGeneratorService/变体 reselect 从原子片段选取, 同片段单视频不重复、同素材多片段可用、跨视频/跨变体原子片段级避让 - atom_clips 未就绪时内存 3-6s 均匀兜底切片,再回退整条素材旧路径 - 85 个单元测试(切片算法/选片器/resolver/service 端到端) --- alembic/versions/079_asset_atom_clips.py | 58 ++++ .../080_edit_plan_clips_atom_clip_id.py | 37 +++ apps/api/app/services/edit_plan_service.py | 73 ++++- .../app/services/plan_generator_service.py | 148 +++++++++- apps/worker/worker_app/celery_app.py | 1 + apps/worker/worker_app/tasks/__init__.py | 5 + apps/worker/worker_app/tasks/atom_clips.py | 82 ++++++ apps/worker/worker_app/tasks/ingest.py | 15 + .../asset_atom_clip_repository.py | 124 ++++++++ .../edit_plan_clip_repository.py | 53 ++++ packages/adapters/sqlalchemy_impl/models.py | 32 ++- packages/domain/__init__.py | 22 ++ packages/domain/asset_atom_clip.py | 85 ++++++ packages/domain/atom_clip_resolver.py | 104 +++++++ packages/domain/atom_clip_selector.py | 264 ++++++++++++++++++ packages/domain/atom_clip_service.py | 215 ++++++++++++++ packages/domain/edit_plan_clip.py | 14 +- packages/ports/asset_atom_clip_repository.py | 57 ++++ tests/unit/test_1970_atom_clip_resolver.py | 120 ++++++++ tests/unit/test_1970_atom_clip_selector.py | 208 ++++++++++++++ tests/unit/test_1970_atom_clip_service.py | 161 +++++++++++ tests/unit/test_1970_atom_plan_generation.py | 189 +++++++++++++ 22 files changed, 2041 insertions(+), 26 deletions(-) create mode 100644 alembic/versions/079_asset_atom_clips.py create mode 100644 alembic/versions/080_edit_plan_clips_atom_clip_id.py create mode 100644 apps/worker/worker_app/tasks/atom_clips.py create mode 100644 packages/adapters/sqlalchemy_impl/asset_atom_clip_repository.py create mode 100644 packages/domain/asset_atom_clip.py create mode 100644 packages/domain/atom_clip_resolver.py create mode 100644 packages/domain/atom_clip_selector.py create mode 100644 packages/domain/atom_clip_service.py create mode 100644 packages/ports/asset_atom_clip_repository.py create mode 100644 tests/unit/test_1970_atom_clip_resolver.py create mode 100644 tests/unit/test_1970_atom_clip_selector.py create mode 100644 tests/unit/test_1970_atom_clip_service.py create mode 100644 tests/unit/test_1970_atom_plan_generation.py diff --git a/alembic/versions/079_asset_atom_clips.py b/alembic/versions/079_asset_atom_clips.py new file mode 100644 index 000000000..8e7e1b697 --- /dev/null +++ b/alembic/versions/079_asset_atom_clips.py @@ -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") diff --git a/alembic/versions/080_edit_plan_clips_atom_clip_id.py b/alembic/versions/080_edit_plan_clips_atom_clip_id.py new file mode 100644 index 000000000..190029c93 --- /dev/null +++ b/alembic/versions/080_edit_plan_clips_atom_clip_id.py @@ -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") diff --git a/apps/api/app/services/edit_plan_service.py b/apps/api/app/services/edit_plan_service.py index 24944060a..2d1660550 100755 --- a/apps/api/app/services/edit_plan_service.py +++ b/apps/api/app/services/edit_plan_service.py @@ -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, @@ -608,18 +609,68 @@ 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 load_atom_clips_for_assets, flatten_candidates + 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=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: diff --git a/apps/api/app/services/plan_generator_service.py b/apps/api/app/services/plan_generator_service.py index 04418ae2b..38e6422af 100755 --- a/apps/api/app/services/plan_generator_service.py +++ b/apps/api/app/services/plan_generator_service.py @@ -35,6 +35,11 @@ from packages.domain.plan_generator_utils import ( map_clip_types_for_mode, ) from packages.domain.smart_match import SCORE_RANDOM_NOISE_MAX, score_asset +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.template_clip_config import TemplateClipConfig logger = logging.getLogger(__name__) @@ -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,105 @@ 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]] = {} diff --git a/apps/worker/worker_app/celery_app.py b/apps/worker/worker_app/celery_app.py index 77a2218a4..915d43a36 100755 --- a/apps/worker/worker_app/celery_app.py +++ b/apps/worker/worker_app/celery_app.py @@ -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", diff --git a/apps/worker/worker_app/tasks/__init__.py b/apps/worker/worker_app/tasks/__init__.py index 74c785861..e2a03307c 100755 --- a/apps/worker/worker_app/tasks/__init__.py +++ b/apps/worker/worker_app/tasks/__init__.py @@ -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", diff --git a/apps/worker/worker_app/tasks/atom_clips.py b/apps/worker/worker_app/tasks/atom_clips.py new file mode 100644 index 000000000..5e0eb2891 --- /dev/null +++ b/apps/worker/worker_app/tasks/atom_clips.py @@ -0,0 +1,82 @@ +"""素材原子切片 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() diff --git a/apps/worker/worker_app/tasks/ingest.py b/apps/worker/worker_app/tasks/ingest.py index 870f843e3..9410d74f8 100755 --- a/apps/worker/worker_app/tasks/ingest.py +++ b/apps/worker/worker_app/tasks/ingest.py @@ -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, diff --git a/packages/adapters/sqlalchemy_impl/asset_atom_clip_repository.py b/packages/adapters/sqlalchemy_impl/asset_atom_clip_repository.py new file mode 100644 index 000000000..b5052ce64 --- /dev/null +++ b/packages/adapters/sqlalchemy_impl/asset_atom_clip_repository.py @@ -0,0 +1,124 @@ +"""素材原子片段仓储 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, + ) diff --git a/packages/adapters/sqlalchemy_impl/edit_plan_clip_repository.py b/packages/adapters/sqlalchemy_impl/edit_plan_clip_repository.py index b76a46211..0b3e5d177 100755 --- a/packages/adapters/sqlalchemy_impl/edit_plan_clip_repository.py +++ b/packages/adapters/sqlalchemy_impl/edit_plan_clip_repository.py @@ -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 diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index dc84bc3d4..eb7f44f69 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -1,7 +1,7 @@ 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 +234,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 +803,34 @@ 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)""" diff --git a/packages/domain/__init__.py b/packages/domain/__init__.py index 6ed2f89ce..17f277349 100755 --- a/packages/domain/__init__.py +++ b/packages/domain/__init__.py @@ -1,5 +1,18 @@ """Domain package for core business entities and rules.""" +from .asset_atom_clip import AssetAtomClip +from . import atom_clip_resolver +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", diff --git a/packages/domain/asset_atom_clip.py b/packages/domain/asset_atom_clip.py new file mode 100644 index 000000000..42db07c33 --- /dev/null +++ b/packages/domain/asset_atom_clip.py @@ -0,0 +1,85 @@ +"""素材原子片段(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, + ) diff --git a/packages/domain/atom_clip_resolver.py b/packages/domain/atom_clip_resolver.py new file mode 100644 index 000000000..7baad0809 --- /dev/null +++ b/packages/domain/atom_clip_resolver.py @@ -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 diff --git a/packages/domain/atom_clip_selector.py b/packages/domain/atom_clip_selector.py new file mode 100644 index 000000000..eaf1ba6fc --- /dev/null +++ b/packages/domain/atom_clip_selector.py @@ -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 diff --git a/packages/domain/atom_clip_service.py b/packages/domain/atom_clip_service.py new file mode 100644 index 000000000..1f669c5e3 --- /dev/null +++ b/packages/domain/atom_clip_service.py @@ -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 diff --git a/packages/domain/edit_plan_clip.py b/packages/domain/edit_plan_clip.py index d62ed13fd..5a2eca7af 100755 --- a/packages/domain/edit_plan_clip.py +++ b/packages/domain/edit_plan_clip.py @@ -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) diff --git a/packages/ports/asset_atom_clip_repository.py b/packages/ports/asset_atom_clip_repository.py new file mode 100644 index 000000000..81f9959c0 --- /dev/null +++ b/packages/ports/asset_atom_clip_repository.py @@ -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 diff --git a/tests/unit/test_1970_atom_clip_resolver.py b/tests/unit/test_1970_atom_clip_resolver.py new file mode 100644 index 000000000..b47876232 --- /dev/null +++ b/tests/unit/test_1970_atom_clip_resolver.py @@ -0,0 +1,120 @@ +"""#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" diff --git a/tests/unit/test_1970_atom_clip_selector.py b/tests/unit/test_1970_atom_clip_selector.py new file mode 100644 index 000000000..d5c9b738b --- /dev/null +++ b/tests/unit/test_1970_atom_clip_selector.py @@ -0,0 +1,208 @@ +"""#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} diff --git a/tests/unit/test_1970_atom_clip_service.py b/tests/unit/test_1970_atom_clip_service.py new file mode 100644 index 000000000..3b81436b2 --- /dev/null +++ b/tests/unit/test_1970_atom_clip_service.py @@ -0,0 +1,161 @@ +"""#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) diff --git a/tests/unit/test_1970_atom_plan_generation.py b/tests/unit/test_1970_atom_plan_generation.py new file mode 100644 index 000000000..8788bd711 --- /dev/null +++ b/tests/unit/test_1970_atom_plan_generation.py @@ -0,0 +1,189 @@ +"""#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 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 +from app.services.plan_generator_service import PlanGeneratorService + + +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"} -- 2.54.0 From 2339385d2c0724c2ea988f723f6a0061916ef56e Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Fri, 18 Sep 2026 02:11:50 +0800 Subject: [PATCH 2/4] =?UTF-8?q?feat(#1970):=20=E6=89=B9=E9=87=8F=E5=8F=98?= =?UTF-8?q?=E4=BD=93=E8=B7=A8=E5=8F=98=E4=BD=93=E5=8E=9F=E5=AD=90=E7=89=87?= =?UTF-8?q?=E6=AE=B5=E7=BA=A7=E7=A1=AC=E9=81=BF=E8=AE=A9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - reselect_plan_for_variant 新增 batch_used_atom_ids 参数 - generation_common 新增 collect_plan_atom_clip_ids - 批量生成循环累积各变体已用 atom_clip_id 传入下一变体 --- apps/api/app/api/routes/generation_tasks.py | 12 ++++++ apps/api/app/services/edit_plan_service.py | 5 ++- apps/api/app/services/generation_common.py | 27 +++++++++++++ tests/unit/test_generation_common.py | 43 +++++++++++++++++++++ 4 files changed, 86 insertions(+), 1 deletion(-) diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index ac4ae0d6b..a638b9006 100755 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -477,8 +477,15 @@ def create_generation_task( # #1855 P0:批次区间避让表,从变体0实际clips构建初始值(公共函数) from app.services.generation_common import collect_plan_segments as _collect_segments + from app.services.generation_common import ( + collect_plan_atom_clip_ids as _collect_atom_ids, + ) _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 +500,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 +537,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) diff --git a/apps/api/app/services/edit_plan_service.py b/apps/api/app/services/edit_plan_service.py index 2d1660550..76a0a8919 100755 --- a/apps/api/app/services/edit_plan_service.py +++ b/apps/api/app/services/edit_plan_service.py @@ -475,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)。 @@ -653,7 +654,9 @@ class EditPlanService: source_clips_data, atom_candidates, historical_atom_ids=historical_atom_ids, - batch_used_atom_ids=None, + batch_used_atom_ids=( + set(batch_used_atom_ids) if batch_used_atom_ids else None + ), rng=rng, ) except Exception: diff --git a/apps/api/app/services/generation_common.py b/apps/api/app/services/generation_common.py index 7f4ac971a..93bb4dbc5 100644 --- a/apps/api/app/services/generation_common.py +++ b/apps/api/app/services/generation_common.py @@ -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, *, diff --git a/tests/unit/test_generation_common.py b/tests/unit/test_generation_common.py index 56ca89c0e..acc10a4c4 100644 --- a/tests/unit/test_generation_common.py +++ b/tests/unit/test_generation_common.py @@ -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"] -- 2.54.0 From 7ba0d122bd5c1c09f882585fbbf70238b8c7bd6a Mon Sep 17 00:00:00 2001 From: CI Bot Date: Thu, 17 Sep 2026 18:21:03 +0000 Subject: [PATCH 3/4] style: auto-format with black + isort + ruff + prettier [skip ci-format-check] --- apps/api/app/api/routes/generation_tasks.py | 4 +- apps/api/app/services/edit_plan_service.py | 9 ++-- .../app/services/plan_generator_service.py | 14 +++---- apps/worker/worker_app/tasks/atom_clips.py | 1 - .../asset_atom_clip_repository.py | 20 ++------- packages/adapters/sqlalchemy_impl/models.py | 19 +++++++-- packages/domain/__init__.py | 2 +- packages/domain/asset_atom_clip.py | 8 +--- tests/unit/test_1970_atom_clip_resolver.py | 20 +++------ tests/unit/test_1970_atom_clip_selector.py | 41 +++++++++---------- tests/unit/test_1970_atom_clip_service.py | 20 +++------ tests/unit/test_1970_atom_plan_generation.py | 18 +++----- 12 files changed, 67 insertions(+), 109 deletions(-) diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index a638b9006..1fdc848f6 100755 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -476,10 +476,8 @@ 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 - from app.services.generation_common import ( - collect_plan_atom_clip_ids as _collect_atom_ids, - ) _batch_segments = _collect_segments(_plan0.id, _plan_svc._clip_repo) # #1970:批次内原子片段硬避让集合 diff --git a/apps/api/app/services/edit_plan_service.py b/apps/api/app/services/edit_plan_service.py index 76a0a8919..3d795a195 100755 --- a/apps/api/app/services/edit_plan_service.py +++ b/apps/api/app/services/edit_plan_service.py @@ -616,10 +616,11 @@ class EditPlanService: from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import ( SQLAlchemyAssetAtomClipRepository, ) - from packages.domain.atom_clip_resolver import load_atom_clips_for_assets, flatten_candidates + 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 atom_repo = SQLAlchemyAssetAtomClipRepository(db) + # 兜底切片只需要时长;本方法已查出 durations,封装一个只读假素材仓储 class _DurationOnlyAssetRepo: def __init__(self, durations_map: dict[str, float]) -> None: @@ -654,9 +655,7 @@ class EditPlanService: 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 - ), + batch_used_atom_ids=(set(batch_used_atom_ids) if batch_used_atom_ids else None), rng=rng, ) except Exception: @@ -673,7 +672,7 @@ class EditPlanService: batch_segments=batch_segments_resolved, target_durations=target_durations, rng=rng, - ) # 片段区间写回素材 metadata(与落库同事务;replace_all_clips_transactional 内 commit) + ) # 片段区间写回素材 metadata(与落库同事务;replace_all_clips_transactional 内 commit) for item in clips_data: aid = item.get("asset_id", "") if aid: diff --git a/apps/api/app/services/plan_generator_service.py b/apps/api/app/services/plan_generator_service.py index 38e6422af..a25fdc26f 100755 --- a/apps/api/app/services/plan_generator_service.py +++ b/apps/api/app/services/plan_generator_service.py @@ -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 @@ -35,11 +40,6 @@ from packages.domain.plan_generator_utils import ( map_clip_types_for_mode, ) from packages.domain.smart_match import SCORE_RANDOM_NOISE_MAX, score_asset -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.template_clip_config import TemplateClipConfig logger = logging.getLogger(__name__) @@ -313,9 +313,7 @@ class PlanGeneratorService: 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) - ) + recently_used = set(self._clip_repo.list_recent_atom_clip_ids_by_user(user_id, limit=200)) except Exception: logger.warning("跨视频原子片段避让查询失败", exc_info=True) diff --git a/apps/worker/worker_app/tasks/atom_clips.py b/apps/worker/worker_app/tasks/atom_clips.py index 5e0eb2891..e18c3a7c0 100644 --- a/apps/worker/worker_app/tasks/atom_clips.py +++ b/apps/worker/worker_app/tasks/atom_clips.py @@ -8,7 +8,6 @@ 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 diff --git a/packages/adapters/sqlalchemy_impl/asset_atom_clip_repository.py b/packages/adapters/sqlalchemy_impl/asset_atom_clip_repository.py index b5052ce64..81f8cefbb 100644 --- a/packages/adapters/sqlalchemy_impl/asset_atom_clip_repository.py +++ b/packages/adapters/sqlalchemy_impl/asset_atom_clip_repository.py @@ -40,9 +40,7 @@ class SQLAlchemyAssetAtomClipRepository: 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() + model = self.session.query(AssetAtomClipModel).filter(AssetAtomClipModel.id == clip_id).first() if model is None: return None return self._to_domain(model) @@ -50,11 +48,7 @@ class SQLAlchemyAssetAtomClipRepository: 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() - ) + 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: @@ -67,11 +61,7 @@ class SQLAlchemyAssetAtomClipRepository: return count def count_by_asset(self, asset_id: str) -> int: - return ( - self.session.query(AssetAtomClipModel) - .filter(AssetAtomClipModel.asset_id == asset_id) - .count() - ) + return self.session.query(AssetAtomClipModel).filter(AssetAtomClipModel.asset_id == asset_id).count() def find_candidates_for_selection( self, @@ -82,9 +72,7 @@ class SQLAlchemyAssetAtomClipRepository: limit: int = 100, ) -> list[AssetAtomClip]: """按筛选条件查找候选原子片段,按时长排序。用于选片逻辑。""" - query = self.session.query(AssetAtomClipModel).filter( - AssetAtomClipModel.asset_id.in_(asset_ids) - ) + 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: diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index eb7f44f69..5285a5d82 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -1,7 +1,20 @@ from datetime import UTC, datetime from typing import Any -from sqlalchemy import JSON, Boolean, Column, DateTime, Float, ForeignKey, 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() @@ -810,9 +823,7 @@ class AssetAtomClipModel(Base): """ __tablename__ = "asset_atom_clips" - __table_args__ = ( - UniqueConstraint("asset_id", "clip_index", name="uq_asset_atom_clips_asset_index"), - ) + __table_args__ = (UniqueConstraint("asset_id", "clip_index", name="uq_asset_atom_clips_asset_index"),) id = Column(String(36), primary_key=True) asset_id = Column( diff --git a/packages/domain/__init__.py b/packages/domain/__init__.py index 17f277349..ce0d1ee05 100755 --- a/packages/domain/__init__.py +++ b/packages/domain/__init__.py @@ -1,7 +1,7 @@ """Domain package for core business entities and rules.""" -from .asset_atom_clip import AssetAtomClip from . import atom_clip_resolver +from .asset_atom_clip import AssetAtomClip from .atom_clip_selector import ( ScoredAtomClip, clips_to_segments, diff --git a/packages/domain/asset_atom_clip.py b/packages/domain/asset_atom_clip.py index 42db07c33..404d7436b 100644 --- a/packages/domain/asset_atom_clip.py +++ b/packages/domain/asset_atom_clip.py @@ -46,15 +46,11 @@ class AssetAtomClip: 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}" - ) + 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}" - ) + 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: diff --git a/tests/unit/test_1970_atom_clip_resolver.py b/tests/unit/test_1970_atom_clip_resolver.py index b47876232..1db49e5a2 100644 --- a/tests/unit/test_1970_atom_clip_resolver.py +++ b/tests/unit/test_1970_atom_clip_resolver.py @@ -64,9 +64,7 @@ class TestLoadAtomClips: """老素材没有 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 - ) + 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 @@ -76,9 +74,7 @@ class TestLoadAtomClips: 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 - ) + 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): @@ -89,9 +85,7 @@ class TestLoadAtomClips: 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 - ) + 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"]) @@ -101,9 +95,7 @@ class TestLoadAtomClips: 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 - ) + 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"]) @@ -113,8 +105,6 @@ class TestLoadAtomClips: 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)]} - ) + 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" diff --git a/tests/unit/test_1970_atom_clip_selector.py b/tests/unit/test_1970_atom_clip_selector.py index d5c9b738b..95de1bdd8 100644 --- a/tests/unit/test_1970_atom_clip_selector.py +++ b/tests/unit/test_1970_atom_clip_selector.py @@ -16,18 +16,22 @@ 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, + 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, + ) ) @@ -117,9 +121,7 @@ class TestSelect: 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 - ) + 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): @@ -159,10 +161,7 @@ class TestClipsToSegments: class TestReselectFromAtoms: def _src(self, n): - return [ - {"order": i, "clip_type": "main", "duration": 4.0, "start_time": 0.0} - for i in range(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( @@ -201,8 +200,6 @@ class TestReselectFromAtoms: 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) - ) + 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} diff --git a/tests/unit/test_1970_atom_clip_service.py b/tests/unit/test_1970_atom_clip_service.py index 3b81436b2..0e0de9a5f 100644 --- a/tests/unit/test_1970_atom_clip_service.py +++ b/tests/unit/test_1970_atom_clip_service.py @@ -78,9 +78,7 @@ class TestComputeAtomClips: """切点 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) - ) + 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) @@ -90,9 +88,7 @@ class TestComputeAtomClips: 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) - ) + 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 @@ -100,22 +96,16 @@ class TestComputeAtomClips: 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) - ) + 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) - ) + 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) - ) + 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): diff --git a/tests/unit/test_1970_atom_plan_generation.py b/tests/unit/test_1970_atom_plan_generation.py index 8788bd711..fb5774972 100644 --- a/tests/unit/test_1970_atom_plan_generation.py +++ b/tests/unit/test_1970_atom_plan_generation.py @@ -15,6 +15,7 @@ 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 @@ -26,7 +27,6 @@ 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 -from app.services.plan_generator_service import PlanGeneratorService class _FakeAsset: @@ -93,12 +93,8 @@ 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) - ] + 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() @@ -152,9 +148,7 @@ class TestAtomClipPlanGeneration: atom_repo = SQLAlchemyAssetAtomClipRepository(db_session) svc = PlanGeneratorService( db_session, - asset_repo=FakeAssetRepo( - {"asset-a": 30.0, "asset-b": 30.0, "asset-c": 30.0} - ), + 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( @@ -174,9 +168,7 @@ class TestAtomClipPlanGeneration: """未注入 atom_clip_repo(旧调用方)时行为不变。""" svc = PlanGeneratorService( db_session, - asset_repo=FakeAssetRepo( - {"asset-a": 30.0, "asset-b": 30.0, "asset-c": 30.0} - ), + asset_repo=FakeAssetRepo({"asset-a": 30.0, "asset-b": 30.0, "asset-c": 30.0}), ) result = svc.generate_from_template( template=_template(), -- 2.54.0 From fbb229bd8347713eaf8838bacefa78d9142319a6 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Fri, 18 Sep 2026 03:38:37 +0800 Subject: [PATCH 4/4] =?UTF-8?q?fix:=20=E4=BF=AE=E5=A4=8D=20develop=20?= =?UTF-8?q?=E9=81=97=E7=95=99=20ruff=20B904/B005=EF=BC=88scripts=5Fai=20ra?= =?UTF-8?q?ise=20from=20None=E3=80=81douyin=20rstrip=20noqa=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- apps/api/app/api/routes/scripts_ai.py | 18 ++++++++++++++---- apps/api/app/services/douyin_resolver.py | 2 +- 2 files changed, 15 insertions(+), 5 deletions(-) diff --git a/apps/api/app/api/routes/scripts_ai.py b/apps/api/app/api/routes/scripts_ai.py index 6420cb490..77d639271 100644 --- a/apps/api/app/api/routes/scripts_ai.py +++ b/apps/api/app/api/routes/scripts_ai.py @@ -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) diff --git a/apps/api/app/services/douyin_resolver.py b/apps/api/app/services/douyin_resolver.py index 89b6ff992..0167837ce 100644 --- a/apps/api/app/services/douyin_resolver.py +++ b/apps/api/app/services/douyin_resolver.py @@ -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: -- 2.54.0