"""Temporal activities wrapping the existing async generation service.""" from __future__ import annotations import logging from datetime import timedelta from app.config import settings from app.services.generation import ( _resolve_generation_sections, _run_section_job, mark_report_generation_failed, run_generation, ) from app.db.database import get_session_factory from app.db.models import Report from app.workflows.models import ReportWorkflowRequest, SectionGenerationInput logger = logging.getLogger(__name__) _ACTIVITY_TIMEOUT = timedelta(minutes=15) try: from temporalio import activity TEMPORAL_AVAILABLE = True except ImportError: # pragma: no cover TEMPORAL_AVAILABLE = False def activity(defn=None, *, name=None): # type: ignore[no-redef] def _wrap(fn): return fn return _wrap if defn is None else _wrap(defn) @activity.defn(name="fetch_sources") async def fetch_sources(request: ReportWorkflowRequest) -> dict[str, str | list[str] | None]: """Load report metadata required before section generation.""" activity.heartbeat("fetch_sources") # type: ignore[attr-defined] factory = get_session_factory() async with factory() as db: report = await db.get(Report, request.report_id) if report is None or report.tenant_id != request.tenant_id: await mark_report_generation_failed( request.report_id, request.tenant_id, "Report not found for workflow.", ) raise ValueError("Report not found for workflow") return { "report_id": request.report_id, "tenant_id": request.tenant_id, "primary_document_id": report.document_id, } @activity.defn(name="retrieve_context") async def retrieve_context(request: ReportWorkflowRequest) -> dict[str, str]: """Prefetch report metadata and warm style/vector paths before section LLM work.""" activity.heartbeat("retrieve_context") # type: ignore[attr-defined] sections = _resolve_generation_sections(request.template_id, request.template_ids) out: dict[str, str] = { "sections": ",".join(sections), "retrieval_level": request.retrieval_level, "section_count": str(len(sections)), } factory = get_session_factory() async with factory() as db: report = await db.get(Report, request.report_id) if report is None or report.tenant_id != request.tenant_id: out["warning"] = "report_not_found" return out out["primary_document_id"] = str(report.document_id or "") out["survey_level"] = str(report.survey_level or "") try: from app.services.generation import _get_or_build_style_profile await _get_or_build_style_profile(request.tenant_id) out["style_profile"] = "warmed" except Exception as exc: # noqa: BLE001 logger.warning("retrieve_context style warm failed: %s", exc) out["style_profile"] = "skipped" if sections and out.get("primary_document_id"): try: from app.agentic.tools import retrieve_tenant_evidence_async seed_q = f"{sections[0]} RICS inspection evidence" hits = await retrieve_tenant_evidence_async( query=seed_q, tenant_id=request.tenant_id, primary_document_id=str(out["primary_document_id"]), secondary_document_ids=list(request.reference_document_ids or []), k=8, rerank_top_n=4, ) out["seed_hits"] = str(len(hits)) except Exception as exc: # noqa: BLE001 logger.warning("retrieve_context seed retrieval failed: %s", exc) out["seed_hits"] = "0" return out @activity.defn(name="generate_section") async def generate_section_activity(payload: SectionGenerationInput) -> str: """Run generation for a single section (full existing pipeline).""" import asyncio async def _heartbeat_loop() -> None: while True: activity.heartbeat(f"generate_section:{payload.section_code}") # type: ignore[attr-defined] await asyncio.sleep(25) hb = asyncio.create_task(_heartbeat_loop()) try: factory = get_session_factory() async with factory() as db: report = await db.get(Report, payload.report_id) if report is None or report.tenant_id != payload.tenant_id: await mark_report_generation_failed( payload.report_id, payload.tenant_id, "Report not found for section activity.", ) raise ValueError("Report not found for section activity") await _run_section_job( mode=payload.mode, db=db, report=report, tenant_id=payload.tenant_id, section_code=payload.section_code, sec_bullets=list(payload.bullets or []), ai_level=payload.ai_level, ai_percent=payload.ai_percent, retrieval_level=payload.retrieval_level, force_regenerate=payload.force_regenerate, strict_uploaded_only=payload.strict_uploaded_only, reference_document_ids=payload.reference_document_ids, draft_paragraph=payload.draft_paragraph, interference_level=payload.interference_level, ) await db.commit() finally: hb.cancel() try: await hb except asyncio.CancelledError: pass return payload.section_code @activity.defn(name="validate_output") async def validate_output( report_id: str, tenant_id: str, section_codes: list[str], ) -> dict[str, str | list[str]]: """Verify expected sections have non-empty persisted text.""" from sqlalchemy import select from app.db.models import ReportSection activity.heartbeat("validate_output") # type: ignore[attr-defined] missing: list[str] = [] factory = get_session_factory() async with factory() as db: result = await db.execute( select(ReportSection).where(ReportSection.report_id == report_id) ) by_code = {row.section_code: row for row in result.scalars().all()} for code in section_codes: row = by_code.get(code) if row is None or not (row.text or "").strip(): missing.append(code) status = "ok" if not missing else "partial" return { "report_id": report_id, "tenant_id": tenant_id, "validation": status, "missing_sections": missing, } @activity.defn(name="assemble_report") async def assemble_report( report_id: str, tenant_id: str, failed_sections: list[str], total_sections: int, section_codes: list[str], ) -> str: """Finalize workflow — set report status after all section activities.""" from sqlalchemy import select from app.db.models import ReportSection from app.services.generation import _finalize_multi_section_report activity.heartbeat("assemble_report") # type: ignore[attr-defined] failed_set = {str(c) for c in (failed_sections or []) if str(c).strip()} factory = get_session_factory() async with factory() as db: report = await db.get(Report, report_id) if report is None or report.tenant_id != tenant_id: await mark_report_generation_failed( report_id, tenant_id, "Report not found during workflow finalization.", ) return "failed" if section_codes: result = await db.execute( select(ReportSection).where(ReportSection.report_id == report_id) ) by_code = {row.section_code: row for row in result.scalars().all()} for code in section_codes: row = by_code.get(code) if row is None or not (row.text or "").strip(): failed_set.add(code) total = max(1, int(total_sections)) _finalize_multi_section_report( report, failure_count=len(failed_set), total=total, phase="temporal", ) await db.commit() return str(report.status.value) def activity_schedule_options() -> dict[str, timedelta]: return {"schedule_to_close_timeout": _ACTIVITY_TIMEOUT} @activity.defn(name="execute_report_generation") async def execute_report_generation(request: ReportWorkflowRequest) -> list[str]: """Run the full generation pipeline (single or multi-section) in one activity.""" import asyncio sections = _resolve_generation_sections(request.template_id, request.template_ids) async def _heartbeat_loop() -> None: while True: activity.heartbeat(f"execute_report_generation:{request.report_id}") # type: ignore[attr-defined] await asyncio.sleep(25) hb = asyncio.create_task(_heartbeat_loop()) try: await run_generation( report_id=request.report_id, tenant_id=request.tenant_id, template_id=request.template_id, bullets=request.bullets, 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, template_ids=request.template_ids, bullets_by_section=request.bullets_by_section, ) finally: hb.cancel() try: await hb except asyncio.CancelledError: pass return sections def temporal_activities_enabled() -> bool: return bool(settings.enable_temporal_workflow and TEMPORAL_AVAILABLE)