mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
3e1479c052
commit
e5dec5e809
13 changed files with 1542 additions and 0 deletions
|
|
@ -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);
|
||||
|
|
@ -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])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
0
litellm/proxy/scheduled_tasks/__init__.py
Normal file
0
litellm/proxy/scheduled_tasks/__init__.py
Normal file
317
litellm/proxy/scheduled_tasks/endpoints.py
Normal file
317
litellm/proxy/scheduled_tasks/endpoints.py
Normal 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)
|
||||
92
litellm/proxy/scheduled_tasks/schedule.py
Normal file
92
litellm/proxy/scheduled_tasks/schedule.py
Normal 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}")
|
||||
246
litellm/proxy/scheduled_tasks/store.py
Normal file
246
litellm/proxy/scheduled_tasks/store.py
Normal 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
|
||||
]
|
||||
84
litellm/proxy/scheduled_tasks/types.py
Normal file
84
litellm/proxy/scheduled_tasks/types.py
Normal 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)
|
||||
|
|
@ -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])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
}
|
||||
|
|
|
|||
0
tests/test_litellm/proxy/scheduled_tasks/__init__.py
Normal file
0
tests/test_litellm/proxy/scheduled_tasks/__init__.py
Normal file
542
tests/test_litellm/proxy/scheduled_tasks/test_endpoints.py
Normal file
542
tests/test_litellm/proxy/scheduled_tasks/test_endpoints.py
Normal 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
|
||||
82
tests/test_litellm/proxy/scheduled_tasks/test_schedule.py
Normal file
82
tests/test_litellm/proxy/scheduled_tasks/test_schedule.py
Normal 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)
|
||||
Loading…
Add table
Reference in a new issue