RandomZ / app /workflows /report_workflow.py
StormShadow308's picture
feat: async pipeline, job queue, generation hardening, and docs
732b14f
Raw
History Blame Contribute Delete
5.85 kB
"""Temporal workflow for durable report generation."""
from __future__ import annotations
import asyncio
from datetime import timedelta
try:
from temporalio import workflow
TEMPORAL_AVAILABLE = True
except ImportError: # pragma: no cover
TEMPORAL_AVAILABLE = False
class _WorkflowStub:
@staticmethod
def defn(cls):
return cls
@staticmethod
def run(fn):
return fn
class unsafe:
@staticmethod
def imports_passed_through():
from contextlib import contextmanager
@contextmanager
def _cm():
yield
return _cm()
workflow = _WorkflowStub() # type: ignore[misc, assignment]
if TEMPORAL_AVAILABLE:
with workflow.unsafe.imports_passed_through():
from app.services.generation import _resolve_generation_sections
from app.workflows.models import (
ReportWorkflowRequest,
SectionGenerationInput,
WorkflowResult,
)
from app.workflows.report_activities import (
activity_schedule_options,
assemble_report,
execute_report_generation,
fetch_sources,
generate_section_activity,
retrieve_context,
validate_output,
)
else:
from app.services.generation import _resolve_generation_sections
from app.workflows.models import (
ReportWorkflowRequest,
SectionGenerationInput,
WorkflowResult,
)
from app.workflows.report_activities import (
activity_schedule_options,
assemble_report,
execute_report_generation,
fetch_sources,
generate_section_activity,
retrieve_context,
validate_output,
)
@workflow.defn(name="ReportGenerationWorkflow")
class ReportGenerationWorkflow:
"""Durable orchestration for POST /reports/{id}/generate."""
@workflow.run
async def run(self, request: ReportWorkflowRequest) -> WorkflowResult:
opts = activity_schedule_options()
timeout = opts["schedule_to_close_timeout"]
await workflow.execute_activity( # type: ignore[attr-defined]
fetch_sources,
request,
start_to_close_timeout=timeout,
schedule_to_close_timeout=timeout,
)
await workflow.execute_activity( # type: ignore[attr-defined]
retrieve_context,
request,
start_to_close_timeout=timedelta(minutes=5),
schedule_to_close_timeout=timedelta(minutes=5),
)
sections = _resolve_generation_sections(request.template_id, request.template_ids)
failed: list[str] = []
if len(sections) <= 1 or not request.parallel_sections:
try:
await workflow.execute_activity( # type: ignore[attr-defined]
execute_report_generation,
request,
start_to_close_timeout=timeout,
schedule_to_close_timeout=timeout,
)
except Exception: # noqa: BLE001
failed = list(sections)
else:
bullets_map = request.bullets_by_section or {}
async def _one(code: str) -> str:
sec_bullets = bullets_map.get(code) if bullets_map else request.bullets
payload = SectionGenerationInput(
report_id=request.report_id,
tenant_id=request.tenant_id,
section_code=code,
bullets=list(sec_bullets or []),
mode=request.mode,
ai_level=request.ai_level,
ai_percent=request.ai_percent,
retrieval_level=request.retrieval_level,
force_regenerate=request.force_regenerate,
strict_uploaded_only=request.strict_uploaded_only,
reference_document_ids=request.reference_document_ids,
draft_paragraph=request.draft_paragraph,
interference_level=request.interference_level,
)
return await workflow.execute_activity( # type: ignore[attr-defined]
generate_section_activity,
payload,
start_to_close_timeout=timeout,
schedule_to_close_timeout=timeout,
)
outcomes = await asyncio.gather(
*[_one(code) for code in sections],
return_exceptions=True,
)
for code, outcome in zip(sections, outcomes, strict=True):
if isinstance(outcome, BaseException):
failed.append(code)
validation = await workflow.execute_activity( # type: ignore[attr-defined]
validate_output,
args=[request.report_id, request.tenant_id, sections],
start_to_close_timeout=timedelta(minutes=2),
schedule_to_close_timeout=timedelta(minutes=2),
)
missing = validation.get("missing_sections") if isinstance(validation, dict) else []
if isinstance(missing, list):
failed_set = set(failed) | {str(c) for c in missing}
failed = [c for c in sections if c in failed_set]
status = await workflow.execute_activity( # type: ignore[attr-defined]
assemble_report,
args=[request.report_id, request.tenant_id, failed, len(sections), sections],
start_to_close_timeout=timedelta(minutes=2),
schedule_to_close_timeout=timedelta(minutes=2),
)
return WorkflowResult(
report_id=request.report_id,
status=status,
failed_sections=failed,
)