"""T4.1 通用长任务 runner 单测(ARCH §7.4)。 直接 `await run_job(...)`(不起后台线程、不连真 DB),注入 fake session 工厂 + fake job repo + fake work,断言: - 成功路:set_running → work → complete(result) → commit; - 失败路:work 抛 → 回滚 → 新 session 里 fail(error) → commit,且异常不冒泡。 独立 session 纪律 = `run_job` 自建 session(经 session_factory),与 `run_overdue_scan` 先例一致。这里的 fake session 工厂记录开了几次 session + 各次 commit。 """ from __future__ import annotations import uuid from collections.abc import AsyncIterator from contextlib import asynccontextmanager from dataclasses import dataclass, field from typing import Any from sqlalchemy.ext.asyncio import AsyncSession from ww_api.services.job_runner import run_job from ww_core.domain.job_repo import ( PROGRESS_COMPLETE, STATUS_DONE, STATUS_FAILED, STATUS_RUNNING, JobView, ) JOB_ID = uuid.UUID("00000000-0000-0000-0000-0000000000aa") class _FakeSession: """最小 fake:记录 commit 次数(提交边界断言)。""" def __init__(self) -> None: self.commits = 0 async def commit(self) -> None: self.commits += 1 class _FakeSessionFactory: """独立 session 工厂替身:`()` → async-CM 产新 `_FakeSession`,记录开了几次。""" def __init__(self) -> None: self.sessions: list[_FakeSession] = [] def __call__(self) -> Any: session = _FakeSession() self.sessions.append(session) @asynccontextmanager async def _cm() -> AsyncIterator[_FakeSession]: yield session return _cm() @dataclass class _FakeJobRepo: """内存 job repo:记录状态流转 + result/error(不触 DB)。""" status: str = "queued" progress: int = 0 result: dict[str, Any] | None = None error: str | None = None calls: list[str] = field(default_factory=list) def _view(self, job_id: uuid.UUID) -> JobView: return JobView( id=job_id, kind="style_learn", status=self.status, progress=self.progress, result=self.result, error=self.error, ) async def set_running(self, job_id: uuid.UUID) -> JobView: self.calls.append("set_running") self.status = STATUS_RUNNING return self._view(job_id) async def complete(self, job_id: uuid.UUID, result: dict[str, Any]) -> JobView: self.calls.append("complete") self.status = STATUS_DONE self.progress = PROGRESS_COMPLETE self.result = dict(result) return self._view(job_id) async def fail(self, job_id: uuid.UUID, error: str) -> JobView: self.calls.append("fail") self.status = STATUS_FAILED self.error = error return self._view(job_id) # ---- success path ---- async def test_run_job_success_sets_done_and_commits() -> None: factory = _FakeSessionFactory() repo = _FakeJobRepo() received: list[AsyncSession] = [] async def work(session: AsyncSession) -> dict[str, Any]: received.append(session) return {"version": 1, "dims_count": 16} await run_job( factory, JOB_ID, work, repo_factory=lambda _s: repo, ) assert repo.calls == ["set_running", "complete"] assert repo.status == STATUS_DONE assert repo.progress == PROGRESS_COMPLETE assert repo.result == {"version": 1, "dims_count": 16} # 自建了一个独立 session 且提交了一次 assert len(factory.sessions) == 1 assert factory.sessions[0].commits == 1 # work 拿到的就是 run_job 自建的 session assert len(received) == 1 # ---- failure path ---- async def test_run_job_failure_sets_failed_and_does_not_raise() -> None: factory = _FakeSessionFactory() repo = _FakeJobRepo() async def work(_session: AsyncSession) -> dict[str, Any]: raise RuntimeError("extraction blew up") # 异常被吞(后台任务边界),不冒泡 await run_job( factory, JOB_ID, work, repo_factory=lambda _s: repo, ) assert "complete" not in repo.calls assert "fail" in repo.calls assert repo.status == STATUS_FAILED assert repo.error == "extraction blew up" # 失败置态在一个**全新** session 里完成并 commit(原事务作废) assert len(factory.sessions) == 2 assert factory.sessions[1].commits == 1