diff --git a/alembic/versions/060_migrate_template_segments_to_clip_configs.py b/alembic/versions/060_migrate_template_segments_to_clip_configs.py new file mode 100644 index 000000000..65641ca23 --- /dev/null +++ b/alembic/versions/060_migrate_template_segments_to_clip_configs.py @@ -0,0 +1,67 @@ +"""migrate template_segments data to template_clip_configs + +将 template_segments 中的历史数据一次性迁移到 template_clip_configs。 +只迁移 template_clip_configs 中尚不存在对应 template_id 的记录,避免覆盖编辑器已发布数据。 + +Revision ID: 060_migrate_segments +Revises: 059_duplicate_rate +Create Date: 2026-08-31 +""" + +import sqlalchemy as sa + +from alembic import op + +revision = "060_migrate_segments" +down_revision = "059_duplicate_rate" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # 从 template_segments 迁移到 template_clip_configs + # 字段映射: + # segment_order -> order + # duration_min -> min_duration + # duration_max -> max_duration + # material_type -> config JSON {"material_type": ...} + # clip_type -> "main" (默认值) + # 只迁移 template_clip_configs 中没有对应 template_id 的记录 + op.execute( + sa.text( + """ + INSERT INTO template_clip_configs + (id, template_id, clip_type, "order", min_duration, max_duration, + text_template, material_requirements, transition_effect, config, + created_at, updated_at) + SELECT + s.id, + s.template_id, + 'main', + s.segment_order, + s.duration_min, + s.duration_max, + '', + '{}', + 'cut', + CASE + WHEN s.material_type IS NOT NULL AND s.material_type != '' + THEN JSON_OBJECT('material_type', s.material_type) + ELSE '{}' + END, + s.created_at, + s.updated_at + FROM template_segments s + WHERE NOT EXISTS ( + SELECT 1 FROM template_clip_configs c + WHERE c.template_id = s.template_id + ) + """ + ) + ) + + +def downgrade() -> None: + # 不可逆迁移:无法确定哪些 clip_configs 是从旧表迁移来的 + # 保留空实现,不回滚 + pass diff --git a/packages/adapters/sqlalchemy_impl/template_repository.py b/packages/adapters/sqlalchemy_impl/template_repository.py index e1860cc33..eb830f977 100755 --- a/packages/adapters/sqlalchemy_impl/template_repository.py +++ b/packages/adapters/sqlalchemy_impl/template_repository.py @@ -1,4 +1,9 @@ -"""SQLAlchemy implementation of TemplateRepository.""" +"""SQLAlchemy implementation of TemplateRepository. + +模板 segments 数据源已统一为 template_clip_configs 表。 +读取时优先 template_clip_configs,回退 template_segments(兼容历史数据)。 +写入全部走 template_clip_configs。 +""" from __future__ import annotations @@ -10,6 +15,7 @@ from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import ( EditPlanModel, TemplateCategoryModel, + TemplateClipConfigModel, TemplateModel, TemplateSegmentModel, ) @@ -47,27 +53,38 @@ class SQLAlchemyTemplateRepository: like_pattern = f"%{keyword}%" query = query.filter(TemplateModel.name.like(like_pattern)) if tag: - # JSON 数组包含指定标签(MySQL JSON_CONTAINS / SQLite json_each 兼容写法用 LIKE) - query = query.filter(TemplateModel.tags.like(f'%"{tag}"%')) + query = query.filter(TemplateModel.tags.like(f'"%{tag}"%')) models = query.order_by(TemplateModel.created_at.desc()).offset(skip).limit(limit).all() templates = [self._model_to_entity(m) for m in models] - # 批量加载所有 segments,避免 N+1 查询 + # 批量加载 segments —— 优先 template_clip_configs if templates: template_ids = [t.id for t in templates] - seg_models = ( - self.session.query(TemplateSegmentModel) - .filter(TemplateSegmentModel.template_id.in_(template_ids)) - .order_by(TemplateSegmentModel.segment_order) + clip_models = ( + self.session.query(TemplateClipConfigModel) + .filter(TemplateClipConfigModel.template_id.in_(template_ids)) + .order_by(TemplateClipConfigModel.order) .all() ) - # 按 template_id 分组 - seg_map: dict[str, list] = {} - for sm in seg_models: - seg_map.setdefault(sm.template_id, []).append( - self._segment_model_to_entity(sm), + clip_map: dict[str, list] = {} + for cm in clip_models: + clip_map.setdefault(cm.template_id, []).append( + self._clip_config_to_segment(cm), ) + # 对没有 clip_configs 的模板,回退读 template_segments + missing_ids = [t.id for t in templates if t.id not in clip_map] + if missing_ids: + old_models = ( + self.session.query(TemplateSegmentModel) + .filter(TemplateSegmentModel.template_id.in_(missing_ids)) + .order_by(TemplateSegmentModel.segment_order) + .all() + ) + for om in old_models: + clip_map.setdefault(om.template_id, []).append( + self._segment_model_to_entity(om), + ) for t in templates: - t.segments = seg_map.get(t.id, []) + t.segments = clip_map.get(t.id, []) return templates def get(self, template_id: str, user_id: str) -> Optional[Template]: @@ -100,7 +117,6 @@ class SQLAlchemyTemplateRepository: is_active=template.is_active, ) self.session.add(model) - # flush 而非 commit,让 create + create_segments 在同一事务中提交 self.session.flush() self.session.refresh(model) result = self._model_to_entity(model) @@ -145,7 +161,10 @@ class SQLAlchemyTemplateRepository: if model is None: return False model.is_active = False - # 级联清理关联的 segments,避免孤儿数据 + # 级联清理两张表的关联数据 + self.session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == template_id, + ).delete(synchronize_session=False) self.session.query(TemplateSegmentModel).filter( TemplateSegmentModel.template_id == template_id, ).delete(synchronize_session=False) @@ -172,7 +191,7 @@ class SQLAlchemyTemplateRepository: if keyword: query = query.filter(TemplateModel.name.like(f"%{keyword}%")) if tag: - query = query.filter(TemplateModel.tags.like(f'%"{tag}"%')) + query = query.filter(TemplateModel.tags.like(f'"%{tag}"%')) return query.count() def copy_template(self, template_id: str, user_id: str, new_name: str) -> Template: @@ -181,9 +200,8 @@ class SQLAlchemyTemplateRepository: if source is None: raise ValueError(f"Template {template_id} not found") - new_id = str(uuid.uuid4()) new_template = Template( - id=new_id, + id=str(uuid.uuid4()), user_id=user_id, name=new_name, mode=source.mode, @@ -197,28 +215,20 @@ class SQLAlchemyTemplateRepository: ) created = self.create(new_template) - # 复制 segments + # 复用 create_segments 写入 template_clip_configs new_segments: List[TemplateSegment] = [] for seg in source.segments: - new_seg = TemplateSegment( + new_segments.append(TemplateSegment( id=str(uuid.uuid4()), - template_id=new_id, + template_id=created.id, segment_order=seg.segment_order, duration_min=seg.duration_min, duration_max=seg.duration_max, material_type=seg.material_type, - ) - new_segments.append(new_seg) - model = TemplateSegmentModel( - id=new_seg.id, - template_id=new_seg.template_id, - segment_order=new_seg.segment_order, - duration_min=new_seg.duration_min, - duration_max=new_seg.duration_max, - material_type=new_seg.material_type, - ) - self.session.add(model) + )) if new_segments: + self.create_segments(new_segments) + else: self.session.commit() created.segments = new_segments @@ -227,34 +237,58 @@ class SQLAlchemyTemplateRepository: # ── Segments ── def list_segments(self, template_id: str) -> List[TemplateSegment]: - models = ( + """优先从 template_clip_configs 读取,回退读 template_segments。""" + clips = ( + self.session.query(TemplateClipConfigModel) + .filter(TemplateClipConfigModel.template_id == template_id) + .order_by(TemplateClipConfigModel.order) + .all() + ) + if clips: + return [self._clip_config_to_segment(m) for m in clips] + # 回退:旧表 + old = ( self.session.query(TemplateSegmentModel) .filter(TemplateSegmentModel.template_id == template_id) .order_by(TemplateSegmentModel.segment_order) .all() ) - return [self._segment_model_to_entity(m) for m in models] + return [self._segment_model_to_entity(m) for m in old] def create_segments(self, segments: List[TemplateSegment]) -> List[TemplateSegment]: + """写入 template_clip_configs 表。material_type 存入 config JSON。""" for seg in segments: - model = TemplateSegmentModel( + config = {"material_type": seg.material_type} if seg.material_type else {} + model = TemplateClipConfigModel( id=seg.id, template_id=seg.template_id, - segment_order=seg.segment_order, - duration_min=seg.duration_min, - duration_max=seg.duration_max, - material_type=seg.material_type, + clip_type="main", + order=seg.segment_order, + min_duration=seg.duration_min, + max_duration=seg.duration_max, + text_template="", + material_requirements={}, + transition_effect="cut", + config=config, ) self.session.add(model) self.session.commit() return segments def delete_segments_by_template(self, template_id: str) -> int: - count = ( - self.session.query(TemplateSegmentModel).filter(TemplateSegmentModel.template_id == template_id).delete() + """删除两张表中的 segments 数据,返回删除总数。""" + c1 = ( + self.session.query(TemplateClipConfigModel) + .filter(TemplateClipConfigModel.template_id == template_id) + .delete() + ) + c2 = ( + self.session.query(TemplateSegmentModel) + .filter(TemplateSegmentModel.template_id == template_id) + .delete() ) self.session.commit() - return count + return c1 + c2 # ── Categories ── @@ -366,6 +400,23 @@ class SQLAlchemyTemplateRepository: updated_at=model.updated_at, ) + @staticmethod + def _clip_config_to_segment(model: TemplateClipConfigModel) -> TemplateSegment: + """将 TemplateClipConfigModel 转换为 TemplateSegment 域实体。""" + material_type = None + if model.config and isinstance(model.config, dict): + material_type = model.config.get("material_type") + return TemplateSegment( + id=model.id, + template_id=model.template_id, + segment_order=model.order, + duration_min=model.min_duration, + duration_max=model.max_duration, + material_type=material_type, + created_at=model.created_at, + updated_at=model.updated_at, + ) + @staticmethod def _category_model_to_entity(model: TemplateCategoryModel) -> TemplateCategory: return TemplateCategory( diff --git a/tests/unit/test_unify_template_segments.py b/tests/unit/test_unify_template_segments.py new file mode 100644 index 000000000..495f802e6 --- /dev/null +++ b/tests/unit/test_unify_template_segments.py @@ -0,0 +1,302 @@ +"""统一模板 segments 数据源单元测试。 + +验证 template_repository 从 template_clip_configs 读取 segments, +写入走 template_clip_configs,回退兼容 template_segments。 +""" + +from __future__ import annotations + +import sys +import uuid +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) + +import pytest +from sqlalchemy import create_engine, text +from sqlalchemy.orm import sessionmaker + +from packages.adapters.sqlalchemy_impl.models import ( + Base, + TemplateClipConfigModel, + TemplateModel, + TemplateSegmentModel, +) +from packages.adapters.sqlalchemy_impl.template_repository import ( + SQLAlchemyTemplateRepository, +) +from packages.domain.template import Template, TemplateSegment + + +@pytest.fixture() +def session(): + engine = create_engine("sqlite:///:memory:") + Base.metadata.create_all(engine) + Session = sessionmaker(bind=engine) + s = Session() + try: + yield s + finally: + s.close() + + +@pytest.fixture() +def repo(session): + return SQLAlchemyTemplateRepository(session) + + +def _make_template( + template_id=None, + user_id="u1", + name="测试模板", + mode="one_take", + segments=None, +): + tid = template_id or str(uuid.uuid4()) + return Template( + id=tid, + user_id=user_id, + name=name, + mode=mode, + category="", + tags=[], + estimated_duration=30.0, + is_active=True, + segments=segments or [], + ) + + +def _make_segment(template_id, order=1, material_type=None): + return TemplateSegment( + id=str(uuid.uuid4()), + template_id=template_id, + segment_order=order, + duration_min=5.0, + duration_max=10.0, + material_type=material_type, + ) + + +# ── 写入测试 ───────────────────────────────────────────────────────────── + + +class TestCreateSegments: + def test_writes_to_clip_configs(self, repo, session): + """create_segments 写入 template_clip_configs 表""" + tpl = _make_template() + repo.create(tpl) + seg = _make_segment(tpl.id, order=1) + repo.create_segments([seg]) + + clips = session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == tpl.id + ).all() + assert len(clips) == 1 + assert clips[0].clip_type == "main" + assert clips[0].order == 1 + assert clips[0].min_duration == 5.0 + assert clips[0].max_duration == 10.0 + assert clips[0].config.get("material_type") is None + + def test_material_type_stored_in_config(self, repo, session): + """material_type 存入 config JSON""" + tpl = _make_template() + repo.create(tpl) + seg = _make_segment(tpl.id, order=1, material_type="voiceover") + repo.create_segments([seg]) + + clip = session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == tpl.id + ).first() + assert clip.config["material_type"] == "voiceover" + + +# ── 读取测试 ───────────────────────────────────────────────────────────── + + +class TestListSegments: + def test_reads_from_clip_configs(self, repo, session): + """list_segments 优先从 template_clip_configs 读取""" + tpl = _make_template() + repo.create(tpl) + seg = _make_segment(tpl.id, order=1, material_type="voiceover") + repo.create_segments([seg]) + + result = repo.list_segments(tpl.id) + assert len(result) == 1 + assert result[0].segment_order == 1 + assert result[0].duration_min == 5.0 + assert result[0].duration_max == 10.0 + assert result[0].material_type == "voiceover" + + def test_fallback_to_old_table(self, repo, session): + """clip_configs 为空时回退读 template_segments""" + tpl = _make_template() + repo.create(tpl) + + # 直接往旧表写入 + old = TemplateSegmentModel( + id=str(uuid.uuid4()), + template_id=tpl.id, + segment_order=1, + duration_min=3.0, + duration_max=8.0, + material_type="场景", + ) + session.add(old) + session.commit() + + result = repo.list_segments(tpl.id) + assert len(result) == 1 + assert result[0].material_type == "场景" + assert result[0].duration_min == 3.0 + + def test_clip_configs_takes_priority(self, repo, session): + """两张表都有数据时,clip_configs 优先""" + tpl = _make_template() + repo.create(tpl) + + # 写入 clip_configs + seg = _make_segment(tpl.id, order=1) + repo.create_segments([seg]) + + # 也写入旧表 + old = TemplateSegmentModel( + id=str(uuid.uuid4()), + template_id=tpl.id, + segment_order=1, + duration_min=1.0, + duration_max=2.0, + material_type=None, + ) + session.add(old) + session.commit() + + result = repo.list_segments(tpl.id) + assert len(result) == 1 + # 应该返回 clip_configs 的数据(min_duration=5.0),不是旧表的(1.0) + assert result[0].duration_min == 5.0 + + +# ── list_by_user 批量加载 ──────────────────────────────────────────────── + + +class TestListByUser: + def test_batch_loads_from_clip_configs(self, repo, session): + """list_by_user 批量从 clip_configs 加载 segments""" + tpl = _make_template() + repo.create(tpl) + seg = _make_segment(tpl.id, order=1, material_type="人物") + repo.create_segments([seg]) + + result = repo.list_by_user("u1") + assert len(result) == 1 + assert len(result[0].segments) == 1 + assert result[0].segments[0].material_type == "人物" + + def test_fallback_for_old_data(self, repo, session): + """list_by_user 对无 clip_configs 的模板回退读旧表""" + tpl = _make_template() + repo.create(tpl) + + old = TemplateSegmentModel( + id=str(uuid.uuid4()), + template_id=tpl.id, + segment_order=1, + duration_min=2.0, + duration_max=6.0, + material_type=None, + ) + session.add(old) + session.commit() + + result = repo.list_by_user("u1") + assert len(result) == 1 + assert len(result[0].segments) == 1 + assert result[0].segments[0].duration_min == 2.0 + + +# ── copy_template ──────────────────────────────────────────────────────── + + +class TestCopyTemplate: + def test_copy_writes_to_clip_configs(self, repo, session): + """复制模板的 segments 写入 clip_configs""" + tpl = _make_template() + repo.create(tpl) + seg = _make_segment(tpl.id, order=1, material_type="voiceover") + repo.create_segments([seg]) + + copied = repo.copy_template(tpl.id, "u1", "副本模板") + assert copied.id != tpl.id + assert len(copied.segments) == 1 + + # 验证写入的是 clip_configs 表 + clips = session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == copied.id + ).all() + assert len(clips) == 1 + assert clips[0].config["material_type"] == "voiceover" + + def test_copy_empty_segments(self, repo, session): + """复制无 segments 的模板不报错""" + tpl = _make_template() + repo.create(tpl) + + copied = repo.copy_template(tpl.id, "u1", "空副本") + assert len(copied.segments) == 0 + + +# ── delete ─────────────────────────────────────────────────────────────── + + +class TestDelete: + def test_delete_cleans_both_tables(self, repo, session): + """删除模板时清理两张表的关联数据""" + tpl = _make_template() + repo.create(tpl) + seg = _make_segment(tpl.id, order=1) + repo.create_segments([seg]) + + # 也写入旧表 + old = TemplateSegmentModel( + id=str(uuid.uuid4()), + template_id=tpl.id, + segment_order=1, + duration_min=1.0, + duration_max=2.0, + ) + session.add(old) + session.commit() + + repo.delete(tpl.id, "u1") + + # 两张表都应该被清理 + c1 = session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == tpl.id + ).count() + c2 = session.query(TemplateSegmentModel).filter( + TemplateSegmentModel.template_id == tpl.id + ).count() + assert c1 == 0 + assert c2 == 0 + + def test_delete_segments_by_template(self, repo, session): + """delete_segments_by_template 清理两张表并返回总数""" + tpl = _make_template() + repo.create(tpl) + seg = _make_segment(tpl.id, order=1) + repo.create_segments([seg]) + + old = TemplateSegmentModel( + id=str(uuid.uuid4()), + template_id=tpl.id, + segment_order=2, + duration_min=1.0, + duration_max=2.0, + ) + session.add(old) + session.commit() + + count = repo.delete_segments_by_template(tpl.id) + assert count == 2 # 1 from clip_configs + 1 from old table