address greptile review feedback (greploop iteration 1)

- Add feature flag guard (LITELLM_SCHEDULED_TASKS_ENABLED) to all endpoints
- Validate expires_at in PATCH path to reject past timestamps
- Fix TOCTOU race in cancel/update by using update_many with status guard
- Default fire_once to False for cron/interval, True for once
- Fix agent_id=NULL bypass in claim_due SQL to prevent intra-owner leakage

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Krrish Dholakia 2026-05-01 12:25:45 -07:00
parent c6f21d0d84
commit 933f38362c
3 changed files with 53 additions and 21 deletions

View file

@ -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.

View file

@ -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

View file

@ -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):