"""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/``, 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