"""SQLAlchemy store for enterprise-research tasks.""" from __future__ import annotations import json from datetime import UTC, datetime from typing import Any from sqlalchemy import delete, select, update from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from deerflow.persistence.enterprise_research.base import EnterpriseResearchReportJobStore, EnterpriseResearchTaskStore from deerflow.persistence.enterprise_research.model import EnterpriseResearchReportJobRow, EnterpriseResearchTaskRow def _to_dict(row: EnterpriseResearchTaskRow) -> dict[str, Any]: try: sources = json.loads(row.sources_json or "[]") except (TypeError, ValueError): sources = [] data: dict[str, Any] = { "id": row.id, "user_id": row.user_id, "subject": row.subject, "focus": row.focus, "template_id": row.template_id, "knowledge_base_ids": list(row.knowledge_base_ids or []), "queries": list(row.queries or []), "selected_source_ids": list(row.selected_source_ids or []), "sources": sources if isinstance(sources, list) else [], "plan_markdown": row.plan_markdown, "report_markdown": row.report_markdown, "report_summary": row.report_summary, "report_model_name": row.report_model_name, "active_report_job_id": row.active_report_job_id, "status": row.status, "warning": row.warning, "created_at": row.created_at.isoformat() if row.created_at else None, "updated_at": row.updated_at.isoformat() if row.updated_at else None, } return data class EnterpriseResearchTaskRepository(EnterpriseResearchTaskStore): def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory async def list_tasks(self, user_id: str, *, limit: int = 20) -> list[dict[str, Any]]: stmt = ( select(EnterpriseResearchTaskRow) .where(EnterpriseResearchTaskRow.user_id == user_id) .order_by(EnterpriseResearchTaskRow.updated_at.desc()) .limit(limit) ) async with self._sf() as session: result = await session.execute(stmt) return [_to_dict(row) for row in result.scalars()] async def get_task(self, task_id: str, user_id: str) -> dict[str, Any] | None: stmt = select(EnterpriseResearchTaskRow).where( EnterpriseResearchTaskRow.id == task_id, EnterpriseResearchTaskRow.user_id == user_id, ) async with self._sf() as session: row = (await session.execute(stmt)).scalar_one_or_none() return _to_dict(row) if row is not None else None async def create_task(self, data: dict[str, Any]) -> dict[str, Any]: now = datetime.now(UTC) row = EnterpriseResearchTaskRow( id=str(data["id"]), user_id=str(data["user_id"]), subject=str(data["subject"]), focus=str(data.get("focus") or ""), template_id=str(data["template_id"]), knowledge_base_ids=list(data.get("knowledge_base_ids") or []), queries=list(data.get("queries") or []), selected_source_ids=list(data.get("selected_source_ids") or []), sources_json=json.dumps(list(data.get("sources") or []), ensure_ascii=False), plan_markdown=data.get("plan_markdown"), report_markdown=data.get("report_markdown"), report_summary=data.get("report_summary"), report_model_name=data.get("report_model_name"), active_report_job_id=data.get("active_report_job_id"), status=str(data.get("status") or "draft"), warning=data.get("warning"), created_at=now, updated_at=now, ) async with self._sf() as session: session.add(row) # Do not refresh after commit: a MySQL read/write split can route it # to a lagging replica. The flushed in-memory row is authoritative. await session.flush() payload = _to_dict(row) await session.commit() return payload async def update_task(self, task_id: str, user_id: str, data: dict[str, Any]) -> dict[str, Any] | None: stmt = select(EnterpriseResearchTaskRow).where( EnterpriseResearchTaskRow.id == task_id, EnterpriseResearchTaskRow.user_id == user_id, ) async with self._sf() as session: row = (await session.execute(stmt)).scalar_one_or_none() if row is None: return None if "status" in data: row.status = str(data["status"]) if "warning" in data: row.warning = data["warning"] if "sources" in data: row.sources_json = json.dumps(list(data["sources"] or []), ensure_ascii=False) if "selected_source_ids" in data: row.selected_source_ids = list(data["selected_source_ids"] or []) if "plan_markdown" in data: row.plan_markdown = data["plan_markdown"] for key in ("report_markdown", "report_summary", "report_model_name", "active_report_job_id"): if key in data: setattr(row, key, data[key]) row.updated_at = datetime.now(UTC) await session.flush() payload = _to_dict(row) await session.commit() return payload async def delete_task(self, task_id: str, user_id: str) -> bool: async with self._sf() as session: result = await session.execute( delete(EnterpriseResearchTaskRow).where( EnterpriseResearchTaskRow.id == task_id, EnterpriseResearchTaskRow.user_id == user_id, ) ) await session.commit() return bool(result.rowcount) def _job_to_dict(row: EnterpriseResearchReportJobRow) -> dict[str, Any]: try: sources = json.loads(row.sources_json or "[]") except (TypeError, ValueError): sources = [] return { "id": row.id, "task_id": row.task_id, "user_id": row.user_id, "status": row.status, "phase": row.phase, "model_name": row.model_name, "action": row.action or "initial", "instruction": row.instruction, "style": row.style, "target_length": row.target_length, "plan_snapshot": row.plan_snapshot, "sources": sources if isinstance(sources, list) else [], "report_markdown": row.report_markdown, "original_report": row.original_report, "report_summary": row.report_summary, "error_message": row.error_message, "created_at": row.created_at.isoformat() if row.created_at else None, "updated_at": row.updated_at.isoformat() if row.updated_at else None, } class EnterpriseResearchReportJobRepository(EnterpriseResearchReportJobStore): def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory async def create_job(self, data: dict[str, Any]) -> dict[str, Any]: now = datetime.now(UTC) status = str(data.get("status") or "queued") active_dedupe_key = ( f"{data['user_id']}:{data['task_id']}" if status in {"queued", "running"} else None ) row = EnterpriseResearchReportJobRow( id=str(data["id"]), task_id=str(data["task_id"]), user_id=str(data["user_id"]), status=status, phase=str(data.get("phase") or "queued"), model_name=data.get("model_name"), action=str(data.get("action") or "initial"), active_dedupe_key=active_dedupe_key, instruction=data.get("instruction"), style=data.get("style"), target_length=data.get("target_length"), plan_snapshot=str(data["plan_snapshot"]), sources_json=json.dumps(list(data.get("sources") or []), ensure_ascii=False), original_report=data.get("original_report"), report_markdown=str(data.get("report_markdown") or ""), report_summary=data.get("report_summary"), error_message=data.get("error_message"), created_at=now, updated_at=now, ) async with self._sf() as session: session.add(row) try: await session.flush() payload = _job_to_dict(row) await session.commit() return payload except IntegrityError: await session.rollback() if not active_dedupe_key: raise existing = (await session.execute( select(EnterpriseResearchReportJobRow) .where(EnterpriseResearchReportJobRow.active_dedupe_key == active_dedupe_key) .limit(1) )).scalar_one_or_none() if existing is None: raise return _job_to_dict(existing) async def get_job(self, job_id: str, user_id: str) -> dict[str, Any] | None: async with self._sf() as session: row = (await session.execute(select(EnterpriseResearchReportJobRow).where( EnterpriseResearchReportJobRow.id == job_id, EnterpriseResearchReportJobRow.user_id == user_id, ))).scalar_one_or_none() return _job_to_dict(row) if row is not None else None async def get_active_for_task(self, task_id: str, user_id: str) -> dict[str, Any] | None: async with self._sf() as session: row = (await session.execute(select(EnterpriseResearchReportJobRow).where( EnterpriseResearchReportJobRow.task_id == task_id, EnterpriseResearchReportJobRow.user_id == user_id, EnterpriseResearchReportJobRow.status.in_(["queued", "running"]), ).order_by(EnterpriseResearchReportJobRow.updated_at.desc()).limit(1))).scalar_one_or_none() return _job_to_dict(row) if row is not None else None async def claim_job(self, job_id: str, user_id: str) -> bool: async with self._sf() as session: result = await session.execute( update(EnterpriseResearchReportJobRow) .where( EnterpriseResearchReportJobRow.id == job_id, EnterpriseResearchReportJobRow.user_id == user_id, EnterpriseResearchReportJobRow.status == "queued", ) .values(status="running", updated_at=datetime.now(UTC)) ) await session.commit() return bool(result.rowcount) async def requeue_job(self, job_id: str, user_id: str, *, expected_updated_at: str | None) -> bool: if not expected_updated_at: return False try: expected = datetime.fromisoformat(expected_updated_at) except ValueError: return False async with self._sf() as session: result = await session.execute( update(EnterpriseResearchReportJobRow) .where( EnterpriseResearchReportJobRow.id == job_id, EnterpriseResearchReportJobRow.user_id == user_id, EnterpriseResearchReportJobRow.status == "running", EnterpriseResearchReportJobRow.updated_at == expected, ) .values( status="queued", phase="queued", report_markdown="", error_message=None, updated_at=datetime.now(UTC), ) ) await session.commit() return bool(result.rowcount) async def transition_job( self, job_id: str, user_id: str, *, expected_statuses: set[str], data: dict[str, Any] ) -> dict[str, Any] | None: allowed = { "status", "phase", "report_markdown", "report_summary", "error_message", "model_name", "action", "instruction", "style", "target_length", "original_report", } values = {key: value for key, value in data.items() if key in allowed} if data.get("status") in {"completed", "failed", "cancelled"}: values["active_dedupe_key"] = None values["updated_at"] = datetime.now(UTC) async with self._sf() as session: result = await session.execute( update(EnterpriseResearchReportJobRow) .where( EnterpriseResearchReportJobRow.id == job_id, EnterpriseResearchReportJobRow.user_id == user_id, EnterpriseResearchReportJobRow.status.in_(expected_statuses), ) .values(**values) ) if not result.rowcount: await session.rollback() return None await session.commit() return await self.get_job(job_id, user_id) async def update_job(self, job_id: str, user_id: str, data: dict[str, Any]) -> dict[str, Any] | None: async with self._sf() as session: row = (await session.execute(select(EnterpriseResearchReportJobRow).where( EnterpriseResearchReportJobRow.id == job_id, EnterpriseResearchReportJobRow.user_id == user_id, ))).scalar_one_or_none() if row is None: return None for key in ( "status", "phase", "report_markdown", "report_summary", "error_message", "model_name", "action", "instruction", "style", "target_length", "original_report", ): if key in data: setattr(row, key, data[key]) if data.get("status") in {"completed", "failed", "cancelled"}: row.active_dedupe_key = None row.updated_at = datetime.now(UTC) await session.flush() payload = _job_to_dict(row) await session.commit() return payload async def list_jobs(self, task_id: str, user_id: str, *, limit: int = 10) -> list[dict[str, Any]]: async with self._sf() as session: rows = (await session.execute(select(EnterpriseResearchReportJobRow).where( EnterpriseResearchReportJobRow.task_id == task_id, EnterpriseResearchReportJobRow.user_id == user_id, ).order_by(EnterpriseResearchReportJobRow.updated_at.desc()).limit(limit))).scalars() return [_job_to_dict(row) for row in rows] async def list_recoverable(self, *, limit: int = 100) -> list[dict[str, Any]]: async with self._sf() as session: rows = (await session.execute( select(EnterpriseResearchReportJobRow) .where(EnterpriseResearchReportJobRow.status.in_(["queued", "running"])) .order_by(EnterpriseResearchReportJobRow.updated_at.asc()) .limit(limit) )).scalars() return [_job_to_dict(row) for row in rows]