From 14580f536a483fc67eea79b3b40e427456769c9c Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Mon, 14 Sep 2026 03:15:00 +0800 Subject: [PATCH] =?UTF-8?q?refactor(backend):=20=E6=99=BA=E8=83=BD?= =?UTF-8?q?=E5=89=AA=E8=BE=91=E5=90=8E=E7=AB=AF=E5=85=AC=E5=85=B1=E9=80=BB?= =?UTF-8?q?=E8=BE=91=E6=8A=BD=E5=8F=96+=E8=A7=A3=E8=80=A6=20(#1884)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/api/app/api/routes/generation_tasks.py | 164 ++--------- .../api/routes/generation_variant_plans.py | 21 +- apps/api/app/services/edit_plan_service.py | 23 +- apps/api/app/services/generation_common.py | 175 ++++++++++++ tests/unit/test_generation_common.py | 259 ++++++++++++++++++ 5 files changed, 465 insertions(+), 177 deletions(-) create mode 100644 apps/api/app/services/generation_common.py create mode 100644 tests/unit/test_generation_common.py diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index 152974ac7..d1e301ae1 100755 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -58,41 +58,10 @@ def _variant_value(values: list[str], index: int, fallback: str = "") -> str: def _query_voice_durations(db: Session, voice_ids: list[str]) -> list[float]: - """批量查询配音素材时长(秒),#1749 配音时长分配用。 + """[已下沉] 路由层兼容别名 → app.services.generation_common.query_voice_durations。""" + from app.services.generation_common import query_voice_durations - 逐项 try/float 硬化:MagicMock/异常/缺失 → 0.0(无配音不分配,不阻断)。 - - #1855 P0修复:不再对 voice_ids 去重,保持与调用方传入顺序/长度一致, - 允许同配音id多次出现时返回相同时长(支持"同配音N变体"的时长对齐)。 - """ - # 先去重查询(IN 查询性能优化),但最终按原始 voice_ids 顺序返回 - raw_ids = list(voice_ids or []) - if not raw_ids: - return [] - # 去重且保序,用于 SQL IN 查询;空字符串/None 视为无效id → 0.0 - unique_ids: list[str] = [] - _seen: set[str] = set() - for v in raw_ids: - if v and v not in _seen: - _seen.add(v) - unique_ids.append(v) - if not unique_ids: - return [0.0 for _ in raw_ids] - try: - from packages.adapters.sqlalchemy_impl.models import AssetModel - - rows = db.query(AssetModel.id, AssetModel.duration).filter(AssetModel.id.in_(unique_ids)).all() - dur_map: dict[str, float] = {} - for row in rows: - try: - dur_map[row[0]] = float(row[1] or 0.0) - except (TypeError, ValueError): - dur_map[row[0]] = 0.0 - # 按原始 voice_ids 顺序返回,保持长度一致;空/None/未查到 → 0.0 - return [dur_map.get(v, 0.0) if v else 0.0 for v in raw_ids] - except Exception: - logger.warning("[生成任务] 配音时长查询失败(按无配音处理,不阻断)", exc_info=True) - return [0.0 for _ in raw_ids] + return query_voice_durations(db, voice_ids) def _to_generation_task_response(task) -> GenerationTaskResponse: @@ -198,61 +167,10 @@ def _writeback_edit_plan_config( title_config: dict | None, db: Session, ) -> None: - """任务入队成功后,回写 EditPlan.config:generation_task_id + title_config。 + """[已下沉] 路由层兼容别名 → app.services.generation_common.writeback_edit_plan_config。""" + from app.services.generation_common import writeback_edit_plan_config - 用 merge 方式更新,不整体覆盖 config,避免丢失其他字段。 - 失败只记日志,不影响任务创建。 - """ - if not plan_id: - return - try: - from packages.adapters.sqlalchemy_impl.models import EditPlanModel - - plan_model = db.query(EditPlanModel).filter(EditPlanModel.id == plan_id).first() - if plan_model is None: - logger.warning("[生成任务] 回写plan.config失败: plan不存在 plan_id=%s", plan_id) - return - - current_config = plan_model.config if isinstance(plan_model.config, dict) else {} - merged = dict(current_config) - merged["generation_task_id"] = task_id - - # 检查标题是否发生变化,如果变化则清除 cover 字段强制重新生成封面 - if title_config: - old_title_config = merged.get("title_config", {}) or {} - old_title_text = (old_title_config.get("text") or "").strip() - new_title_text = (title_config.get("text") or "").strip() - if old_title_text != new_title_text: - # 标题变化,清除旧封面 - if "cover" in merged: - del merged["cover"] - logger.info( - "[生成任务] 标题变化,清除旧封面: plan_id=%s old_title=%s new_title=%s", - plan_id, - old_title_text, - new_title_text, - ) - merged["title_config"] = title_config - - plan_model.config = merged - db.commit() - logger.info( - "[生成任务] 回写plan.config成功: plan_id=%s task_id=%s keys=%s", - plan_id, - task_id, - list(merged.keys()), - ) - except Exception as e: - logger.warning( - "[生成任务] 回写plan.config异常(不影响任务创建): plan_id=%s error=%s", - plan_id, - e, - exc_info=True, - ) - try: - db.rollback() - except Exception: - pass + return writeback_edit_plan_config(plan_id, task_id, title_config, db) def _resolve_project_and_library( @@ -503,25 +421,14 @@ def create_generation_task( # 各变体配音时长(查询硬化:异常 → 0.0 不阻断) voice_durations = _query_voice_durations(db, variant_voices) - # 解析批量源 plan:优先前端传入;否则按 template_id + user 查最新(与单任务兜底同源) - batch_source_plan_id = request.source_edit_plan_id - if not batch_source_plan_id and request.template_id: - try: - from packages.adapters.sqlalchemy_impl.models import EditPlanModel + # 解析批量源 plan:优先前端传入;否则按 template_id + user 查最新(公共函数) + from app.services.generation_common import resolve_latest_plan_by_template - _latest = ( - db.query(EditPlanModel) - .filter( - EditPlanModel.template_id == request.template_id, - EditPlanModel.created_by_user_id == user_id, - ) - .order_by(EditPlanModel.created_at.desc()) - .first() - ) - if _latest: - batch_source_plan_id = _latest.id - except Exception: - logger.warning("[生成任务] 批量源 plan 解析失败", exc_info=True) + batch_source_plan_id = ( + request.source_edit_plan_id + or resolve_latest_plan_by_template(db, template_id=request.template_id, user_id=user_id) + or "" + ) if not batch_source_plan_id and not request.variant_plan_ids: # 无任何可用源 plan:批量变体无从选片,明确报错,严禁静默共用/同源 @@ -566,24 +473,10 @@ def create_generation_task( ) from clone_err variant_plan_ids.append(_plan0.id) - # #1855 P0:批次区间避让表,从变体0实际clips构建初始值 - def _collect_segments(pid): - segs = {} - _sk, _pg = 0, 500 - while True: - _b = _plan_svc._clip_repo.list_by_plan(pid, skip=_sk, limit=_pg) - if not _b: - break - for _c in _b: - if _c.asset_id and float(_c.duration or 0) > 0: - _st = float(_c.start_time or 0.0) - segs.setdefault(_c.asset_id, []).append((_st, _st + float(_c.duration))) - if len(_b) < _pg: - break - _sk += _pg - return segs + # #1855 P0:批次区间避让表,从变体0实际clips构建初始值(公共函数) + from app.services.generation_common import collect_plan_segments as _collect_segments - _batch_segments = _collect_segments(_plan0.id) + _batch_segments = _collect_segments(_plan0.id, _plan_svc._clip_repo) # 变体 1..N-1 独立选片(传入累积batch_segments做素材区间避让) for task_index in range(1, count): @@ -631,7 +524,7 @@ def create_generation_task( # #1855 P0:把新变体的clips区间追加到batch_segments,供下一变体避让 try: - _new_segs = _collect_segments(variant.id) + _new_segs = _collect_segments(variant.id, _plan_svc._clip_repo) for _aid, _ivs in _new_segs.items(): _batch_segments.setdefault(_aid, []).extend(_ivs) except Exception: @@ -659,24 +552,13 @@ def create_generation_task( ) _single_vd: list[float] = _query_voice_durations(db, _voices) _single_dur = _single_vd[0] if _single_vd else 0.0 - _single_plan = request.source_edit_plan_id - if not _single_plan and request.template_id: - try: - from packages.adapters.sqlalchemy_impl.models import EditPlanModel + from app.services.generation_common import resolve_latest_plan_by_template - _latest = ( - db.query(EditPlanModel) - .filter( - EditPlanModel.template_id == request.template_id, - EditPlanModel.created_by_user_id == user_id, - ) - .order_by(EditPlanModel.created_at.desc()) - .first() - ) - if _latest: - _single_plan = _latest.id - except Exception: - logger.warning("[生成任务] 单任务源 plan 解析失败", exc_info=True) + _single_plan = ( + request.source_edit_plan_id + or resolve_latest_plan_by_template(db, template_id=request.template_id, user_id=user_id) + or "" + ) if _single_dur > 0 and _single_plan: from app.services.edit_plan_service import EditPlanService diff --git a/apps/api/app/api/routes/generation_variant_plans.py b/apps/api/app/api/routes/generation_variant_plans.py index 7c1de1123..bd62fd804 100644 --- a/apps/api/app/api/routes/generation_variant_plans.py +++ b/apps/api/app/api/routes/generation_variant_plans.py @@ -90,25 +90,12 @@ def create_variant_plans( except VariantVoiceError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc - # 解析源 plan:显式传入优先;否则按 template_id + user 查最新 + # 解析源 plan:显式传入优先;否则按 template_id + user 查最新(公共函数) + from app.services.generation_common import resolve_latest_plan_by_template + source_plan_id = request.source_edit_plan_id.strip() if not source_plan_id and request.template_id.strip(): - try: - from packages.adapters.sqlalchemy_impl.models import EditPlanModel - - _latest = ( - db.query(EditPlanModel) - .filter( - EditPlanModel.template_id == request.template_id.strip(), - EditPlanModel.created_by_user_id == user_id, - ) - .order_by(EditPlanModel.created_at.desc()) - .first() - ) - if _latest: - source_plan_id = _latest.id - except Exception: - logger.exception("[variant-plans] 源 plan 解析失败") + source_plan_id = resolve_latest_plan_by_template(db, template_id=request.template_id, user_id=user_id) or "" if not source_plan_id: raise HTTPException( diff --git a/apps/api/app/services/edit_plan_service.py b/apps/api/app/services/edit_plan_service.py index 6f2a4949d..f9f752404 100755 --- a/apps/api/app/services/edit_plan_service.py +++ b/apps/api/app/services/edit_plan_service.py @@ -1003,25 +1003,10 @@ class EditPlanService: logger.exception("变体0 配音分配失败(不阻断): plan=%s", plan0.id) plan_ids.append(plan0.id) - # #1855 P0:批次内素材区间避让表——从变体0实际落库的clips构建初始值 - def _collect_plan_segments(pid: str) -> dict[str, list[tuple[float, float]]]: - """分页读取 plan 所有 clips,构建 {asset_id: [(start, end), ...]} 区间表。""" - segs: dict[str, list[tuple[float, float]]] = {} - _sk2, _pg2 = 0, 500 - while True: - _b2 = self._clip_repo.list_by_plan(pid, skip=_sk2, limit=_pg2) - if not _b2: - break - for _c in _b2: - if _c.asset_id and float(_c.duration or 0) > 0: - _st = float(_c.start_time or 0.0) - segs.setdefault(_c.asset_id, []).append((_st, _st + float(_c.duration))) - if len(_b2) < _pg2: - break - _sk2 += _pg2 - return segs + # #1855 P0:批次内素材区间避让表——从变体0实际落库的clips构建初始值(公共函数) + from app.services.generation_common import collect_plan_segments as _collect_plan_segments - batch_segments_acc: dict[str, list[tuple[float, float]]] = _collect_plan_segments(plan0.id) + batch_segments_acc: dict[str, list[tuple[float, float]]] = _collect_plan_segments(plan0.id, self._clip_repo) # 变体 1..N-1:独立选片(传入累积的 batch_segments 做区间避让) for i in range(1, count): @@ -1061,7 +1046,7 @@ class EditPlanService: # #1855 P0:把当前新变体的 clips 区间追加到 batch_segments,供下一变体避让 try: - _new_segs = _collect_plan_segments(variant.id) + _new_segs = _collect_plan_segments(variant.id, self._clip_repo) for _aid, _ivs in _new_segs.items(): batch_segments_acc.setdefault(_aid, []).extend(_ivs) except Exception: diff --git a/apps/api/app/services/generation_common.py b/apps/api/app/services/generation_common.py new file mode 100644 index 000000000..fd39fbe6b --- /dev/null +++ b/apps/api/app/services/generation_common.py @@ -0,0 +1,175 @@ +"""智能剪辑公共服务辅助函数(从 route 层下沉)。 + +集中管理: +- query_voice_durations:批量查询配音素材时长 +- writeback_edit_plan_config:任务入队后回写 EditPlan.config +- collect_plan_segments:分页读取 plan clips 构建素材区间表(变体避让用) +- resolve_latest_plan_by_template:按 template_id + user_id 查最新 EditPlan + +设计原则: +- 无副作用的纯查询 / 幂等写回;失败一律不阻断主流程(记日志 + 返回安全默认值) +- 不依赖 FastAPI / HTTPException,便于 service 层和 worker 复用 +""" + +from __future__ import annotations + +import logging +from typing import Any, Optional + +from sqlalchemy.orm import Session + +logger = logging.getLogger(__name__) + + +def query_voice_durations(db: Session, voice_ids: list[str]) -> list[float]: + """批量查询配音素材时长(秒),#1749 配音时长分配用。 + + 逐项 try/float 硬化:MagicMock/异常/缺失 → 0.0(无配音不分配,不阻断)。 + + #1855 P0修复:不再对 voice_ids 去重,保持与调用方传入顺序/长度一致, + 允许同配音id多次出现时返回相同时长(支持"同配音N变体"的时长对齐)。 + """ + raw_ids = list(voice_ids or []) + if not raw_ids: + return [] + unique_ids: list[str] = [] + _seen: set[str] = set() + for v in raw_ids: + if v and v not in _seen: + _seen.add(v) + unique_ids.append(v) + if not unique_ids: + return [0.0 for _ in raw_ids] + try: + from packages.adapters.sqlalchemy_impl.models import AssetModel + + rows = db.query(AssetModel.id, AssetModel.duration).filter(AssetModel.id.in_(unique_ids)).all() + dur_map: dict[str, float] = {} + for row in rows: + try: + dur_map[row[0]] = float(row[1] or 0.0) + except (TypeError, ValueError): + dur_map[row[0]] = 0.0 + return [dur_map.get(v, 0.0) if v else 0.0 for v in raw_ids] + except Exception: + logger.warning("[generation_common] 配音时长查询失败(按无配音处理,不阻断)", exc_info=True) + return [0.0 for _ in raw_ids] + + +def writeback_edit_plan_config( + plan_id: str, + task_id: str, + title_config: dict | None, + db: Session, +) -> None: + """任务入队成功后,回写 EditPlan.config:generation_task_id + title_config。 + + 用 merge 方式更新,不整体覆盖 config,避免丢失其他字段。 + 失败只记日志,不影响任务创建。 + """ + if not plan_id: + return + try: + from packages.adapters.sqlalchemy_impl.models import EditPlanModel + + plan_model = db.query(EditPlanModel).filter(EditPlanModel.id == plan_id).first() + if plan_model is None: + logger.warning("[generation_common] 回写plan.config失败: plan不存在 plan_id=%s", plan_id) + return + + current_config = plan_model.config if isinstance(plan_model.config, dict) else {} + merged = dict(current_config) + merged["generation_task_id"] = task_id + + if title_config: + old_title_config = merged.get("title_config", {}) or {} + old_title_text = (old_title_config.get("text") or "").strip() + new_title_text = (title_config.get("text") or "").strip() + if old_title_text != new_title_text: + if "cover" in merged: + del merged["cover"] + logger.info( + "[generation_common] 标题变化,清除旧封面: plan_id=%s old_title=%s new_title=%s", + plan_id, + old_title_text, + new_title_text, + ) + merged["title_config"] = title_config + + plan_model.config = merged + db.commit() + logger.info( + "[generation_common] 回写plan.config成功: plan_id=%s task_id=%s keys=%s", + plan_id, + task_id, + list(merged.keys()), + ) + except Exception as e: + logger.warning( + "[generation_common] 回写plan.config异常(不影响任务创建): plan_id=%s error=%s", + plan_id, + e, + exc_info=True, + ) + try: + db.rollback() + except Exception: + pass + + +def collect_plan_segments( + plan_id: str, + clip_repo: Any, + *, + page_size: int = 500, +) -> dict[str, list[tuple[float, float]]]: + """分页读取 plan 所有 clips,构建 {asset_id: [(start, end), ...]} 素材区间表。 + + 用于 #1855 P0 批次内素材区间避让(变体间素材片段重叠控制)。 + """ + segs: dict[str, list[tuple[float, float]]] = {} + 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: + if c.asset_id and float(c.duration or 0) > 0: + st = float(c.start_time or 0.0) + segs.setdefault(c.asset_id, []).append((st, st + float(c.duration))) + if len(batch) < pg: + break + sk += pg + return segs + + +def resolve_latest_plan_by_template( + db: Session, + *, + template_id: str, + user_id: str, +) -> Optional[str]: + """按 template_id + user_id 查找最新的 EditPlan.id(模板兜底用)。找不到返回 None。""" + if not (template_id or "").strip(): + return None + try: + from packages.adapters.sqlalchemy_impl.models import EditPlanModel + + latest = ( + db.query(EditPlanModel) + .filter( + EditPlanModel.template_id == template_id.strip(), + EditPlanModel.created_by_user_id == user_id, + ) + .order_by(EditPlanModel.created_at.desc()) + .first() + ) + return latest.id if latest else None + except Exception: + logger.warning( + "[generation_common] 按template查找最新plan失败: template=%s user=%s", + template_id, + user_id, + exc_info=True, + ) + return None diff --git a/tests/unit/test_generation_common.py b/tests/unit/test_generation_common.py new file mode 100644 index 000000000..baa1b9d38 --- /dev/null +++ b/tests/unit/test_generation_common.py @@ -0,0 +1,259 @@ +"""generation_common 公共服务辅助函数单元测试。 + +覆盖 query_voice_durations / writeback_edit_plan_config / collect_plan_segments / +resolve_latest_plan_by_template 四个下沉函数的主路径、边界与容错路径。 +""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest + +# ═══════════════════════════════════════════════════════════════════════════════ +# query_voice_durations +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestQueryVoiceDurations: + def _make_db_with_rows(self, rows): + """构造 MagicMock db,query().filter().all() 返回 rows。""" + db = MagicMock() + db.query.return_value.filter.return_value.all.return_value = list(rows) + return db + + def test_empty_input_returns_empty_list(self): + from app.services.generation_common import query_voice_durations + + db = MagicMock() + assert query_voice_durations(db, []) == [] + assert query_voice_durations(db, None) == [] + db.query.assert_not_called() + + def test_all_empty_or_falsy_ids_returns_zero_list(self): + from app.services.generation_common import query_voice_durations + + db = MagicMock() + assert query_voice_durations(db, ["", None, ""]) == [0.0, 0.0, 0.0] + + def test_normal_lookup_returns_durations_in_input_order(self): + from app.services.generation_common import query_voice_durations + + db = self._make_db_with_rows([("v1", 3.5), ("v2", 7.2)]) + result = query_voice_durations(db, ["v1", "v2", "v-missing"]) + assert result == [3.5, 7.2, 0.0] + + def test_duplicate_ids_returns_consistent_durations_preserves_order(self): + """#1855:同配音 id 多次出现应返回相同时长,保持输入顺序/长度。""" + from app.services.generation_common import query_voice_durations + + db = self._make_db_with_rows([("v1", 4.0)]) + result = query_voice_durations(db, ["v1", "v1", "v1"]) + assert result == [4.0, 4.0, 4.0] + + def test_non_numeric_duration_coerced_to_zero(self): + from app.services.generation_common import query_voice_durations + + db = self._make_db_with_rows([("v1", None), ("v2", "not-a-number"), ("v3", 2.0)]) + result = query_voice_durations(db, ["v1", "v2", "v3"]) + assert result == [0.0, 0.0, 2.0] + + def test_db_exception_returns_zeros_and_logs(self, caplog): + from app.services.generation_common import query_voice_durations + + db = MagicMock() + db.query.side_effect = RuntimeError("DB boom") + with caplog.at_level("WARNING"): + result = query_voice_durations(db, ["v1", "v2"]) + assert result == [0.0, 0.0] + assert any("配音时长查询失败" in rec.message for rec in caplog.records) + + +# ═══════════════════════════════════════════════════════════════════════════════ +# writeback_edit_plan_config +# ═══════════════════════════════════════════════════════════════════════════════ + + +def _make_plan_model(config=None): + plan = MagicMock() + plan.config = config if config is not None else {} + return plan + + +class TestWritebackEditPlanConfig: + def test_empty_plan_id_returns_immediately(self): + from app.services.generation_common import writeback_edit_plan_config + + db = MagicMock() + writeback_edit_plan_config("", "task1", None, db) + db.query.assert_not_called() + + def test_plan_not_found_logs_and_returns(self, caplog): + from app.services.generation_common import writeback_edit_plan_config + + db = MagicMock() + db.query.return_value.filter.return_value.first.return_value = None + with caplog.at_level("WARNING"): + writeback_edit_plan_config("p999", "task1", None, db) + db.commit.assert_not_called() + assert any("plan不存在" in rec.message for rec in caplog.records) + + def test_writes_task_id_preserves_existing_config(self): + from app.services.generation_common import writeback_edit_plan_config + + plan = _make_plan_model({"other": "keep-me"}) + db = MagicMock() + db.query.return_value.filter.return_value.first.return_value = plan + writeback_edit_plan_config("p1", "task-xyz", None, db) + assert plan.config["generation_task_id"] == "task-xyz" + assert plan.config["other"] == "keep-me" + assert "title_config" not in plan.config + db.commit.assert_called_once() + + def test_merges_title_config_without_title_change(self): + from app.services.generation_common import writeback_edit_plan_config + + plan = _make_plan_model({"title_config": {"text": "old"}, "cover": "x"}) + db = MagicMock() + db.query.return_value.filter.return_value.first.return_value = plan + writeback_edit_plan_config("p1", "t1", {"text": "old"}, db) + assert plan.config["title_config"] == {"text": "old"} + # 标题未变 → cover 保留 + assert plan.config.get("cover") == "x" + + def test_title_change_clears_cover(self): + from app.services.generation_common import writeback_edit_plan_config + + plan = _make_plan_model({"title_config": {"text": "old"}, "cover": "x"}) + db = MagicMock() + db.query.return_value.filter.return_value.first.return_value = plan + writeback_edit_plan_config("p1", "t1", {"text": "new-title"}, db) + assert "cover" not in plan.config + assert plan.config["title_config"] == {"text": "new-title"} + + def test_config_not_dict_treated_as_empty(self): + from app.services.generation_common import writeback_edit_plan_config + + plan = _make_plan_model(config=None) + db = MagicMock() + db.query.return_value.filter.return_value.first.return_value = plan + writeback_edit_plan_config("p1", "t1", {"text": "hi"}, db) + assert plan.config["generation_task_id"] == "t1" + assert plan.config["title_config"] == {"text": "hi"} + + def test_exception_triggers_rollback_and_logs(self, caplog): + from app.services.generation_common import writeback_edit_plan_config + + db = MagicMock() + db.query.return_value.filter.return_value.first.side_effect = RuntimeError("fail") + with caplog.at_level("WARNING"): + writeback_edit_plan_config("p1", "t1", None, db) + db.rollback.assert_called_once() + assert any("回写plan.config异常" in rec.message for rec in caplog.records) + + def test_exception_with_rollback_also_failing_is_safe(self, caplog): + """外层异常后,db.rollback() 自己也抛异常时也不应中断(pass 兜底)。""" + from app.services.generation_common import writeback_edit_plan_config + + db = MagicMock() + db.query.return_value.filter.return_value.first.side_effect = RuntimeError("fail") + db.rollback.side_effect = RuntimeError("rollback boom") + with caplog.at_level("WARNING"): + # 不应抛出异常 + writeback_edit_plan_config("p1", "t1", None, db) + assert any("回写plan.config异常" in rec.message for rec in caplog.records) + + +# ═══════════════════════════════════════════════════════════════════════════════ +# collect_plan_segments +# ═══════════════════════════════════════════════════════════════════════════════ + + +def _make_clip(asset_id, start, duration): + c = MagicMock() + c.asset_id = asset_id + c.start_time = start + c.duration = duration + return c + + +class TestCollectPlanSegments: + def test_empty_plan_returns_empty(self): + from app.services.generation_common import collect_plan_segments + + repo = MagicMock() + repo.list_by_plan.return_value = [] + assert collect_plan_segments("p1", repo) == {} + + def test_single_page_collects_segments(self): + from app.services.generation_common import collect_plan_segments + + repo = MagicMock() + repo.list_by_plan.side_effect = [ + [_make_clip("a1", 0.0, 5.0), _make_clip("a1", 10.0, 3.0), _make_clip("a2", 2.0, 4.0)], + [], + ] + segs = collect_plan_segments("p1", repo, page_size=500) + assert segs["a1"] == [(0.0, 5.0), (10.0, 13.0)] + assert segs["a2"] == [(2.0, 6.0)] + + def test_pagination_walks_all_batches(self): + from app.services.generation_common import collect_plan_segments + + repo = MagicMock() + page1 = [_make_clip("a1", 0.0, 1.0)] * 2 + page2 = [_make_clip("a2", 0.0, 2.0)] * 2 + page3 = [_make_clip("a3", 0.0, 1.0)] # short final batch → stop + repo.list_by_plan.side_effect = [page1, page2, page3] + segs = collect_plan_segments("p1", repo, page_size=2) + assert set(segs.keys()) == {"a1", "a2", "a3"} + assert repo.list_by_plan.call_count == 3 + + def test_skips_zero_or_negative_duration_clips(self): + from app.services.generation_common import collect_plan_segments + + repo = MagicMock() + repo.list_by_plan.side_effect = [ + [_make_clip(None, 0.0, 5.0), _make_clip("a1", 0.0, 0.0), _make_clip("a1", 1.0, -1.0)], + [], + ] + assert collect_plan_segments("p1", repo) == {} + + +# ═══════════════════════════════════════════════════════════════════════════════ +# resolve_latest_plan_by_template +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestResolveLatestPlanByTemplate: + @pytest.mark.parametrize("tid", ["", None, " "]) + def test_empty_template_returns_none(self, tid): + from app.services.generation_common import resolve_latest_plan_by_template + + db = MagicMock() + assert resolve_latest_plan_by_template(db, template_id=tid, user_id="u1") is None + db.query.assert_not_called() + + def test_returns_latest_plan_id(self): + from app.services.generation_common import resolve_latest_plan_by_template + + db = MagicMock() + latest = MagicMock(id="plan-xyz") + db.query.return_value.filter.return_value.order_by.return_value.first.return_value = latest + assert resolve_latest_plan_by_template(db, template_id=" tpl1 ", user_id="u1") == "plan-xyz" + + def test_no_plan_returns_none(self): + from app.services.generation_common import resolve_latest_plan_by_template + + db = MagicMock() + db.query.return_value.filter.return_value.order_by.return_value.first.return_value = None + assert resolve_latest_plan_by_template(db, template_id="tpl", user_id="u") is None + + def test_db_exception_returns_none_and_logs(self, caplog): + from app.services.generation_common import resolve_latest_plan_by_template + + db = MagicMock() + db.query.side_effect = RuntimeError("boom") + 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)