321 lines
15 KiB
Python
321 lines
15 KiB
Python
"""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]
|