mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
d174c43518
commit
65dcf43257
13 changed files with 738 additions and 22 deletions
|
|
@ -0,0 +1 @@
|
|||
ALTER TABLE "LiteLLM_Lens" ADD COLUMN IF NOT EXISTS "due_at" TIMESTAMP(3) DEFAULT '1970-01-01 00:00:00';
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_Lens_due_at_idx"
|
||||
ON "LiteLLM_Lens" ("due_at");
|
||||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)})
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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)))
|
||||
|
||||
|
|
|
|||
199
tests/integration/database/test_lens_scheduler_load.py
Normal file
199
tests/integration/database/test_lens_scheduler_load.py
Normal 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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue