"""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, )