380 lines
16 KiB
Python
380 lines
16 KiB
Python
"""SQLAlchemy repository for scheduled tasks."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from datetime import UTC, datetime
|
|
from typing import Any
|
|
|
|
import uuid
|
|
|
|
from sqlalchemy import delete, select, update
|
|
from sqlalchemy.exc import IntegrityError
|
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
|
|
|
from deerflow.persistence.scheduled_tasks.base import ScheduledTaskStore
|
|
from deerflow.persistence.scheduled_tasks.model import (
|
|
ScheduledTaskDeliveryProfileRow,
|
|
ScheduledTaskRow,
|
|
ScheduledTaskRunRow,
|
|
ScheduledTaskSubscriptionRow,
|
|
SchedulerThreadRow,
|
|
)
|
|
|
|
|
|
def _to_dict(row) -> dict[str, Any]:
|
|
data = row.to_dict()
|
|
for key in ("created_at", "updated_at", "next_run_at", "last_run_at", "scheduled_for", "started_at", "finished_at", "published_at"):
|
|
value = data.get(key)
|
|
if isinstance(value, datetime):
|
|
data[key] = value.isoformat()
|
|
if "execution_context_json" in data:
|
|
data["execution_context"] = data.pop("execution_context_json") or {}
|
|
if "result_json" in data:
|
|
data["result"] = data.pop("result_json") or None
|
|
return data
|
|
|
|
|
|
class ScheduledTaskRepository(ScheduledTaskStore):
|
|
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
|
|
self._sf = session_factory
|
|
|
|
async def get_or_create_scheduler_thread(self, user_id: str, thread_id_factory) -> dict:
|
|
async with self._sf() as session:
|
|
row = await session.get(SchedulerThreadRow, user_id)
|
|
if row is not None:
|
|
return _to_dict(row)
|
|
row = SchedulerThreadRow(user_id=user_id, thread_id=thread_id_factory())
|
|
session.add(row)
|
|
try:
|
|
await session.commit()
|
|
except IntegrityError:
|
|
await session.rollback()
|
|
row = await session.get(SchedulerThreadRow, user_id)
|
|
if row is None:
|
|
raise
|
|
await session.refresh(row)
|
|
return _to_dict(row)
|
|
|
|
async def list_tasks(self, user_id: str) -> list[dict[str, Any]]:
|
|
async with self._sf() as session:
|
|
result = await session.execute(select(ScheduledTaskRow).where(ScheduledTaskRow.user_id == user_id).order_by(ScheduledTaskRow.created_at.desc()))
|
|
return [_to_dict(row) for row in result.scalars()]
|
|
|
|
async def get_task(self, task_id: str, user_id: str) -> dict[str, Any] | None:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskRow, task_id)
|
|
if row is None or row.user_id != user_id:
|
|
return None
|
|
return _to_dict(row)
|
|
|
|
async def create_task(self, data: dict[str, Any]) -> dict[str, Any]:
|
|
row = ScheduledTaskRow(
|
|
task_id=data["task_id"],
|
|
user_id=data["user_id"],
|
|
scheduler_thread_id=data["scheduler_thread_id"],
|
|
source_thread_id=data.get("source_thread_id"),
|
|
source_agent_name=data.get("source_agent_name"),
|
|
execution_agent_name=data.get("execution_agent_name"),
|
|
name=data["name"],
|
|
description=data.get("description"),
|
|
prompt=data["prompt"],
|
|
schedule_text=data["schedule_text"],
|
|
cron_expr=data["cron_expr"],
|
|
timezone=data["timezone"],
|
|
enabled=data.get("enabled", True),
|
|
next_run_at=data.get("next_run_at"),
|
|
last_status=data.get("last_status"),
|
|
execution_context_json=data.get("execution_context") or {},
|
|
)
|
|
async with self._sf() as session:
|
|
session.add(row)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return _to_dict(row)
|
|
|
|
async def update_task(self, task_id: str, user_id: str, data: dict[str, Any]) -> dict[str, Any] | None:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskRow, task_id)
|
|
if row is None or row.user_id != user_id:
|
|
return None
|
|
for key, value in data.items():
|
|
attr = "execution_context_json" if key == "execution_context" else key
|
|
if hasattr(row, attr):
|
|
setattr(row, attr, value)
|
|
row.updated_at = datetime.now(UTC)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return _to_dict(row)
|
|
|
|
async def delete_task(self, task_id: str, user_id: str) -> bool:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskRow, task_id)
|
|
if row is None or row.user_id != user_id:
|
|
return False
|
|
await session.delete(row)
|
|
await session.commit()
|
|
return True
|
|
|
|
async def due_tasks(self, now: datetime, *, limit: int = 20) -> list[dict[str, Any]]:
|
|
async with self._sf() as session:
|
|
result = await session.execute(
|
|
select(ScheduledTaskRow)
|
|
.where(ScheduledTaskRow.enabled.is_(True), ScheduledTaskRow.next_run_at.is_not(None), ScheduledTaskRow.next_run_at <= now)
|
|
.order_by(ScheduledTaskRow.next_run_at.asc())
|
|
.limit(limit)
|
|
)
|
|
return [_to_dict(row) for row in result.scalars()]
|
|
|
|
# --- Admin (cross-user) ---
|
|
|
|
async def list_all_tasks(self) -> list[dict[str, Any]]:
|
|
async with self._sf() as session:
|
|
result = await session.execute(select(ScheduledTaskRow).order_by(ScheduledTaskRow.created_at.desc()))
|
|
return [_to_dict(row) for row in result.scalars()]
|
|
|
|
async def update_task_any(self, task_id: str, data: dict[str, Any]) -> dict[str, Any] | None:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskRow, task_id)
|
|
if row is None:
|
|
return None
|
|
for key, value in data.items():
|
|
attr = "execution_context_json" if key == "execution_context" else key
|
|
if hasattr(row, attr):
|
|
setattr(row, attr, value)
|
|
row.updated_at = datetime.now(UTC)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return _to_dict(row)
|
|
|
|
async def delete_task_any(self, task_id: str) -> bool:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskRow, task_id)
|
|
if row is None:
|
|
return False
|
|
await session.delete(row)
|
|
await session.commit()
|
|
return True
|
|
|
|
async def list_runs_any(self, task_id: str, *, limit: int = 50) -> list[dict[str, Any]]:
|
|
async with self._sf() as session:
|
|
result = await session.execute(
|
|
select(ScheduledTaskRunRow)
|
|
.where(ScheduledTaskRunRow.task_id == task_id)
|
|
.order_by(ScheduledTaskRunRow.scheduled_for.desc())
|
|
.limit(limit)
|
|
)
|
|
return [_to_dict(row) for row in result.scalars()]
|
|
|
|
async def list_stuck_runs(self, cutoff: datetime, *, limit: int = 50) -> list[dict[str, Any]]:
|
|
async with self._sf() as session:
|
|
result = await session.execute(
|
|
select(ScheduledTaskRunRow)
|
|
.where(
|
|
ScheduledTaskRunRow.status.in_(("running", "pending")),
|
|
ScheduledTaskRunRow.started_at.is_not(None),
|
|
ScheduledTaskRunRow.started_at <= cutoff,
|
|
)
|
|
.order_by(ScheduledTaskRunRow.started_at.asc())
|
|
.limit(limit)
|
|
)
|
|
return [_to_dict(row) for row in result.scalars()]
|
|
|
|
async def create_run_once(self, data: dict[str, Any]) -> dict[str, Any] | None:
|
|
row = ScheduledTaskRunRow(
|
|
id=data["id"],
|
|
task_id=data["task_id"],
|
|
user_id=data["user_id"],
|
|
scheduler_thread_id=data["scheduler_thread_id"],
|
|
scheduled_for=data["scheduled_for"],
|
|
status=data.get("status", "pending"),
|
|
started_at=data.get("started_at"),
|
|
result_json=data.get("result") or None,
|
|
)
|
|
async with self._sf() as session:
|
|
session.add(row)
|
|
try:
|
|
await session.commit()
|
|
except IntegrityError:
|
|
await session.rollback()
|
|
return None
|
|
await session.refresh(row)
|
|
return _to_dict(row)
|
|
|
|
async def update_run(self, run_id: str, data: dict[str, Any]) -> None:
|
|
async with self._sf() as session:
|
|
values = dict(data)
|
|
if "result" in values:
|
|
values["result_json"] = values.pop("result")
|
|
values["updated_at"] = datetime.now(UTC)
|
|
await session.execute(update(ScheduledTaskRunRow).where(ScheduledTaskRunRow.id == run_id).values(**values))
|
|
await session.commit()
|
|
|
|
async def list_runs(self, task_id: str, user_id: str, *, limit: int = 50) -> list[dict[str, Any]]:
|
|
async with self._sf() as session:
|
|
result = await session.execute(
|
|
select(ScheduledTaskRunRow)
|
|
.where(ScheduledTaskRunRow.task_id == task_id, ScheduledTaskRunRow.user_id == user_id)
|
|
.order_by(ScheduledTaskRunRow.scheduled_for.desc())
|
|
.limit(limit)
|
|
)
|
|
return [_to_dict(row) for row in result.scalars()]
|
|
|
|
async def get_run(self, run_id: str, user_id: str) -> dict[str, Any] | None:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskRunRow, run_id)
|
|
if row is None or row.user_id != user_id:
|
|
return None
|
|
return _to_dict(row)
|
|
|
|
async def delete_run(self, run_id: str, user_id: str) -> bool:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskRunRow, run_id)
|
|
if row is None or row.user_id != user_id:
|
|
return False
|
|
await session.delete(row)
|
|
await session.commit()
|
|
return True
|
|
|
|
async def delete_run_any(self, run_id: str) -> bool:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskRunRow, run_id)
|
|
if row is None:
|
|
return False
|
|
await session.delete(row)
|
|
await session.commit()
|
|
return True
|
|
|
|
async def get_delivery_profile(self, user_id: str) -> dict[str, Any]:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskDeliveryProfileRow, user_id)
|
|
if row is None:
|
|
return {"user_id": user_id, "element_user_id": None}
|
|
return _to_dict(row)
|
|
|
|
async def update_delivery_profile(self, user_id: str, data: dict[str, Any]) -> dict[str, Any]:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskDeliveryProfileRow, user_id)
|
|
if row is None:
|
|
row = ScheduledTaskDeliveryProfileRow(user_id=user_id)
|
|
session.add(row)
|
|
if "element_user_id" in data:
|
|
element_user_id = data["element_user_id"]
|
|
row.element_user_id = element_user_id.strip() if isinstance(element_user_id, str) and element_user_id.strip() else None
|
|
if "element_room_id" in data:
|
|
row.element_room_id = data["element_room_id"] or None
|
|
row.updated_at = datetime.now(UTC)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return _to_dict(row)
|
|
|
|
# --- Publish / subscribe ---
|
|
|
|
async def get_task_any(self, task_id: str) -> dict[str, Any] | None:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskRow, task_id)
|
|
return _to_dict(row) if row is not None else None
|
|
|
|
async def get_run_any(self, run_id: str) -> dict[str, Any] | None:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskRunRow, run_id)
|
|
return _to_dict(row) if row is not None else None
|
|
|
|
async def set_task_published(self, task_id: str, owner_user_id: str, *, published: bool) -> dict[str, Any] | None:
|
|
async with self._sf() as session:
|
|
row = await session.get(ScheduledTaskRow, task_id)
|
|
if row is None or row.user_id != owner_user_id:
|
|
return None
|
|
row.published = bool(published)
|
|
row.published_at = datetime.now(UTC) if published else None
|
|
row.updated_at = datetime.now(UTC)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return _to_dict(row)
|
|
|
|
async def list_published_tasks(self, *, limit: int = 50, offset: int = 0) -> list[dict[str, Any]]:
|
|
async with self._sf() as session:
|
|
result = await session.execute(
|
|
select(ScheduledTaskRow)
|
|
.where(ScheduledTaskRow.published.is_(True))
|
|
# ``NULLS LAST`` is unsupported by MySQL — sort by an
|
|
# ``IS NULL`` flag instead (False<True ⇒ non-null rows first),
|
|
# which is portable across mysql/sqlite/postgres.
|
|
.order_by(
|
|
ScheduledTaskRow.published_at.is_(None),
|
|
ScheduledTaskRow.published_at.desc(),
|
|
ScheduledTaskRow.created_at.desc(),
|
|
)
|
|
.offset(max(0, offset))
|
|
.limit(max(1, limit))
|
|
)
|
|
return [_to_dict(row) for row in result.scalars()]
|
|
|
|
async def list_subscribed_tasks(self, user_id: str) -> list[dict[str, Any]]:
|
|
async with self._sf() as session:
|
|
result = await session.execute(
|
|
select(ScheduledTaskRow, ScheduledTaskSubscriptionRow)
|
|
.join(ScheduledTaskSubscriptionRow, ScheduledTaskSubscriptionRow.task_id == ScheduledTaskRow.task_id)
|
|
.where(ScheduledTaskSubscriptionRow.user_id == user_id)
|
|
.order_by(ScheduledTaskSubscriptionRow.created_at.desc())
|
|
)
|
|
tasks: list[dict[str, Any]] = []
|
|
for task_row, sub_row in result.all():
|
|
task = _to_dict(task_row)
|
|
task["_subscription"] = _to_dict(sub_row)
|
|
tasks.append(task)
|
|
return tasks
|
|
|
|
async def get_subscription(self, task_id: str, user_id: str) -> dict[str, Any] | None:
|
|
async with self._sf() as session:
|
|
result = await session.execute(
|
|
select(ScheduledTaskSubscriptionRow)
|
|
.where(ScheduledTaskSubscriptionRow.task_id == task_id, ScheduledTaskSubscriptionRow.user_id == user_id)
|
|
.limit(1)
|
|
)
|
|
row = result.scalar_one_or_none()
|
|
return _to_dict(row) if row is not None else None
|
|
|
|
async def upsert_subscription(self, task_id: str, user_id: str, *, notify_element: bool) -> dict[str, Any]:
|
|
async with self._sf() as session:
|
|
result = await session.execute(
|
|
select(ScheduledTaskSubscriptionRow)
|
|
.where(ScheduledTaskSubscriptionRow.task_id == task_id, ScheduledTaskSubscriptionRow.user_id == user_id)
|
|
.limit(1)
|
|
)
|
|
row = result.scalar_one_or_none()
|
|
if row is None:
|
|
row = ScheduledTaskSubscriptionRow(
|
|
id=str(uuid.uuid4()),
|
|
task_id=task_id,
|
|
user_id=user_id,
|
|
notify_element=bool(notify_element),
|
|
)
|
|
session.add(row)
|
|
else:
|
|
row.notify_element = bool(notify_element)
|
|
row.updated_at = datetime.now(UTC)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return _to_dict(row)
|
|
|
|
async def delete_subscription(self, task_id: str, user_id: str) -> bool:
|
|
async with self._sf() as session:
|
|
result = await session.execute(
|
|
delete(ScheduledTaskSubscriptionRow).where(
|
|
ScheduledTaskSubscriptionRow.task_id == task_id,
|
|
ScheduledTaskSubscriptionRow.user_id == user_id,
|
|
)
|
|
)
|
|
await session.commit()
|
|
return (result.rowcount or 0) > 0
|
|
|
|
async def list_notify_subscribers(self, task_id: str) -> list[dict[str, Any]]:
|
|
async with self._sf() as session:
|
|
result = await session.execute(
|
|
select(ScheduledTaskSubscriptionRow).where(
|
|
ScheduledTaskSubscriptionRow.task_id == task_id,
|
|
ScheduledTaskSubscriptionRow.notify_element.is_(True),
|
|
)
|
|
)
|
|
return [_to_dict(row) for row in result.scalars()]
|