aaxaxax's picture
initial commit
2aa8b3a
Raw
History Blame
21.4 kB
"""Bounded, thread-safe job admission and FIFO scheduling.
The queue deliberately does not start or stop threads. It owns state and
cooperative cancellation tokens; :mod:`distinct_agent.worker` owns execution.
"""
from __future__ import annotations
import math
import threading
import time
from collections import deque
from collections.abc import Callable, Mapping
from dataclasses import dataclass, field
from typing import Any
from distinct_protocol import (
MAX_LIVE_STEP_TEXT,
MAX_LIVE_STEPS_PER_JOB,
PHASE_DONE,
PHASE_QUEUED_ON_WORKER,
PHASE_WORKING,
STEP_PHASE,
AgentSnapshot,
AgentStatus,
JobSpec,
JobStatus,
normalise_phase,
)
# The single cadence: the agent's sync request carries the snapshot up and
# brings claimable work down, so the poll *is* the heartbeat. The server
# advertises this value at pairing and measures liveness against it.
DEFAULT_POLL_INTERVAL_SECONDS = 5.0
# Deprecated alias; there was never a separate snapshot cadence in practice.
DEFAULT_SNAPSHOT_INTERVAL_SECONDS = DEFAULT_POLL_INTERVAL_SECONDS
TERMINAL_STATUSES = frozenset(
{
JobStatus.COMPLETED,
JobStatus.FAILED,
JobStatus.CANCELLED,
JobStatus.EXPIRED,
}
)
@dataclass(frozen=True)
class OfferDecision:
"""The local admission decision returned for a server job offer."""
job_id: str
accepted: bool
status: JobStatus
queue_position: int | None
reason: str | None = None
def to_dict(self) -> dict:
return {
"job_id": self.job_id,
"accepted": self.accepted,
"status": self.status.value,
"queue_position": self.queue_position,
"reason": self.reason,
}
@dataclass(frozen=True)
class CancellationDecision:
"""Outcome of a cooperative cancellation request."""
job_id: str
found: bool
requested: bool
immediate: bool
status: JobStatus | None
reason: str | None = None
@dataclass(frozen=True)
class JobView:
"""A safe copy of one queue record for UI or diagnostics."""
job_id: str
status: JobStatus
queue_position: int | None
progress: float
progress_message: str
cancel_requested: bool
terminal_reason: str | None
phase: str = PHASE_QUEUED_ON_WORKER
@dataclass
class _Entry:
# ``job_id`` is kept separately so the entry stays usable after ``job`` has
# been dropped. A terminal entry retains only scheduling metadata: the
# JobSpec, and therefore the user's prompt, is released immediately.
job_id: str
job: JobSpec | None
status: JobStatus = JobStatus.QUEUED
progress: float = 0.0
progress_message: str = "Queued"
#: Machine-readable phase from the protocol vocabulary. Distinct from
#: ``progress_message``, which is prose: the phase drives the interface's
#: own copy and its spinner, the message is detail beside it.
phase: str = PHASE_QUEUED_ON_WORKER
#: Bounded ring of live steps for this job. Oldest is dropped rather than
#: refused, because a run that stops streaming when it gets interesting is
#: worse than one that forgets its opening moves.
live_steps: list[dict] = field(default_factory=list)
cancel_event: threading.Event = field(default_factory=threading.Event)
offered_monotonic: float = 0.0
started_monotonic: float | None = None
started_wall: float | None = None
finished_monotonic: float | None = None
terminal_reason: str | None = None
class InMemoryJobQueue:
"""A bounded FIFO with a configurable number of active jobs.
``capacity`` is the total number of non-terminal jobs, including active
work. This makes admission deterministic and prevents a fast producer
from exceeding the advertised agent capacity.
"""
def __init__(
self,
*,
capacity: int = 4,
max_active: int = 1,
agent_id: str | None = None,
history_limit: int = 256,
default_job_seconds: float = 0.0,
monotonic: Callable[[], float] = time.monotonic,
wall_clock: Callable[[], float] = time.time,
) -> None:
if isinstance(capacity, bool) or not 1 <= capacity <= 64:
raise ValueError("capacity must be an integer from 1 to 64")
if isinstance(max_active, bool) or not 1 <= max_active <= capacity:
raise ValueError("max_active must be between 1 and capacity")
if isinstance(history_limit, bool) or history_limit < 1:
raise ValueError("history_limit must be a positive integer")
if not math.isfinite(default_job_seconds) or default_job_seconds < 0:
raise ValueError("default_job_seconds must be finite and non-negative")
self.capacity = capacity
self.max_active = max_active
self.agent_id = agent_id
self.history_limit = history_limit
self.default_job_seconds = float(default_job_seconds)
self._monotonic = monotonic
self._wall_clock = wall_clock
self._waiting: deque[str] = deque()
self._active: dict[str, _Entry] = {}
self._entries: dict[str, _Entry] = {}
self._terminal_order: deque[str] = deque()
self._durations: deque[float] = deque(maxlen=32)
self._closed = False
self._condition = threading.Condition(threading.RLock())
@property
def closed(self) -> bool:
with self._condition:
return self._closed
def __len__(self) -> int:
with self._condition:
return len(self._waiting) + len(self._active)
def offer(self, job: JobSpec) -> OfferDecision:
"""Admit ``job`` or return an explicit, non-throwing rejection."""
if not isinstance(job, JobSpec):
raise TypeError("job must be a distinct_protocol.JobSpec")
with self._condition:
if self._closed:
return self._reject(job.id, "queue_closed")
if self.agent_id and job.target_agent_id != self.agent_id:
return self._reject(job.id, "wrong_agent")
if job.id in self._entries:
return self._reject(job.id, "duplicate_job")
if len(self._waiting) + len(self._active) >= self.capacity:
return self._reject(job.id, "queue_full")
entry = _Entry(job_id=job.id, job=job, offered_monotonic=self._monotonic())
self._entries[job.id] = entry
self._waiting.append(job.id)
position = len(self._waiting)
self._condition.notify_all()
return OfferDecision(
job_id=job.id,
accepted=True,
status=JobStatus.QUEUED,
queue_position=position,
)
def _reject(self, job_id: str, reason: str) -> OfferDecision:
return OfferDecision(
job_id=job_id,
accepted=False,
status=JobStatus.OFFERED,
queue_position=None,
reason=reason,
)
def claim_next(self) -> JobSpec | None:
"""Move the oldest waiting job to running, if a slot is available."""
with self._condition:
if len(self._active) >= self.max_active or not self._waiting:
return None
job_id = self._waiting.popleft()
entry = self._entries[job_id]
entry.status = JobStatus.RUNNING
entry.started_monotonic = self._monotonic()
entry.started_wall = time.time()
entry.progress_message = "Starting"
entry.phase = PHASE_WORKING
self._active[job_id] = entry
self._condition.notify_all()
return entry.job
def wait_and_claim(self, timeout: float | None = None) -> JobSpec | None:
"""Wait until work and an execution slot are available."""
if timeout is not None and timeout < 0:
raise ValueError("timeout cannot be negative")
deadline = None if timeout is None else self._monotonic() + timeout
with self._condition:
while True:
job = self.claim_next()
if job is not None:
return job
if self._closed and not self._waiting:
return None
remaining = None if deadline is None else deadline - self._monotonic()
if remaining is not None and remaining <= 0:
return None
self._condition.wait(remaining)
def update_progress(
self,
job_id: str,
fraction: float,
message: str = "",
*,
phase: str | None = None,
step: Mapping[str, Any] | None = None,
) -> JobView:
if isinstance(fraction, bool) or not isinstance(fraction, int | float):
raise ValueError("progress must be numeric")
fraction = float(fraction)
if not math.isfinite(fraction) or not 0.0 <= fraction <= 1.0:
raise ValueError("progress must be finite and between zero and one")
if len(message) > 512:
raise ValueError("progress message exceeds 512 characters")
with self._condition:
entry = self._require_entry(job_id)
if entry.status not in {JobStatus.QUEUED, JobStatus.RUNNING}:
raise ValueError("cannot update progress for a terminal job")
# Progress is monotonic. A stale callback cannot move a UI bar back.
entry.progress = max(entry.progress, fraction)
if message:
entry.progress_message = message
if phase is not None:
# Normalised here rather than trusted: the harness is local code
# today, but the phase reaches a user's screen and the boundary
# belongs where the value enters the shared structure.
entry.phase = normalise_phase(phase)
if step is not None:
self._append_step_unlocked(entry, step)
return self._view_unlocked(job_id)
def record_step(self, job_id: str, step: Mapping[str, Any]) -> None:
"""Append one live step without touching the progress fraction."""
with self._condition:
entry = self._entries.get(job_id)
if entry is None or entry.status not in {JobStatus.QUEUED, JobStatus.RUNNING}:
return
self._append_step_unlocked(entry, step)
def _append_step_unlocked(self, entry: _Entry, step: Mapping[str, Any]) -> None:
record = dict(step)
record.setdefault("phase", entry.phase)
record.setdefault("kind", STEP_PHASE)
record.setdefault("at", time.time())
record["text"] = str(record.get("text") or "")[:MAX_LIVE_STEP_TEXT]
entry.live_steps.append(record)
# Keep the tail, not the head: what the run is doing now matters more
# than how it opened, and the cap is a wire limit rather than a policy.
if len(entry.live_steps) > MAX_LIVE_STEPS_PER_JOB:
del entry.live_steps[: len(entry.live_steps) - MAX_LIVE_STEPS_PER_JOB]
def cancel(self, job_id: str, reason: str = "cancelled") -> CancellationDecision:
"""Cancel queued work immediately or signal active work cooperatively."""
with self._condition:
entry = self._entries.get(job_id)
if entry is None:
return CancellationDecision(job_id, False, False, False, None, "unknown_job")
if entry.status in TERMINAL_STATUSES:
return CancellationDecision(
job_id, True, False, True, entry.status, "already_terminal"
)
entry.cancel_event.set()
entry.progress_message = "Cancellation requested"
if entry.status == JobStatus.QUEUED:
self._waiting.remove(job_id)
self._finish_unlocked(entry, JobStatus.CANCELLED, reason)
self._condition.notify_all()
return CancellationDecision(
job_id, True, True, True, JobStatus.CANCELLED
)
return CancellationDecision(job_id, True, True, False, entry.status)
def cancellation_event(self, job_id: str) -> threading.Event:
with self._condition:
return self._require_entry(job_id).cancel_event
def complete(self, job_id: str) -> JobView:
with self._condition:
entry = self._require_active(job_id)
entry.progress = 1.0
self._finish_unlocked(entry, JobStatus.COMPLETED, None)
self._condition.notify_all()
return self._view_unlocked(job_id)
def fail(self, job_id: str, reason: str) -> JobView:
with self._condition:
entry = self._require_active(job_id)
self._finish_unlocked(entry, JobStatus.FAILED, reason[:4096])
self._condition.notify_all()
return self._view_unlocked(job_id)
def finish_cancelled(self, job_id: str, reason: str = "cancelled") -> JobView:
with self._condition:
entry = self._require_active(job_id)
self._finish_unlocked(entry, JobStatus.CANCELLED, reason)
self._condition.notify_all()
return self._view_unlocked(job_id)
def expire(self, job_id: str, reason: str = "lease_expired") -> JobView:
with self._condition:
entry = self._require_entry(job_id)
if entry.status == JobStatus.QUEUED:
self._waiting.remove(job_id)
elif entry.status == JobStatus.RUNNING:
entry.cancel_event.set()
entry.progress_message = "Expiration requested"
entry.terminal_reason = reason
# A Python thread cannot be killed safely. Keep the active
# slot occupied until the runner cooperatively returns.
return self._view_unlocked(job_id)
else:
return self._view_unlocked(job_id)
self._finish_unlocked(entry, JobStatus.EXPIRED, reason)
self._condition.notify_all()
return self._view_unlocked(job_id)
def view(self, job_id: str) -> JobView | None:
with self._condition:
if job_id not in self._entries:
return None
return self._view_unlocked(job_id)
def queued_jobs(self) -> tuple[JobSpec, ...]:
with self._condition:
return tuple(
self._entries[job_id].job
for job_id in self._waiting
if self._entries[job_id].job is not None
)
def active_jobs(self) -> tuple[JobSpec, ...]:
with self._condition:
return tuple(entry.job for entry in self._active.values() if entry.job is not None)
def retains_prompt(self, job_id: str) -> bool:
"""True while this queue still holds the prompt text for ``job_id``.
Exposed so the erasure guarantee can be asserted from outside rather
than taken on trust.
"""
with self._condition:
entry = self._entries.get(job_id)
return entry is not None and entry.job is not None
def snapshot(self, *, agent_id: str | None = None, energy_available: bool = False) -> AgentSnapshot:
"""Build the serializable snapshot sent with every poll."""
resolved_agent_id = agent_id or self.agent_id
if not resolved_agent_id:
raise ValueError("agent_id is required to create a snapshot")
with self._condition:
active_ids = tuple(self._active)
queued_ids = tuple(self._waiting)
outstanding = len(active_ids) + len(queued_ids)
if self._closed:
status = AgentStatus.DRAINING if outstanding else AgentStatus.OFFLINE
elif outstanding >= self.capacity:
status = AgentStatus.OVERLOADED
elif outstanding:
status = AgentStatus.BUSY
else:
status = AgentStatus.ONLINE
progress = {
job_id: self._entries[job_id].progress
for job_id in active_ids + queued_ids
}
# The live trace: the latest harness step per job ("Local model
# step 2", "Calling allowed tool calculate"). Carried on the same
# poll the progress fraction already rides, so the user sees what
# the run is doing rather than only how far along it claims to be.
progress_notes = {
job_id: self._entries[job_id].progress_message
for job_id in active_ids + queued_ids
if self._entries[job_id].progress_message
}
progress_phase = {
job_id: self._entries[job_id].phase
for job_id in active_ids + queued_ids
}
live_steps = {
job_id: tuple(dict(step) for step in self._entries[job_id].live_steps)
for job_id in active_ids + queued_ids
if self._entries[job_id].live_steps
}
started_at = {
job_id: float(self._entries[job_id].started_wall)
for job_id in active_ids + queued_ids
if self._entries[job_id].started_wall is not None
}
duration = (
sum(self._durations) / len(self._durations)
if self._durations
else self.default_job_seconds
)
batches = math.ceil(outstanding / self.max_active) if outstanding else 0
return AgentSnapshot(
agent_id=resolved_agent_id,
status=status,
active_job_ids=active_ids,
queued_job_ids=queued_ids,
queue_capacity=self.capacity,
estimated_wait_s=float(batches * duration),
progress=progress,
progress_notes=progress_notes,
progress_phase=progress_phase,
live_steps=live_steps,
started_at=started_at,
energy_available=bool(energy_available),
last_seen=self._wall_clock(),
)
def close(self, *, cancel_waiting: bool = False) -> None:
"""Stop new admission and optionally cancel all waiting work."""
with self._condition:
self._closed = True
if cancel_waiting:
for job_id in tuple(self._waiting):
entry = self._entries[job_id]
entry.cancel_event.set()
self._finish_unlocked(entry, JobStatus.CANCELLED, "queue_closed")
self._waiting.clear()
self._condition.notify_all()
def _finish_unlocked(
self, entry: _Entry, status: JobStatus, reason: str | None
) -> None:
job_id = entry.job_id
now = self._monotonic()
if entry.started_monotonic is not None and job_id in self._active:
self._durations.append(max(0.0, now - entry.started_monotonic))
self._active.pop(job_id, None)
entry.status = status
entry.finished_monotonic = now
entry.terminal_reason = reason
entry.progress_message = status.value.replace("_", " ").title()
entry.phase = PHASE_DONE
# The live step stream is erased with the prompt, for the same reason:
# it quotes the run. A terminal job keeps scheduling metadata only.
entry.live_steps = []
# Release the JobSpec, and with it the prompt and conversation context,
# at the moment the job becomes terminal. History is bounded at
# ``history_limit`` entries, but those entries are now scheduling
# metadata only -- previously each one pinned a full prompt in memory
# for up to 256 jobs.
entry.job = None
self._terminal_order.append(job_id)
self._trim_history_unlocked()
def _trim_history_unlocked(self) -> None:
while len(self._terminal_order) > self.history_limit:
old_id = self._terminal_order.popleft()
old = self._entries.get(old_id)
if old is not None and old.status in TERMINAL_STATUSES:
self._entries.pop(old_id, None)
def _require_entry(self, job_id: str) -> _Entry:
entry = self._entries.get(job_id)
if entry is None:
raise KeyError(job_id)
return entry
def _require_active(self, job_id: str) -> _Entry:
entry = self._active.get(job_id)
if entry is None:
raise ValueError(f"job {job_id!r} is not active")
return entry
def _view_unlocked(self, job_id: str) -> JobView:
entry = self._entries[job_id]
if job_id in self._active:
position: int | None = 0
else:
try:
position = tuple(self._waiting).index(job_id) + 1
except ValueError:
position = None
return JobView(
job_id=job_id,
status=entry.status,
queue_position=position,
progress=entry.progress,
progress_message=entry.progress_message,
cancel_requested=entry.cancel_event.is_set(),
terminal_reason=entry.terminal_reason,
phase=entry.phase,
)