RandomZ / app /workflows /report_activities.py
StormShadow308's picture
feat: async pipeline, job queue, generation hardening, and docs
732b14f
Raw
History Blame Contribute Delete
10.2 kB
"""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)