add /v1/tasks scheduled-task scheduler endpoints

Add Postgres-backed scheduler primitive to the proxy so external agents
can offload "fire this work at time X" without running their own
APScheduler + tasks table. Six endpoints under /v1/tasks behind
LITELLM_SCHEDULED_TASKS_ENABLED:

- POST   /v1/tasks            create
- GET    /v1/tasks            list
- GET    /v1/tasks/{id}       fetch
- PATCH  /v1/tasks/{id}       update (pending only)
- DELETE /v1/tasks/{id}       cancel / stop early
- GET    /v1/tasks/due        atomic claim + schedule advance

Identity (user_id, team_id, agent_id) and ownership (owner_token, FK to
LiteLLM_VerificationToken) are stamped from auth — never accepted from
the request body. /due reads agent_id from the calling key.

claim_due() uses one approved raw-SQL exception per CLAUDE.md
(SELECT FOR UPDATE SKIP LOCKED is required for multi-pod safety and is
not expressible via Prisma model methods). All other writes go through
Prisma model methods, batched inside the same transaction.

Migration adds:
- Partial index on (next_run_at) WHERE status='pending' so the ticker
  scans only candidate rows.
- CHECK constraints for schedule_kind, status, and the
  "action='check' implies check_prompt IS NOT NULL" invariant.

Schedule kinds: interval ('5m'/'2h'/'1d'), cron (5-field crontab with
optional IANA tz), once. compute_next_run is ported from the original
test_local_agent_2/task_runner.py.

Behind LITELLM_SCHEDULED_TASKS_ENABLED=false default — router only
mounts when the flag is on.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Krrish Dholakia 2026-04-29 12:59:34 -07:00
parent 3e1479c052
commit e5dec5e809
13 changed files with 1542 additions and 0 deletions

View file

@ -0,0 +1,66 @@
-- CreateTable
CREATE TABLE "LiteLLM_ScheduledTaskTable" (
"task_id" TEXT NOT NULL,
"owner_token" TEXT NOT NULL,
"user_id" TEXT,
"team_id" TEXT,
"agent_id" TEXT,
"title" TEXT NOT NULL,
"action" TEXT NOT NULL,
"action_args" JSONB,
"check_prompt" TEXT,
"format_prompt" TEXT,
"schedule_kind" TEXT NOT NULL,
"schedule_spec" TEXT NOT NULL,
"schedule_tz" TEXT,
"next_run_at" TIMESTAMP(3) NOT NULL,
"expires_at" TIMESTAMP(3) NOT NULL,
"fire_once" BOOLEAN NOT NULL DEFAULT true,
"status" TEXT NOT NULL DEFAULT 'pending',
"last_fired_at" TIMESTAMP(3),
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL,
CONSTRAINT "LiteLLM_ScheduledTaskTable_pkey" PRIMARY KEY ("task_id")
);
-- CreateIndex
CREATE INDEX "LiteLLM_ScheduledTaskTable_owner_token_status_idx"
ON "LiteLLM_ScheduledTaskTable" ("owner_token", "status");
-- CreateIndex
CREATE INDEX "LiteLLM_ScheduledTaskTable_user_id_idx"
ON "LiteLLM_ScheduledTaskTable" ("user_id");
-- CreateIndex
CREATE INDEX "LiteLLM_ScheduledTaskTable_team_id_idx"
ON "LiteLLM_ScheduledTaskTable" ("team_id");
-- CreateIndex
CREATE INDEX "LiteLLM_ScheduledTaskTable_agent_id_status_idx"
ON "LiteLLM_ScheduledTaskTable" ("agent_id", "status");
-- AddForeignKey
ALTER TABLE "LiteLLM_ScheduledTaskTable"
ADD CONSTRAINT "LiteLLM_ScheduledTaskTable_owner_token_fkey"
FOREIGN KEY ("owner_token")
REFERENCES "LiteLLM_VerificationToken" ("token")
ON DELETE CASCADE ON UPDATE CASCADE;
-- Hand-appended (Prisma cannot express partial indexes or CHECK constraints):
-- Partial index over rows the ticker actually scans.
CREATE INDEX "LiteLLM_ScheduledTaskTable_next_run_due_idx"
ON "LiteLLM_ScheduledTaskTable" ("next_run_at")
WHERE status = 'pending';
-- Constrain enums at DB. Bad writes bounce.
ALTER TABLE "LiteLLM_ScheduledTaskTable"
ADD CONSTRAINT "LiteLLM_ScheduledTaskTable_schedule_kind_check"
CHECK (schedule_kind IN ('interval','cron','once'));
ALTER TABLE "LiteLLM_ScheduledTaskTable"
ADD CONSTRAINT "LiteLLM_ScheduledTaskTable_status_check"
CHECK (status IN ('pending','fired','expired','cancelled'));
ALTER TABLE "LiteLLM_ScheduledTaskTable"
ADD CONSTRAINT "LiteLLM_ScheduledTaskTable_action_prompt_check"
CHECK (action <> 'check' OR check_prompt IS NOT NULL);

View file

@ -410,6 +410,7 @@ model LiteLLM_VerificationToken {
litellm_project_table LiteLLM_ProjectTable? @relation(fields: [project_id], references: [project_id])
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
jwt_key_mappings LiteLLM_JWTKeyMapping[]
scheduled_tasks LiteLLM_ScheduledTaskTable[]
// SELECT COUNT(*) FROM (SELECT "public"."LiteLLM_VerificationToken"."token" FROM "public"."LiteLLM_VerificationToken" WHERE ("public"."LiteLLM_VerificationToken"."user_id" = $1 AND ("public"."LiteLLM_VerificationToken"."team_id" IS NULL OR "public"."LiteLLM_VerificationToken"."team_id" <> $2)) OFFSET $3 ) AS "sub"
// SELECT ... FROM "public"."LiteLLM_VerificationToken" WHERE "public"."LiteLLM_VerificationToken"."user_id" = $1 OFFSET $2
@ -1290,3 +1291,38 @@ model LiteLLM_AdaptiveRouterSession {
@@id([session_id, router_name, model_name])
@@index([last_activity_at], map: "idx_adaptive_router_session_activity")
}
model LiteLLM_ScheduledTaskTable {
task_id String @id @default(uuid())
owner_token String
owner_key LiteLLM_VerificationToken @relation(fields: [owner_token], references: [token], onDelete: Cascade)
user_id String?
team_id String?
agent_id String?
title String
action String
action_args Json?
check_prompt String?
format_prompt String?
schedule_kind String
schedule_spec String
schedule_tz String?
next_run_at DateTime
expires_at DateTime
fire_once Boolean @default(true)
status String @default("pending")
last_fired_at DateTime?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@index([owner_token, status])
@@index([user_id])
@@index([team_id])
@@index([agent_id, status])
}

View file

@ -511,6 +511,9 @@ from litellm.proxy.utils import (
prefetch_config_params,
update_spend,
)
from litellm.proxy.scheduled_tasks.endpoints import (
router as scheduled_tasks_router,
)
from litellm.proxy.vector_store_endpoints.endpoints import router as vector_store_router
from litellm.proxy.vector_store_endpoints.management_endpoints import (
router as vector_store_management_router,
@ -14277,6 +14280,8 @@ app.include_router(model_access_group_management_router)
app.include_router(tag_management_router)
app.include_router(tool_management_router)
app.include_router(memory_router)
if str(os.getenv("LITELLM_SCHEDULED_TASKS_ENABLED", "false")).lower() == "true":
app.include_router(scheduled_tasks_router)
app.include_router(cost_tracking_settings_router)
app.include_router(router_settings_router)
app.include_router(fallback_management_router)

View file

@ -0,0 +1,317 @@
"""
FastAPI endpoints for scheduled tasks.
Six endpoints under /v1/tasks. Auth via user_api_key_auth — owner_token,
user_id, team_id, agent_id are all stamped from the calling key, never
accepted from the client.
"""
from __future__ import annotations
from datetime import datetime, timezone
from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
import litellm.proxy.scheduled_tasks.store as store
from litellm.proxy.scheduled_tasks.schedule import (
compute_next_run,
validate_schedule,
)
from litellm.proxy.scheduled_tasks.types import (
CreateScheduledTaskRequest,
DueTaskResponse,
DueTasksResponse,
ListScheduledTasksResponse,
ScheduledTaskResponse,
UpdateScheduledTaskRequest,
)
router = APIRouter()
def _require_token(user_api_key_dict: UserAPIKeyAuth) -> str:
"""
Resolve the FK target for owner_token. We deliberately prefer `.token`
(the hashed value stored in LiteLLM_VerificationToken) over `.api_key`
(the raw key as sent by the client), so the FK constraint matches.
"""
token = user_api_key_dict.token or user_api_key_dict.api_key
if not token:
raise HTTPException(status_code=401, detail="missing api key")
return token
def _row_to_response(row) -> ScheduledTaskResponse:
"""Convert a Prisma row (model object) into the response schema."""
return ScheduledTaskResponse(
task_id=row.task_id,
owner_token=row.owner_token,
user_id=row.user_id,
team_id=row.team_id,
agent_id=row.agent_id,
title=row.title,
action=row.action,
action_args=row.action_args,
check_prompt=row.check_prompt,
format_prompt=row.format_prompt,
schedule_kind=row.schedule_kind,
schedule_spec=row.schedule_spec,
schedule_tz=row.schedule_tz,
next_run_at=row.next_run_at,
expires_at=row.expires_at,
fire_once=row.fire_once,
status=row.status,
last_fired_at=row.last_fired_at,
created_at=row.created_at,
updated_at=row.updated_at,
)
def _get_prisma_client():
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not initialised")
return prisma_client
@router.post(
"/v1/tasks",
tags=["scheduled tasks"],
response_model=ScheduledTaskResponse,
)
async def create_scheduled_task(
data: CreateScheduledTaskRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
if data.action == "check" and not data.check_prompt:
raise HTTPException(
status_code=400,
detail="check_prompt is required when action='check'",
)
try:
validate_schedule(
kind=data.schedule_kind,
spec=data.schedule_spec,
tz=data.schedule_tz,
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
owner_token = _require_token(user_api_key_dict)
prisma_client = _get_prisma_client()
active = await store.count_active_for_owner(prisma_client, owner_token)
if active >= store.MAX_ACTIVE_TASKS_PER_KEY:
raise HTTPException(
status_code=429,
detail=(
f"too many active scheduled tasks "
f"(max {store.MAX_ACTIVE_TASKS_PER_KEY})"
),
)
now = datetime.now(timezone.utc)
if data.expires_at <= now:
raise HTTPException(status_code=400, detail="expires_at must be in the future")
next_run_at = compute_next_run(
kind=data.schedule_kind,
spec=data.schedule_spec,
tz=data.schedule_tz,
from_time=now,
)
row = await store.create_task(
prisma_client,
owner_token=owner_token,
user_id=user_api_key_dict.user_id,
team_id=user_api_key_dict.team_id,
agent_id=user_api_key_dict.agent_id,
title=data.title,
action=data.action,
action_args=data.action_args,
check_prompt=data.check_prompt,
format_prompt=data.format_prompt,
schedule_kind=data.schedule_kind,
schedule_spec=data.schedule_spec,
schedule_tz=data.schedule_tz,
next_run_at=next_run_at,
expires_at=data.expires_at,
fire_once=data.fire_once,
)
return _row_to_response(row)
@router.get(
"/v1/tasks",
tags=["scheduled tasks"],
response_model=ListScheduledTasksResponse,
)
async def list_scheduled_tasks(
include_terminal: bool = Query(False),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
owner_token = _require_token(user_api_key_dict)
prisma_client = _get_prisma_client()
rows = await store.list_tasks_for_owner(
prisma_client,
owner_token=owner_token,
include_terminal=include_terminal,
)
return ListScheduledTasksResponse(tasks=[_row_to_response(r) for r in rows])
@router.get(
"/v1/tasks/due",
tags=["scheduled tasks"],
response_model=DueTasksResponse,
)
async def get_due_tasks(
actions: Optional[str] = Query(
None,
description="Comma-separated list of action names this worker handles.",
),
limit: int = Query(20, ge=1, le=100),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Atomic claim of due tasks. Returns immediately, possibly with empty list.
Schedule advance happens server-side: recurring rows get a fresh
next_run_at, fire_once rows flip to 'fired'.
"""
_require_token(user_api_key_dict)
prisma_client = _get_prisma_client()
parsed_actions: Optional[List[str]] = None
if actions:
parsed_actions = [a.strip() for a in actions.split(",") if a.strip()]
if not parsed_actions:
parsed_actions = None
rows = await store.claim_due(
prisma_client,
agent_id=user_api_key_dict.agent_id,
actions=parsed_actions,
limit=limit,
)
return DueTasksResponse(
tasks=[
DueTaskResponse(
task_id=r["task_id"],
title=r["title"],
action=r["action"],
action_args=r["action_args"],
check_prompt=r["check_prompt"],
format_prompt=r["format_prompt"],
scheduled_for=r["scheduled_for"],
)
for r in rows
]
)
@router.get(
"/v1/tasks/{task_id}",
tags=["scheduled tasks"],
response_model=ScheduledTaskResponse,
)
async def get_scheduled_task(
task_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
owner_token = _require_token(user_api_key_dict)
prisma_client = _get_prisma_client()
row = await store.get_task_for_owner(
prisma_client,
task_id=task_id,
owner_token=owner_token,
)
if row is None:
raise HTTPException(status_code=404, detail="task not found")
return _row_to_response(row)
@router.patch(
"/v1/tasks/{task_id}",
tags=["scheduled tasks"],
response_model=ScheduledTaskResponse,
)
async def update_scheduled_task(
task_id: str,
data: UpdateScheduledTaskRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
owner_token = _require_token(user_api_key_dict)
prisma_client = _get_prisma_client()
fields = data.model_dump(exclude_unset=True)
if not fields:
raise HTTPException(status_code=400, detail="no fields to update")
existing = await store.get_task_for_owner(
prisma_client,
task_id=task_id,
owner_token=owner_token,
)
if existing is None:
raise HTTPException(status_code=404, detail="task not found")
if existing.status != "pending":
raise HTTPException(
status_code=400,
detail=f"cannot update task in status '{existing.status}'",
)
merged_kind = fields.get("schedule_kind", existing.schedule_kind)
merged_spec = fields.get("schedule_spec", existing.schedule_spec)
merged_tz = fields.get("schedule_tz", existing.schedule_tz)
if any(k in fields for k in ("schedule_kind", "schedule_spec", "schedule_tz")):
try:
validate_schedule(kind=merged_kind, spec=merged_spec, tz=merged_tz)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
merged_action = fields.get("action", existing.action)
merged_check = fields.get("check_prompt", existing.check_prompt)
if merged_action == "check" and not merged_check:
raise HTTPException(
status_code=400,
detail="check_prompt is required when action='check'",
)
try:
row = await store.update_task_for_owner(
prisma_client,
task_id=task_id,
owner_token=owner_token,
fields=fields,
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
if row is None:
raise HTTPException(status_code=404, detail="task not found")
return _row_to_response(row)
@router.delete(
"/v1/tasks/{task_id}",
tags=["scheduled tasks"],
response_model=ScheduledTaskResponse,
)
async def cancel_scheduled_task(
task_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
owner_token = _require_token(user_api_key_dict)
prisma_client = _get_prisma_client()
row = await store.cancel_task_for_owner(
prisma_client,
task_id=task_id,
owner_token=owner_token,
)
if row is None:
raise HTTPException(status_code=404, detail="task not found")
return _row_to_response(row)

View file

@ -0,0 +1,92 @@
"""
Schedule parsing + next-run computation for LiteLLM_ScheduledTaskTable.
Ported from test_local_agent_2/task_runner.py:299-340. Three schedule kinds:
interval: '30s', '5m', '2h', '1d'
cron: standard 5-field crontab; honours schedule_tz (IANA)
once: parked at year 9999; fire_once=True flips status='fired' after
first fire so we never consult next_run_at again.
All returned datetimes are UTC, timezone-aware.
"""
from __future__ import annotations
import re
from datetime import datetime, timedelta, timezone
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
from apscheduler.triggers.cron import ( # type: ignore[import-not-found, import-untyped]
CronTrigger,
)
_INTERVAL_RE = re.compile(r"^\s*(\d+)\s*([smhd])\s*$", re.IGNORECASE)
_INTERVAL_UNITS = {
"s": "seconds",
"m": "minutes",
"h": "hours",
"d": "days",
}
_FAR_FUTURE = datetime(9999, 1, 1, tzinfo=timezone.utc)
def _parse_interval(spec: str) -> timedelta:
m = _INTERVAL_RE.match(spec or "")
if not m:
raise ValueError(f"invalid interval schedule_spec: {spec!r}")
n = int(m.group(1))
unit = _INTERVAL_UNITS[m.group(2).lower()]
return timedelta(**{unit: n})
def validate_schedule(
*,
kind: str,
spec: str,
tz: str | None,
) -> None:
"""Raise ValueError if any of (kind, spec, tz) is malformed."""
if kind == "interval":
_parse_interval(spec)
elif kind == "cron":
try:
CronTrigger.from_crontab(spec, timezone=tz or "UTC")
except Exception as e:
raise ValueError(f"invalid cron schedule_spec: {spec!r} ({e})") from e
if tz is not None:
try:
ZoneInfo(tz)
except ZoneInfoNotFoundError as e:
raise ValueError(f"invalid schedule_tz: {tz!r}") from e
elif kind == "once":
return
else:
raise ValueError(f"unknown schedule_kind: {kind!r}")
def compute_next_run(
*,
kind: str,
spec: str,
tz: str | None,
from_time: datetime,
) -> datetime:
"""
Returns UTC datetime. For 'once', returns far-future sentinel.
`from_time` should be timezone-aware; if naive, treated as UTC.
"""
if from_time.tzinfo is None:
from_time = from_time.replace(tzinfo=timezone.utc)
if kind == "interval":
return from_time + _parse_interval(spec)
if kind == "cron":
trigger = CronTrigger.from_crontab(spec, timezone=tz or "UTC")
nxt = trigger.get_next_fire_time(None, from_time)
if nxt is None:
return _FAR_FUTURE
return nxt.astimezone(timezone.utc)
if kind == "once":
return _FAR_FUTURE
raise ValueError(f"unknown schedule_kind: {kind!r}")

View file

@ -0,0 +1,246 @@
"""
Storage layer for LiteLLM_ScheduledTaskTable.
CRUD via Prisma model methods. Single approved raw-SQL exception in
claim_due() — Prisma cannot express SELECT FOR UPDATE SKIP LOCKED, which
is required for multi-pod tick safety.
"""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, List, Optional
from litellm.proxy.scheduled_tasks.schedule import compute_next_run
# Whitelist of fields PATCH can touch. Anything outside this set would let
# a buggy caller rewrite scheduling state (status, fired flags, ...).
UPDATABLE_FIELDS = frozenset(
{
"title",
"check_prompt",
"schedule_kind",
"schedule_spec",
"schedule_tz",
"next_run_at",
"expires_at",
"fire_once",
"action",
"action_args",
"format_prompt",
}
)
MAX_ACTIVE_TASKS_PER_KEY = 10
TERMINAL_STATUSES = ("fired", "expired", "cancelled")
async def count_active_for_owner(prisma_client: Any, owner_token: str) -> int:
return await prisma_client.db.litellm_scheduledtasktable.count(
where={"owner_token": owner_token, "status": "pending"},
)
async def create_task(
prisma_client: Any,
*,
owner_token: str,
user_id: Optional[str],
team_id: Optional[str],
agent_id: Optional[str],
title: str,
action: str,
action_args: Optional[Dict[str, Any]],
check_prompt: Optional[str],
format_prompt: Optional[str],
schedule_kind: str,
schedule_spec: str,
schedule_tz: Optional[str],
next_run_at: datetime,
expires_at: datetime,
fire_once: bool,
) -> Any:
return await prisma_client.db.litellm_scheduledtasktable.create(
data={
"owner_token": owner_token,
"user_id": user_id,
"team_id": team_id,
"agent_id": agent_id,
"title": title,
"action": action,
"action_args": action_args,
"check_prompt": check_prompt,
"format_prompt": format_prompt,
"schedule_kind": schedule_kind,
"schedule_spec": schedule_spec,
"schedule_tz": schedule_tz,
"next_run_at": next_run_at,
"expires_at": expires_at,
"fire_once": fire_once,
}
)
async def list_tasks_for_owner(
prisma_client: Any,
*,
owner_token: str,
include_terminal: bool,
) -> List[Any]:
where: Dict[str, Any] = {"owner_token": owner_token}
if not include_terminal:
where["status"] = "pending"
return await prisma_client.db.litellm_scheduledtasktable.find_many(
where=where,
order={"created_at": "desc"},
)
async def get_task_for_owner(
prisma_client: Any,
*,
task_id: str,
owner_token: str,
) -> Optional[Any]:
return await prisma_client.db.litellm_scheduledtasktable.find_first(
where={"task_id": task_id, "owner_token": owner_token},
)
async def update_task_for_owner(
prisma_client: Any,
*,
task_id: str,
owner_token: str,
fields: Dict[str, Any],
) -> Optional[Any]:
"""
Update one or more whitelisted fields. Only when status='pending'.
Returns updated row, or None if not found / not owned / not pending.
"""
bad = [k for k in fields if k not in UPDATABLE_FIELDS]
if bad:
raise ValueError(f"cannot update fields: {', '.join(sorted(bad))}")
if not fields:
raise ValueError("no fields to update")
existing = await prisma_client.db.litellm_scheduledtasktable.find_first(
where={
"task_id": task_id,
"owner_token": owner_token,
"status": "pending",
},
)
if existing is None:
return None
return await prisma_client.db.litellm_scheduledtasktable.update(
where={"task_id": task_id},
data=fields,
)
async def cancel_task_for_owner(
prisma_client: Any,
*,
task_id: str,
owner_token: str,
) -> Optional[Any]:
existing = await prisma_client.db.litellm_scheduledtasktable.find_first(
where={
"task_id": task_id,
"owner_token": owner_token,
"status": "pending",
},
)
if existing is None:
return None
return await prisma_client.db.litellm_scheduledtasktable.update(
where={"task_id": task_id},
data={"status": "cancelled"},
)
async def claim_due(
prisma_client: Any,
*,
agent_id: Optional[str],
actions: Optional[List[str]],
limit: int,
) -> List[Dict[str, Any]]:
"""
Atomically claim due tasks and advance their schedule.
APPROVED RAW-SQL EXCEPTION (CLAUDE.md): SELECT FOR UPDATE SKIP LOCKED
is required for multi-pod safety and is not expressible via Prisma
model methods. Per-row writes use Prisma inside the same transaction.
"""
now = datetime.now(timezone.utc)
async with prisma_client.db.tx(timeout=timedelta(seconds=30)) as tx:
rows = await tx.query_raw(
"""
SELECT task_id, schedule_kind, schedule_spec, schedule_tz,
fire_once, expires_at, next_run_at,
action, action_args, check_prompt, format_prompt, title
FROM "LiteLLM_ScheduledTaskTable"
WHERE status = 'pending'
AND next_run_at <= now()
AND ($1::text IS NULL OR agent_id = $1)
AND ($2::text[] IS NULL OR action = ANY($2))
ORDER BY next_run_at
LIMIT $3
FOR UPDATE SKIP LOCKED
""",
agent_id,
actions,
limit,
)
if not rows:
return []
async with tx.batch_() as batcher:
for r in rows:
expires_at = r["expires_at"]
if isinstance(expires_at, str):
expires_at = datetime.fromisoformat(
expires_at.replace("Z", "+00:00")
)
if expires_at.tzinfo is None:
expires_at = expires_at.replace(tzinfo=timezone.utc)
if expires_at <= now:
new_status = "expired"
new_next = r["next_run_at"]
elif r["fire_once"]:
new_status = "fired"
new_next = r["next_run_at"]
else:
new_status = "pending"
new_next = compute_next_run(
kind=r["schedule_kind"],
spec=r["schedule_spec"],
tz=r["schedule_tz"],
from_time=now,
)
batcher.litellm_scheduledtasktable.update(
where={"task_id": r["task_id"]},
data={
"status": new_status,
"next_run_at": new_next,
"last_fired_at": now,
},
)
return [
{
"task_id": r["task_id"],
"title": r["title"],
"action": r["action"],
"action_args": r["action_args"],
"check_prompt": r["check_prompt"],
"format_prompt": r["format_prompt"],
"scheduled_for": r["next_run_at"],
}
for r in rows
]

View file

@ -0,0 +1,84 @@
from datetime import datetime
from typing import Any, Dict, List, Literal, Optional
from pydantic import BaseModel, Field
ScheduleKind = Literal["interval", "cron", "once"]
TaskStatus = Literal["pending", "fired", "expired", "cancelled"]
class CreateScheduledTaskRequest(BaseModel):
title: str
action: str
action_args: Optional[Dict[str, Any]] = None
check_prompt: Optional[str] = None
format_prompt: Optional[str] = None
schedule_kind: ScheduleKind
schedule_spec: str
schedule_tz: Optional[str] = None
expires_at: datetime
fire_once: bool = True
class UpdateScheduledTaskRequest(BaseModel):
title: Optional[str] = None
check_prompt: Optional[str] = None
schedule_kind: Optional[ScheduleKind] = None
schedule_spec: Optional[str] = None
schedule_tz: Optional[str] = None
next_run_at: Optional[datetime] = None
expires_at: Optional[datetime] = None
fire_once: Optional[bool] = None
action: Optional[str] = None
action_args: Optional[Dict[str, Any]] = None
format_prompt: Optional[str] = None
class ScheduledTaskResponse(BaseModel):
task_id: str
owner_token: str
user_id: Optional[str]
team_id: Optional[str]
agent_id: Optional[str]
title: str
action: str
action_args: Optional[Dict[str, Any]]
check_prompt: Optional[str]
format_prompt: Optional[str]
schedule_kind: str
schedule_spec: str
schedule_tz: Optional[str]
next_run_at: datetime
expires_at: datetime
fire_once: bool
status: str
last_fired_at: Optional[datetime]
created_at: datetime
updated_at: datetime
class DueTaskResponse(BaseModel):
"""Trimmed shape returned by /due — only fields the agent needs to dispatch."""
task_id: str
title: str
action: str
action_args: Optional[Dict[str, Any]]
check_prompt: Optional[str]
format_prompt: Optional[str]
scheduled_for: datetime
class ListScheduledTasksResponse(BaseModel):
tasks: List[ScheduledTaskResponse]
class DueTasksResponse(BaseModel):
tasks: List[DueTaskResponse] = Field(default_factory=list)

View file

@ -410,6 +410,7 @@ model LiteLLM_VerificationToken {
litellm_project_table LiteLLM_ProjectTable? @relation(fields: [project_id], references: [project_id])
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
jwt_key_mappings LiteLLM_JWTKeyMapping[]
scheduled_tasks LiteLLM_ScheduledTaskTable[]
// SELECT COUNT(*) FROM (SELECT "public"."LiteLLM_VerificationToken"."token" FROM "public"."LiteLLM_VerificationToken" WHERE ("public"."LiteLLM_VerificationToken"."user_id" = $1 AND ("public"."LiteLLM_VerificationToken"."team_id" IS NULL OR "public"."LiteLLM_VerificationToken"."team_id" <> $2)) OFFSET $3 ) AS "sub"
// SELECT ... FROM "public"."LiteLLM_VerificationToken" WHERE "public"."LiteLLM_VerificationToken"."user_id" = $1 OFFSET $2
@ -1290,3 +1291,38 @@ model LiteLLM_AdaptiveRouterSession {
@@id([session_id, router_name, model_name])
@@index([last_activity_at], map: "idx_adaptive_router_session_activity")
}
model LiteLLM_ScheduledTaskTable {
task_id String @id @default(uuid())
owner_token String
owner_key LiteLLM_VerificationToken @relation(fields: [owner_token], references: [token], onDelete: Cascade)
user_id String?
team_id String?
agent_id String?
title String
action String
action_args Json?
check_prompt String?
format_prompt String?
schedule_kind String
schedule_spec String
schedule_tz String?
next_run_at DateTime
expires_at DateTime
fire_once Boolean @default(true)
status String @default("pending")
last_fired_at DateTime?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@index([owner_token, status])
@@index([user_id])
@@index([team_id])
@@index([agent_id, status])
}

View file

@ -410,6 +410,7 @@ model LiteLLM_VerificationToken {
litellm_project_table LiteLLM_ProjectTable? @relation(fields: [project_id], references: [project_id])
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
jwt_key_mappings LiteLLM_JWTKeyMapping[]
scheduled_tasks LiteLLM_ScheduledTaskTable[]
// SELECT COUNT(*) FROM (SELECT "public"."LiteLLM_VerificationToken"."token" FROM "public"."LiteLLM_VerificationToken" WHERE ("public"."LiteLLM_VerificationToken"."user_id" = $1 AND ("public"."LiteLLM_VerificationToken"."team_id" IS NULL OR "public"."LiteLLM_VerificationToken"."team_id" <> $2)) OFFSET $3 ) AS "sub"
// SELECT ... FROM "public"."LiteLLM_VerificationToken" WHERE "public"."LiteLLM_VerificationToken"."user_id" = $1 OFFSET $2
@ -1290,3 +1291,38 @@ model LiteLLM_AdaptiveRouterSession {
@@id([session_id, router_name, model_name])
@@index([last_activity_at], map: "idx_adaptive_router_session_activity")
}
model LiteLLM_ScheduledTaskTable {
task_id String @id @default(uuid())
owner_token String
owner_key LiteLLM_VerificationToken @relation(fields: [owner_token], references: [token], onDelete: Cascade)
user_id String?
team_id String?
agent_id String?
title String
action String
action_args Json?
check_prompt String?
format_prompt String?
schedule_kind String
schedule_spec String
schedule_tz String?
next_run_at DateTime
expires_at DateTime
fire_once Boolean @default(true)
status String @default("pending")
last_fired_at DateTime?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@index([owner_token, status])
@@index([user_id])
@@index([team_id])
@@index([agent_id, status])
}

View file

@ -0,0 +1,542 @@
"""
Endpoint tests for /v1/tasks. FastAPI TestClient + in-memory fake of the
prisma scheduled-tasks table. The /due test path uses a fake of query_raw +
tx() that exercises the same Python-side schedule advance logic.
"""
import os
import sys
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, List, Optional
from unittest.mock import MagicMock, patch
from fastapi import FastAPI
from fastapi.testclient import TestClient
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy._types import UserAPIKeyAuth, hash_token
from litellm.proxy.scheduled_tasks.endpoints import router
# UserAPIKeyAuth's model_validator rewrites api_key + token to the hashed
# form when api_key starts with "sk-". The endpoint reads `.token` for the
# FK target, so we derive the same hash here for assertions/seeding.
_TEST_API_KEY = "sk-test"
_TEST_OWNER_TOKEN = hash_token(_TEST_API_KEY)
def _now() -> datetime:
return datetime.now(timezone.utc)
def _make_row(**kwargs) -> MagicMock:
"""Build a Prisma-row-like object."""
now = _now()
defaults: Dict[str, Any] = {
"task_id": "task-1",
"owner_token": _TEST_OWNER_TOKEN,
"user_id": "user-a",
"team_id": "team-a",
"agent_id": "agent-a",
"title": "default",
"action": "check",
"action_args": None,
"check_prompt": "is it done?",
"format_prompt": None,
"schedule_kind": "interval",
"schedule_spec": "5m",
"schedule_tz": None,
"next_run_at": now + timedelta(minutes=5),
"expires_at": now + timedelta(days=1),
"fire_once": True,
"status": "pending",
"last_fired_at": None,
"created_at": now,
"updated_at": now,
}
defaults.update(kwargs)
row = MagicMock()
for k, v in defaults.items():
setattr(row, k, v)
return row
class _FakeScheduledTaskTable:
"""In-memory fake of prisma_client.db.litellm_scheduledtasktable."""
def __init__(self):
self.rows: List[MagicMock] = []
self._counter = 0
def _matches(self, row: MagicMock, where: Dict[str, Any]) -> bool:
for k, v in where.items():
if getattr(row, k, None) != v:
return False
return True
def _filter(self, where: Optional[Dict[str, Any]]) -> List[MagicMock]:
if not where:
return list(self.rows)
return [r for r in self.rows if self._matches(r, where)]
async def create(self, data: Dict[str, Any]) -> MagicMock:
self._counter += 1
row = _make_row(task_id=f"task-{self._counter}", **data)
self.rows.append(row)
return row
async def count(self, where: Optional[Dict[str, Any]] = None) -> int:
return len(self._filter(where))
async def find_first(
self, where: Optional[Dict[str, Any]] = None
) -> Optional[MagicMock]:
rows = self._filter(where)
return rows[0] if rows else None
async def find_many(
self,
where: Optional[Dict[str, Any]] = None,
order: Optional[Dict[str, str]] = None,
) -> List[MagicMock]:
_ = order
return self._filter(where)
async def update(self, where: Dict[str, Any], data: Dict[str, Any]) -> MagicMock:
for r in self.rows:
if r.task_id == where["task_id"]:
for k, v in data.items():
setattr(r, k, v)
return r
raise Exception("Not found")
class _FakeBatcher:
"""Capture batched updates and flush them on context exit."""
def __init__(self, table: _FakeScheduledTaskTable):
self._table = table
self._ops: List = []
@property
def litellm_scheduledtasktable(self):
outer = self
class _Proxy:
def update(self, where, data):
outer._ops.append((where, data))
return _Proxy()
async def __aenter__(self):
return self
async def __aexit__(self, *exc):
for where, data in self._ops:
await self._table.update(where=where, data=data)
class _FakeTx:
"""Tx context — exposes query_raw + batch_ + the table."""
def __init__(self, table: _FakeScheduledTaskTable):
self._table = table
self.litellm_scheduledtasktable = table
async def query_raw(self, query: str, *args) -> List[Dict[str, Any]]:
# Model the production claim query: rows where status='pending',
# next_run_at <= now, optional agent_id and actions filters, ordered
# by next_run_at, limited.
agent_id, actions, limit = args
now = _now()
out: List[Dict[str, Any]] = []
for r in sorted(self._table.rows, key=lambda x: x.next_run_at):
if r.status != "pending":
continue
if r.next_run_at > now:
continue
if agent_id is not None and r.agent_id != agent_id:
continue
if actions is not None and r.action not in actions:
continue
out.append(
{
"task_id": r.task_id,
"schedule_kind": r.schedule_kind,
"schedule_spec": r.schedule_spec,
"schedule_tz": r.schedule_tz,
"fire_once": r.fire_once,
"expires_at": r.expires_at,
"next_run_at": r.next_run_at,
"action": r.action,
"action_args": r.action_args,
"check_prompt": r.check_prompt,
"format_prompt": r.format_prompt,
"title": r.title,
}
)
if len(out) >= limit:
break
return out
def batch_(self) -> _FakeBatcher:
return _FakeBatcher(self._table)
class _FakeTxFactory:
def __init__(self, table: _FakeScheduledTaskTable):
self._table = table
def __call__(self, *args, **kwargs):
return self
async def __aenter__(self):
return _FakeTx(self._table)
async def __aexit__(self, *exc):
return None
def _make_prisma() -> MagicMock:
client = MagicMock()
client.db = MagicMock()
table = _FakeScheduledTaskTable()
client.db.litellm_scheduledtasktable = table
client.db.tx = _FakeTxFactory(table)
return client
def _make_app(prisma: Any) -> TestClient:
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
def _auth_override():
return UserAPIKeyAuth(
api_key=_TEST_API_KEY,
user_id="user-a",
team_id="team-a",
agent_id="agent-a",
)
app = FastAPI()
app.include_router(router)
app.dependency_overrides[user_api_key_auth] = _auth_override
return TestClient(app, raise_server_exceptions=True)
def _patch_prisma(prisma):
return patch(
"litellm.proxy.scheduled_tasks.endpoints._get_prisma_client",
return_value=prisma,
)
def _create_payload(**overrides) -> Dict[str, Any]:
payload: Dict[str, Any] = {
"title": "watch PR 123",
"action": "check",
"check_prompt": "is PR 123 merged?",
"schedule_kind": "interval",
"schedule_spec": "5m",
"schedule_tz": None,
"expires_at": (_now() + timedelta(days=1)).isoformat(),
"fire_once": True,
}
payload.update(overrides)
return payload
class TestCreate:
def setup_method(self):
self.prisma = _make_prisma()
self.client = _make_app(self.prisma)
def test_happy_path_stamps_identity_from_auth(self):
with _patch_prisma(self.prisma):
r = self.client.post("/v1/tasks", json=_create_payload())
assert r.status_code == 200, r.text
body = r.json()
assert body["user_id"] == "user-a"
assert body["team_id"] == "team-a"
assert body["agent_id"] == "agent-a"
# owner_token must match the auth-resolved hashed token (FK to
# LiteLLM_VerificationToken.token), not the raw "sk-..." form.
assert body["owner_token"]
assert body["owner_token"] != "sk-test"
assert body["status"] == "pending"
def test_check_action_requires_check_prompt(self):
payload = _create_payload(action="check", check_prompt=None)
with _patch_prisma(self.prisma):
r = self.client.post("/v1/tasks", json=payload)
assert r.status_code == 400
assert "check_prompt" in r.json()["detail"]
def test_action_other_than_check_does_not_require_check_prompt(self):
payload = _create_payload(
action="pr_digest",
check_prompt=None,
action_args={"channel_id": "C456"},
format_prompt="bullet list",
schedule_kind="cron",
schedule_spec="0 9 * * 1-5",
schedule_tz="America/Los_Angeles",
fire_once=False,
)
with _patch_prisma(self.prisma):
r = self.client.post("/v1/tasks", json=payload)
assert r.status_code == 200, r.text
def test_invalid_cron_rejected(self):
payload = _create_payload(
schedule_kind="cron",
schedule_spec="not a cron",
)
with _patch_prisma(self.prisma):
r = self.client.post("/v1/tasks", json=payload)
assert r.status_code == 400
assert "cron" in r.json()["detail"].lower()
def test_invalid_tz_rejected(self):
payload = _create_payload(
action="pr_digest",
check_prompt=None,
schedule_kind="cron",
schedule_spec="0 9 * * *",
schedule_tz="Europe/Atlantis",
)
with _patch_prisma(self.prisma):
r = self.client.post("/v1/tasks", json=payload)
assert r.status_code == 400
def test_eleventh_task_rejected(self):
with _patch_prisma(self.prisma):
for i in range(10):
r = self.client.post(
"/v1/tasks",
json=_create_payload(title=f"t{i}"),
)
assert r.status_code == 200, r.text
r = self.client.post("/v1/tasks", json=_create_payload(title="overflow"))
assert r.status_code == 429
def test_expires_in_past_rejected(self):
payload = _create_payload(
expires_at=(_now() - timedelta(seconds=1)).isoformat(),
)
with _patch_prisma(self.prisma):
r = self.client.post("/v1/tasks", json=payload)
assert r.status_code == 400
class TestListAndGet:
def setup_method(self):
self.prisma = _make_prisma()
self.client = _make_app(self.prisma)
def test_list_excludes_terminal_by_default(self):
with _patch_prisma(self.prisma):
r1 = self.client.post("/v1/tasks", json=_create_payload(title="active"))
assert r1.status_code == 200
# Manually flip a row to 'fired' to simulate post-fire state
self.prisma.db.litellm_scheduledtasktable.rows[0].status = "pending"
self.prisma.db.litellm_scheduledtasktable.rows[0] = _make_row(
task_id="task-1", title="active", status="pending"
)
self.prisma.db.litellm_scheduledtasktable.rows.append(
_make_row(task_id="task-2", title="done", status="fired")
)
r = self.client.get("/v1/tasks")
assert r.status_code == 200
titles = [t["title"] for t in r.json()["tasks"]]
assert "done" not in titles
assert "active" in titles
def test_list_with_terminal_includes_all(self):
with _patch_prisma(self.prisma):
self.prisma.db.litellm_scheduledtasktable.rows.append(
_make_row(task_id="task-2", title="done", status="fired")
)
self.prisma.db.litellm_scheduledtasktable.rows.append(
_make_row(task_id="task-3", title="active", status="pending")
)
r = self.client.get("/v1/tasks?include_terminal=true")
assert r.status_code == 200
titles = [t["title"] for t in r.json()["tasks"]]
assert "done" in titles
assert "active" in titles
def test_get_owned(self):
with _patch_prisma(self.prisma):
self.client.post("/v1/tasks", json=_create_payload(title="mine"))
tid = self.prisma.db.litellm_scheduledtasktable.rows[0].task_id
r = self.client.get(f"/v1/tasks/{tid}")
assert r.status_code == 200
assert r.json()["title"] == "mine"
def test_get_foreign_404(self):
# Insert a row owned by another token.
self.prisma.db.litellm_scheduledtasktable.rows.append(
_make_row(task_id="foreign-1", owner_token="sk-other")
)
with _patch_prisma(self.prisma):
r = self.client.get("/v1/tasks/foreign-1")
assert r.status_code == 404
class TestUpdate:
def setup_method(self):
self.prisma = _make_prisma()
self.client = _make_app(self.prisma)
def test_update_pending(self):
with _patch_prisma(self.prisma):
self.client.post("/v1/tasks", json=_create_payload(title="orig"))
tid = self.prisma.db.litellm_scheduledtasktable.rows[0].task_id
r = self.client.patch(f"/v1/tasks/{tid}", json={"title": "renamed"})
assert r.status_code == 200
assert r.json()["title"] == "renamed"
def test_update_invalid_field_rejected(self):
# Pydantic strips unknown fields during model construction. The
# whitelist guard inside store still bounces explicit attempts to
# touch status/owner_token if they ever leak through.
with _patch_prisma(self.prisma):
self.client.post("/v1/tasks", json=_create_payload())
tid = self.prisma.db.litellm_scheduledtasktable.rows[0].task_id
# Empty body — no updatable fields → 400
r = self.client.patch(f"/v1/tasks/{tid}", json={})
assert r.status_code == 400
def test_update_terminal_rejected(self):
self.prisma.db.litellm_scheduledtasktable.rows.append(
_make_row(task_id="t-fired", status="fired")
)
with _patch_prisma(self.prisma):
r = self.client.patch("/v1/tasks/t-fired", json={"title": "nope"})
assert r.status_code == 400
def test_update_foreign_404(self):
self.prisma.db.litellm_scheduledtasktable.rows.append(
_make_row(task_id="foreign-1", owner_token="sk-other")
)
with _patch_prisma(self.prisma):
r = self.client.patch("/v1/tasks/foreign-1", json={"title": "hijack"})
assert r.status_code == 404
def test_update_invalid_schedule_rejected(self):
with _patch_prisma(self.prisma):
self.client.post("/v1/tasks", json=_create_payload())
tid = self.prisma.db.litellm_scheduledtasktable.rows[0].task_id
r = self.client.patch(
f"/v1/tasks/{tid}",
json={"schedule_kind": "interval", "schedule_spec": "banana"},
)
assert r.status_code == 400
class TestCancel:
def setup_method(self):
self.prisma = _make_prisma()
self.client = _make_app(self.prisma)
def test_cancel_owned(self):
with _patch_prisma(self.prisma):
self.client.post("/v1/tasks", json=_create_payload())
tid = self.prisma.db.litellm_scheduledtasktable.rows[0].task_id
r = self.client.delete(f"/v1/tasks/{tid}")
assert r.status_code == 200
assert r.json()["status"] == "cancelled"
def test_cancel_foreign_404(self):
self.prisma.db.litellm_scheduledtasktable.rows.append(
_make_row(task_id="foreign-1", owner_token="sk-other")
)
with _patch_prisma(self.prisma):
r = self.client.delete("/v1/tasks/foreign-1")
assert r.status_code == 404
class TestDue:
def setup_method(self):
self.prisma = _make_prisma()
self.client = _make_app(self.prisma)
def _seed_due_row(self, **kwargs):
defaults: Dict[str, Any] = {
"task_id": f"due-{len(self.prisma.db.litellm_scheduledtasktable.rows) + 1}",
"next_run_at": _now() - timedelta(seconds=10),
"status": "pending",
"agent_id": "agent-a",
}
defaults.update(kwargs)
self.prisma.db.litellm_scheduledtasktable.rows.append(_make_row(**defaults))
def test_due_claims_pending_rows(self):
self._seed_due_row(action="check")
self._seed_due_row(action="pr_digest", fire_once=False, schedule_spec="5m")
with _patch_prisma(self.prisma):
r = self.client.get("/v1/tasks/due")
assert r.status_code == 200
body = r.json()
assert len(body["tasks"]) == 2
def test_due_advances_recurring_and_fires_once(self):
self._seed_due_row(action="check", fire_once=True)
self._seed_due_row(action="pr_digest", fire_once=False, schedule_spec="5m")
with _patch_prisma(self.prisma):
self.client.get("/v1/tasks/due")
rows = self.prisma.db.litellm_scheduledtasktable.rows
once_row = next(r for r in rows if r.fire_once)
recurring_row = next(r for r in rows if not r.fire_once)
assert once_row.status == "fired"
assert recurring_row.status == "pending"
assert recurring_row.next_run_at > _now()
def test_due_second_call_empty(self):
self._seed_due_row(action="check", fire_once=True)
with _patch_prisma(self.prisma):
r1 = self.client.get("/v1/tasks/due")
assert len(r1.json()["tasks"]) == 1
r2 = self.client.get("/v1/tasks/due")
assert len(r2.json()["tasks"]) == 0
def test_due_skips_other_agents(self):
self._seed_due_row(action="check", agent_id="agent-other")
self._seed_due_row(action="check", agent_id="agent-a")
with _patch_prisma(self.prisma):
r = self.client.get("/v1/tasks/due")
assert len(r.json()["tasks"]) == 1
assert r.json()["tasks"][0]["task_id"] != "due-1" or (
self.prisma.db.litellm_scheduledtasktable.rows[0].agent_id == "agent-a"
)
def test_due_actions_filter(self):
self._seed_due_row(action="check")
self._seed_due_row(action="pr_digest")
self._seed_due_row(action="other_thing")
with _patch_prisma(self.prisma):
r = self.client.get("/v1/tasks/due?actions=check,pr_digest")
assert r.status_code == 200
actions = sorted(t["action"] for t in r.json()["tasks"])
assert actions == ["check", "pr_digest"]
def test_due_expires_past_flips_to_expired(self):
self._seed_due_row(
action="check",
fire_once=False,
schedule_spec="5m",
expires_at=_now() - timedelta(seconds=1),
)
with _patch_prisma(self.prisma):
self.client.get("/v1/tasks/due")
row = self.prisma.db.litellm_scheduledtasktable.rows[0]
assert row.status == "expired"
def test_due_skips_future_rows(self):
self._seed_due_row(
action="check",
next_run_at=_now() + timedelta(minutes=5),
)
with _patch_prisma(self.prisma):
r = self.client.get("/v1/tasks/due")
assert len(r.json()["tasks"]) == 0

View file

@ -0,0 +1,82 @@
"""
Pure-function tests for schedule parsing and next-run computation.
No DB, no FastAPI — just exercise the parsing branches.
"""
from datetime import datetime, timedelta, timezone
import pytest
from litellm.proxy.scheduled_tasks.schedule import (
_FAR_FUTURE,
compute_next_run,
validate_schedule,
)
class TestValidateSchedule:
def test_interval_ok(self):
validate_schedule(kind="interval", spec="30s", tz=None)
validate_schedule(kind="interval", spec="5m", tz=None)
validate_schedule(kind="interval", spec="2h", tz=None)
validate_schedule(kind="interval", spec="1d", tz=None)
def test_interval_bad_spec(self):
with pytest.raises(ValueError, match="invalid interval"):
validate_schedule(kind="interval", spec="banana", tz=None)
with pytest.raises(ValueError, match="invalid interval"):
validate_schedule(kind="interval", spec="5x", tz=None)
def test_cron_ok(self):
validate_schedule(kind="cron", spec="0 9 * * 1-5", tz="America/Los_Angeles")
validate_schedule(kind="cron", spec="*/15 * * * *", tz=None)
def test_cron_bad_spec(self):
with pytest.raises(ValueError, match="invalid cron"):
validate_schedule(kind="cron", spec="not a cron", tz=None)
def test_cron_bad_tz(self):
with pytest.raises(ValueError, match="invalid"):
validate_schedule(kind="cron", spec="0 9 * * *", tz="Europe/Atlantis")
def test_once_ok(self):
validate_schedule(kind="once", spec="ignored", tz=None)
def test_unknown_kind(self):
with pytest.raises(ValueError, match="unknown schedule_kind"):
validate_schedule(kind="weekly", spec="x", tz=None)
class TestComputeNextRun:
def test_interval_advances(self):
now = datetime(2026, 4, 29, 12, 0, 0, tzinfo=timezone.utc)
nxt = compute_next_run(kind="interval", spec="5m", tz=None, from_time=now)
assert nxt == now + timedelta(minutes=5)
def test_interval_naive_input_treated_utc(self):
naive = datetime(2026, 4, 29, 12, 0, 0)
nxt = compute_next_run(kind="interval", spec="1h", tz=None, from_time=naive)
assert nxt.tzinfo is not None
assert nxt.utcoffset() == timedelta(0)
def test_cron_returns_utc(self):
now = datetime(2026, 4, 29, 0, 0, 0, tzinfo=timezone.utc)
nxt = compute_next_run(
kind="cron",
spec="0 9 * * *",
tz="America/Los_Angeles",
from_time=now,
)
assert nxt.tzinfo is not None
assert nxt.utcoffset() == timedelta(0)
assert nxt > now
def test_once_returns_far_future(self):
now = datetime(2026, 4, 29, 12, 0, 0, tzinfo=timezone.utc)
nxt = compute_next_run(kind="once", spec="x", tz=None, from_time=now)
assert nxt == _FAR_FUTURE
def test_unknown_kind_raises(self):
now = datetime(2026, 4, 29, 12, 0, 0, tzinfo=timezone.utc)
with pytest.raises(ValueError):
compute_next_run(kind="weekly", spec="x", tz=None, from_time=now)