File size: 4,981 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
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
"""Redis LIST queue for report generation jobs."""

from __future__ import annotations

import asyncio
import logging
from collections.abc import Awaitable, Callable

from app.config import settings
from app.jobs.models import GenerationJob, JobType
from app.redis_client import get_redis, job_queue_enabled
from app.services.generation import mark_report_generation_failed, run_agentic_full_report_job, run_generation

logger = logging.getLogger(__name__)


def job_queue_active() -> bool:
    return job_queue_enabled()


async def enqueue_generation_job(job: GenerationJob) -> None:
    """Push a job onto the Redis list (left = newest; worker BRPOP from right)."""
    redis = await get_redis()
    await redis.lpush(settings.job_queue_key, job.model_dump_json())
    logger.info(
        "Enqueued %s job report=%s tenant=%s",
        job.job_type.value,
        job.report_id,
        job.tenant_id,
    )


async def _brpop_job() -> GenerationJob | None:
    redis = await get_redis()
    result = await redis.brpop(
        settings.job_queue_key,
        timeout=int(settings.job_queue_block_seconds),
    )
    if not result:
        return None
    _, raw = result
    return GenerationJob.model_validate_json(raw)


async def _execute_job(job: GenerationJob) -> None:
    if job.job_type == JobType.generate:
        p = job.payload
        await run_generation(
            report_id=job.report_id,
            tenant_id=job.tenant_id,
            template_id=str(p.get("template_id") or ""),
            bullets=list(p.get("bullets") or []),
            mode=str(p.get("mode") or "generate"),
            ai_level=int(p.get("ai_level") or 3),
            ai_percent=p.get("ai_percent"),
            retrieval_level=str(p.get("retrieval_level") or "paragraph"),
            force_regenerate=bool(p.get("force_regenerate", False)),
            strict_uploaded_only=bool(p.get("strict_uploaded_only", False)),
            reference_document_ids=p.get("reference_document_ids"),
            draft_paragraph=p.get("draft_paragraph"),
            interference_level=p.get("interference_level"),
            template_ids=p.get("template_ids"),
            bullets_by_section=p.get("bullets_by_section"),
        )
        return

    if job.job_type == JobType.agentic_full:
        p = job.payload
        await run_agentic_full_report_job(
            job.report_id,
            job.tenant_id,
            bullets_by_section=dict(p.get("bullets_by_section") or {}),
            ai_percent=int(p.get("ai_percent") or 50),
            retrieval_level=str(p.get("retrieval_level") or "paragraph"),
            reference_document_ids=p.get("reference_document_ids"),
            similarity_scan=bool(p.get("similarity_scan", False)),
            peer_sections=dict(p.get("peer_sections") or {}),
            similarity_exclude_document_ids=p.get("similarity_exclude_document_ids"),
            interference_level=p.get("interference_level"),
        )
        return

    raise ValueError(f"Unknown job type: {job.job_type}")


async def _run_job_safe(job: GenerationJob) -> None:
    try:
        await _execute_job(job)
    except Exception as exc:  # noqa: BLE001
        logger.exception(
            "Job failed type=%s report=%s",
            job.job_type.value,
            job.report_id,
        )
        await mark_report_generation_failed(job.report_id, job.tenant_id, str(exc))


async def process_jobs_forever() -> None:
    """Blocking worker loop — run from ``jobs_worker.py``."""
    sem = asyncio.Semaphore(int(settings.job_queue_max_concurrent))
    _tasks: set[asyncio.Task] = set()  # type: ignore[type-arg]
    logger.info(
        "Jobs worker listening key=%s concurrency=%s",
        settings.job_queue_key,
        settings.job_queue_max_concurrent,
    )

    async def _worker(job: GenerationJob) -> None:
        async with sem:
            await _run_job_safe(job)

    while True:
        job = await _brpop_job()
        if job is None:
            continue
        from app.api.background_tasks import _log_task_outcome

        t = asyncio.create_task(_worker(job))
        _tasks.add(t)
        t.add_done_callback(_tasks.discard)
        t.add_done_callback(_log_task_outcome)


async def dispatch_or_enqueue(
    *,
    job: GenerationJob,
    inline_factory: Callable[[], Awaitable[None]],
) -> str:
    """Enqueue to Redis or run inline; returns ``redis`` | ``inline``.

    Falls back to in-process execution when Redis is configured but unreachable.
    """
    from app.api.background_tasks import spawn_background_task

    if job_queue_active():
        try:
            await enqueue_generation_job(job)
            return "redis"
        except Exception as exc:  # noqa: BLE001
            logger.warning(
                "Redis enqueue failed for report=%s (%s); falling back to in-process task",
                job.report_id,
                exc,
            )
    spawn_background_task(inline_factory())
    return "inline"