Spaces:
Sleeping
Sleeping
File size: 5,852 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 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 | """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,
)
|