deer-flow/backend/tests/schedule_fakes.py
rayhpeng c38d291505 refactor(schedule): standardize the module on the hexagonal architecture
Replaces the pre-hexagonal scheduled-task implementation with a slice
built to the layering spec: a pure domain (two aggregates, two state
machines, the policy value object), output ports it declares itself,
SQL/launcher/thread adapters implementing them under `app/adapters/`,
and a composition root that is the one place any of them is
instantiated.

The old implementation mixed all of that into `app/scheduler/service.py`
and a router that reached straight into repositories, so the rules that
matter -- overlap policy, lease handling, which write owns which
timestamp -- were only reachable through a live database. They are now
unit-assertable on in-memory fakes, with the contract suite running each
port against both the fake and real sqlite, and the concurrency
invariants pinned by dedicated race tests.

Two bugs the old shape hid are fixed on the way: a completion hook that
replayed a stale snapshot and rolled back the launch write, and a
corrupt stored row surfacing to the client as a 4xx.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-07-31 11:16:46 +08:00

294 lines
11 KiB
Python

"""Shared zero-IO doubles for the schedule context.
Lives in its own module rather than inside a test file so the port contract
suite and the service tests can both import it as an ordinary module, without
relying on the tests directory happening to be on ``sys.path``.
**Concurrency is explicitly out of scope.** These doubles model the
single-threaded semantics of each port -- which rows `claim_due` selects, that
`add` refuses a second active run -- and provide no atomicity whatsoever. A
green run here says nothing about two dispatchers racing; that is covered
against a real database in ``test_schedule_dispatch_race.py``. Do not
read a passing contract suite as licence to run more than one scheduler.
"""
from __future__ import annotations
from dataclasses import dataclass, field, replace
from datetime import datetime, timedelta
from deerflow.domain.schedule.exceptions import ActiveRunConflictError
from deerflow.domain.schedule.model import (
ACTIVE_RUN_STATUSES,
TERMINAL_RUN_STATUSES,
TERMINAL_TASK_STATUSES,
RunStatus,
ScheduledRun,
ScheduledTask,
ScheduleType,
TaskStatus,
)
from deerflow.domain.schedule.ports import LaunchedRun
@dataclass
class _TaskRow:
"""One stored task: the aggregate plus the columns the domain has no
business knowing about.
The lease lives here rather than on ``ScheduledTask`` for the same reason
it lives in the table and not in the domain -- it is a storage-level
concurrency control, not a business fact about the task.
"""
task: ScheduledTask
lease_owner: str | None = None
lease_expires_at: datetime | None = None
class InMemoryScheduledTaskRepository:
"""ScheduledTaskRepository double backed by a dict."""
def __init__(self) -> None:
self._rows: dict[str, _TaskRow] = {}
# -- inspection helpers (tests only, not part of the port) -----------
def lease_of(self, task_id: str) -> tuple[str | None, datetime | None]:
row = self._rows[task_id]
return row.lease_owner, row.lease_expires_at
def seed(self, task: ScheduledTask, *, lease_owner: str | None = None, lease_expires_at: datetime | None = None) -> ScheduledTask:
"""Install a task directly, optionally already claimed."""
self._rows[task.task_id] = _TaskRow(task=task, lease_owner=lease_owner, lease_expires_at=lease_expires_at)
return task
# -- port ------------------------------------------------------------
async def add(self, task: ScheduledTask) -> ScheduledTask:
self._rows[task.task_id] = _TaskRow(task=task)
return task
async def get(self, task_id: str, *, user_id: str) -> ScheduledTask | None:
row = self._rows.get(task_id)
if row is None or row.task.user_id != user_id:
return None
return row.task
async def list_by_user(self, user_id: str) -> list[ScheduledTask]:
rows = [row.task for row in self._rows.values() if row.task.user_id == user_id]
return sorted(rows, key=lambda t: (t.created_at, t.task_id), reverse=True)
async def list_by_user_and_thread(self, user_id: str, thread_id: str) -> list[ScheduledTask]:
tasks = await self.list_by_user(user_id)
return [task for task in tasks if task.thread_id == thread_id]
async def save(self, task: ScheduledTask) -> ScheduledTask | None:
row = self._rows.get(task.task_id)
if row is None or row.task.user_id != task.user_id:
return None
row.task = task
return task
async def delete(self, task_id: str, *, user_id: str) -> bool:
row = self._rows.get(task_id)
if row is None or row.task.user_id != user_id:
return False
del self._rows[task_id]
return True
async def claim_due(self, *, now: datetime, lease_seconds: int, limit: int) -> list[ScheduledTask]:
def claimable(row: _TaskRow) -> bool:
task = row.task
if task.next_run_at is None or task.next_run_at > now:
return False
expired = row.lease_expires_at is not None and row.lease_expires_at < now
if task.status is TaskStatus.ENABLED:
return row.lease_expires_at is None or expired
# Stuck mid-dispatch: the claimer died between claim and launch.
return task.status is TaskStatus.RUNNING and expired
due = sorted(
(row for row in self._rows.values() if claimable(row)),
key=lambda row: (row.task.next_run_at, row.task.task_id),
)[:limit]
claimed = []
for row in due:
# Stands in for whatever identity a real adapter records.
row.lease_owner = "fake-worker"
row.lease_expires_at = now + timedelta(seconds=lease_seconds)
row.task = replace(row.task, status=TaskStatus.RUNNING)
claimed.append(row.task)
return claimed
async def record_launch(
self,
task_id: str,
*,
status: TaskStatus,
next_run_at: datetime | None,
last_run_at: datetime | None,
last_run_id: str | None,
last_thread_id: str | None,
last_error: str | None,
increment_run_count: bool,
protect_terminal: bool = False,
) -> None:
row = self._rows.get(task_id)
if row is None:
return
task = row.task
if not (protect_terminal and task.status in TERMINAL_TASK_STATUSES):
task = replace(task, status=status, last_error=last_error)
task = replace(
task,
next_run_at=next_run_at,
last_run_at=last_run_at,
last_run_id=last_run_id,
last_thread_id=last_thread_id,
)
if increment_run_count:
task = replace(task, run_count=task.run_count + 1)
row.task = task
row.lease_owner = None
row.lease_expires_at = None
async def record_completion(
self,
task_id: str,
*,
user_id: str,
status: TaskStatus | None,
error: str | None,
) -> None:
row = self._rows.get(task_id)
if row is None or row.task.user_id != user_id:
return
# Only the verdict; every scheduling field belongs to record_launch.
row.task = replace(row.task, last_error=error)
if status is not None:
row.task = replace(row.task, status=status)
async def cancel_stuck_once_tasks(self, *, error: str) -> int:
cancelled = 0
for row in self._rows.values():
stuck = row.task.status is TaskStatus.RUNNING and row.task.schedule.schedule_type is ScheduleType.ONCE and row.lease_expires_at is None
if stuck:
row.task = replace(row.task, status=TaskStatus.CANCELLED, last_error=error)
cancelled += 1
return cancelled
class InMemoryScheduledRunRepository:
"""ScheduledRunRepository double that also models the active-slot rule.
``add`` refusing a second active record for one task is not an
implementation detail to be skipped here: the service collapses that
rejection to the same outcome as its own fast path, and a double that never
raises would let that collapse go untested.
"""
def __init__(self) -> None:
self._rows: dict[str, ScheduledRun] = {}
def all_runs(self) -> list[ScheduledRun]:
"""Inspection helper (tests only, not part of the port)."""
return list(self._rows.values())
async def add(self, run: ScheduledRun) -> ScheduledRun:
if run.is_active and any(other.task_id == run.task_id and other.is_active for other in self._rows.values()):
raise ActiveRunConflictError(f"scheduled task {run.task_id!r} already has an active run")
self._rows[run.record_id] = run
return run
async def list_by_task(self, task_id: str, *, limit: int = 50, offset: int = 0) -> list[ScheduledRun]:
rows = [run for run in self._rows.values() if run.task_id == task_id]
rows.sort(key=lambda run: (run.created_at, run.record_id), reverse=True)
return rows[offset : offset + limit]
async def count_active(self) -> int:
return sum(1 for run in self._rows.values() if run.is_active)
async def has_active(self, task_id: str) -> bool:
return any(run.task_id == task_id and run.is_active for run in self._rows.values())
async def update_status(
self,
record_id: str,
*,
status: RunStatus,
run_id: str | None = None,
error: str | None = None,
started_at: datetime | None = None,
finished_at: datetime | None = None,
protect_terminal: bool = False,
) -> None:
run = self._rows.get(record_id)
if run is None:
return
if protect_terminal and run.status in TERMINAL_RUN_STATUSES:
# Keep the terminal verdict; only backfill what it could not know.
if run.run_id is None and run_id is not None:
run = replace(run, run_id=run_id)
if run.started_at is None and started_at is not None:
run = replace(run, started_at=started_at)
self._rows[record_id] = run
return
run = replace(run, status=status, run_id=run_id, error=error)
if started_at is not None:
run = replace(run, started_at=started_at)
if finished_at is not None:
run = replace(run, finished_at=finished_at)
self._rows[record_id] = run
async def mark_stale_active(self, *, error: str) -> int:
stale = [run for run in self._rows.values() if run.status in ACTIVE_RUN_STATUSES]
for run in stale:
self._rows[run.record_id] = replace(run, status=RunStatus.INTERRUPTED, error=error)
return len(stale)
@dataclass
class FakeRunLauncher:
"""RunLauncher double.
`fail_with` makes every launch raise instead -- set it to a ThreadBusyError
or a LaunchFailedError to drive the two branches the domain distinguishes.
"""
fail_with: Exception | None = None
calls: list[dict] = field(default_factory=list)
async def launch(
self,
*,
thread_id: str,
assistant_id: str | None,
prompt: str,
owner_user_id: str | None,
metadata: dict[str, str],
) -> LaunchedRun:
self.calls.append(
{
"thread_id": thread_id,
"assistant_id": assistant_id,
"prompt": prompt,
"owner_user_id": owner_user_id,
"metadata": metadata,
}
)
if self.fail_with is not None:
raise self.fail_with
return LaunchedRun(run_id=f"run-{len(self.calls)}", thread_id=thread_id)
class FakeThreadLookup:
"""ThreadLookup double backed by a thread_id -> owner mapping."""
def __init__(self, threads: dict[str, str] | None = None) -> None:
self._threads = dict(threads or {})
async def exists_for_user(self, thread_id: str, user_id: str) -> bool:
return self._threads.get(thread_id) == user_id