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

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]