"""SQLAlchemy-backed TaskCOP compatibility task store.""" from __future__ import annotations from datetime import datetime from typing import Any from sqlalchemy import func, or_, select from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from deerflow.persistence.taskcop_tasks.base import TASK_STATUS_ANALYSIS_COMPLETED, TaskCopTaskStore from deerflow.persistence.taskcop_tasks.model import TaskCopTaskRow def _as_legacy(row: TaskCopTaskRow) -> dict[str, Any]: return { "id": row.id, "taskId": row.id, "taskName": row.task_name, "taskContent": row.overview, "overview": row.overview, "taskDirection": row.task_direction, "taskLevel": row.task_level, "taskStatus": row.task_status, "trendSucc": row.task_status >= TASK_STATUS_ANALYSIS_COMPLETED, "taskType": row.task_type, "taskSuperiorFlag": row.task_superior_flag, "pTaskId": row.parent_task_id, "issuingDept": row.issuing_dept, "taskNo": row.task_no, "taskStartDate": row.task_start_date.isoformat() if row.task_start_date else None, "taskEndDate": row.task_end_date.isoformat() if row.task_end_date else None, "sourceType": 0, "ifCreateTask": 0, "sumAll": 0, "endCount": 0, "createTime": row.create_time.isoformat(), "updateTime": row.update_time.isoformat(), } class TaskCopTaskRepository(TaskCopTaskStore): def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory async def create_task(self, data: dict[str, Any]) -> dict[str, Any]: for _ in range(10): try: return await self._create_task_once(data) except IntegrityError: if data.get("id") is not None: raise raise RuntimeError("Failed to allocate a TaskCOP task id after multiple attempts") async def _create_task_once(self, data: dict[str, Any]) -> dict[str, Any]: task_id = str(data.get("id") or await self._next_task_id()) row = TaskCopTaskRow( id=task_id, task_name=str(data["task_name"]), overview=str(data.get("overview") or ""), task_level=str(data.get("task_level") or "1"), task_direction=str(data.get("task_direction") or "W方向"), task_status=int(data.get("task_status", 0)), task_type=int(data.get("task_type", 1)), task_superior_flag=int(data.get("task_superior_flag", 0)), parent_task_id=data.get("parent_task_id") or None, issuing_dept=str(data.get("issuing_dept") or ""), task_no=str(data.get("task_no") or ""), created_by=str(data.get("created_by") or "default"), task_start_date=data.get("task_start_date"), task_end_date=data.get("task_end_date"), ) async with self._sf() as session: session.add(row) await session.commit() await session.refresh(row) return _as_legacy(row) async def upsert_task(self, data: dict[str, Any]) -> tuple[dict[str, Any], bool]: """Import a caller-supplied task id, fully replacing an existing row.""" task_id = str(data["id"]) async with self._sf() as session: row = await session.get(TaskCopTaskRow, task_id) created = row is None if row is None: row = TaskCopTaskRow(id=task_id, task_name=str(data["task_name"])) session.add(row) row.task_name = str(data["task_name"]) row.overview = str(data.get("overview") or "") row.task_level = str(data.get("task_level") or "1") row.task_direction = str(data.get("task_direction") or "W方向") row.task_status = int(data.get("task_status", 0)) row.task_type = int(data.get("task_type", 1)) row.task_superior_flag = int(data.get("task_superior_flag", 0)) row.parent_task_id = data.get("parent_task_id") or None row.issuing_dept = str(data.get("issuing_dept") or "") row.task_no = str(data.get("task_no") or "") row.created_by = str(data.get("created_by") or "default") row.task_start_date = data.get("task_start_date") row.task_end_date = data.get("task_end_date") await session.commit() await session.refresh(row) return _as_legacy(row), created async def _next_task_id(self) -> str: """Return the next four-digit id, skipping historical UUID task rows.""" async with self._sf() as session: task_ids = (await session.scalars(select(TaskCopTaskRow.id))).all() numeric_ids = (int(task_id) for task_id in task_ids if task_id.isdecimal()) return f"{max(numeric_ids, default=0) + 1:04d}" async def list_tasks( self, *, created_by: str | None = None, keyword: str = "", task_direction: str | None = None, task_statuses: set[int] | None = None, start_time: datetime | None = None, end_time: datetime | None = None, page_num: int = 1, page_size: int = 10, ) -> tuple[int, list[dict[str, Any]]]: filters = [] if created_by is not None: filters.append(TaskCopTaskRow.created_by == created_by) if keyword.strip(): pattern = f"%{keyword.strip()}%" filters.append( or_( TaskCopTaskRow.task_name.ilike(pattern), TaskCopTaskRow.overview.ilike(pattern), TaskCopTaskRow.task_no.ilike(pattern), ) ) if task_direction: filters.append(TaskCopTaskRow.task_direction == task_direction) if task_statuses: filters.append(TaskCopTaskRow.task_status.in_(task_statuses)) if start_time: filters.append(TaskCopTaskRow.create_time >= start_time) if end_time: filters.append(TaskCopTaskRow.create_time <= end_time) offset = max(page_num - 1, 0) * page_size async with self._sf() as session: total = await session.scalar(select(func.count()).select_from(TaskCopTaskRow).where(*filters)) stmt = ( select(TaskCopTaskRow) .where(*filters) .order_by(TaskCopTaskRow.create_time.desc()) .offset(offset) .limit(page_size) ) result = await session.execute(stmt) return int(total or 0), [_as_legacy(row) for row in result.scalars()] async def get_task(self, task_id: str) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(TaskCopTaskRow, task_id) return _as_legacy(row) if row is not None else None async def delete_task(self, task_id: str) -> bool: async with self._sf() as session: row = await session.get(TaskCopTaskRow, task_id) if row is None: return False await session.delete(row) await session.commit() return True async def mark_analysis_completed(self, task_id: str) -> bool: async with self._sf() as session: row = await session.get(TaskCopTaskRow, task_id) if row is None: return False if row.task_status < TASK_STATUS_ANALYSIS_COMPLETED: row.task_status = TASK_STATUS_ANALYSIS_COMPLETED await session.commit() return True