RandomZ / app /tests /test_generation_abort_stub.py
StormShadow308's picture
feat: async pipeline, job queue, generation hardening, and docs
732b14f
Raw
History Blame Contribute Delete
3.61 kB
"""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()