"""GenerationTaskRepository - cleanup_stale_pending 超时 pending 清理单元测试。""" import sys from datetime import datetime, timedelta, timezone from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) from sqlalchemy import create_engine, text from sqlalchemy.orm import sessionmaker from packages.adapters.sqlalchemy_impl.generation_task_repository import ( SQLAlchemyGenerationTaskRepository, ) from packages.adapters.sqlalchemy_impl.models import Base from packages.domain import GenerationTask, GenerationTaskStatus def _repository(): engine = create_engine("sqlite:///:memory:") Base.metadata.create_all(engine) session = sessionmaker(bind=engine)() return SQLAlchemyGenerationTaskRepository(session), session, engine def _make_task(**kwargs) -> GenerationTask: defaults = dict( project_id="proj-1", asset_library_id="lib-1", created_by_user_id="user-1", ) defaults.update(kwargs) return GenerationTask.create(**defaults) # --------------------------------------------------------------------------- # cleanup_stale_pending 基本测试 # --------------------------------------------------------------------------- def test_cleanup_stale_pending_no_tasks_returns_zero(): """没有任务时返回 0。""" repo, _, _ = _repository() count = repo.cleanup_stale_pending(timeout_minutes=30) assert count == 0 def test_cleanup_stale_pending_recent_pending_not_cleaned(): """30 分钟内的 pending 任务不被清理。""" repo, _, _ = _repository() task = _make_task() repo.create(task) # 刚创建的 pending 任务不应被清理 count = repo.cleanup_stale_pending(timeout_minutes=30) assert count == 0 assert repo.get(task.id).status == GenerationTaskStatus.PENDING def test_cleanup_stale_pending_old_pending_marked_failed(): """超过 30 分钟的 pending 任务被标记为 failed。""" repo, _, engine = _repository() task = _make_task() repo.create(task) # 手动把 created_at 改到 1 小时前 with engine.connect() as conn: conn.execute( text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"), {"ts": datetime.now(timezone.utc) - timedelta(hours=1), "id": task.id}, ) conn.commit() count = repo.cleanup_stale_pending(timeout_minutes=30) assert count == 1 saved = repo.get(task.id) assert saved.status == GenerationTaskStatus.FAILED assert saved.error_message == "pending timeout: auto cleanup" assert saved.error_info.get("error_type") == "PendingTimeout" assert "30" in saved.error_info["message"] assert "failed_at" in saved.error_info assert saved.completed_at is not None def test_cleanup_stale_pending_running_not_touched(): """running 任务不受影响,只清理 pending。""" repo, _, engine = _repository() task = _make_task() repo.create(task) task.mark_processing() repo.update(task) # 回写 created_at 到 1 小时前 with engine.connect() as conn: conn.execute( text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"), {"ts": datetime.now(timezone.utc) - timedelta(hours=1), "id": task.id}, ) conn.commit() count = repo.cleanup_stale_pending(timeout_minutes=30) assert count == 0 assert repo.get(task.id).status == GenerationTaskStatus.RUNNING def test_cleanup_stale_pending_custom_timeout(): """自定义超时时间生效。""" repo, _, engine = _repository() task = _make_task() repo.create(task) # 回写 created_at 到 20 分钟前 with engine.connect() as conn: conn.execute( text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"), {"ts": datetime.now(timezone.utc) - timedelta(minutes=20), "id": task.id}, ) conn.commit() # 30 分钟超时:不清理 count_30 = repo.cleanup_stale_pending(timeout_minutes=30) assert count_30 == 0 # 15 分钟超时:清理 count_15 = repo.cleanup_stale_pending(timeout_minutes=15) assert count_15 == 1 assert repo.get(task.id).status == GenerationTaskStatus.FAILED def test_cleanup_stale_pending_multiple(): """批量清理多个超时的 pending 任务。""" repo, _, engine = _repository() tasks = [] for i in range(5): t = _make_task(project_id=f"proj-{i}") repo.create(t) tasks.append(t) # 全部回写 created_at 到 2 小时前 with engine.connect() as conn: for t in tasks: conn.execute( text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"), {"ts": datetime.now(timezone.utc) - timedelta(hours=2), "id": t.id}, ) conn.commit() count = repo.cleanup_stale_pending(timeout_minutes=30) assert count == 5 for t in tasks: assert repo.get(t.id).status == GenerationTaskStatus.FAILED