diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000000_lens_due_at/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000000_lens_due_at/migration.sql new file mode 100644 index 00000000000..676d9b124c9 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000000_lens_due_at/migration.sql @@ -0,0 +1 @@ +ALTER TABLE "LiteLLM_Lens" ADD COLUMN IF NOT EXISTS "due_at" TIMESTAMP(3) DEFAULT '1970-01-01 00:00:00'; diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000100_lens_due_at_index/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000100_lens_due_at_index/migration.sql new file mode 100644 index 00000000000..36311c4e749 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000100_lens_due_at_index/migration.sql @@ -0,0 +1,2 @@ +CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_Lens_due_at_idx" +ON "LiteLLM_Lens" ("due_at"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 514df905866..ceb127c31bd 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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 { diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index 21fde94edfe..d0b91aeea2f 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -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) diff --git a/litellm/proxy/lens/repository.py b/litellm/proxy/lens/repository.py index ddbc1aad44a..cbb66f338fa 100644 --- a/litellm/proxy/lens/repository.py +++ b/litellm/proxy/lens/repository.py @@ -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 diff --git a/litellm/proxy/lens/state.py b/litellm/proxy/lens/state.py index 38efc7f23bb..599dca8f1c6 100644 --- a/litellm/proxy/lens/state.py +++ b/litellm/proxy/lens/state.py @@ -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)}) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 514df905866..ceb127c31bd 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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 { diff --git a/schema.prisma b/schema.prisma index 514df905866..ceb127c31bd 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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 { diff --git a/tests/integration/database/test_lens_repository.py b/tests/integration/database/test_lens_repository.py index e29e92a2505..cd2f4996657 100644 --- a/tests/integration/database/test_lens_repository.py +++ b/tests/integration/database/test_lens_repository.py @@ -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))) diff --git a/tests/integration/database/test_lens_scheduler_load.py b/tests/integration/database/test_lens_scheduler_load.py new file mode 100644 index 00000000000..0afa41bea74 --- /dev/null +++ b/tests/integration/database/test_lens_scheduler_load.py @@ -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) diff --git a/tests/proxy_behavior/lens/test_lifecycle.py b/tests/proxy_behavior/lens/test_lifecycle.py index e3027474baa..5f30a11cf70 100644 --- a/tests/proxy_behavior/lens/test_lifecycle.py +++ b/tests/proxy_behavior/lens/test_lifecycle.py @@ -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: diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index 34d51110761..5ac5b756dd2 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -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 diff --git a/tests/unit/proxy/lens/test_state.py b/tests/unit/proxy/lens/test_state.py index 9db69557e28..80eb638a610 100644 --- a/tests/unit/proxy/lens/test_state.py +++ b/tests/unit/proxy/lens/test_state.py @@ -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"), (