183 lines
7.6 KiB
Python
183 lines
7.6 KiB
Python
"""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
|