Spaces:
Sleeping
Sleeping
| """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: | |
| def defn(cls): | |
| return cls | |
| def run(fn): | |
| return fn | |
| class unsafe: | |
| def imports_passed_through(): | |
| from contextlib import 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, | |
| ) | |
| class ReportGenerationWorkflow: | |
| """Durable orchestration for POST /reports/{id}/generate.""" | |
| 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, | |
| ) | |