mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
* chore(lens): remove deployment screenshots * refactor(lens)!: rename internal engine package and API * fix(lens): pin worker image for renamed API * test(lens): cover fresh and populated rename migrations * fix(lens): protect db-push upgrades and restore routing and CI * fix(lens): resolve migration tables across schemas and include database driver
227 lines
10 KiB
Python
227 lines
10 KiB
Python
import asyncio
|
|
import hashlib
|
|
import os
|
|
from collections.abc import AsyncIterator
|
|
from datetime import datetime, timedelta, timezone
|
|
from typing import Final
|
|
from uuid import uuid4
|
|
|
|
import pytest
|
|
import pytest_asyncio
|
|
from fastapi import HTTPException, Request
|
|
from fastapi.security import HTTPAuthorizationCredentials
|
|
|
|
from litellm import Router
|
|
from litellm.proxy import proxy_server
|
|
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
|
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|
from litellm.proxy.lens import endpoints
|
|
from litellm.proxy.lens.models import Check, Coverage, LensSettings, ModelRequest, Progress, Result, RunRequest
|
|
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
|
|
|
|
|
@pytest_asyncio.fixture(loop_scope="function")
|
|
async def lens_database() -> AsyncIterator[PrismaClient]:
|
|
original_db: Final = proxy_server.prisma_client
|
|
original_router: Final = proxy_server.llm_router
|
|
original_settings: Final = proxy_server.general_settings
|
|
proxy_server.general_settings = {
|
|
**original_settings,
|
|
"allowed_ips": ["127.0.0.1"],
|
|
"use_x_forwarded_for": True,
|
|
"mcp_trusted_proxy_ranges": ["192.0.2.100/32"],
|
|
"mcp_xff_num_trusted_hops": 1,
|
|
}
|
|
client: Final = PrismaClient(os.environ["DATABASE_URL"], ProxyLogging(UserApiKeyCache()))
|
|
await client.connect()
|
|
proxy_server.prisma_client = client
|
|
proxy_server.llm_router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "lens-test-analysis",
|
|
"litellm_params": {
|
|
"model": "openai/lens-test-analysis",
|
|
"api_key": "test-only",
|
|
"mock_response": '{"observations":[]}',
|
|
"input_cost_per_token": 0.000001,
|
|
"output_cost_per_token": 0.000002,
|
|
},
|
|
}
|
|
]
|
|
)
|
|
try:
|
|
yield client
|
|
finally:
|
|
proxy_server.general_settings = original_settings
|
|
proxy_server.prisma_client = original_db
|
|
proxy_server.llm_router = original_router
|
|
await client.disconnect()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database: PrismaClient) -> None:
|
|
admin: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
|
settings: Final = LensSettings(
|
|
name="Lifecycle regression",
|
|
model="lens-test-analysis",
|
|
enabled=False,
|
|
checks=(Check(id="retries", instruction="Find unrecovered retries"),),
|
|
)
|
|
lens: Final = await endpoints.create_lens(settings, admin)
|
|
key_id: Final = hashlib.sha256(uuid4().bytes).hexdigest()
|
|
await lens_database.db.litellm_verificationtoken.create(data={"token": key_id, "models": ["lens-test-analysis"]})
|
|
registration: Final = await endpoints.register_worker(
|
|
endpoints.WorkerName(name="Test analyzer", analysis_key_id=key_id), admin
|
|
)
|
|
credentials: Final = HTTPAuthorizationCredentials(scheme="Bearer", credentials=registration.token)
|
|
worker: Final = await endpoints.worker_auth(credentials)
|
|
try:
|
|
assert lens.jobs[0].status == "queued"
|
|
stored_worker: Final = await endpoints.repository().worker(
|
|
hashlib.sha256(registration.token.encode()).hexdigest()
|
|
)
|
|
assert stored_worker is not None and stored_worker.id == worker.id
|
|
assert worker.id == registration.worker.id
|
|
listing: Final = await endpoints.list_lenses(admin)
|
|
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))
|
|
)
|
|
winners: Final = tuple(claim for claim in claims if claim is not None)
|
|
assert len(winners) == 1
|
|
claimed: Final = winners[0]
|
|
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)
|
|
)
|
|
is None
|
|
)
|
|
assert await endpoints.progress(
|
|
lens.id, claimed.job.id, Progress(stage="Reviewing", coverage=Coverage(screened=2)), worker
|
|
)
|
|
assert await endpoints.heartbeat(lens.id, claimed.job.id, worker)
|
|
response: Final = await endpoints.model(
|
|
lens.id,
|
|
claimed.job.id,
|
|
ModelRequest(prompt="Return an empty observations list", purpose="extract"),
|
|
worker,
|
|
Request(
|
|
{
|
|
"type": "http",
|
|
"scheme": "http",
|
|
"path": "/lens/worker/model",
|
|
"headers": [],
|
|
"client": ("127.0.0.1", 1234),
|
|
}
|
|
),
|
|
)
|
|
assert '"observations"' in response.content
|
|
with pytest.raises(HTTPException) as denied_ip:
|
|
await endpoints.model(
|
|
lens.id,
|
|
claimed.job.id,
|
|
ModelRequest(prompt="Must not run", purpose="extract"),
|
|
worker,
|
|
Request(
|
|
{
|
|
"type": "http",
|
|
"scheme": "http",
|
|
"path": "/lens/worker/model",
|
|
"headers": [(b"x-forwarded-for", b"127.0.0.1")],
|
|
"client": ("192.0.2.1", 1234),
|
|
}
|
|
),
|
|
)
|
|
assert denied_ip.value.status_code == 403
|
|
forwarded: Final = await endpoints.model(
|
|
lens.id,
|
|
claimed.job.id,
|
|
ModelRequest(prompt="Return an empty observations list", purpose="extract"),
|
|
worker,
|
|
Request(
|
|
{
|
|
"type": "http",
|
|
"scheme": "http",
|
|
"path": "/lens/worker/model",
|
|
"headers": [(b"x-forwarded-for", b"127.0.0.1")],
|
|
"client": ("192.0.2.100", 1234),
|
|
}
|
|
),
|
|
)
|
|
assert '"observations"' in forwarded.content
|
|
with pytest.raises(HTTPException) as spoofed_chain:
|
|
await endpoints.model(
|
|
lens.id,
|
|
claimed.job.id,
|
|
ModelRequest(prompt="Must not run", purpose="extract"),
|
|
worker,
|
|
Request(
|
|
{
|
|
"type": "http",
|
|
"scheme": "http",
|
|
"path": "/lens/worker/model",
|
|
"headers": [(b"x-forwarded-for", b"127.0.0.1, 192.0.2.1")],
|
|
"client": ("192.0.2.100", 1234),
|
|
}
|
|
),
|
|
)
|
|
assert spoofed_chain.value.status_code == 403
|
|
charged: Final = await endpoints.get_lens(lens.id, worker.scope)
|
|
assert charged.spent == pytest.approx(response.cost + forwarded.cost)
|
|
assert charged.jobs[0].cost == pytest.approx(response.cost + forwarded.cost)
|
|
legacy: Final = worker.model_copy(update={"analysis_key_id": None})
|
|
await endpoints.repository().save_worker(legacy)
|
|
authenticated_legacy: Final = await endpoints.worker_auth(credentials)
|
|
assert authenticated_legacy.analysis_key_id is None
|
|
with pytest.raises(HTTPException) as needs_billing:
|
|
await endpoints.claim(authenticated_legacy, protocol_version=2)
|
|
assert needs_billing.value.status_code == 409
|
|
assert await endpoints.heartbeat(lens.id, claimed.job.id, authenticated_legacy)
|
|
finished: Final = await endpoints.result(
|
|
lens.id, claimed.job.id, Result(coverage=Coverage(screened=2)), authenticated_legacy
|
|
)
|
|
assert finished.jobs[0].status == "completed"
|
|
assert finished.jobs[0].coverage.screened == 2
|
|
assert finished.last_scan_at == claimed.job.end
|
|
assert finished.next_run_at > finished.jobs[0].finished_at
|
|
assert await endpoints.result(lens.id, claimed.job.id, Result(coverage=Coverage()), worker) == finished
|
|
with pytest.raises(HTTPException) as stale:
|
|
await endpoints.heartbeat(lens.id, claimed.job.id, worker)
|
|
assert stale.value.status_code == 409
|
|
edited: Final = await endpoints.update_lens(
|
|
lens.id, settings.model_copy(update={"interval_minutes": 7}), admin
|
|
)
|
|
assert edited.revision == lens.revision + 1
|
|
rerun: Final = await endpoints.run_lens(lens.id, RunRequest(lookback_hours=3), admin)
|
|
assert rerun.jobs[0].settings.interval_minutes == 7
|
|
assert rerun.jobs[0].created_at - rerun.jobs[0].start == timedelta(hours=3)
|
|
history: Final = await endpoints.list_runs(lens.id, admin, offset=0)
|
|
assert {job.id for job in history} == {claimed.job.id, rerun.jobs[0].id}
|
|
archived: Final = await endpoints.read_run(lens.id, claimed.job.id, admin)
|
|
assert archived == finished.jobs[0]
|
|
assert archived.settings.interval_minutes == 15
|
|
assert archived.findings == ()
|
|
with pytest.raises(HTTPException) as foreign_history:
|
|
await endpoints.read_run(lens.id, claimed.job.id, UserAPIKeyAuth(team_id="other"))
|
|
assert foreign_history.value.status_code == 403
|
|
cancelled: Final = await endpoints.cancel_lens(lens.id, admin)
|
|
assert cancelled.jobs[0].status == "cancelled"
|
|
assert await endpoints.cancel_lens(lens.id, admin) == cancelled
|
|
assert await endpoints.revoke_worker(worker.id, admin)
|
|
assert await endpoints.repository().set_worker_billing(worker.id, key_id) is None
|
|
with pytest.raises(HTTPException) as revoked_billing:
|
|
await endpoints.set_worker_billing(worker.id, endpoints.WorkerBilling(analysis_key_id=key_id), admin)
|
|
assert revoked_billing.value.status_code == 409
|
|
with pytest.raises(HTTPException) as revoked:
|
|
await endpoints.worker_auth(credentials)
|
|
assert revoked.value.status_code == 401
|
|
with pytest.raises(HTTPException) as foreign:
|
|
await endpoints.get_lens(lens.id, endpoints.Scope(team_id="other"))
|
|
assert foreign.value.status_code == 404
|
|
finally:
|
|
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_LensRun" WHERE lens_id=$1', lens.id)
|
|
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', lens.id)
|
|
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_LensWorker" WHERE id=$1', worker.id)
|
|
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_VerificationToken" WHERE token=$1', key_id)
|