"""素材片段使用记录追踪服务测试(asset_segment_tracker). 覆盖: - get_used_segments 聚合 metadata 中持久化的区间 - record_used_segments 追加记录(不 commit,保留原有 metadata 字段) - remove_used_segment 匹配删除(tolerance + plan_id) - reset_used_segments 清空轮回(其他字段不动) - make_reset_callback 同时清持久化和内存 - _calc_random_start_time 的 on_exhausted 轮回回调 """ from __future__ import annotations import json import os import sys from pathlib import Path os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing") os.environ.setdefault("DATABASE_URL", "sqlite:///test.db") sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) import pytest from app.services import asset_segment_tracker as ast from app.services.asset_segment_tracker import ( get_used_segments, make_reset_callback, record_used_segments, remove_used_segment, reset_used_segments, ) from packages.domain.plan_generator_utils import _calc_random_start_time class FakeModel: """模拟 AssetModel:id + classification_result(JSON Text)+ updated_at。""" def __init__(self, asset_id: str, meta: dict | None = None): self.id = asset_id self.classification_result = json.dumps(meta, ensure_ascii=False) if meta else None self.updated_at = None def meta(self) -> dict: return json.loads(self.classification_result) if self.classification_result else {} class _InExpr: def __init__(self, ids, models): self._ids = ids self._models = models def all(self): return [self._models[i] for i in self._ids if i in self._models] class _EqExpr: def __init__(self, target_id, models): self._target_id = target_id self._models = models def first(self): return self._models.get(self._target_id) class FakeSession: """模拟 db:db.query(Model).filter(Model.id.in_(ids)).all() / .filter(Model.id == id).first()。 tracker 模块里的 AssetModel 被 monkeypatch 为 FakeModel 类, 这里用挂在类上的伪 column 对象接住 in_ / __eq__。 """ class _Col: def __init__(self, models): self._models = models def in_(self, ids): return _InExpr(list(ids), self._models) def __eq__(self, other): return _EqExpr(other, self._models) def __init__(self, models: dict[str, FakeModel]): self._models = models self.commits = 0 def query(self, _model): col = self._Col(self._models) class _Q: def filter(self_inner, expr): return expr q = _Q() # 让 tracker 里 AssetModel.id 能取到伪 column _model.id = col return q def commit(self): self.commits += 1 @pytest.fixture def patched_model(monkeypatch): """把 tracker 模块内的 AssetModel 替换为 FakeModel(供 FakeSession 挂伪 column)。""" monkeypatch.setattr(ast, "AssetModel", FakeModel) @pytest.fixture def models(): return {} def _db(models): return FakeSession(models) # ── get_used_segments ───────────────────────────────────────────────────────── def test_get_used_segments_aggregates_ranges(patched_model): models = { "a1": FakeModel( "a1", { "used_time_ranges": [ {"start": 1.0, "end": 5.0, "plan_id": "p1"}, {"start": 9.0, "end": 12.0, "plan_id": "p2"}, ] }, ), "a2": FakeModel("a2", {"other": 1}), # 无区间记录 "a3": FakeModel("a3"), # metadata 为空 } db = _db(models) result = get_used_segments(db, ["a1", "a2", "a3", "missing"]) assert result == {"a1": [(1.0, 5.0), (9.0, 12.0)]} def test_get_used_segments_empty_input(patched_model): assert get_used_segments(_db({}), []) == {} # ── record_used_segments ────────────────────────────────────────────────────── def test_record_appends_and_no_commit(patched_model): models = {"a1": FakeModel("a1", {"generation_use_count": 3})} db = _db(models) record_used_segments(db, "a1", 2.0, 6.5, "plan-x") meta = models["a1"].meta() assert meta["generation_use_count"] == 3 # 原有字段保留 ranges = meta["used_time_ranges"] assert len(ranges) == 1 assert ranges[0]["start"] == 2.0 assert ranges[0]["end"] == 6.5 assert ranges[0]["plan_id"] == "plan-x" assert "created_at" in ranges[0] assert db.commits == 0 # 不自行 commit(事务由调用方控制) def test_record_multiple_appends_in_order(patched_model): models = {"a1": FakeModel("a1")} db = _db(models) record_used_segments(db, "a1", 0.0, 4.0, "p1") record_used_segments(db, "a1", 10.0, 14.0, "p1") ranges = models["a1"].meta()["used_time_ranges"] assert [r["start"] for r in ranges] == [0.0, 10.0] def test_record_missing_asset_is_noop(patched_model): db = _db({}) record_used_segments(db, "ghost", 0.0, 1.0, "p1") # 不抛异常 # ── remove_used_segment ─────────────────────────────────────────────────────── def test_remove_matching_segment(patched_model): models = {"a1": FakeModel("a1")} db = _db(models) record_used_segments(db, "a1", 0.0, 4.0, "p1") record_used_segments(db, "a1", 10.0, 14.0, "p1") removed = remove_used_segment(db, "a1", 0.0, 4.0, plan_id="p1") assert removed is True ranges = models["a1"].meta()["used_time_ranges"] assert len(ranges) == 1 assert ranges[0]["start"] == 10.0 def test_remove_not_found_returns_false(patched_model): models = {"a1": FakeModel("a1")} db = _db(models) record_used_segments(db, "a1", 0.0, 4.0, "p1") assert remove_used_segment(db, "a1", 99.0, 100.0, plan_id="p1") is False def test_remove_respects_tolerance(patched_model): models = { "a1": FakeModel("a1", {"used_time_ranges": [{"start": 5.0, "end": 9.0, "plan_id": "p1"}]}), "a2": FakeModel("a2", {"used_time_ranges": [{"start": 5.0, "end": 9.0, "plan_id": "p1"}]}), } db = _db(models) # 偏差 0.3 秒,在 tolerance=0.5 内 → 删除成功 assert remove_used_segment(db, "a1", 5.3, 8.7, plan_id="p1") is True # 偏差 2 秒,超出 tolerance → 删除失败 assert remove_used_segment(db, "a2", 7.0, 11.0, plan_id="p1") is False def test_remove_plan_id_must_match(patched_model): models = {"a1": FakeModel("a1", {"used_time_ranges": [{"start": 5.0, "end": 9.0, "plan_id": "plan-A"}]})} db = _db(models) # 时间匹配但 plan_id 不同 → 不删除 assert remove_used_segment(db, "a1", 5.0, 9.0, plan_id="plan-B") is False assert len(models["a1"].meta()["used_time_ranges"]) == 1 def test_remove_legacy_record_without_plan_id(patched_model): """旧数据记录没有 plan_id 字段时,MediaKit 移动片段仍能按时间匹配删除(防容量泄漏)。""" models = { "a1": FakeModel( "a1", {"used_time_ranges": [{"start": 5.0, "end": 9.0}]}, # 旧记录无 plan_id ) } db = _db(models) # 传入 plan_id,但记录本身无 plan_id → 按时间匹配,允许删除 assert remove_used_segment(db, "a1", 5.0, 9.0, plan_id="plan-new") is True assert models["a1"].meta()["used_time_ranges"] == [] # ── reset_used_segments ─────────────────────────────────────────────────────── def test_reset_clears_ranges_keeps_other_fields(patched_model): models = { "a1": FakeModel( "a1", { "generation_use_count": 9, "used_time_ranges": [{"start": 1, "end": 2}], }, ) } db = _db(models) reset_used_segments(db, "a1") meta = models["a1"].meta() assert meta["used_time_ranges"] == [] assert meta["generation_use_count"] == 9 assert db.commits == 0 # ── make_reset_callback ─────────────────────────────────────────────────────── def test_reset_callback_clears_persisted_and_memory(patched_model): models = {"a1": FakeModel("a1", {"used_time_ranges": [{"start": 0, "end": 30}]})} db = _db(models) used_segments = {"a1": [(0.0, 30.0)], "a2": [(1.0, 2.0)]} cb = make_reset_callback(db, used_segments) cb("a1") assert "a1" not in used_segments # 内存清空 assert "a2" in used_segments # 其他素材不受影响 assert models["a1"].meta()["used_time_ranges"] == [] # ── _calc_random_start_time 轮回回调 ────────────────────────────────────────── def test_calc_random_start_invokes_reset_when_exhausted(): """素材区间被占满(100 次随机必重叠)→ 触发 on_exhausted,重置后重试成功。""" durations = {"a1": 30.0} used = {"a1": [(0.0, 10.0), (10.0, 20.0), (20.0, 30.0)]} reset_called = [] def _on_exhausted(asset_id): reset_called.append(asset_id) used.pop(asset_id, None) # 模拟轮回清空 result = _calc_random_start_time("a1", 10.0, durations, used, on_exhausted=_on_exhausted) assert reset_called == ["a1"] assert result is not None assert 0.0 <= result <= 20.0 # max_start = 30 - 10 def test_calc_random_start_no_callback_keeps_legacy_fallback(): """不传 on_exhausted 时保持旧降级行为,不报错。""" durations = {"a1": 30.0} used = {"a1": [(0.0, 10.0), (10.0, 20.0), (20.0, 30.0)]} result = _calc_random_start_time("a1", 10.0, durations, used) assert result is not None def test_calc_random_start_with_space_does_not_reset(): """有充足空闲区间时不触发 reset。""" durations = {"a1": 100.0} used = {"a1": [(0.0, 50.0)]} reset_called = [] result = _calc_random_start_time("a1", 5.0, durations, used, on_exhausted=lambda aid: reset_called.append(aid)) assert reset_called == [] assert result is not None