File size: 3,614 Bytes
732b14f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
"""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()