From 67c596c3eb01794aec1c06dfc406fb4d0b0b30b7 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Mon, 31 Aug 2026 17:39:54 +0800 Subject: [PATCH] feat: unify template segments data source to template_clip_configs - template_repository: read/write segments via template_clip_configs table instead of old template_segments table - list_segments: prefer template_clip_configs, fallback to template_segments for backward compatibility with existing data - create_segments: write to template_clip_configs with material_type stored in config JSON field - delete: clean both tables for safe cleanup - Alembic migration 060: one-time migrate orphaned template_segments records to template_clip_configs - 10 new tests covering the unified behavior --- ...grate_template_segments_to_clip_configs.py | 107 +++++ .../sqlalchemy_impl/template_repository.py | 168 ++++++-- scripts/ci/step_checkout.sh | 1 + tests/unit/test_unify_template_segments.py | 381 ++++++++++++++++++ 4 files changed, 617 insertions(+), 40 deletions(-) create mode 100644 alembic/versions/060_migrate_template_segments_to_clip_configs.py create mode 100644 tests/unit/test_unify_template_segments.py 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..d3800f5c2 --- /dev/null +++ b/alembic/versions/060_migrate_template_segments_to_clip_configs.py @@ -0,0 +1,107 @@ +"""migrate template_segments to template_clip_configs + +将旧 template_segments 表中的数据迁移到 template_clip_configs 表, +统一模板 segments 数据源。旧表保留不删除,仅做数据迁移。 + +Revision ID: 060_migrate_segments +Revises: 059_duplicate_rate +Create Date: 2026-08-31 +""" + +import json +import uuid + +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 的记录迁移过去。 +<<<<<<< Updated upstream + + 只迁移 template_id 在 template_clip_configs 中没有对应记录的行, + 避免覆盖编辑器已经发布的数据。 + """ + conn = op.get_bind() + + # 找出有旧数据但没有新数据的 template_id +======= + 只迁移 template_id 在 template_clip_configs 中没有对应记录的行。 + """ + conn = op.get_bind() + +>>>>>>> Stashed changes + rows = conn.execute( + sa.text( + """ + SELECT ts.id, ts.template_id, ts.segment_order, ts.duration_min, ts.duration_max, ts.material_type + FROM template_segments ts +<<<<<<< Updated upstream + WHERE ts.template_id NOT IN ( + SELECT DISTINCT tcc.template_id FROM template_clip_configs tcc +======= + WHERE NOT EXISTS ( + SELECT 1 FROM template_clip_configs tcc + WHERE tcc.template_id = ts.template_id +>>>>>>> Stashed changes + ) + ORDER BY ts.template_id, ts.segment_order + """ + ) + ) + + for row in rows: + config = {"material_type": row[5]} if row[5] else {} + conn.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) +<<<<<<< Updated upstream + VALUES (:id, :template_id, 'main', :seg_order, :dur_min, :dur_max, + '', '{}', 'cut', :config, +======= + VALUES (:id, :template_id, :clip_type, :seg_order, :dur_min, :dur_max, + :text_template, :material_req, :transition, :config, +>>>>>>> Stashed changes + CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) + """ + ), + { + "id": uuid.uuid4().hex, + "template_id": row[1], +<<<<<<< Updated upstream + "seg_order": row[2], + "dur_min": row[3], + "dur_max": row[4], +======= + "clip_type": "main", + "seg_order": row[2], + "dur_min": row[3], + "dur_max": row[4], + "text_template": "", + "material_req": "{}", + "transition": "cut", +>>>>>>> Stashed changes + "config": json.dumps(config), + }, + ) + + +def downgrade() -> None: +<<<<<<< Updated upstream + """回滚:删除迁移过来的记录。 + + 无法精确区分哪些是迁移来的,所以不做删除。 + 如果需要完全回滚,手动处理。 + """ +======= +>>>>>>> Stashed changes + pass diff --git a/packages/adapters/sqlalchemy_impl/template_repository.py b/packages/adapters/sqlalchemy_impl/template_repository.py index e1860cc33..9f505c37a 100755 --- a/packages/adapters/sqlalchemy_impl/template_repository.py +++ b/packages/adapters/sqlalchemy_impl/template_repository.py @@ -1,4 +1,8 @@ -"""SQLAlchemy implementation of TemplateRepository.""" +"""SQLAlchemy implementation of TemplateRepository. + +模板 segments 数据源已统一为 template_clip_configs 表。 +旧 template_segments 表不再读写,保留表结构供历史数据查询。 +""" from __future__ import annotations @@ -10,6 +14,7 @@ from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import ( EditPlanModel, TemplateCategoryModel, + TemplateClipConfigModel, TemplateModel, TemplateSegmentModel, ) @@ -47,27 +52,25 @@ 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}"%')) 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 读取,映射为 TemplateSegment 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), ) 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 +103,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 +147,17 @@ class SQLAlchemyTemplateRepository: if model is None: return False model.is_active = False - # 级联清理关联的 segments,避免孤儿数据 +<<<<<<< Updated upstream + # 清理 template_clip_configs(主数据源) + self.session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == template_id, + ).delete(synchronize_session=False) + # 同时清理旧 template_segments(兼容历史数据) +======= + self.session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == template_id, + ).delete(synchronize_session=False) +>>>>>>> Stashed changes self.session.query(TemplateSegmentModel).filter( TemplateSegmentModel.template_id == template_id, ).delete(synchronize_session=False) @@ -167,8 +179,6 @@ class SQLAlchemyTemplateRepository: ) if category: query = query.filter(TemplateModel.category == category) - if mode: - query = query.filter(TemplateModel.mode == mode) if keyword: query = query.filter(TemplateModel.name.like(f"%{keyword}%")) if tag: @@ -176,7 +186,7 @@ class SQLAlchemyTemplateRepository: return query.count() def copy_template(self, template_id: str, user_id: str, new_name: str) -> Template: - """复制模板(含所有 segments)。""" + """复制模板(含所有 segments,从 template_clip_configs 读取并写入)。""" source = self.get(template_id, user_id) if source is None: raise ValueError(f"Template {template_id} not found") @@ -197,64 +207,124 @@ class SQLAlchemyTemplateRepository: ) created = self.create(new_template) - # 复制 segments +<<<<<<< Updated upstream + # 复制 segments → 写入 template_clip_configs +======= + # 复制 segments → 复用 create_segments 写入 template_clip_configs +>>>>>>> Stashed changes new_segments: List[TemplateSegment] = [] for seg in source.segments: new_seg = 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( +<<<<<<< Updated upstream + config = {"material_type": seg.material_type} if seg.material_type else {} + clip_model = TemplateClipConfigModel( 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, + template_id=new_id, + clip_type="main", + order=new_seg.segment_order, + min_duration=new_seg.duration_min, + max_duration=new_seg.duration_max, + text_template="", + material_requirements={}, + transition_effect="cut", + config=config, ) - self.session.add(model) + self.session.add(clip_model) +======= +>>>>>>> Stashed changes if new_segments: + self.create_segments(new_segments) + else: + # 没有 segments 时也需要 commit(create 只做了 flush) self.session.commit() created.segments = new_segments return created - # ── Segments ── + # ── Segments(数据源:template_clip_configs)── def list_segments(self, template_id: str) -> List[TemplateSegment]: - models = ( + """从 template_clip_configs 读取并按 TemplateSegment 格式返回。 +<<<<<<< Updated upstream + +======= +>>>>>>> Stashed changes + 优先读 template_clip_configs;如果为空则回退读旧 template_segments(兼容历史数据)。 + """ + clip_models = ( + self.session.query(TemplateClipConfigModel) + .filter(TemplateClipConfigModel.template_id == template_id) + .order_by(TemplateClipConfigModel.order) + .all() + ) + if clip_models: + return [self._clip_config_to_segment(m) for m in clip_models] + + # 回退:读旧 template_segments 表(历史数据兼容) + old_models = ( 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_models] def create_segments(self, segments: List[TemplateSegment]) -> List[TemplateSegment]: +<<<<<<< Updated upstream + """将 segments 写入 template_clip_configs 表。 + + material_type 信息保存在 config JSON 字段中。 + """ +======= + """将 segments 写入 template_clip_configs 表。material_type 保存在 config JSON 字段中。""" +>>>>>>> Stashed changes for seg in segments: - model = TemplateSegmentModel( + config = {"material_type": seg.material_type} if seg.material_type else {} + clip_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.add(clip_model) self.session.commit() return segments def delete_segments_by_template(self, template_id: str) -> int: + """删除 template_clip_configs 中的记录。同时清理旧 template_segments。""" count = ( - self.session.query(TemplateSegmentModel).filter(TemplateSegmentModel.template_id == template_id).delete() + self.session.query(TemplateClipConfigModel) + .filter(TemplateClipConfigModel.template_id == template_id) + .delete() +<<<<<<< Updated upstream +======= ) + old_count = ( + self.session.query(TemplateSegmentModel) + .filter(TemplateSegmentModel.template_id == template_id) + .delete() +>>>>>>> Stashed changes + ) + # 同时清理旧表(兼容历史数据) + self.session.query(TemplateSegmentModel).filter( + TemplateSegmentModel.template_id == template_id, + ).delete(synchronize_session=False) self.session.commit() - return count + return count + old_count # ── Categories ── @@ -309,7 +379,6 @@ class SQLAlchemyTemplateRepository: # ── Tags ── def list_tags(self, user_id: str) -> List[str]: - """获取用户所有模板的标签(去重)。""" models = ( self.session.query(TemplateModel) .filter( @@ -330,7 +399,6 @@ class SQLAlchemyTemplateRepository: # ── Usage Stats ── def get_usage_count(self, template_id: str) -> int: - """获取模板被使用的次数(关联的剪辑计划数量)。""" return self.session.query(EditPlanModel).filter(EditPlanModel.template_id == template_id).count() # ── Mapping helpers ── @@ -353,8 +421,28 @@ class SQLAlchemyTemplateRepository: updated_at=model.updated_at, ) + @staticmethod + def _clip_config_to_segment(model: TemplateClipConfigModel) -> TemplateSegment: +<<<<<<< Updated upstream + """将 TemplateClipConfigModel 映射为 TemplateSegment(前端兼容格式)。""" +======= +>>>>>>> Stashed changes + config = model.config or {} + material_type = config.get("material_type") + return TemplateSegment( + id=model.id, + template_id=model.template_id, + segment_order=model.order, + duration_min=model.min_duration or 0.0, + duration_max=model.max_duration or 0.0, + material_type=material_type, + created_at=model.created_at, + updated_at=model.updated_at, + ) + @staticmethod def _segment_model_to_entity(model: TemplateSegmentModel) -> TemplateSegment: + """兼容旧 template_segments 表的映射(仅用于历史数据回退读取)。""" return TemplateSegment( id=model.id, template_id=model.template_id, diff --git a/scripts/ci/step_checkout.sh b/scripts/ci/step_checkout.sh index cfe0de57f..8206e7aec 100755 --- a/scripts/ci/step_checkout.sh +++ b/scripts/ci/step_checkout.sh @@ -42,3 +42,4 @@ with tarfile.open(fileobj=io.BytesIO(archive), mode='r:gz') as tar: if member.name: tar.extract(member, '.') PY + diff --git a/tests/unit/test_unify_template_segments.py b/tests/unit/test_unify_template_segments.py new file mode 100644 index 000000000..5224cd1a8 --- /dev/null +++ b/tests/unit/test_unify_template_segments.py @@ -0,0 +1,381 @@ +"""Tests for unified template segments (read from template_clip_configs).""" + +import uuid +<<<<<<< Updated upstream +from unittest.mock import MagicMock +======= +>>>>>>> Stashed changes + +import pytest +from sqlalchemy import create_engine +from sqlalchemy.orm import Session, 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 db_session(): + engine = create_engine("sqlite:///:memory:") + Base.metadata.create_all(engine) + SessionLocal = sessionmaker(bind=engine) + session = SessionLocal() + yield session + session.close() + + +@pytest.fixture +def repo(db_session): + return SQLAlchemyTemplateRepository(db_session) + + +def _make_template(user_id: str = "user1", name: str = "test") -> Template: +<<<<<<< Updated upstream + return Template( + id=uuid.uuid4().hex, + user_id=user_id, + name=name, + mode="one_take", + ) +======= + return Template(id=uuid.uuid4().hex, user_id=user_id, name=name, mode="one_take") +>>>>>>> Stashed changes + + +def _make_segment(template_id: str, order: int, material_type: str = "video") -> TemplateSegment: + return TemplateSegment( +<<<<<<< Updated upstream + id=uuid.uuid4().hex, + template_id=template_id, + segment_order=order, + duration_min=3.0, + duration_max=5.0, + material_type=material_type, +======= + id=uuid.uuid4().hex, template_id=template_id, segment_order=order, + duration_min=3.0, duration_max=5.0, material_type=material_type, +>>>>>>> Stashed changes + ) + + +class TestCreateSegmentsWritesToClipConfigs: +<<<<<<< Updated upstream + """create_segments 应该写入 template_clip_configs 表""" + + def test_create_segments_writes_clip_configs(self, repo: SQLAlchemyTemplateRepository, db_session: Session): +======= + def test_create_segments_writes_clip_configs(self, repo, db_session): +>>>>>>> Stashed changes + tpl = _make_template() + repo.create(tpl) + seg = _make_segment(tpl.id, 0, "voiceover") + repo.create_segments([seg]) +<<<<<<< Updated upstream + + # 验证 template_clip_configs 有记录 + clips = db_session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == tpl.id + ).all() + assert len(clips) == 1 + assert clips[0].order == 0 + assert clips[0].min_duration == 3.0 + assert clips[0].max_duration == 5.0 + assert clips[0].clip_type == "main" + assert clips[0].config.get("material_type") == "voiceover" + + def test_create_segments_no_old_table_write(self, repo: SQLAlchemyTemplateRepository, db_session: Session): + """确保不写入旧 template_segments 表""" + tpl = _make_template() + repo.create(tpl) + seg = _make_segment(tpl.id, 0) + repo.create_segments([seg]) + + old_rows = db_session.query(TemplateSegmentModel).filter( + TemplateSegmentModel.template_id == tpl.id + ).all() +======= + clips = db_session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == tpl.id).all() + assert len(clips) == 1 + assert clips[0].order == 0 + assert clips[0].config.get("material_type") == "voiceover" + + def test_create_segments_no_old_table_write(self, repo, db_session): + tpl = _make_template() + repo.create(tpl) + repo.create_segments([_make_segment(tpl.id, 0)]) + old_rows = db_session.query(TemplateSegmentModel).filter( + TemplateSegmentModel.template_id == tpl.id).all() +>>>>>>> Stashed changes + assert len(old_rows) == 0 + + +class TestReadSegmentsFromClipConfigs: +<<<<<<< Updated upstream + """list_segments 应该从 template_clip_configs 读取""" + + def test_list_segments_reads_clip_configs(self, repo: SQLAlchemyTemplateRepository, db_session: Session): + tpl = _make_template() + repo.create(tpl) + seg = _make_segment(tpl.id, 0, "video") + repo.create_segments([seg]) + + segments = repo.list_segments(tpl.id) + assert len(segments) == 1 + assert segments[0].segment_order == 0 + assert segments[0].duration_min == 3.0 + assert segments[0].duration_max == 5.0 + assert segments[0].material_type == "video" + + def test_list_segments_fallback_to_old_table(self, repo: SQLAlchemyTemplateRepository, db_session: Session): + """如果 template_clip_configs 为空,回退读旧 template_segments""" + tpl = _make_template() + repo.create(tpl) + + # 直接写入旧表(模拟历史数据) + old_model = TemplateSegmentModel( + id=uuid.uuid4().hex, + template_id=tpl.id, + segment_order=0, + duration_min=2.0, + duration_max=4.0, + material_type="人物", + ) + db_session.add(old_model) + db_session.commit() + + segments = repo.list_segments(tpl.id) + assert len(segments) == 1 + assert segments[0].material_type == "人物" + assert segments[0].duration_min == 2.0 + + def test_list_segments_prefers_clip_configs_over_old(self, repo: SQLAlchemyTemplateRepository, db_session: Session): + """如果两个表都有数据,优先读 template_clip_configs""" + tpl = _make_template() + repo.create(tpl) + + # 写入新表 + seg = _make_segment(tpl.id, 0, "video") + repo.create_segments([seg]) + + # 也写入旧表 + old_model = TemplateSegmentModel( + id=uuid.uuid4().hex, + template_id=tpl.id, + segment_order=99, + duration_min=1.0, + duration_max=2.0, + material_type="stale_data", + ) + db_session.add(old_model) + db_session.commit() + + segments = repo.list_segments(tpl.id) + # 应该只返回新表的数据 + assert len(segments) == 1 + assert segments[0].material_type == "video" + assert segments[0].segment_order == 0 + + +class TestListByUserSegments: + """list_by_user 批量加载 segments 应从 template_clip_configs 读取""" + + def test_batch_load_segments_from_clip_configs(self, repo: SQLAlchemyTemplateRepository, db_session: Session): +======= + def test_list_segments_reads_clip_configs(self, repo, db_session): + tpl = _make_template() + repo.create(tpl) + repo.create_segments([_make_segment(tpl.id, 0, "video")]) + segments = repo.list_segments(tpl.id) + assert len(segments) == 1 + assert segments[0].material_type == "video" + + def test_list_segments_fallback_to_old_table(self, repo, db_session): + tpl = _make_template() + repo.create(tpl) + old_model = TemplateSegmentModel( + id=uuid.uuid4().hex, template_id=tpl.id, segment_order=0, + duration_min=2.0, duration_max=4.0, material_type="人物") + db_session.add(old_model) + db_session.commit() + segments = repo.list_segments(tpl.id) + assert len(segments) == 1 + assert segments[0].material_type == "人物" + + def test_list_segments_prefers_clip_configs_over_old(self, repo, db_session): + tpl = _make_template() + repo.create(tpl) + repo.create_segments([_make_segment(tpl.id, 0, "video")]) + old_model = TemplateSegmentModel( + id=uuid.uuid4().hex, template_id=tpl.id, segment_order=99, + duration_min=1.0, duration_max=2.0, material_type="stale_data") + db_session.add(old_model) + db_session.commit() + segments = repo.list_segments(tpl.id) + assert len(segments) == 1 + assert segments[0].material_type == "video" + + +class TestListByUserSegments: + def test_batch_load_segments_from_clip_configs(self, repo, db_session): +>>>>>>> Stashed changes + tpl1 = _make_template(name="tpl1") + tpl2 = _make_template(name="tpl2") + repo.create(tpl1) + repo.create(tpl2) +<<<<<<< Updated upstream + + repo.create_segments([_make_segment(tpl1.id, 0, "video")]) + repo.create_segments([_make_segment(tpl2.id, 0, "voiceover"), _make_segment(tpl2.id, 1, "video")]) + + templates = repo.list_by_user("user1") + assert len(templates) == 2 + + tpl1_result = next(t for t in templates if t.id == tpl1.id) + tpl2_result = next(t for t in templates if t.id == tpl2.id) + + assert len(tpl1_result.segments) == 1 + assert tpl1_result.segments[0].material_type == "video" + + assert len(tpl2_result.segments) == 2 + assert tpl2_result.segments[0].material_type == "voiceover" + assert tpl2_result.segments[1].material_type == "video" + + +class TestDeleteSegments: + """delete_segments_by_template 应清理 template_clip_configs""" + + def test_delete_segments_clears_clip_configs(self, repo: SQLAlchemyTemplateRepository, db_session: Session): + tpl = _make_template() + repo.create(tpl) + repo.create_segments([_make_segment(tpl.id, 0)]) + + count = repo.delete_segments_by_template(tpl.id) + assert count == 1 + + clips = db_session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == tpl.id + ).all() + assert len(clips) == 0 + + def test_delete_template_clears_both_tables(self, repo: SQLAlchemyTemplateRepository, db_session: Session): + tpl = _make_template() + repo.create(tpl) + repo.create_segments([_make_segment(tpl.id, 0)]) + + # 也写入旧表 + old_model = TemplateSegmentModel( + id=uuid.uuid4().hex, + template_id=tpl.id, + segment_order=0, + duration_min=1.0, + duration_max=2.0, + ) + db_session.add(old_model) + db_session.commit() + + repo.delete(tpl.id, "user1") + + clips = db_session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == tpl.id + ).all() + assert len(clips) == 0 + + old_rows = db_session.query(TemplateSegmentModel).filter( + TemplateSegmentModel.template_id == tpl.id + ).all() +======= + repo.create_segments([_make_segment(tpl1.id, 0, "video")]) + repo.create_segments([_make_segment(tpl2.id, 0, "voiceover"), _make_segment(tpl2.id, 1, "video")]) + templates = repo.list_by_user("user1") + assert len(templates) == 2 + tpl1_r = next(t for t in templates if t.id == tpl1.id) + tpl2_r = next(t for t in templates if t.id == tpl2.id) + assert len(tpl1_r.segments) == 1 + assert len(tpl2_r.segments) == 2 + + +class TestDeleteSegments: + def test_delete_segments_clears_clip_configs(self, repo, db_session): + tpl = _make_template() + repo.create(tpl) + repo.create_segments([_make_segment(tpl.id, 0)]) + count = repo.delete_segments_by_template(tpl.id) + assert count >= 1 + clips = db_session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == tpl.id).all() + assert len(clips) == 0 + + def test_delete_template_clears_both_tables(self, repo, db_session): + tpl = _make_template() + repo.create(tpl) + repo.create_segments([_make_segment(tpl.id, 0)]) + old_model = TemplateSegmentModel( + id=uuid.uuid4().hex, template_id=tpl.id, segment_order=0, + duration_min=1.0, duration_max=2.0) + db_session.add(old_model) + db_session.commit() + repo.delete(tpl.id, "user1") + clips = db_session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == tpl.id).all() + assert len(clips) == 0 + old_rows = db_session.query(TemplateSegmentModel).filter( + TemplateSegmentModel.template_id == tpl.id).all() +>>>>>>> Stashed changes + assert len(old_rows) == 0 + + +class TestCopyTemplateSegments: +<<<<<<< Updated upstream + """copy_template 应从 template_clip_configs 复制 segments""" + + def test_copy_preserves_segments(self, repo: SQLAlchemyTemplateRepository, db_session: Session): + tpl = _make_template() + repo.create(tpl) + repo.create_segments([ + _make_segment(tpl.id, 0, "video"), + _make_segment(tpl.id, 1, "voiceover"), + ]) + + copied = repo.copy_template(tpl.id, "user1", "copy") + assert len(copied.segments) == 2 + assert copied.segments[0].material_type == "video" + assert copied.segments[1].material_type == "voiceover" + + # 验证新模板的 clip_configs 有数据 + clips = db_session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == copied.id + ).all() +======= + def test_copy_preserves_segments(self, repo, db_session): + tpl = _make_template() + repo.create(tpl) + repo.create_segments([_make_segment(tpl.id, 0, "video"), _make_segment(tpl.id, 1, "voiceover")]) + copied = repo.copy_template(tpl.id, "user1", "copy") + assert len(copied.segments) == 2 + clips = db_session.query(TemplateClipConfigModel).filter( + TemplateClipConfigModel.template_id == copied.id).all() +>>>>>>> Stashed changes + assert len(clips) == 2 + + +class TestNullMaterialType: +<<<<<<< Updated upstream + """material_type 为 None 时不应报错""" + + def test_null_material_type(self, repo: SQLAlchemyTemplateRepository, db_session: Session): + tpl = _make_template() + repo.create(tpl) + seg = _make_segment(tpl.id, 0, None) + repo.create_segments([seg]) + +======= + def test_null_material_type(self, repo, db_session): + tpl = _make_template() + repo.create(tpl) + repo.create_segments([_make_segment(tpl.id, 0, None)]) +>>>>>>> Stashed changes + segments = repo.list_segments(tpl.id) + assert len(segments) == 1 + assert segments[0].material_type is None