mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
add error tracking + lazy expiry to scheduled tasks
Three additions to match the local task-scheduler semantics that the
proxy was missing:
1. consecutive_errors + last_error columns on LiteLLM_ScheduledTaskTable.
New POST /v1/tasks/{task_id}/report endpoint:
{"result": "success" | "error", "reason": "..."}
success → resets counter, clears last_error.
error → increments counter. On the Nth consecutive error
(store.MAX_CONSECUTIVE_ERRORS = 3), status flips to 'failed'
and /due stops emitting the task. Caller's next list() shows
it as failed with last_error populated; agent renders its
own user-facing notification — proxy stays out of the
notification channel.
2. Lazy expiry sweep on every list/get. Previously a task that expired
before its next fire window sat in 'pending' indefinitely (claim_due
only flips status when the row is also otherwise due). Added
sweep_expired_for_owner() called from list_tasks_for_owner and
get_task_for_owner — flips pending rows past expires_at to 'expired'
before any read returns them. Cheap update_many gated on the
(owner_token, status, expires_at) index.
3. Status CHECK constraint widened to include 'failed'.
Tests:
- 6 new tests in TestReport / TestLazyExpiry
- Existing 39 still pass
- /Users/krrishdholakia/Documents/temp_py_folder/test_tasks.py smoke
script extended with sections 4b (report success/error/failure flip),
4c (10-task cap), and 6 (lazy expiry)
The per-key 10-task cap was already enforced server-side via
store.MAX_ACTIVE_TASKS_PER_KEY → 429; smoke test now exercises it.
Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
cda2a9cc80
commit
8587f3b770
8 changed files with 241 additions and 3 deletions
|
|
@ -19,6 +19,8 @@ CREATE TABLE "LiteLLM_ScheduledTaskTable" (
|
|||
"fire_once" BOOLEAN NOT NULL DEFAULT true,
|
||||
"status" TEXT NOT NULL DEFAULT 'pending',
|
||||
"last_fired_at" TIMESTAMP(3),
|
||||
"consecutive_errors" INTEGER NOT NULL DEFAULT 0,
|
||||
"last_error" TEXT,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"updated_at" TIMESTAMP(3) NOT NULL,
|
||||
|
||||
|
|
@ -54,7 +56,7 @@ ALTER TABLE "LiteLLM_ScheduledTaskTable"
|
|||
CHECK (schedule_kind IN ('interval','cron','once'));
|
||||
ALTER TABLE "LiteLLM_ScheduledTaskTable"
|
||||
ADD CONSTRAINT "LiteLLM_ScheduledTaskTable_status_check"
|
||||
CHECK (status IN ('pending','fired','expired','cancelled'));
|
||||
CHECK (status IN ('pending','fired','expired','cancelled','failed'));
|
||||
ALTER TABLE "LiteLLM_ScheduledTaskTable"
|
||||
ADD CONSTRAINT "LiteLLM_ScheduledTaskTable_action_prompt_check"
|
||||
CHECK (action <> 'check' OR check_prompt IS NOT NULL);
|
||||
|
|
|
|||
|
|
@ -1321,6 +1321,8 @@ model LiteLLM_ScheduledTaskTable {
|
|||
|
||||
status String @default("pending")
|
||||
last_fired_at DateTime?
|
||||
consecutive_errors Int @default(0)
|
||||
last_error String?
|
||||
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from litellm.proxy.scheduled_tasks.types import (
|
|||
DueTaskResponse,
|
||||
DueTasksResponse,
|
||||
ListScheduledTasksResponse,
|
||||
ReportTaskResultRequest,
|
||||
ScheduledTaskResponse,
|
||||
UpdateScheduledTaskRequest,
|
||||
)
|
||||
|
|
@ -66,6 +67,8 @@ def _row_to_response(row) -> ScheduledTaskResponse:
|
|||
fire_once=row.fire_once,
|
||||
status=row.status,
|
||||
last_fired_at=row.last_fired_at,
|
||||
consecutive_errors=getattr(row, "consecutive_errors", 0) or 0,
|
||||
last_error=getattr(row, "last_error", None),
|
||||
created_at=row.created_at,
|
||||
updated_at=row.updated_at,
|
||||
)
|
||||
|
|
@ -318,3 +321,36 @@ async def cancel_scheduled_task(
|
|||
if row is None:
|
||||
raise HTTPException(status_code=404, detail="task not found")
|
||||
return _row_to_response(row)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/tasks/{task_id}/report",
|
||||
tags=["scheduled tasks"],
|
||||
response_model=ScheduledTaskResponse,
|
||||
)
|
||||
async def report_task_result(
|
||||
task_id: str,
|
||||
data: ReportTaskResultRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Agent reports outcome of one dispatch attempt.
|
||||
|
||||
success → resets the consecutive-error counter.
|
||||
error → increments it. On the Nth consecutive error
|
||||
(store.MAX_CONSECUTIVE_ERRORS), status flips to 'failed' and
|
||||
/due stops returning the task. Caller's next list() shows it
|
||||
as failed with last_error set.
|
||||
"""
|
||||
owner_token = _require_token(user_api_key_dict)
|
||||
prisma_client = _get_prisma_client()
|
||||
row = await store.report_task_result(
|
||||
prisma_client,
|
||||
task_id=task_id,
|
||||
owner_token=owner_token,
|
||||
result=data.result,
|
||||
reason=data.reason,
|
||||
)
|
||||
if row is None:
|
||||
raise HTTPException(status_code=404, detail="task not found")
|
||||
return _row_to_response(row)
|
||||
|
|
|
|||
|
|
@ -49,7 +49,8 @@ UPDATABLE_FIELDS = frozenset(
|
|||
JSON_FIELDS = frozenset({"action_args", "metadata"})
|
||||
|
||||
MAX_ACTIVE_TASKS_PER_KEY = 10
|
||||
TERMINAL_STATUSES = ("fired", "expired", "cancelled")
|
||||
MAX_CONSECUTIVE_ERRORS = 3
|
||||
TERMINAL_STATUSES = ("fired", "expired", "cancelled", "failed")
|
||||
|
||||
|
||||
async def count_active_for_owner(prisma_client: Any, owner_token: str) -> int:
|
||||
|
|
@ -101,12 +102,32 @@ async def create_task(
|
|||
return await prisma_client.db.litellm_scheduledtasktable.create(data=data)
|
||||
|
||||
|
||||
async def sweep_expired_for_owner(prisma_client: Any, owner_token: str) -> int:
|
||||
"""
|
||||
Lazy expiry: flip pending rows past expires_at to 'expired' before any
|
||||
read returns them. Without this, a task that never fires after its
|
||||
expires_at sits in 'pending' until the next claim that happens to scan
|
||||
it — which may be never if no other rows are ever due.
|
||||
|
||||
Returns count of rows flipped.
|
||||
"""
|
||||
return await prisma_client.db.litellm_scheduledtasktable.update_many(
|
||||
where={
|
||||
"owner_token": owner_token,
|
||||
"status": "pending",
|
||||
"expires_at": {"lte": datetime.now(timezone.utc)},
|
||||
},
|
||||
data={"status": "expired"},
|
||||
)
|
||||
|
||||
|
||||
async def list_tasks_for_owner(
|
||||
prisma_client: Any,
|
||||
*,
|
||||
owner_token: str,
|
||||
include_terminal: bool,
|
||||
) -> List[Any]:
|
||||
await sweep_expired_for_owner(prisma_client, owner_token)
|
||||
where: Dict[str, Any] = {"owner_token": owner_token}
|
||||
if not include_terminal:
|
||||
where["status"] = "pending"
|
||||
|
|
@ -122,11 +143,55 @@ async def get_task_for_owner(
|
|||
task_id: str,
|
||||
owner_token: str,
|
||||
) -> Optional[Any]:
|
||||
await sweep_expired_for_owner(prisma_client, owner_token)
|
||||
return await prisma_client.db.litellm_scheduledtasktable.find_first(
|
||||
where={"task_id": task_id, "owner_token": owner_token},
|
||||
)
|
||||
|
||||
|
||||
async def report_task_result(
|
||||
prisma_client: Any,
|
||||
*,
|
||||
task_id: str,
|
||||
owner_token: str,
|
||||
result: str,
|
||||
reason: Optional[str],
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
Agent reports outcome of one dispatch attempt.
|
||||
|
||||
success → reset consecutive_errors, clear last_error.
|
||||
error → bump consecutive_errors. If >= MAX_CONSECUTIVE_ERRORS, flip
|
||||
status to 'failed' so /due stops re-emitting it.
|
||||
|
||||
Scoped by (task_id, owner_token). Returns updated row, or None if not
|
||||
found / not owned.
|
||||
"""
|
||||
existing = await prisma_client.db.litellm_scheduledtasktable.find_first(
|
||||
where={"task_id": task_id, "owner_token": owner_token},
|
||||
)
|
||||
if existing is None:
|
||||
return None
|
||||
|
||||
if result == "success":
|
||||
return await prisma_client.db.litellm_scheduledtasktable.update(
|
||||
where={"task_id": task_id},
|
||||
data={"consecutive_errors": 0, "last_error": None},
|
||||
)
|
||||
|
||||
new_count = (existing.consecutive_errors or 0) + 1
|
||||
update: Dict[str, Any] = {
|
||||
"consecutive_errors": new_count,
|
||||
"last_error": reason,
|
||||
}
|
||||
if new_count >= MAX_CONSECUTIVE_ERRORS and existing.status == "pending":
|
||||
update["status"] = "failed"
|
||||
return await prisma_client.db.litellm_scheduledtasktable.update(
|
||||
where={"task_id": task_id},
|
||||
data=update,
|
||||
)
|
||||
|
||||
|
||||
async def update_task_for_owner(
|
||||
prisma_client: Any,
|
||||
*,
|
||||
|
|
|
|||
|
|
@ -4,7 +4,8 @@ from typing import Any, Dict, List, Literal, Optional
|
|||
from pydantic import BaseModel, Field
|
||||
|
||||
ScheduleKind = Literal["interval", "cron", "once"]
|
||||
TaskStatus = Literal["pending", "fired", "expired", "cancelled"]
|
||||
TaskStatus = Literal["pending", "fired", "expired", "cancelled", "failed"]
|
||||
ReportResult = Literal["success", "error"]
|
||||
|
||||
|
||||
class CreateScheduledTaskRequest(BaseModel):
|
||||
|
|
@ -62,11 +63,20 @@ class ScheduledTaskResponse(BaseModel):
|
|||
|
||||
status: str
|
||||
last_fired_at: Optional[datetime]
|
||||
consecutive_errors: int = 0
|
||||
last_error: Optional[str] = None
|
||||
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class ReportTaskResultRequest(BaseModel):
|
||||
"""Agent reports the outcome of one dispatch attempt."""
|
||||
|
||||
result: ReportResult
|
||||
reason: Optional[str] = None
|
||||
|
||||
|
||||
class DueTaskResponse(BaseModel):
|
||||
"""Trimmed shape returned by /due — only fields the agent needs to dispatch."""
|
||||
|
||||
|
|
|
|||
|
|
@ -1321,6 +1321,8 @@ model LiteLLM_ScheduledTaskTable {
|
|||
|
||||
status String @default("pending")
|
||||
last_fired_at DateTime?
|
||||
consecutive_errors Int @default(0)
|
||||
last_error String?
|
||||
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
|
|
|
|||
|
|
@ -1321,6 +1321,8 @@ model LiteLLM_ScheduledTaskTable {
|
|||
|
||||
status String @default("pending")
|
||||
last_fired_at DateTime?
|
||||
consecutive_errors Int @default(0)
|
||||
last_error String?
|
||||
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
|
|
|
|||
|
|
@ -53,6 +53,8 @@ def _make_row(**kwargs) -> MagicMock:
|
|||
"fire_once": True,
|
||||
"status": "pending",
|
||||
"last_fired_at": None,
|
||||
"consecutive_errors": 0,
|
||||
"last_error": None,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
}
|
||||
|
|
@ -120,6 +122,29 @@ class _FakeScheduledTaskTable:
|
|||
_ = order
|
||||
return self._filter(where)
|
||||
|
||||
async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> int:
|
||||
# Minimal Prisma-style filter: support {"lte": <dt>} for expires_at.
|
||||
count = 0
|
||||
for r in self.rows:
|
||||
ok = True
|
||||
for k, v in where.items():
|
||||
if isinstance(v, dict):
|
||||
actual = getattr(r, k, None)
|
||||
if "lte" in v:
|
||||
if actual is None or actual > v["lte"]:
|
||||
ok = False
|
||||
break
|
||||
else:
|
||||
if getattr(r, k, None) != v:
|
||||
ok = False
|
||||
break
|
||||
if not ok:
|
||||
continue
|
||||
for k, v in data.items():
|
||||
setattr(r, k, v)
|
||||
count += 1
|
||||
return count
|
||||
|
||||
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"]:
|
||||
|
|
@ -578,3 +603,97 @@ class TestDue:
|
|||
with _patch_prisma(self.prisma):
|
||||
r = self.client.get("/v1/tasks/due")
|
||||
assert len(r.json()["tasks"]) == 0
|
||||
|
||||
|
||||
class TestReport:
|
||||
def setup_method(self):
|
||||
self.prisma = _make_prisma()
|
||||
self.client = _make_app(self.prisma)
|
||||
|
||||
def _seed(self, **kwargs):
|
||||
defaults: Dict[str, Any] = {
|
||||
"task_id": "rep-1",
|
||||
"status": "pending",
|
||||
}
|
||||
defaults.update(kwargs)
|
||||
self.prisma.db.litellm_scheduledtasktable.rows.append(_make_row(**defaults))
|
||||
|
||||
def test_success_clears_counter(self):
|
||||
self._seed(consecutive_errors=2, last_error="prev")
|
||||
with _patch_prisma(self.prisma):
|
||||
r = self.client.post(
|
||||
"/v1/tasks/rep-1/report",
|
||||
json={"result": "success"},
|
||||
)
|
||||
assert r.status_code == 200
|
||||
body = r.json()
|
||||
assert body["consecutive_errors"] == 0
|
||||
assert body["last_error"] is None
|
||||
|
||||
def test_first_error_bumps_counter_no_flip(self):
|
||||
self._seed()
|
||||
with _patch_prisma(self.prisma):
|
||||
r = self.client.post(
|
||||
"/v1/tasks/rep-1/report",
|
||||
json={"result": "error", "reason": "boom"},
|
||||
)
|
||||
assert r.status_code == 200
|
||||
body = r.json()
|
||||
assert body["consecutive_errors"] == 1
|
||||
assert body["last_error"] == "boom"
|
||||
assert body["status"] == "pending"
|
||||
|
||||
def test_third_error_flips_to_failed(self):
|
||||
self._seed(consecutive_errors=2)
|
||||
with _patch_prisma(self.prisma):
|
||||
r = self.client.post(
|
||||
"/v1/tasks/rep-1/report",
|
||||
json={"result": "error", "reason": "still broken"},
|
||||
)
|
||||
assert r.status_code == 200
|
||||
body = r.json()
|
||||
assert body["consecutive_errors"] == 3
|
||||
assert body["status"] == "failed"
|
||||
assert body["last_error"] == "still broken"
|
||||
|
||||
def test_report_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.post(
|
||||
"/v1/tasks/foreign-1/report",
|
||||
json={"result": "error"},
|
||||
)
|
||||
assert r.status_code == 404
|
||||
|
||||
def test_failed_task_not_returned_by_due(self):
|
||||
self._seed(
|
||||
status="failed",
|
||||
next_run_at=_now() - timedelta(seconds=10),
|
||||
)
|
||||
with _patch_prisma(self.prisma):
|
||||
r = self.client.get("/v1/tasks/due")
|
||||
assert r.json()["tasks"] == []
|
||||
|
||||
|
||||
class TestLazyExpiry:
|
||||
def setup_method(self):
|
||||
self.prisma = _make_prisma()
|
||||
self.client = _make_app(self.prisma)
|
||||
|
||||
def test_list_flips_expired_pending_rows(self):
|
||||
self.prisma.db.litellm_scheduledtasktable.rows.append(
|
||||
_make_row(
|
||||
task_id="exp-1",
|
||||
status="pending",
|
||||
next_run_at=_now() + timedelta(hours=1), # not yet due
|
||||
expires_at=_now() - timedelta(seconds=1), # already expired
|
||||
)
|
||||
)
|
||||
with _patch_prisma(self.prisma):
|
||||
r = self.client.get("/v1/tasks?include_terminal=true")
|
||||
assert r.status_code == 200
|
||||
body = r.json()
|
||||
statuses = {t["task_id"]: t["status"] for t in body["tasks"]}
|
||||
assert statuses["exp-1"] == "expired"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue