diff --git a/litellm/proxy/scheduled_tasks/endpoints.py b/litellm/proxy/scheduled_tasks/endpoints.py index 35a77597d88..1c6c0af1268 100644 --- a/litellm/proxy/scheduled_tasks/endpoints.py +++ b/litellm/proxy/scheduled_tasks/endpoints.py @@ -8,6 +8,7 @@ accepted from the client. from __future__ import annotations +import os from datetime import datetime, timezone from typing import List, Optional @@ -33,6 +34,13 @@ from litellm.proxy.scheduled_tasks.types import ( router = APIRouter() +def _require_feature_enabled() -> None: + if os.environ.get("LITELLM_SCHEDULED_TASKS_ENABLED", "false").lower() != "true": + raise HTTPException( + status_code=404, detail="scheduled tasks feature is not enabled" + ) + + def _require_token(user_api_key_dict: UserAPIKeyAuth) -> str: """ Resolve the FK target for owner_token. We deliberately prefer `.token` @@ -90,6 +98,7 @@ def _get_prisma_client(): async def create_scheduled_task( data: CreateScheduledTaskRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + _feature: None = Depends(_require_feature_enabled), ): if data.action == "check" and not data.check_prompt: raise HTTPException( @@ -122,6 +131,8 @@ async def create_scheduled_task( if data.expires_at <= now: raise HTTPException(status_code=400, detail="expires_at must be in the future") + fire_once = data.fire_once if data.fire_once is not None else (data.schedule_kind == "once") + next_run_at = compute_next_run( kind=data.schedule_kind, spec=data.schedule_spec, @@ -146,7 +157,7 @@ async def create_scheduled_task( schedule_tz=data.schedule_tz, next_run_at=next_run_at, expires_at=data.expires_at, - fire_once=data.fire_once, + fire_once=fire_once, ) return _row_to_response(row) @@ -159,6 +170,7 @@ async def create_scheduled_task( async def list_scheduled_tasks( include_terminal: bool = Query(False), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + _feature: None = Depends(_require_feature_enabled), ): owner_token = _require_token(user_api_key_dict) prisma_client = _get_prisma_client() @@ -182,6 +194,7 @@ async def get_due_tasks( ), limit: int = Query(20, ge=1, le=100), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + _feature: None = Depends(_require_feature_enabled), ): """ Atomic claim of due tasks. Returns immediately, possibly with empty list. @@ -229,6 +242,7 @@ async def get_due_tasks( async def get_scheduled_task( task_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + _feature: None = Depends(_require_feature_enabled), ): owner_token = _require_token(user_api_key_dict) prisma_client = _get_prisma_client() @@ -251,6 +265,7 @@ async def update_scheduled_task( task_id: str, data: UpdateScheduledTaskRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + _feature: None = Depends(_require_feature_enabled), ): owner_token = _require_token(user_api_key_dict) prisma_client = _get_prisma_client() @@ -289,6 +304,16 @@ async def update_scheduled_task( detail="check_prompt is required when action='check'", ) + if "expires_at" in fields: + now = datetime.now(timezone.utc) + new_expires = fields["expires_at"] + if new_expires.tzinfo is None: + new_expires = new_expires.replace(tzinfo=timezone.utc) + if new_expires <= now: + raise HTTPException( + status_code=400, detail="expires_at must be in the future" + ) + try: row = await store.update_task_for_owner( prisma_client, @@ -311,6 +336,7 @@ async def update_scheduled_task( async def cancel_scheduled_task( task_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + _feature: None = Depends(_require_feature_enabled), ): owner_token = _require_token(user_api_key_dict) prisma_client = _get_prisma_client() @@ -333,6 +359,7 @@ async def report_task_result( task_id: str, data: ReportTaskResultRequest, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + _feature: None = Depends(_require_feature_enabled), ): """ Agent reports outcome of one dispatch attempt. diff --git a/litellm/proxy/scheduled_tasks/store.py b/litellm/proxy/scheduled_tasks/store.py index a652ff4923d..e2078765f4c 100644 --- a/litellm/proxy/scheduled_tasks/store.py +++ b/litellm/proxy/scheduled_tasks/store.py @@ -209,16 +209,6 @@ async def update_task_for_owner( 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 - # Mirror the create path: Json columns must be json-encoded for # prisma-client-python. Drop the key entirely when caller sent null; # passing None to a Json? column is rejected. @@ -234,10 +224,19 @@ async def update_task_for_owner( raise ValueError("no fields to update") fields = cleaned - return await prisma_client.db.litellm_scheduledtasktable.update( - where={"task_id": task_id}, + affected = await prisma_client.db.litellm_scheduledtasktable.update_many( + where={ + "task_id": task_id, + "owner_token": owner_token, + "status": "pending", + }, data=fields, ) + if affected == 0: + return None + return await prisma_client.db.litellm_scheduledtasktable.find_first( + where={"task_id": task_id, "owner_token": owner_token}, + ) async def cancel_task_for_owner( @@ -246,19 +245,19 @@ async def cancel_task_for_owner( task_id: str, owner_token: str, ) -> Optional[Any]: - existing = await prisma_client.db.litellm_scheduledtasktable.find_first( + affected = await prisma_client.db.litellm_scheduledtasktable.update_many( 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"}, ) + if affected == 0: + return None + return await prisma_client.db.litellm_scheduledtasktable.find_first( + where={"task_id": task_id, "owner_token": owner_token}, + ) async def claim_due( @@ -297,7 +296,7 @@ async def claim_due( WHERE status = 'pending' AND next_run_at <= now() AND owner_token = $1 - AND ($2::text IS NULL OR agent_id = $2) + AND (($2::text IS NULL AND agent_id IS NULL) OR agent_id = $2) AND ($3::text[] IS NULL OR action = ANY($3)) ORDER BY next_run_at LIMIT $4 diff --git a/litellm/proxy/scheduled_tasks/types.py b/litellm/proxy/scheduled_tasks/types.py index 0209a8749f1..ff6973c87f0 100644 --- a/litellm/proxy/scheduled_tasks/types.py +++ b/litellm/proxy/scheduled_tasks/types.py @@ -21,7 +21,13 @@ class CreateScheduledTaskRequest(BaseModel): schedule_tz: Optional[str] = None expires_at: datetime - fire_once: bool = True + fire_once: Optional[bool] = Field( + default=None, + description=( + "Defaults to True for kind='once', False for 'interval'/'cron'. " + "Set explicitly to override." + ), + ) class UpdateScheduledTaskRequest(BaseModel):