"""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 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()]