diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260429000000_add_scheduled_tasks/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260429000000_add_scheduled_tasks/migration.sql new file mode 100644 index 00000000000..919ae5ad431 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260429000000_add_scheduled_tasks/migration.sql @@ -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); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 8f07c5afa3f..fc6cb6b04d9 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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]) +} diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 8f676df04cd..2ca071ba3c8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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) diff --git a/litellm/proxy/scheduled_tasks/__init__.py b/litellm/proxy/scheduled_tasks/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/proxy/scheduled_tasks/endpoints.py b/litellm/proxy/scheduled_tasks/endpoints.py new file mode 100644 index 00000000000..fe0506cd1cd --- /dev/null +++ b/litellm/proxy/scheduled_tasks/endpoints.py @@ -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) diff --git a/litellm/proxy/scheduled_tasks/schedule.py b/litellm/proxy/scheduled_tasks/schedule.py new file mode 100644 index 00000000000..c66cdd7b0b9 --- /dev/null +++ b/litellm/proxy/scheduled_tasks/schedule.py @@ -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}") diff --git a/litellm/proxy/scheduled_tasks/store.py b/litellm/proxy/scheduled_tasks/store.py new file mode 100644 index 00000000000..f88b239ac3c --- /dev/null +++ b/litellm/proxy/scheduled_tasks/store.py @@ -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 + ] diff --git a/litellm/proxy/scheduled_tasks/types.py b/litellm/proxy/scheduled_tasks/types.py new file mode 100644 index 00000000000..dcb60e2641e --- /dev/null +++ b/litellm/proxy/scheduled_tasks/types.py @@ -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) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 8f07c5afa3f..fc6cb6b04d9 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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]) +} diff --git a/schema.prisma b/schema.prisma index 8f07c5afa3f..fc6cb6b04d9 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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]) +} diff --git a/tests/test_litellm/proxy/scheduled_tasks/__init__.py b/tests/test_litellm/proxy/scheduled_tasks/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/scheduled_tasks/test_endpoints.py b/tests/test_litellm/proxy/scheduled_tasks/test_endpoints.py new file mode 100644 index 00000000000..d3a735fda17 --- /dev/null +++ b/tests/test_litellm/proxy/scheduled_tasks/test_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/scheduled_tasks/test_schedule.py b/tests/test_litellm/proxy/scheduled_tasks/test_schedule.py new file mode 100644 index 00000000000..c0083ea0327 --- /dev/null +++ b/tests/test_litellm/proxy/scheduled_tasks/test_schedule.py @@ -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)