feat(scheduler,conversations,api): cron and schedule in the gateway tz, job runs persisted

This commit is contained in:
hh
2026-09-01 23:13:17 +02:00
parent 9b1127acff
commit 1954a6e9f2
10 changed files with 326 additions and 43 deletions
+1
View File
@@ -210,6 +210,7 @@ async def _async_main() -> None:
rotation=Rotation(
conversations, gateway.rotation or RotationPolicy(tz=gateway.tz)
),
tz=gateway.tz,
)
conversations.scheduler = scheduler
+8 -3
View File
@@ -23,7 +23,7 @@ import logging
import re
import uuid
from dataclasses import dataclass, field
from datetime import UTC, date, datetime, timedelta
from datetime import UTC, date, datetime, timedelta, tzinfo
from typing import TYPE_CHECKING, Any, cast
from claude_agent_sdk import (
@@ -298,6 +298,10 @@ class Conversations:
self._tasks: set[asyncio.Task[None]] = set()
self._idle_task: asyncio.Task[None] | None = None
@property
def db(self) -> Database:
return self._db
@property
def queue(self) -> InjectQueue:
return self._queue
@@ -964,6 +968,7 @@ class Conversations:
if result.text.strip():
await self.inject(parent, result.text, urgency="normal", origin="слив")
await self.set_status(conv, "merged")
await self.mark_closed(conv)
self._bus.publish(
"conversation.merged",
conversation_id=conv.external_id,
@@ -1946,7 +1951,7 @@ class Conversations:
)
def parse_at(at: str) -> datetime:
def parse_at(at: str, tz: tzinfo = UTC) -> datetime:
raw = at.strip()
match = _RELATIVE.match(raw.replace(" ", ""))
if match:
@@ -1954,7 +1959,7 @@ def parse_at(at: str) -> datetime:
return datetime.now(UTC) + timedelta(seconds=int(amount) * _UNITS[unit])
parsed = datetime.fromisoformat(raw)
if parsed.tzinfo is None:
parsed = parsed.astimezone()
parsed = parsed.replace(tzinfo=tz)
return parsed.astimezone(UTC)
+119 -26
View File
@@ -14,19 +14,27 @@ from __future__ import annotations
import asyncio
import contextlib
import functools
import json
import logging
import traceback
from dataclasses import dataclass
from datetime import UTC, datetime, timedelta
from datetime import UTC, datetime, timedelta, tzinfo
from typing import TYPE_CHECKING, Any, cast
from zoneinfo import ZoneInfo
from croniter import croniter
from pgqueuer import PgQueuer, Queries
from pgqueuer.core.executors import ScheduleExecutor
from pgqueuer.models import JobId
from sqlalchemy import func
from sqlmodel import col, select
from starlette.applications import Starlette
from starlette.responses import JSONResponse
from starlette.routing import Route
from beaver_gateway.core.conversations import parse_at
from beaver_gateway.storage.models import JobRunRecord
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable, Coroutine, Sequence
@@ -42,12 +50,27 @@ if TYPE_CHECKING:
from beaver_gateway.core.rotation import Rotation
from beaver_gateway.storage.models import Conversation
__all__ = ["INJECT", "Budget", "Job", "JobRun", "Scheduler"]
__all__ = ["INJECT", "Budget", "Job", "JobRun", "LocalCron", "Scheduler", "next_run"]
_log = logging.getLogger("beaver_gateway.core.scheduler")
INJECT = "inject"
RETRY = timedelta(minutes=15)
ERROR_MAX = 2000
def next_run(expression: str, tz: tzinfo, now: datetime | None = None) -> datetime:
"""The cron's next fire time, read in ``tz``, returned in UTC."""
start = (now or datetime.now(UTC)).astimezone(tz)
return croniter(expression, start_time=start).get_next(datetime).astimezone(UTC)
@dataclass
class LocalCron(ScheduleExecutor):
tz: tzinfo = UTC
def get_next(self) -> datetime:
return next_run(self.parameters.expression, self.tz)
@dataclass(frozen=True, slots=True)
@@ -165,9 +188,11 @@ class Scheduler:
budget: Budget | None = None,
rotation: Rotation | None = None,
heartbeat: timedelta = timedelta(seconds=30),
tz: str = "UTC",
) -> None:
self.conversations = conversations
self.rotation = rotation
self.tz = ZoneInfo(tz)
self.budget = budget or Budget()
self._jobs = {job.name: job for job in jobs}
self._driver = driver
@@ -229,7 +254,12 @@ class Scheduler:
if job.cron:
@pgq.schedule(job.name, job.cron, clean_old=True)
@pgq.schedule(
job.name,
job.cron,
clean_old=True,
executor_factory=functools.partial(LocalCron, tz=self.tz),
)
async def cron(_schedule: PgSchedule) -> None:
await self._dispatch(job, trigger="cron", payload={})
@@ -257,19 +287,33 @@ class Scheduler:
return
started = datetime.now(UTC)
bus.publish("job.start", job=job.name, trigger=trigger)
status = "done"
status, error = "done", None
try:
await job.run(JobRun(job, trigger, payload, self))
except Exception: # noqa: BLE001
status = "failed"
except Exception as exc: # noqa: BLE001
status, error = "failed", _error_text(exc)
_log.exception("job %s (%s) failed", job.name, trigger)
self._runs[job.name] = {
"trigger": trigger,
"started_at": started.isoformat(timespec="seconds"),
"status": status,
}
record = JobRunRecord(
job=job.name,
trigger=trigger,
started_at=started,
finished_at=datetime.now(UTC),
status=status,
error=error,
payload=payload,
)
await self._record(record)
self._runs[job.name] = self._run_public(record)
bus.publish("job.end", job=job.name, trigger=trigger, status=status)
async def _record(self, record: JobRunRecord) -> None:
try:
async with self.conversations.db.session() as session:
session.add(record)
await session.commit()
except Exception: # noqa: BLE001
_log.exception("job %s: run not recorded", record.job)
async def trigger(
self,
job: Job,
@@ -325,14 +369,14 @@ class Scheduler:
if self._queries is None:
msg = "scheduler needs postgres; `schedule` is unavailable"
raise RuntimeError(msg)
when = parse_at(at)
when = parse_at(at, self.tz).astimezone(self.tz)
delay = max(when - datetime.now(UTC), timedelta(0))
payload = _encode(
{
"conversation": conv.external_id,
"text": text,
"urgency": urgency,
"at": when.isoformat(timespec="seconds"),
"at": when.isoformat(timespec="minutes"),
}
)
ids = await self._queries.enqueue(
@@ -347,7 +391,7 @@ class Scheduler:
"schedule.created",
conversation_id=conv.external_id,
job=job_id,
execute_at=when.isoformat(timespec="seconds"),
execute_at=when.isoformat(timespec="minutes"),
)
return job_id, when
@@ -411,16 +455,22 @@ class Scheduler:
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)
_job_public(row, self.tz)
for row in await self._queries.browse_queue(limit=200)
]
utilization = await self.utilization()
last = await self._last_runs()
return {
"enabled": self.enabled,
"tz": str(self.tz),
"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],
"jobs": [
self._job_public(job, crons.get(job.name), last.get(job.name))
for job in self.jobs
],
"queue": queue,
}
@@ -428,24 +478,60 @@ class Scheduler:
if self._queries is None:
return []
rows = await self._queries.browse_queue(limit=500, entrypoints=[INJECT])
out = [_job_public(row) for row in rows]
out = [_job_public(row, self.tz) 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]:
async def runs(self, name: str, limit: int = 50) -> list[dict[str, Any]]:
stmt = (
select(JobRunRecord)
.where(JobRunRecord.job == name)
.order_by(col(JobRunRecord.id).desc())
.limit(limit)
)
async with self.conversations.db.session() as session:
rows = (await session.exec(stmt)).all()
return [self._run_public(row) for row in rows]
async def _last_runs(self) -> dict[str, dict[str, Any]]:
newest = select(func.max(JobRunRecord.id)).group_by(JobRunRecord.job)
stmt = select(JobRunRecord).where(col(JobRunRecord.id).in_(newest))
try:
async with self.conversations.db.session() as session:
rows = (await session.exec(stmt)).all()
except Exception: # noqa: BLE001
_log.exception("job runs unreadable, using memory")
return dict(self._runs)
return {row.job: self._run_public(row) for row in rows}
def _job_public(
self, job: Job, cron: PgSchedule | None, run: dict[str, Any] | 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,
"next_run": _iso(cron.next_run, self.tz) if cron is not None else None,
"last_run": _iso(cron.last_run, self.tz) if cron is not None else None,
"status": str(cron.status) if cron is not None else None,
"run": self._runs.get(job.name),
"run": run or self._runs.get(job.name),
}
def _run_public(self, row: JobRunRecord) -> dict[str, Any]:
return {
"id": row.id,
"job": row.job,
"trigger": row.trigger,
"started_at": _iso(row.started_at, self.tz),
"finished_at": _iso(row.finished_at, self.tz),
"status": row.status,
"error": row.error,
"payload": row.payload,
}
# ---- http ----------------------------------------------------------
@@ -506,21 +592,28 @@ def _decode(payload: bytes | None) -> dict[str, Any]:
return data if isinstance(data, dict) else {"raw": data}
def _job_public(row: PgJob) -> dict[str, Any]:
def _job_public(row: PgJob, tz: tzinfo) -> 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),
"execute_after": _iso(row.execute_after, tz),
"created": _iso(row.created, tz),
"attempts": row.attempts,
"payload": _decode(row.payload),
}
def _error_text(exc: BaseException) -> str:
head = f"{type(exc).__name__}: {exc}"
return f"{head}\n{''.join(traceback.format_exception(exc))}"[:ERROR_MAX]
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
def _iso(value: datetime | None, tz: tzinfo = UTC) -> str | None:
if value is None:
return None
return _aware(value).astimezone(tz).isoformat(timespec="seconds")
+3 -1
View File
@@ -22,7 +22,9 @@ from datetime import UTC, datetime
from typing import TYPE_CHECKING, Any
try:
from claude_agent_sdk._cli_version import __cli_version__ as _cli_version
from claude_agent_sdk._cli_version import __cli_version__
_cli_version = str(__cli_version__)
except ImportError: # pragma: no cover
_cli_version = "2.1.248"
+16 -4
View File
@@ -62,7 +62,7 @@ if TYPE_CHECKING:
from pathlib import Path
from beaver_gateway.core.conversations import Conversations
from beaver_gateway.core.scheduler import Scheduler
from beaver_gateway.core.scheduler import Job, Scheduler
from beaver_gateway.frontends.base import GatewayRuntime
_log = logging.getLogger("beaver_gateway.frontends.api")
@@ -580,15 +580,27 @@ def build_app(runtime: GatewayRuntime, *, memory_root: Path | None = None) -> Fa
await require_token(request, runtime, scope=SCOPE)
return await scheduler_of().snapshot()
def job_of(scheduler: Scheduler, name: str) -> Job:
job = scheduler.job(name)
if job is None:
raise HTTPException(status.HTTP_404_NOT_FOUND, f"no job {name!r}")
return job
@app.post("/jobs/{name}/run", status_code=status.HTTP_202_ACCEPTED)
async def run_job(name: str, request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
scheduler = scheduler_of()
job = scheduler.job(name)
if job is None:
raise HTTPException(status.HTTP_404_NOT_FOUND, f"no job {name!r}")
job = job_of(scheduler, name)
return {"job": await scheduler.trigger(job, await body_of(request))}
@app.get("/jobs/{name}/runs")
async def job_runs(name: str, request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
scheduler = scheduler_of()
job = job_of(scheduler, name)
limit = min(int(request.query_params.get("limit", 50)), 500)
return {"job": job.name, "runs": await scheduler.runs(job.name, limit)}
@app.delete("/jobs/queue/{job_id}")
async def cancel_job(job_id: int, request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
+27
View File
@@ -360,6 +360,32 @@ class RateLimit(SQLModel, table=True):
)
class JobRunRecord(SQLModel, table=True):
"""One finished run of a scheduler job (§3.6): what fired it, how it ended.
``error`` keeps the head of the traceback of a failed run; ``payload``
is the trigger's data (``{}`` for cron) and should stay small.
"""
__tablename__ = "job_runs"
id: int | None = Field(default=None, primary_key=True)
job: str = Field(index=True)
trigger: str
started_at: datetime
finished_at: datetime
status: str
error: str | None = Field(default=None)
payload: dict[str, Any] = Field(
default_factory=dict,
sa_column=Column(
JSON().with_variant(JSONB(), "postgresql"),
nullable=False,
server_default=text("'{}'"),
),
)
__all__ = [
"AuditLog",
"Conversation",
@@ -367,6 +393,7 @@ __all__ = [
"ConversationMessage",
"Delivery",
"InjectQueueItem",
"JobRunRecord",
"RateLimit",
"TelegramUpdate",
"Token",