"""Report validation and terminal status transitions (no stuck ``generating``).""" from __future__ import annotations from unittest.mock import AsyncMock, MagicMock, patch import pytest from app.db.models import ReportStatus from app.services.generation import ( abort_generation_if_report_invalid, mark_report_generation_complete_if_still_generating, mark_report_generation_failed, run_generation, ) @pytest.mark.asyncio async def test_abort_generation_false_when_report_missing() -> None: mock_session = MagicMock() mock_session.get = AsyncMock(return_value=None) mock_ctx = MagicMock() mock_ctx.__aenter__ = AsyncMock(return_value=mock_session) mock_ctx.__aexit__ = AsyncMock(return_value=False) with patch("app.services.generation.get_session_factory") as mock_factory: mock_factory.return_value.return_value = mock_ctx ok = await abort_generation_if_report_invalid( "missing-report-id", "tenant-a", ) assert ok is False @pytest.mark.asyncio async def test_abort_generation_marks_failed_on_tenant_mismatch() -> None: marked: list[tuple[str, str, str]] = [] async def _mark_failed(report_id: str, tenant_id: str, reason: str) -> None: marked.append((report_id, tenant_id, reason)) mock_report = MagicMock() mock_report.tenant_id = "tenant-a" mock_session = MagicMock() mock_session.get = AsyncMock(return_value=mock_report) mock_ctx = MagicMock() mock_ctx.__aenter__ = AsyncMock(return_value=mock_session) mock_ctx.__aexit__ = AsyncMock(return_value=False) with ( patch( "app.services.generation.mark_report_generation_failed", side_effect=_mark_failed, ), patch("app.services.generation.get_session_factory") as mock_factory, ): mock_factory.return_value.return_value = mock_ctx ok = await abort_generation_if_report_invalid( "r-tenant-mismatch", "tenant-b", reason="Tenant mismatch.", ) assert ok is False assert marked == [("r-tenant-mismatch", "tenant-a", "Tenant mismatch.")] @pytest.mark.asyncio async def test_run_generation_marks_failed_on_invalid_section_codes() -> None: marked: list[str] = [] async def _mark_failed(_rid: str, _tid: str, reason: str) -> None: marked.append(reason) with ( patch( "app.services.generation.mark_report_generation_failed", side_effect=_mark_failed, ), patch("app.services.generation.abort_generation_if_report_invalid", return_value=True), ): await run_generation( report_id="r1", tenant_id="tenant-a", template_id="NOT_REAL", bullets=[], ) assert marked assert "Unknown section" in marked[0] @pytest.mark.asyncio async def test_mark_complete_if_still_generating_uses_atomic_update() -> None: mock_result = MagicMock() mock_result.rowcount = 1 mock_session = MagicMock() mock_session.execute = AsyncMock(return_value=mock_result) mock_session.commit = AsyncMock() mock_ctx = MagicMock() mock_ctx.__aenter__ = AsyncMock(return_value=mock_session) mock_ctx.__aexit__ = AsyncMock(return_value=False) with patch("app.services.generation.get_session_factory") as mock_factory: mock_factory.return_value.return_value = mock_ctx await mark_report_generation_complete_if_still_generating("r1", "tenant-a") mock_session.execute.assert_awaited_once() mock_session.commit.assert_awaited_once()