deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/persistence/taskcop_tasks/sql.py
2026-09-07 18:24:55 +08:00

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