Files
beaver-gateway/src/beaver_gateway/core/scheduler.py
T

527 lines
18 KiB
Python

"""Jobs and deferred injects on pgqueuer (§3.6, §4.5).
A job is a name, a handler and its triggers: a cron expression, the
webhook ``/hooks/<name>``, gateway bus events. The executor is pgqueuer on
the gateway's own Postgres (a dedicated autocommit connection for
LISTEN/NOTIFY), so cron ticks, webhook deliveries and the ``schedule``
tool's one-off injects all live in one table and survive a restart. A
handler only queues work for a conversation and returns; the turn itself
runs in the conversation's worker. Non-critical jobs step aside while the
subscription window is past its threshold.
"""
from __future__ import annotations
import asyncio
import contextlib
import json
import logging
from dataclasses import dataclass
from datetime import UTC, datetime, timedelta
from typing import TYPE_CHECKING, Any, cast
from pgqueuer import PgQueuer, Queries
from pgqueuer.models import JobId
from starlette.applications import Starlette
from starlette.responses import JSONResponse
from starlette.routing import Route
from beaver_gateway.core.conversations import parse_at
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable, Coroutine, Sequence
from pgqueuer.models import Job as PgJob
from pgqueuer.models import Schedule as PgSchedule
from pgqueuer.ports.driver import Driver
from starlette.requests import Request
from beaver_gateway.core.conversations import Conversations, DistillResult
from beaver_gateway.core.distill import LineCap
from beaver_gateway.core.injects import Priority
from beaver_gateway.core.rotation import Rotation
from beaver_gateway.storage.models import Conversation
__all__ = ["INJECT", "Budget", "Job", "JobRun", "Scheduler"]
_log = logging.getLogger("beaver_gateway.core.scheduler")
INJECT = "inject"
RETRY = timedelta(minutes=15)
@dataclass(frozen=True, slots=True)
class Job:
name: str
run: Callable[[JobRun], Awaitable[None]]
cron: str | None = None
webhook: bool = False
events: tuple[str, ...] = ()
critical: bool = True
dedupe: bool = True
@property
def entrypoint(self) -> str:
return f"job:{self.name}"
@dataclass(frozen=True, slots=True)
class Budget:
threshold: float = 0.7
tokens: int | None = None
window: timedelta = timedelta(hours=5)
limit_window: str = "five_hour"
@dataclass(slots=True)
class JobRun:
job: Job
trigger: str
payload: dict[str, Any]
scheduler: Scheduler
@property
def conversations(self) -> Conversations:
return self.scheduler.conversations
async def master(self) -> Conversation | None:
masters = await self.conversations.find(kind="master", status="open", limit=1)
return masters[0] if masters else None
async def inject_master(
self, text: str, *, urgency: Priority = "normal", origin: str | None = None
) -> bool:
master = await self.master()
if master is None:
_log.error(
"job %s: no open master, inject lost: %s", self.job.name, text[:200]
)
return False
await self.conversations.inject(
master, text, urgency=urgency, origin=origin or self.job.name
)
return True
async def spawn_job(
self,
*,
agent: str,
text: str,
title: str | None = None,
line_cap: LineCap | None = None,
) -> Conversation:
"""A headless job turn; ``line_cap`` bounces a rewrite past the cap."""
return await self.conversations.spawn(
kind="job",
agent=agent,
seed="brief",
text=text,
title=title,
origin="job",
flags={"line_cap": line_cap.as_flags()} if line_cap else None,
)
async def close_idle(
self,
*,
kind: str = "deep",
days: int = 2,
limit: int = 3,
since: datetime | None = None,
) -> list[DistillResult]:
"""§4.5: close chats quiet for ``days``, at most ``limit`` per run."""
out: list[DistillResult] = []
for conv in await self.conversations.idle(
kind=kind, days=days, since=since, limit=limit
):
try:
out.append(
await self.conversations.distill(conv, reason=f"idle {days}d")
)
except Exception: # noqa: BLE001
_log.exception("closing idle %s failed", conv.external_id)
return out
async def retry_in(self, delay: timedelta) -> None:
await self.scheduler.trigger(
self.job, self.payload, delay=delay, trigger=self.trigger
)
async def rotate(self) -> list[Conversation]:
rotation = self.scheduler.rotation
return await rotation.tick() if rotation is not None else []
def background(self, coro: Coroutine[Any, Any, Any]) -> None:
self.scheduler.background(coro)
class Scheduler:
def __init__(
self,
*,
conversations: Conversations,
jobs: Sequence[Job] = (),
driver: Driver | None = None,
budget: Budget | None = None,
rotation: Rotation | None = None,
heartbeat: timedelta = timedelta(seconds=30),
) -> None:
self.conversations = conversations
self.rotation = rotation
self.budget = budget or Budget()
self._jobs = {job.name: job for job in jobs}
self._driver = driver
self._queries: Queries | None = None
self._heartbeat = heartbeat
self._tasks: set[asyncio.Task[Any]] = set()
self._runs: dict[str, dict[str, Any]] = {}
@property
def enabled(self) -> bool:
return self._queries is not None
@property
def jobs(self) -> list[Job]:
return list(self._jobs.values())
def job(self, name: str) -> Job | None:
return self._jobs.get(name)
# ---- lifecycle -----------------------------------------------------
async def start(self) -> None:
if self._driver is not None:
queries = Queries(self._driver)
if not await queries.has_table("pgqueuer"):
await queries.install()
pgq = PgQueuer(self._driver, queries=queries)
self._register(pgq)
self._queries = queries
self._spawn(pgq.run(heartbeat_timeout=self._heartbeat))
else:
_log.warning("scheduler: no postgres driver - cron and webhooks are off")
self._spawn(self._events())
async def stop(self) -> None:
for task in list(self._tasks):
task.cancel()
for task in list(self._tasks):
with contextlib.suppress(BaseException):
await task
self._tasks.clear()
self._queries = None
def _register(self, pgq: PgQueuer) -> None:
@pgq.entrypoint(INJECT)
async def deliver(job: PgJob) -> None:
await self._deliver(job)
for job in self._jobs.values():
self._register_job(pgq, job)
def _register_job(self, pgq: PgQueuer, job: Job) -> None:
@pgq.entrypoint(job.entrypoint)
async def queued(pg_job: PgJob) -> None:
data = _decode(pg_job.payload)
await self._dispatch(
job, trigger=str(data.pop("trigger", "queue")), payload=data
)
if job.cron:
@pgq.schedule(job.name, job.cron, clean_old=True)
async def cron(_schedule: PgSchedule) -> None:
await self._dispatch(job, trigger="cron", payload={})
async def _events(self) -> None:
listeners = [job for job in self._jobs.values() if job.events]
if not listeners:
return
async for event in self.conversations.bus.stream():
for job in listeners:
if event.get("type") in job.events:
self._spawn(
self._dispatch(job, trigger="event", payload=dict(event))
)
# ---- dispatch ------------------------------------------------------
async def _dispatch(
self, job: Job, *, trigger: str, payload: dict[str, Any]
) -> None:
bus = self.conversations.bus
if not job.critical and await self.throttled():
bus.publish("job.deferred", job=job.name, trigger=trigger)
if trigger != "cron":
await self.trigger(job, payload, delay=RETRY, trigger=trigger)
return
started = datetime.now(UTC)
bus.publish("job.start", job=job.name, trigger=trigger)
status = "done"
try:
await job.run(JobRun(job, trigger, payload, self))
except Exception: # noqa: BLE001
status = "failed"
_log.exception("job %s (%s) failed", job.name, trigger)
self._runs[job.name] = {
"trigger": trigger,
"started_at": started.isoformat(timespec="seconds"),
"status": status,
}
bus.publish("job.end", job=job.name, trigger=trigger, status=status)
async def trigger(
self,
job: Job,
payload: dict[str, Any] | None = None,
*,
delay: timedelta | None = None,
trigger: str = "manual",
) -> int | None:
"""Queue one run of ``job``; runs inline when there is no executor."""
data = {**(payload or {}), "trigger": trigger}
if self._queries is None:
self._spawn(
self._dispatch(job, trigger=trigger, payload=payload or {}), delay=delay
)
return None
ids = await self._queries.enqueue(
job.entrypoint, _encode(data), execute_after=delay
)
return int(ids[0]) if ids and ids[0] is not None else None
async def hook(self, name: str, payload: dict[str, Any]) -> int | None:
job = self._jobs.get(name)
if job is None or not job.webhook:
msg = f"no webhook job {name!r}"
raise LookupError(msg)
self.conversations.bus.publish("hook", job=name)
if self._queries is None:
return await self.trigger(job, payload, trigger="webhook")
ids = await self._queries.enqueue(
job.entrypoint,
_encode({**payload, "trigger": "webhook"}),
dedupe_key=f"hook:{name}" if job.dedupe else None,
on_conflict="skip" if job.dedupe else "raise",
)
return int(ids[0]) if ids and ids[0] is not None else None
# ---- deferred injects ----------------------------------------------
async def schedule(
self,
conv: Conversation,
at: str,
text: str,
*,
urgency: Priority = "wake",
dedupe_key: str | None = None,
) -> tuple[int | None, datetime]:
"""A deferred inject.
``wake`` by default: it was promised for a time, so it starts a turn
then instead of waiting for the normal window.
"""
if self._queries is None:
msg = "scheduler needs postgres; `schedule` is unavailable"
raise RuntimeError(msg)
when = parse_at(at)
delay = max(when - datetime.now(UTC), timedelta(0))
payload = _encode(
{
"conversation": conv.external_id,
"text": text,
"urgency": urgency,
"at": when.isoformat(timespec="seconds"),
}
)
ids = await self._queries.enqueue(
INJECT,
payload,
execute_after=delay,
dedupe_key=dedupe_key,
on_conflict="skip" if dedupe_key else "raise",
)
job_id = int(ids[0]) if ids and ids[0] is not None else None
self.conversations.bus.publish(
"schedule.created",
conversation_id=conv.external_id,
job=job_id,
execute_at=when.isoformat(timespec="seconds"),
)
return job_id, when
async def cancel(self, job_id: int) -> bool:
if self._queries is None:
return False
row = await self._queries.queue_job_by_id(JobId(job_id))
if row is None:
return False
await self._queries.mark_job_as_cancelled([JobId(job_id)])
self.conversations.bus.publish("schedule.cancelled", job=job_id)
return True
async def _deliver(self, job: PgJob) -> None:
data = _decode(job.payload)
key = str(data.get("conversation", ""))
conv = await self.conversations.resolve(key)
if conv is None:
_log.error(
"scheduled inject #%s lost: conversation %r is gone: %s",
job.id,
key,
str(data.get("text", ""))[:200],
)
return
await self.conversations.inject(
conv,
str(data.get("text", "")),
urgency=cast("Priority", data.get("urgency") or "wake"),
origin="schedule",
)
# ---- budget --------------------------------------------------------
async def utilization(self) -> float | None:
now = datetime.now(UTC)
values: list[float] = []
for row in await self.conversations.rate_limits(limit=100):
if row.window != self.budget.limit_window or row.utilization is None:
continue
fresh = now - _aware(row.ts) < self.budget.window
live = row.resets_at is None or _aware(row.resets_at) > now
if fresh and live:
values.append(float(row.utilization))
break
if self.budget.tokens:
tokens = await self.conversations.usage_tokens(now - self.budget.window)
values.append(tokens / self.budget.tokens)
return max(values) if values else None
async def throttled(self) -> bool:
utilization = await self.utilization()
return utilization is not None and utilization > self.budget.threshold
# ---- introspection -------------------------------------------------
async def snapshot(self) -> dict[str, Any]:
crons: dict[str, PgSchedule] = {}
queue: list[dict[str, Any]] = []
if self._queries is not None:
for row in await self._queries.peek_schedule():
crons[str(row.entrypoint)] = row
queue = [
_job_public(row) for row in await self._queries.browse_queue(limit=200)
]
utilization = await self.utilization()
return {
"enabled": self.enabled,
"utilization": utilization,
"throttled": utilization is not None
and utilization > self.budget.threshold,
"threshold": self.budget.threshold,
"jobs": [self._job_public(job, crons.get(job.name)) for job in self.jobs],
"queue": queue,
}
async def scheduled(self, conv: Conversation | None = None) -> list[dict[str, Any]]:
if self._queries is None:
return []
rows = await self._queries.browse_queue(limit=500, entrypoints=[INJECT])
out = [_job_public(row) for row in rows]
if conv is not None:
out = [
j for j in out if j["payload"].get("conversation") == conv.external_id
]
return out
def _job_public(self, job: Job, cron: PgSchedule | None) -> dict[str, Any]:
return {
"name": job.name,
"cron": job.cron,
"webhook": job.webhook,
"events": list(job.events),
"critical": job.critical,
"next_run": _iso(cron.next_run) if cron is not None else None,
"last_run": _iso(cron.last_run) if cron is not None else None,
"status": str(cron.status) if cron is not None else None,
"run": self._runs.get(job.name),
}
# ---- http ----------------------------------------------------------
def app(self, authorize: Callable[[Request], Awaitable[Any]]) -> Starlette:
async def hook(request: Request) -> JSONResponse:
await authorize(request)
name = request.path_params["name"]
payload = await _payload(request)
try:
job_id = await self.hook(name, payload)
except LookupError as exc:
return JSONResponse({"error": str(exc)}, status_code=404)
return JSONResponse({"job": job_id, "name": name}, status_code=202)
return Starlette(routes=[Route("/{name}", hook, methods=["POST"])])
# ---- internals -----------------------------------------------------
def background(self, coro: Coroutine[Any, Any, Any]) -> None:
self._spawn(coro)
def _spawn(
self, coro: Coroutine[Any, Any, Any], *, delay: timedelta | None = None
) -> None:
async def later() -> None:
if delay is not None:
await asyncio.sleep(delay.total_seconds())
await coro
task = asyncio.create_task(later() if delay is not None else coro)
self._tasks.add(task)
task.add_done_callback(self._tasks.discard)
async def _payload(request: Request) -> dict[str, Any]:
body = await request.body()
if not body:
return dict(request.query_params)
try:
data = json.loads(body)
except ValueError:
return {"raw": body.decode("utf-8", errors="replace")}
return data if isinstance(data, dict) else {"raw": data}
def _encode(data: dict[str, Any]) -> bytes:
return json.dumps(data, ensure_ascii=False, default=str).encode("utf-8")
def _decode(payload: bytes | None) -> dict[str, Any]:
if not payload:
return {}
try:
data = json.loads(payload)
except ValueError:
return {"raw": payload.decode("utf-8", errors="replace")}
return data if isinstance(data, dict) else {"raw": data}
def _job_public(row: PgJob) -> dict[str, Any]:
return {
"id": int(row.id),
"entrypoint": row.entrypoint,
"status": str(row.status),
"execute_after": _iso(row.execute_after),
"created": _iso(row.created),
"attempts": row.attempts,
"payload": _decode(row.payload),
}
def _aware(value: datetime) -> datetime:
return value if value.tzinfo is not None else value.replace(tzinfo=UTC)
def _iso(value: datetime | None) -> str | None:
return _aware(value).isoformat(timespec="seconds") if value is not None else None