Spaces:
Sleeping
Sleeping
| """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) | |
| 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, | |
| } | |
| 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 | |
| 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 | |
| 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, | |
| } | |
| 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} | |
| 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) | |