perf(lens): claim worker jobs from an indexed due queue instead of scanning every lens (#45095)

* perf(lens): claim worker jobs from an indexed due queue instead of scanning every lens

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

* fix(lens): apply the due_at index concurrently on its own and default legacy rows to due

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

* test(lens): pass repository to claim lifecycle tests

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

* fix(lens): page past unsupported due lenses and declare the full due_at index

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

* fix(lens): claim due lenses in a loop instead of recursion

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-07 12:24:01 -07:00 • committed by GitHub
parent d174c43518
commit 65dcf43257
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 738 additions and 22 deletions

View file

@ -0,0 +1 @@
ALTER TABLE "LiteLLM_Lens" ADD COLUMN IF NOT EXISTS "due_at" TIMESTAMP(3) DEFAULT '1970-01-01 00:00:00';

View file

@ -0,0 +1,2 @@
CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_Lens_due_at_idx"
ON "LiteLLM_Lens" ("due_at");

View file

@ -1939,6 +1939,9 @@ model LiteLLM_Lens {
id String @id
version Int @default(0)
data Json
due_at DateTime? @default(dbgenerated("'1970-01-01 00:00:00'::timestamp without time zone"))
@@index([due_at], map: "LiteLLM_Lens_due_at_idx")
}
model LiteLLM_LensRun {

View file

@ -1,10 +1,11 @@
import hashlib
import secrets
from collections.abc import Awaitable, Callable
from datetime import datetime, timedelta, timezone
from functools import reduce
from itertools import chain
from types import MappingProxyType
from typing import Annotated, Final, TypeAlias
from typing import Annotated, Final, Protocol, TypeAlias
from uuid import uuid4
from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response
@ -49,7 +50,7 @@ from litellm.proxy.lens.models import (
WorkerCreated,
)
from litellm.proxy.lens.release import PROTOCOL_VERSION, release_tag, worker_image
from litellm.proxy.lens.repository import LensRepository, WriterDatabase
from litellm.proxy.lens.repository import DueLens, LensRepository, WriterDatabase
from litellm.proxy.lens.reviews import criteria_key
from litellm.proxy.lens.sources import ActivityAvailability, SourceReader, Storage, parse_execution
from litellm.proxy.lens.state import (
@ -71,11 +72,24 @@ from litellm.proxy.tracing_runtime import provide_storage
from litellm.types.llms.base import LiteLLMBaseModel
router: Final = APIRouter(prefix="/lens", tags=["Lens"])
CLAIM_CANDIDATES: Final = 20
_bearer: Final = HTTPBearer()
Auth: TypeAlias = Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)]
StorageDep: TypeAlias = Annotated[Storage | None, Depends(provide_storage)]
class _ClaimRepository(Protocol):
async def due(
self, scope: Scope, now: datetime, limit: int, after: DueLens | None = None
) -> tuple[DueLens, ...]: ...
async def sync_due(self, lens: Lens) -> None: ...
async def update(
self, lens_id: str, transform: Callable[[Lens], Lens], attempts: int, *, changed_only: bool
) -> Lens | None: ...
def repository() -> LensRepository:
from litellm.proxy.proxy_server import prisma_client
@ -501,13 +515,29 @@ async def claim(worker: WorkerAuth, protocol_version: int = 1, worker_release: s
if worker.analysis_key_id is None:
raise HTTPException(409, "Assign an analysis key to this worker in Lens setup")
now: Final = datetime.now(timezone.utc)
await repository().heartbeat(worker.id, now.isoformat())
for candidate in await repository().lenses():
if not can_access(worker.scope, candidate.scope):
continue
if claimed := await claim_candidate(candidate, worker, now):
return claimed
return None
lens_repository: Final = repository()
await lens_repository.heartbeat(worker.id, now.isoformat())
return await claim_due(worker, now, lens_repository)
async def claim_due(
worker: Worker,
now: datetime,
lens_repository: _ClaimRepository,
supports_model: Callable[[Worker, LensSettings], Awaitable[bool]] = worker_supports_model,
) -> Claim | None:
after: DueLens | None = None # rebind-ok: keyset cursor advances one page at a time
while True:
page = await lens_repository.due(worker.scope, now, CLAIM_CANDIDATES, after)
for candidate in page:
if not can_access(worker.scope, candidate.lens.scope):
continue
if claimed := await claim_candidate(candidate.lens, worker, now, lens_repository, supports_model):
return claimed
await lens_repository.sync_due(candidate.lens)
if len(page) < CLAIM_CANDIDATES:
return None
after = page[-1]
@router.post("/worker/{lens_id}/{job_id}/progress", response_model=bool)
@ -752,9 +782,15 @@ async def heartbeat(lens_id: str, job_id: str, worker: WorkerAuth) -> bool:
return await progress(lens_id, job_id, Progress(), worker)
async def claim_candidate(candidate: Lens, worker: Worker, now: datetime) -> Claim | None:
async def claim_candidate(
candidate: Lens,
worker: Worker,
now: datetime,
lens_repository: _ClaimRepository,
supports_model: Callable[[Worker, LensSettings], Awaitable[bool]] = worker_supports_model,
) -> Claim | None:
active: Final = current_job(candidate)
if not await worker_supports_model(worker, active.settings if active else candidate.settings):
if not await supports_model(worker, active.settings if active else candidate.settings):
return None
job_id: Final = str(uuid4())
@ -765,7 +801,7 @@ async def claim_candidate(candidate: Lens, worker: Worker, now: datetime) -> Cla
return e
return claim_job(scheduled, worker, now)
updated: Final = await repository().update(candidate.id, schedule, changed_only=True)
updated: Final = await lens_repository.update(candidate.id, schedule, attempts=1, changed_only=True)
if updated is None:
return None
job: Final = current_job(updated)

View file

@ -3,6 +3,7 @@ import json
import random
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Protocol
@ -24,7 +25,7 @@ from litellm.proxy.lens.models import (
Worker,
)
from litellm.proxy.lens.reviews import criteria_key
from litellm.proxy.lens.state import apply_progress, current_job, replace_job
from litellm.proxy.lens.state import apply_progress, current_job, due_at, replace_job
from litellm.types.llms.base import LiteLLMBaseModel
if TYPE_CHECKING:
@ -39,6 +40,18 @@ class Database(Protocol):
class Row(LiteLLMBaseModel):
data: JsonValue
due_at: datetime | None = None
class DueRow(LiteLLMBaseModel):
data: JsonValue
due_at: datetime
@dataclass(frozen=True, slots=True)
class DueLens:
lens: Lens
due_at: datetime
class FindingRun(LiteLLMBaseModel):
@ -47,6 +60,26 @@ class FindingRun(LiteLLMBaseModel):
_ROWS: Final = TypeAdapter(tuple[Row, ...])
_DUE_ROWS: Final = TypeAdapter(tuple[DueRow, ...])
_DUE_QUERY: Final[LiteralString] = """SELECT data, due_at FROM "LiteLLM_Lens"
WHERE due_at IS NOT NULL AND due_at <= ($4::timestamptz AT TIME ZONE 'UTC')
AND ($1::boolean OR (
COALESCE((data->'scope'->>'all_teams')::boolean, false) IS NOT TRUE
AND COALESCE(data->'scope'->>'team_id', '')=$2
AND ($2 <> '' OR COALESCE(data->'scope'->>'api_key_hash', '')=$3)
))
ORDER BY due_at, id
LIMIT $5"""
_DUE_AFTER_QUERY: Final[LiteralString] = """SELECT data, due_at FROM "LiteLLM_Lens"
WHERE due_at IS NOT NULL AND due_at <= ($4::timestamptz AT TIME ZONE 'UTC')
AND (due_at, id) > ($6::timestamp, $7)
AND ($1::boolean OR (
COALESCE((data->'scope'->>'all_teams')::boolean, false) IS NOT TRUE
AND COALESCE(data->'scope'->>'team_id', '')=$2
AND ($2 <> '' OR COALESCE(data->'scope'->>'api_key_hash', '')=$3)
))
ORDER BY due_at, id
LIMIT $5"""
UPDATE_ATTEMPTS: Final = 40
UPDATE_BACKOFF_SECONDS: Final = 0.02
@ -146,6 +179,30 @@ class LensRepository:
rows: Final = _ROWS.validate_python(await self.db.query_raw('SELECT data FROM "LiteLLM_Lens" ORDER BY id'))
return tuple(Lens.model_validate(row.data) for row in rows)
async def due(self, scope: Scope, now: datetime, limit: int, after: DueLens | None = None) -> tuple[DueLens, ...]:
query: Final[LiteralString] = _DUE_QUERY if after is None else _DUE_AFTER_QUERY
parameters: Final[tuple[object, ...]] = (
(
scope.all_teams,
scope.team_id,
scope.api_key_hash,
now.isoformat(),
limit,
)
if after is None
else (
scope.all_teams,
scope.team_id,
scope.api_key_hash,
now.isoformat(),
limit,
after.due_at,
after.lens.id,
)
)
rows: Final = _DUE_ROWS.validate_python(await self.db.query_raw(query, *parameters), from_attributes=True)
return tuple(DueLens(lens=Lens.model_validate(row.data), due_at=row.due_at) for row in rows)
async def get(self, lens_id: str) -> Lens | None:
rows: Final = _ROWS.validate_python(
await self.db.query_raw(
@ -157,12 +214,25 @@ class LensRepository:
async def create(self, lens: Lens) -> Lens:
await self.db.execute_raw(
'INSERT INTO "LiteLLM_Lens" (id, version, data) VALUES ($1,0,$2::jsonb)',
"""INSERT INTO "LiteLLM_Lens" (id, version, data, due_at)
VALUES ($1,0,$2::jsonb,($3::timestamptz AT TIME ZONE 'UTC'))""",
lens.id,
lens.model_dump_json(),
scheduled_at.isoformat() if (scheduled_at := due_at(lens)) else None,
)
return lens
async def sync_due(self, lens: Lens) -> None:
await self.db.execute_raw(
"""UPDATE "LiteLLM_Lens"
SET due_at=($3::timestamptz AT TIME ZONE 'UTC')
WHERE id=$1 AND version=$2
AND due_at IS DISTINCT FROM ($3::timestamptz AT TIME ZONE 'UTC')""",
lens.id,
lens.version,
scheduled_at.isoformat() if (scheduled_at := due_at(lens)) else None,
)
async def update(
self,
lens_id: str,
@ -193,7 +263,8 @@ class LensRepository:
"""WITH previous AS MATERIALIZED (
SELECT data FROM "LiteLLM_Lens" WHERE id=$2 AND version=$3 FOR UPDATE
), updated AS (
UPDATE "LiteLLM_Lens" SET data=$1::jsonb, version=version+1
UPDATE "LiteLLM_Lens" SET data=$1::jsonb, version=version+1,
due_at=($4::timestamptz AT TIME ZONE 'UTC')
WHERE id=$2 AND version=$3 AND EXISTS (SELECT 1 FROM previous) RETURNING id
)
, archived AS (INSERT INTO "LiteLLM_LensRun" (id, lens_id, created_at, data)
@ -207,6 +278,7 @@ class LensRepository:
updated.model_dump_json(),
lens_id,
previous.version,
scheduled_at.isoformat() if (scheduled_at := due_at(updated)) else None,
)
)
return bool(rows and rows[0].data == 1), updated

View file

@ -37,6 +37,15 @@ def current_job(lens: Lens) -> Job | None:
return next((job for job in lens.jobs if job.status in ("queued", "running")), None)
def due_at(lens: Lens) -> datetime | None:
job: Final = current_job(lens)
if job is None:
return lens.next_run_at if lens.settings.enabled else None
if job.status == "queued":
return job.created_at
return job.lease_until or job.created_at
def replace_job(lens: Lens, job: Job) -> Lens:
return lens.model_copy(
update=MappingProxyType({"jobs": tuple(job if old.id == job.id else old for old in lens.jobs)})

View file

@ -1939,6 +1939,9 @@ model LiteLLM_Lens {
id String @id
version Int @default(0)
data Json
due_at DateTime? @default(dbgenerated("'1970-01-01 00:00:00'::timestamp without time zone"))
@@index([due_at], map: "LiteLLM_Lens_due_at_idx")
}
model LiteLLM_LensRun {

View file

@ -1939,6 +1939,9 @@ model LiteLLM_Lens {
id String @id
version Int @default(0)
data Json
due_at DateTime? @default(dbgenerated("'1970-01-01 00:00:00'::timestamp without time zone"))
@@index([due_at], map: "LiteLLM_Lens_due_at_idx")
}
model LiteLLM_LensRun {

View file

@ -14,6 +14,7 @@ import pytest_asyncio
from fastapi import HTTPException
from prisma import Prisma
from psycopg import sql
from pydantic import TypeAdapter
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.db.prisma_client import PrismaWrapper
@ -35,7 +36,7 @@ from litellm.proxy.lens.models import (
Worker,
)
from litellm.proxy.lens.repository import LensRepository, WriterDatabase
from litellm.proxy.lens.state import claim_job, queue_job
from litellm.proxy.lens.state import cancel_job, claim_job, current_job, due_at, end_job, queue_job, replace_job
@pytest_asyncio.fixture(loop_scope="function")
@ -44,6 +45,274 @@ async def lens_db() -> AsyncIterator[Prisma]:
yield db
def _scheduled_lens(
lens_id: str,
scope: Scope,
now: datetime,
next_run_at: datetime,
*,
enabled: bool = True,
jobs: tuple[Job, ...] = (),
) -> Lens:
return Lens(
id=lens_id,
scope=scope,
settings=LensSettings(
name="Scheduling test",
model="analysis",
context="Find unexpected behavior",
enabled=enabled,
),
created_at=now,
next_run_at=next_run_at,
jobs=jobs,
budget_month=now.strftime("%Y-%m"),
)
def _stored_due_at(lens_id: str) -> datetime | None:
with psycopg.connect(os.environ["DATABASE_URL"]) as connection:
row: Final = connection.execute('SELECT due_at FROM "LiteLLM_Lens" WHERE id=%s', (lens_id,)).fetchone()
return TypeAdapter(datetime | None).validate_python(row[0]) if row else None
async def _assert_due_column(repo: LensRepository, lens_id: str) -> None:
stored: Final = await repo.get(lens_id)
assert stored is not None
expected: Final = due_at(stored)
actual: Final = _stored_due_at(lens_id)
if expected is None:
assert actual is None
return
assert actual is not None
difference: Final = actual.replace(tzinfo=timezone.utc) - expected.astimezone(timezone.utc)
assert abs(difference.total_seconds()) <= 0.001
@pytest.mark.asyncio
async def test_due_filters_by_schedule_and_scope(lens_db: Prisma) -> None:
utc_now: Final = datetime.now(timezone.utc).replace(microsecond=0)
worker_now: Final = utc_now.astimezone(timezone(timedelta(hours=3)))
team_id: Final = uuid4().hex
worker_scope: Final = Scope(team_id=team_id)
worker: Final = Worker(id=uuid4().hex, name="worker", scope=worker_scope, last_seen=worker_now)
repo: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db)))
due_lens: Final = _scheduled_lens(uuid4().hex, worker_scope, utc_now, utc_now - timedelta(minutes=20))
future_lens: Final = _scheduled_lens(uuid4().hex, worker_scope, utc_now, utc_now + timedelta(minutes=20))
disabled_lens: Final = _scheduled_lens(
uuid4().hex, worker_scope, utc_now, utc_now - timedelta(minutes=10), enabled=False
)
live_queued: Final = queue_job(
_scheduled_lens(uuid4().hex, worker_scope, utc_now - timedelta(minutes=5), utc_now - timedelta(minutes=5)),
utc_now - timedelta(minutes=5),
uuid4().hex,
)
live_lens: Final = claim_job(live_queued, worker, worker_now)
expired_queued: Final = queue_job(
_scheduled_lens(uuid4().hex, worker_scope, utc_now - timedelta(minutes=10), utc_now - timedelta(minutes=10)),
utc_now - timedelta(minutes=10),
uuid4().hex,
)
expired_claimed: Final = claim_job(expired_queued, worker, utc_now - timedelta(minutes=10))
expired_job: Final = expired_claimed.jobs[0].model_copy(update={"lease_until": utc_now - timedelta(minutes=5)})
expired_lens: Final = expired_claimed.model_copy(update={"jobs": (expired_job,)})
other_lens: Final = _scheduled_lens(
uuid4().hex, Scope(team_id=uuid4().hex), utc_now, utc_now - timedelta(minutes=3)
)
worker_key: Final = uuid4().hex
key_lens: Final = _scheduled_lens(
uuid4().hex, Scope(api_key_hash=worker_key), utc_now, utc_now - timedelta(minutes=2)
)
candidates: Final = (due_lens, future_lens, disabled_lens, live_lens, expired_lens, other_lens, key_lens)
await asyncio.gather(*(repo.create(candidate) for candidate in candidates))
try:
await lens_db.execute_raw(
"""UPDATE "LiteLLM_Lens"
SET data=jsonb_set(data, '{scope}', jsonb_build_object('team_id', $2))
WHERE id=$1""",
due_lens.id,
team_id,
)
await lens_db.execute_raw(
"""UPDATE "LiteLLM_Lens"
SET data=jsonb_set(data, '{scope}', jsonb_build_object('api_key_hash', $2))
WHERE id=$1""",
key_lens.id,
worker_key,
)
team_due: Final = await repo.due(worker_scope, worker_now, 20)
assert tuple(candidate.lens.id for candidate in team_due) == tuple(
lens.id for lens in sorted((due_lens, expired_lens), key=lambda lens: (due_at(lens), lens.id))
)
assert team_due[0].lens.scope == worker_scope
key_due: Final = await repo.due(Scope(api_key_hash=worker_key), worker_now, 20)
assert tuple(candidate.lens.id for candidate in key_due) == (key_lens.id,)
all_due: Final = await repo.due(Scope(all_teams=True), worker_now, 20)
assert {candidate.lens.id for candidate in all_due} == {
due_lens.id,
expired_lens.id,
other_lens.id,
key_lens.id,
}
finally:
await lens_db.execute_raw(
'DELETE FROM "LiteLLM_Lens" WHERE id=ANY($1::text[])',
tuple(lens.id for lens in candidates),
)
@pytest.mark.asyncio
async def test_due_pages_lenses_with_equal_due_at_without_skipping_or_repeating(lens_db: Prisma) -> None:
now: Final = datetime.now(timezone.utc).replace(microsecond=0)
scope: Final = Scope(team_id=uuid4().hex)
repo: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db)))
lenses: Final = tuple(_scheduled_lens(uuid4().hex, scope, now, now - timedelta(minutes=1)) for _ in range(45))
await asyncio.gather(*(repo.create(lens) for lens in lenses))
try:
await lens_db.execute_raw(
"""UPDATE "LiteLLM_Lens"
SET due_at=$2::timestamp
WHERE id=ANY($1::text[])""",
tuple(lens.id for lens in lenses),
"1970-01-01 00:00:00",
)
first: Final = await repo.due(scope, now, 20)
second: Final = await repo.due(scope, now, 20, first[-1])
third: Final = await repo.due(scope, now, 20, second[-1])
assert tuple(len(page) for page in (first, second, third)) == (20, 20, 5)
ids: Final = tuple(candidate.lens.id for candidate in (*first, *second, *third))
assert ids == tuple(sorted(lens.id for lens in lenses))
finally:
await lens_db.execute_raw(
'DELETE FROM "LiteLLM_Lens" WHERE id=ANY($1::text[])',
tuple(lens.id for lens in lenses),
)
@pytest.mark.asyncio
async def test_due_at_stays_consistent_through_job_lifecycle(lens_db: Prisma) -> None:
now: Final = datetime.now(timezone.utc).replace(microsecond=0)
scope: Final = Scope(team_id=uuid4().hex)
lens: Final = _scheduled_lens(uuid4().hex, scope, now, now)
repo: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db)))
worker: Final = Worker(id=uuid4().hex, name="worker", scope=scope, last_seen=now)
await repo.create(lens)
try:
await _assert_due_column(repo, lens.id)
job_id: Final = uuid4().hex
claimed: Final = await repo.update(
lens.id,
lambda candidate: claim_job(queue_job(candidate, now, job_id), worker, now),
attempts=1,
)
assert claimed is not None
await _assert_due_column(repo, lens.id)
active: Final = current_job(claimed)
assert active is not None
progressed: Final = await repo.progress(lens.id, active, Progress())
assert progressed is not None
await _assert_due_column(repo, lens.id)
result_at: Final = datetime.now(timezone.utc)
def finish(candidate: Lens) -> Lens:
active_job: Final = current_job(candidate)
if active_job is None:
return candidate
return replace_job(candidate, end_job(active_job, "completed", result_at)).model_copy(
update={"next_run_at": result_at + timedelta(minutes=candidate.settings.interval_minutes)}
)
completed: Final = await repo.update(lens.id, finish, attempts=1)
assert completed is not None
await _assert_due_column(repo, lens.id)
cancelled_at: Final = datetime.now(timezone.utc)
cancelled: Final = await repo.update(
lens.id,
lambda candidate: cancel_job(
queue_job(candidate, cancelled_at, uuid4().hex, trigger="manual"),
cancelled_at,
),
attempts=1,
)
assert cancelled is not None
await _assert_due_column(repo, lens.id)
finally:
await lens_db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', lens.id)
@pytest.mark.asyncio
async def test_sync_due_repairs_legacy_rows_and_ignores_stale_versions(lens_db: Prisma) -> None:
now: Final = datetime.now(timezone.utc).replace(microsecond=0)
team_id: Final = uuid4().hex
scope: Final = Scope(team_id=team_id)
worker: Final = Worker(id=uuid4().hex, name="worker", scope=scope, last_seen=now)
repo: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db)))
due_idle: Final = _scheduled_lens(uuid4().hex, scope, now, now - timedelta(minutes=20))
future_idle: Final = _scheduled_lens(uuid4().hex, scope, now, now + timedelta(minutes=20))
disabled_idle: Final = _scheduled_lens(uuid4().hex, scope, now, now - timedelta(minutes=10), enabled=False)
queued_lens: Final = queue_job(
_scheduled_lens(uuid4().hex, scope, now, now + timedelta(minutes=20), enabled=False),
now - timedelta(minutes=3),
uuid4().hex,
trigger="manual",
)
live_lens: Final = claim_job(
queue_job(
_scheduled_lens(uuid4().hex, scope, now, now + timedelta(minutes=20)),
now - timedelta(minutes=10),
uuid4().hex,
),
worker,
now,
)
expired_claimed: Final = claim_job(
queue_job(
_scheduled_lens(uuid4().hex, scope, now, now + timedelta(minutes=20)),
now - timedelta(minutes=10),
uuid4().hex,
),
worker,
now - timedelta(minutes=10),
)
expired_lens: Final = expired_claimed.model_copy(
update={"jobs": (expired_claimed.jobs[0].model_copy(update={"lease_until": now - timedelta(minutes=5)}),)}
)
candidates: Final = (due_idle, future_idle, disabled_idle, queued_lens, live_lens, expired_lens)
await asyncio.gather(*(repo.create(candidate) for candidate in candidates))
try:
past: Final = now - timedelta(hours=1)
await lens_db.execute_raw(
"""UPDATE "LiteLLM_Lens"
SET due_at=($2::timestamptz AT TIME ZONE 'UTC')
WHERE id=ANY($1::text[])""",
tuple(lens.id for lens in candidates),
past.isoformat(),
)
legacy_due: Final = await repo.due(scope, now, 20)
assert {candidate.lens.id for candidate in legacy_due} == {lens.id for lens in candidates}
for candidate in legacy_due:
await repo.sync_due(candidate.lens)
repaired_due: Final = await repo.due(scope, now, 20)
assert {candidate.lens.id for candidate in repaired_due} == {due_idle.id, queued_lens.id, expired_lens.id}
await asyncio.gather(*(_assert_due_column(repo, lens.id) for lens in candidates))
stale: Final = await repo.get(future_idle.id)
assert stale is not None
await lens_db.execute_raw(
"""UPDATE "LiteLLM_Lens"
SET version=version+1, due_at=($2::timestamptz AT TIME ZONE 'UTC')
WHERE id=$1""",
stale.id,
past.isoformat(),
)
await repo.sync_due(stale)
assert _stored_due_at(stale.id) == past.replace(tzinfo=None)
finally:
await lens_db.execute_raw(
'DELETE FROM "LiteLLM_Lens" WHERE id=ANY($1::text[])',
tuple(lens.id for lens in candidates),
)
@pytest.mark.asyncio
async def test_concurrent_workers_cannot_both_acquire_the_same_job(lens_db: Prisma) -> None:
now: Final = datetime.now(timezone.utc)
@ -291,6 +560,19 @@ def test_db_push_creates_fresh_lens_tables_and_preserves_them_on_restart(monkeyp
assert connection.execute(
sql.SQL("SELECT id, data FROM {}").format(sql.Identifier(schema, "LiteLLM_Lens"))
).fetchall() == [("saved", {"keep": True})]
assert (
connection.execute(
sql.SQL("SELECT due_at FROM {} WHERE id='saved'").format(sql.Identifier(schema, "LiteLLM_Lens"))
).fetchone()[0]
is not None
)
due_index: Final = connection.execute(
"""SELECT indexdef FROM pg_indexes
WHERE schemaname=%s AND tablename='LiteLLM_Lens' AND indexname='LiteLLM_Lens_due_at_idx'""",
(schema,),
).fetchone()
assert due_index is not None
assert "WHERE" not in due_index[0]
finally:
connection.execute(sql.SQL("DROP SCHEMA {} CASCADE").format(sql.Identifier(schema)))

View file

@ -0,0 +1,199 @@
import asyncio
import json
import os
import sys
from collections.abc import AsyncIterator
from contextlib import AbstractAsyncContextManager
from datetime import datetime, timedelta, timezone
from time import perf_counter
from typing import Final
from uuid import uuid4
import pytest
import pytest_asyncio
from prisma import Prisma
from pydantic import TypeAdapter
from typing_extensions import LiteralString
from litellm.proxy.db.prisma_client import PrismaWrapper
from litellm.proxy.lens.endpoints import claim_due
from litellm.proxy.lens.models import Evidence, Finding, Lens, LensSettings, Scope, Worker
from litellm.proxy.lens.repository import Database, LensRepository, Row, WriterDatabase
from litellm.proxy.lens.state import current_job
@pytest_asyncio.fixture(loop_scope="function")
async def lens_db() -> AsyncIterator[Prisma]:
async with Prisma(datasource={"url": os.environ["DATABASE_URL"]}) as db:
yield db
class ReadMeter:
def __init__(self) -> None:
self.batches: tuple[tuple[int, ...], ...] = ()
def record(self, document_sizes: tuple[int, ...]) -> None:
self.batches = (*self.batches, document_sizes)
@property
def document_count(self) -> int:
return sum(len(batch) for batch in self.batches)
@property
def total_bytes(self) -> int:
return sum(sum(batch) for batch in self.batches)
class MeasuredDatabase:
def __init__(self, database: WriterDatabase, meter: ReadMeter) -> None:
self.database: Final = database
self.meter: Final = meter
async def query_raw(self, query: LiteralString, *args: object) -> object:
rows: Final = await self.database.query_raw(query, *args)
if 'FROM "LiteLLM_Lens"' in query and "WHERE id" not in query:
documents: Final = TypeAdapter(tuple[Row, ...]).validate_python(rows)
self.meter.record(
tuple(len(json.dumps(row.data, separators=(",", ":")).encode("utf-8")) for row in documents)
)
return rows
async def execute_raw(self, query: LiteralString, *args: object) -> int:
return await self.database.execute_raw(query, *args)
def transaction(self) -> AbstractAsyncContextManager[Database]:
return self.database.transaction()
def _large_lens(lens_id: str, scope: Scope, now: datetime, next_run_at: datetime) -> Lens:
findings: Final = tuple(
Finding(
id=f"f{index}",
title=f"Issue {index}",
description="Repeated operation returns an unexpected result.",
check_id="behavior",
evidence=(
Evidence(
execution_id=f"t{index}",
span_id=f"s{index}",
quote="Unexpected result",
),
),
first_seen=now,
last_seen=now,
revision=1,
)
for index in range(100)
)
return Lens(
id=lens_id,
scope=scope,
settings=LensSettings(
name="Claim scheduler load",
model="analysis",
context="Find unexpected behavior",
enabled=True,
),
created_at=now,
next_run_at=next_run_at,
findings=findings,
budget_month=now.strftime("%Y-%m"),
)
def _due_lens(lens_id: str, scope: Scope, now: datetime, model: str, next_run_at: datetime) -> Lens:
return Lens(
id=lens_id,
scope=scope,
settings=LensSettings(
name="Claim paging test",
model=model,
context="Find unexpected behavior",
enabled=True,
),
created_at=now,
next_run_at=next_run_at,
budget_month=now.strftime("%Y-%m"),
)
async def _supports_model(_worker: Worker, _settings: LensSettings) -> bool:
return True
async def _supports_supported_model(_worker: Worker, settings: LensSettings) -> bool:
return settings.model == "supported"
@pytest.mark.asyncio
async def test_claim_due_reaches_a_supported_lens_behind_a_full_page_of_unsupported_ones(
lens_db: Prisma,
) -> None:
now: Final = datetime.now(timezone.utc).replace(microsecond=0)
scope: Final = Scope(team_id=uuid4().hex)
worker: Final = Worker(id=uuid4().hex, name="paging-test-worker", scope=scope, last_seen=now)
unsupported_at: Final = now - timedelta(minutes=5)
supported_at: Final = now - timedelta(minutes=1)
unsupported: Final = tuple(_due_lens(uuid4().hex, scope, now, "unsupported", unsupported_at) for _ in range(25))
supported: Final = _due_lens(uuid4().hex, scope, now, "supported", supported_at)
candidates: Final = (*unsupported, supported)
repository: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db)))
await asyncio.gather(*(repository.create(candidate) for candidate in candidates))
try:
claim: Final = await claim_due(worker, now, repository, _supports_supported_model)
assert claim is not None
assert claim.lens_id == supported.id
assert claim.job.status == "running"
finally:
await lens_db.execute_raw(
'DELETE FROM "LiteLLM_Lens" WHERE id=ANY($1::text[])',
tuple(candidate.id for candidate in candidates),
)
@pytest.mark.asyncio
async def test_lens_claim_reads_scale_with_due_lenses_not_total_lenses(lens_db: Prisma) -> None:
now: Final = datetime.now(timezone.utc).replace(microsecond=0)
scope: Final = Scope(team_id=uuid4().hex)
worker: Final = Worker(id=uuid4().hex, name="load-test-worker", scope=scope, last_seen=now)
due_lens: Final = _large_lens(uuid4().hex, scope, now, now - timedelta(seconds=1))
initial_future: Final = tuple(_large_lens(uuid4().hex, scope, now, now + timedelta(days=1)) for _ in range(20))
additional_future: Final = tuple(_large_lens(uuid4().hex, scope, now, now + timedelta(days=1)) for _ in range(200))
ids: Final = tuple(lens.id for lens in (due_lens, *initial_future, *additional_future))
writer: Final = WriterDatabase(PrismaWrapper(lens_db))
seed_repository: Final = LensRepository(writer)
await asyncio.gather(*(seed_repository.create(lens) for lens in (due_lens, *initial_future)))
try:
before_meter: Final = ReadMeter()
before_repository: Final = LensRepository(MeasuredDatabase(writer, before_meter))
before_started: Final = perf_counter()
before_claim: Final = await claim_due(worker, now, before_repository, _supports_model)
before_seconds: Final = perf_counter() - before_started
assert before_claim is not None
assert before_claim.lens_id == due_lens.id
assert before_claim.job.status == "running"
claimed_lens: Final = await seed_repository.get(due_lens.id)
assert claimed_lens is not None
assert current_job(claimed_lens) == before_claim.job
await seed_repository.update(
due_lens.id,
lambda lens: lens.model_copy(update={"jobs": (), "next_run_at": now - timedelta(seconds=1)}),
attempts=1,
)
await asyncio.gather(*(seed_repository.create(lens) for lens in additional_future))
after_meter: Final = ReadMeter()
after_repository: Final = LensRepository(MeasuredDatabase(writer, after_meter))
after_started: Final = perf_counter()
after_claim: Final = await claim_due(worker, now, after_repository, _supports_model)
after_seconds: Final = perf_counter() - after_started
assert after_claim is not None
assert after_claim.lens_id == due_lens.id
assert after_claim.job.status == "running"
sys.stdout.write(
f"claim read: before={before_meter.total_bytes} bytes, {before_seconds:.4f}s; "
f"after={after_meter.total_bytes} bytes, {after_seconds:.4f}s\n"
)
assert before_meter.document_count == after_meter.document_count == 1
assert before_meter.total_bytes == after_meter.total_bytes
finally:
await lens_db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=ANY($1::text[])', ids)

View file

@ -197,7 +197,10 @@ async def test_team_route_requires_a_worker_with_matching_model_access(lens_data
)
try:
wrong_team: Final = await endpoints.register_worker(endpoints.WorkerName(analysis_key_id=key_b), admin)
assert await endpoints.claim_candidate(lens, wrong_team.worker, datetime.now(timezone.utc)) is None
assert (
await endpoints.claim_candidate(lens, wrong_team.worker, datetime.now(timezone.utc), endpoints.repository())
is None
)
for operation in (
endpoints.create_lens(settings, admin),
endpoints.run_lens(lens.id, RunRequest(), admin),
@ -212,7 +215,9 @@ async def test_team_route_requires_a_worker_with_matching_model_access(lens_data
assert edited.settings.context == "Use sources"
right_team: Final = await endpoints.register_worker(endpoints.WorkerName(analysis_key_id=key_a), admin)
await endpoints.validate_workers(settings, lens.scope)
claim: Final = await endpoints.claim_candidate(lens, right_team.worker, datetime.now(timezone.utc))
claim: Final = await endpoints.claim_candidate(
lens, right_team.worker, datetime.now(timezone.utc), endpoints.repository()
)
assert claim is not None and claim.job.worker_id == right_team.worker.id
finally:
await lens_database.db.execute_raw(
@ -257,7 +262,10 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
assert lens.id in tuple(e.id for e in listing.lenses)
assert worker.id in tuple(w.id for w in listing.workers)
claims: Final = await asyncio.gather(
*(endpoints.claim_candidate(lens, worker, datetime.now(timezone.utc)) for _ in range(8))
*(
endpoints.claim_candidate(lens, worker, datetime.now(timezone.utc), endpoints.repository())
for _ in range(8)
)
)
winners: Final = tuple(claim for claim in claims if claim is not None)
assert len(winners) == 1
@ -265,7 +273,10 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
assert claimed.job.worker_id == worker.id
assert (
await endpoints.claim_candidate(
await endpoints.get_lens(lens.id, worker.scope), worker, datetime.now(timezone.utc)
await endpoints.get_lens(lens.id, worker.scope),
worker,
datetime.now(timezone.utc),
endpoints.repository(),
)
is None
)
@ -420,7 +431,9 @@ async def test_failed_model_requests_release_lens_budget_reservations(lens_datab
registration: Final = await endpoints.register_worker(endpoints.WorkerName(analysis_key_id=key_id), admin)
worker: Final = registration.worker
try:
claimed: Final = await endpoints.claim_candidate(lens, worker, datetime.now(timezone.utc))
claimed: Final = await endpoints.claim_candidate(
lens, worker, datetime.now(timezone.utc), endpoints.repository()
)
assert claimed is not None
for _ in range(3):
with pytest.raises(HTTPException) as failed:

View file

@ -1,3 +1,4 @@
from collections.abc import Callable
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Final
@ -10,6 +11,7 @@ import litellm
from litellm import Router
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.lens.endpoints import (
claim_due,
list_agents,
read_reviews,
result,
@ -35,8 +37,9 @@ from litellm.proxy.lens.models import (
Scope,
TraceFindingsRequest,
TraceIdentity,
Worker,
)
from litellm.proxy.lens.repository import Row
from litellm.proxy.lens.repository import DueLens, Row
from litellm.proxy.lens.state import claim_job, queue_job, replace_job
from tests.unit.proxy.lens.test_agent_workspace import execution
from tests.unit.proxy.lens.test_state import NOW, lens, worker
@ -624,3 +627,53 @@ async def test_unknown_gateway_release_refuses_registration_and_claims(monkeypat
await claim(worker(), protocol_version=PROTOCOL_VERSION, worker_release="")
assert claim_error.value.status_code == 503
assert claim_error.value.detail == registration_error.value.detail
@pytest.mark.asyncio
async def test_claim_due_pages_through_more_than_a_thousand_full_pages() -> None:
candidate_lens: Final = lens()
def candidate_page(page_number: int, size: int) -> tuple[DueLens, ...]:
return tuple(
DueLens(
lens=candidate_lens.model_copy(update={"id": f"lens-{page_number * 20 + offset:05}"}),
due_at=NOW,
)
for offset in range(size)
)
full_pages: Final = tuple(candidate_page(page_number, 20) for page_number in range(1_200))
pages: Final = (*full_pages, candidate_page(1_200, 1))
assigned_worker: Final = worker()
class PagingRepository:
def __init__(self) -> None:
self.after_calls: tuple[DueLens | None, ...] = ()
async def due(
self, scope: Scope, now: datetime, limit: int, after: DueLens | None = None
) -> tuple[DueLens, ...]:
assert scope == assigned_worker.scope
assert now == NOW
assert limit == 20
self.after_calls = (*self.after_calls, after)
return pages[len(self.after_calls) - 1]
async def sync_due(self, lens: Lens) -> None:
return None
async def update(
self, lens_id: str, transform: Callable[[Lens], Lens], attempts: int, *, changed_only: bool
) -> Lens | None:
raise AssertionError("Unsupported models must not update candidates")
async def reject_model(_worker: Worker, _settings: LensSettings) -> bool:
return False
repository: Final = PagingRepository()
claim: Final = await claim_due(assigned_worker, NOW, repository, reject_model)
expected_after: Final = (None, *(page[-1] for page in pages[:-1]))
assert claim is None
assert len(repository.after_calls) == 1_201
assert repository.after_calls == expected_after

View file

@ -37,6 +37,7 @@ from litellm.proxy.lens.state import (
cancel_job,
claim_job,
current_job,
due_at,
end_job,
merge_finding,
next_scan_start,
@ -89,6 +90,22 @@ def worker(team: str = "alpha", identity: str = "worker") -> Worker:
return Worker(id=identity, name=identity, scope=Scope(team_id=team), last_seen=NOW)
def lens_with_job(
status: Literal["queued", "running", "completed"],
lease_until: datetime | None = None,
*,
enabled: bool = True,
trigger: Literal["schedule", "manual"] = "schedule",
) -> Lens:
original: Final = lens()
configured: Final = original.model_copy(
update={"settings": original.settings.model_copy(update={"enabled": enabled})}
)
queued: Final = queue_job(configured, NOW, "job", trigger=trigger)
job: Final = queued.jobs[0].model_copy(update={"status": status, "lease_until": lease_until})
return queued.model_copy(update={"jobs": (job,)})
def finding(execution: str) -> FindingDraft:
return FindingDraft(
title="Repeated failed searches",
@ -98,6 +115,29 @@ def finding(execution: str) -> FindingDraft:
)
@pytest.mark.parametrize(
("candidate", "expected"),
(
pytest.param(lens(), NOW, id="idle-enabled"),
pytest.param(
lens().model_copy(update={"settings": lens().settings.model_copy(update={"enabled": False})}),
None,
id="idle-disabled",
),
pytest.param(lens_with_job("queued", enabled=False, trigger="manual"), NOW, id="queued-manual-while-disabled"),
pytest.param(
lens_with_job("running", NOW + timedelta(minutes=5)),
NOW + timedelta(minutes=5),
id="running-with-lease",
),
pytest.param(lens_with_job("running"), NOW, id="running-without-lease"),
pytest.param(lens_with_job("completed"), NOW, id="completed-only"),
),
)
def test_due_at_matches_the_current_scheduling_state(candidate: Lens, expected: datetime | None) -> None:
assert due_at(candidate) == expected
@pytest.mark.parametrize(
("viewer", "target", "allowed"),
(