mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(lens): run investigations with configured wildcard models (#44233)
* fix(lens): run investigations with configured wildcard models * fix(lens): validate worker model access and pricing before analysis * fix(lens): bound worker validation and preserve unrelated edits
This commit is contained in:
parent
0238ec9721
commit
9fb327e8c6
7 changed files with 428 additions and 30 deletions
|
|
@ -0,0 +1,3 @@
|
|||
CREATE INDEX IF NOT EXISTS "LiteLLM_LensWorker_active_scope_idx"
|
||||
ON "LiteLLM_LensWorker" USING GIN ((data->'scope') jsonb_path_ops)
|
||||
WHERE data @> '{"revoked": false}'::jsonb;
|
||||
|
|
@ -9,12 +9,15 @@ from uuid import uuid4
|
|||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
from pydantic import AwareDatetime, BaseModel, Field, TypeAdapter
|
||||
from pydantic import AwareDatetime, BaseModel, Field
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy._types import LitellmUserRoles, ModelAccessDeniedProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import can_key_call_model
|
||||
from litellm.proxy.auth.resolvers.exceptions import KeyNotFoundError
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.db.routing_prisma_wrapper import writer_wrapper
|
||||
from litellm.proxy.lens.billing import validate_key
|
||||
from litellm.proxy.lens.inference import Deployment, deployment_prices
|
||||
from litellm.proxy.lens.models import (
|
||||
Claim,
|
||||
Execution,
|
||||
|
|
@ -127,20 +130,60 @@ def validate_selection(settings: LensSettings) -> None:
|
|||
raise HTTPException(422, "Choose execution IDs returned by the activity preview")
|
||||
|
||||
|
||||
def validate_model(settings: LensSettings, auth: UserAPIKeyAuth) -> None:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
async def validate_model(settings: LensSettings, auth: UserAPIKeyAuth) -> None:
|
||||
from litellm.proxy.proxy_server import llm_router, prisma_client
|
||||
|
||||
validate_selection(settings)
|
||||
if llm_router is None or settings.model not in llm_router.get_model_names(team_id=auth.team_id):
|
||||
deployments: Final = (
|
||||
llm_router.get_model_list(model_name=settings.model, team_id=auth.team_id) if llm_router else ()
|
||||
)
|
||||
if not deployments:
|
||||
raise HTTPException(400, "Choose a model configured on this LiteLLM instance")
|
||||
allowed_models: Final = TypeAdapter(tuple[str, ...]).validate_python(auth.model_dump().get("models") or ())
|
||||
if (
|
||||
auth.user_role != LitellmUserRoles.PROXY_ADMIN
|
||||
and allowed_models
|
||||
and settings.model not in allowed_models
|
||||
and "all-proxy-models" not in allowed_models
|
||||
):
|
||||
raise HTTPException(403, "This key does not have access to the analysis model")
|
||||
if auth.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
try:
|
||||
await can_key_call_model(
|
||||
model=settings.model,
|
||||
llm_model_list=deployments,
|
||||
valid_token=auth,
|
||||
llm_router=llm_router,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
except ModelAccessDeniedProxyException as exc:
|
||||
raise HTTPException(403, "This key does not have access to the analysis model") from exc
|
||||
for deployment in deployments:
|
||||
deployment_prices(Deployment.model_validate(deployment))
|
||||
|
||||
|
||||
async def worker_supports_model(worker: Worker, settings: LensSettings) -> bool:
|
||||
if worker.revoked or worker.analysis_key_id is None:
|
||||
return False
|
||||
try:
|
||||
auth: Final = await validate_key(worker.analysis_key_id)
|
||||
if auth is None:
|
||||
return False
|
||||
await validate_model(settings, auth)
|
||||
except KeyNotFoundError:
|
||||
return False
|
||||
except HTTPException as exc:
|
||||
if exc.status_code not in (400, 401, 403):
|
||||
raise
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
async def validate_workers(settings: LensSettings, scope: Scope) -> None:
|
||||
workers: Final = repository().eligible_workers(scope)
|
||||
first: Final = await anext(workers, None)
|
||||
if first is None or await worker_supports_model(first, settings):
|
||||
return
|
||||
async for worker in workers:
|
||||
if await worker_supports_model(worker, settings):
|
||||
return
|
||||
raise HTTPException(
|
||||
400,
|
||||
"No worker can use this analysis model. Choose a model available to the worker's virtual key, "
|
||||
"or update its model access and pricing.",
|
||||
)
|
||||
|
||||
|
||||
@router.get("", response_model=LensList)
|
||||
|
|
@ -156,7 +199,8 @@ async def list_lenses(auth: Auth, storage: StorageDep) -> LensList:
|
|||
@router.post("", response_model=Lens)
|
||||
async def create_lens(settings: LensSettings, auth: Auth) -> Lens:
|
||||
scope: Final = user_scope(auth, write=True)
|
||||
validate_model(settings, auth)
|
||||
await validate_model(settings, auth)
|
||||
await validate_workers(settings, scope)
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
lens: Final = Lens(
|
||||
id=str(uuid4()),
|
||||
|
|
@ -183,8 +227,10 @@ async def list_agents(auth: Auth, storage: StorageDep) -> tuple[str, ...]:
|
|||
|
||||
@router.put("/{lens_id}", response_model=Lens)
|
||||
async def update_lens(lens_id: str, settings: LensSettings, auth: Auth) -> Lens:
|
||||
await get_lens(lens_id, user_scope(auth, write=True))
|
||||
validate_model(settings, auth)
|
||||
lens: Final = await get_lens(lens_id, user_scope(auth, write=True))
|
||||
validate_selection(settings)
|
||||
if settings.model != lens.settings.model or (settings.enabled and not lens.settings.enabled):
|
||||
await validate_model(settings, auth)
|
||||
return required(
|
||||
await repository().update(
|
||||
lens_id,
|
||||
|
|
@ -202,9 +248,10 @@ async def update_lens(lens_id: str, settings: LensSettings, auth: Auth) -> Lens:
|
|||
|
||||
@router.post("/{lens_id}/runs", response_model=Lens)
|
||||
async def run_lens(lens_id: str, body: RunRequest, auth: Auth) -> Lens:
|
||||
await get_lens(lens_id, user_scope(auth, write=True))
|
||||
if body.settings is not None:
|
||||
validate_model(body.settings, auth)
|
||||
lens: Final = await get_lens(lens_id, user_scope(auth, write=True))
|
||||
settings: Final = body.settings or lens.settings
|
||||
await validate_model(settings, auth)
|
||||
await validate_workers(settings, lens.scope)
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
job_id: Final = str(uuid4())
|
||||
return required(
|
||||
|
|
@ -528,10 +575,16 @@ async def heartbeat(lens_id: str, job_id: str, worker: WorkerAuth) -> bool:
|
|||
|
||||
|
||||
async def claim_candidate(candidate: Lens, worker: Worker, now: datetime) -> Claim | None:
|
||||
active: Final = current_job(candidate)
|
||||
if not await worker_supports_model(worker, active.settings if active else candidate.settings):
|
||||
return None
|
||||
job_id: Final = str(uuid4())
|
||||
|
||||
def schedule(e: Lens) -> Lens:
|
||||
scheduled: Final = queue_job(e, now, job_id) if e.settings.enabled and e.next_run_at <= now else e
|
||||
job: Final = current_job(scheduled)
|
||||
if job and job.settings.model != (active.settings.model if active else candidate.settings.model):
|
||||
return e
|
||||
return claim_job(scheduled, worker, now)
|
||||
|
||||
updated: Final = await repository().update(candidate.id, schedule, changed_only=True)
|
||||
|
|
|
|||
|
|
@ -3,9 +3,10 @@ from types import MappingProxyType
|
|||
from typing import Final
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
|
||||
import litellm
|
||||
from litellm.exceptions import ModelNotMappedError
|
||||
from litellm.integrations.clickhouse.context import lens_analysis
|
||||
from litellm.litellm_core_utils.initialize_dynamic_callback_params import inherit_message_logging_privacy
|
||||
from litellm.proxy.lens.billing import complete, validate_key
|
||||
|
|
@ -59,6 +60,17 @@ class Prices(BaseModel):
|
|||
input_cost_per_token_above_128k_tokens: float = 0
|
||||
output_cost_per_token_above_128k_tokens: float = 0
|
||||
|
||||
@field_validator(
|
||||
"input_cost_per_token_above_200k_tokens",
|
||||
"output_cost_per_token_above_200k_tokens",
|
||||
"input_cost_per_token_above_128k_tokens",
|
||||
"output_cost_per_token_above_128k_tokens",
|
||||
mode="before",
|
||||
)
|
||||
@classmethod
|
||||
def missing_tier_rate(cls, value: object) -> object:
|
||||
return 0 if value is None else value
|
||||
|
||||
|
||||
def deployment_prices(deployment: Deployment) -> Prices:
|
||||
params: Final = deployment.litellm_params
|
||||
|
|
@ -66,7 +78,14 @@ def deployment_prices(deployment: Deployment) -> Prices:
|
|||
return Prices(
|
||||
input_cost_per_token=params.input_cost_per_token, output_cost_per_token=params.output_cost_per_token
|
||||
)
|
||||
return Prices.model_validate(litellm.get_model_info(model=params.model))
|
||||
try:
|
||||
return Prices.model_validate(litellm.get_model_info(model=params.model))
|
||||
except (ModelNotMappedError, ValueError) as exc:
|
||||
raise HTTPException(
|
||||
400,
|
||||
f"Pricing is not configured for {params.model}. Set input_cost_per_token and output_cost_per_token "
|
||||
"on its deployment before running an investigation.",
|
||||
) from exc
|
||||
|
||||
|
||||
def quote(deployments: tuple[Deployment, ...], prompt: str) -> float:
|
||||
|
|
|
|||
|
|
@ -1,11 +1,12 @@
|
|||
from collections.abc import Awaitable, Callable
|
||||
import json
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Protocol
|
||||
|
||||
from pydantic import BaseModel, JsonValue, TypeAdapter
|
||||
|
||||
from litellm.proxy.db.prisma_client import PrismaWrapper
|
||||
from litellm.proxy.lens.models import Job, Lens, Worker
|
||||
from litellm.proxy.lens.models import Job, Lens, Scope, Worker
|
||||
|
||||
|
||||
class Database(Protocol):
|
||||
|
|
@ -117,6 +118,33 @@ class LensRepository:
|
|||
rows: Final = _ROWS.validate_python(await self.db.query_raw('SELECT data FROM "LiteLLM_LensWorker"'))
|
||||
return tuple(Worker.model_validate(row.data) for row in rows)
|
||||
|
||||
async def eligible_workers(self, scope: Scope) -> AsyncIterator[Worker]:
|
||||
scoped: Final = (
|
||||
{"all_teams": True}
|
||||
if scope.all_teams
|
||||
else {"team_id": scope.team_id}
|
||||
if scope.team_id
|
||||
else {"team_id": "", "api_key_hash": scope.api_key_hash}
|
||||
)
|
||||
cursor = "" # rebind-ok: advance the keyset cursor after each bounded page
|
||||
while True:
|
||||
rows = _ROWS.validate_python(
|
||||
await self.db.query_raw(
|
||||
"""SELECT data FROM "LiteLLM_LensWorker"
|
||||
WHERE data @> '{"revoked": false}'::jsonb AND id > $1
|
||||
AND (data->'scope' @> '{"all_teams": true}'::jsonb OR data->'scope' @> $2::jsonb)
|
||||
ORDER BY id LIMIT 50""",
|
||||
cursor,
|
||||
json.dumps(scoped),
|
||||
)
|
||||
)
|
||||
workers = tuple(Worker.model_validate(row.data) for row in rows)
|
||||
for worker in workers:
|
||||
yield worker
|
||||
if len(workers) < 50:
|
||||
return
|
||||
cursor = workers[-1].id
|
||||
|
||||
async def worker(self, token_hash: str) -> Worker | None:
|
||||
rows: Final = _ROWS.validate_python(
|
||||
await self.db.query_raw(
|
||||
|
|
|
|||
|
|
@ -10,13 +10,27 @@ import pytest
|
|||
import pytest_asyncio
|
||||
from fastapi import HTTPException, Request
|
||||
from fastapi.security import HTTPAuthorizationCredentials
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
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.lens.models import (
|
||||
Check,
|
||||
Coverage,
|
||||
Lens,
|
||||
LensSettings,
|
||||
ModelRequest,
|
||||
Progress,
|
||||
Result,
|
||||
RunRequest,
|
||||
Scope,
|
||||
Worker,
|
||||
)
|
||||
from litellm.proxy.lens.repository import Database, LensRepository, Row
|
||||
from litellm.proxy.lens.state import can_access
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
|
||||
|
||||
|
|
@ -46,7 +60,18 @@ async def lens_database() -> AsyncIterator[PrismaClient]:
|
|||
"input_cost_per_token": 0.000001,
|
||||
"output_cost_per_token": 0.000002,
|
||||
},
|
||||
}
|
||||
},
|
||||
{
|
||||
"model_name": "lens-team-route",
|
||||
"model_info": {"team_id": "lens-test-team-a", "team_public_model_name": "private/*"},
|
||||
"litellm_params": {
|
||||
"model": "openai/*",
|
||||
"api_key": "test-only",
|
||||
"input_cost_per_token": 0.000001,
|
||||
"output_cost_per_token": 0.000002,
|
||||
},
|
||||
},
|
||||
{"model_name": "unpriced/*", "litellm_params": {"model": "openai/*", "api_key": "test-only"}},
|
||||
]
|
||||
)
|
||||
try:
|
||||
|
|
@ -58,6 +83,138 @@ async def lens_database() -> AsyncIterator[PrismaClient]:
|
|||
await client.disconnect()
|
||||
|
||||
|
||||
class _ObservedDatabase:
|
||||
def __init__(self, db: Database) -> None:
|
||||
self.db: Final = db
|
||||
self.page_sizes: tuple[int, ...] = ()
|
||||
|
||||
async def query_raw(self, query: str, *args: object) -> object:
|
||||
rows: Final = TypeAdapter(tuple[Row, ...]).validate_python(await self.db.query_raw(query, *args))
|
||||
self.page_sizes = (*self.page_sizes, len(rows))
|
||||
return rows
|
||||
|
||||
async def execute_raw(self, query: str, *args: object) -> int:
|
||||
return await self.db.execute_raw(query, *args)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("kind", ("all", "team", "key"))
|
||||
@pytest.mark.asyncio
|
||||
async def test_eligible_workers_filter_before_bounded_pages(lens_database: PrismaClient, kind: str) -> None:
|
||||
prefix: Final = str(uuid4())
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
scopes: Final = {
|
||||
"all": Scope(all_teams=True),
|
||||
"team": Scope(team_id=prefix),
|
||||
"key": Scope(api_key_hash=prefix),
|
||||
}
|
||||
workers: Final = (
|
||||
*(Worker(id=f"{prefix}-{i:03}", name=prefix, scope=scopes["all"], last_seen=now) for i in range(65)),
|
||||
Worker(id=f"{prefix}-team", name=prefix, scope=scopes["team"], last_seen=now),
|
||||
Worker(id=f"{prefix}-key", name=prefix, scope=scopes["key"], last_seen=now),
|
||||
Worker(id=f"{prefix}-foreign", name=prefix, scope=Scope(team_id="other"), last_seen=now),
|
||||
Worker(id=f"{prefix}-other-key", name=prefix, scope=Scope(api_key_hash="other"), last_seen=now),
|
||||
Worker(id=f"{prefix}-revoked", name=prefix, scope=scopes["all"], last_seen=now, revoked=True),
|
||||
)
|
||||
repo: Final = endpoints.repository()
|
||||
try:
|
||||
for worker in workers:
|
||||
await repo.save_worker(worker, hashlib.sha256(worker.id.encode()).hexdigest())
|
||||
observed: Final = _ObservedDatabase(repo.db)
|
||||
eligible: Final = [worker async for worker in LensRepository(observed).eligible_workers(scopes[kind])]
|
||||
expected: Final = tuple(w for w in workers if not w.revoked and can_access(w.scope, scopes[kind]))
|
||||
assert tuple(w.id for w in eligible) == tuple(sorted(w.id for w in expected))
|
||||
assert observed.page_sizes == (50, len(expected) - 50)
|
||||
finally:
|
||||
await lens_database.db.execute_raw("DELETE FROM \"LiteLLM_LensWorker\" WHERE data->>'name'=$1", prefix)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("enabled", (True, False))
|
||||
@pytest.mark.asyncio
|
||||
async def test_unpriced_saved_model_allows_edits_but_not_new_runs(lens_database: PrismaClient, enabled: bool) -> None:
|
||||
admin: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
original: Final = Lens(
|
||||
id=str(uuid4()),
|
||||
scope=Scope(all_teams=True),
|
||||
created_at=now,
|
||||
next_run_at=now,
|
||||
budget_month=now.strftime("%Y-%m"),
|
||||
settings=LensSettings(
|
||||
name="Saved investigation", model="unpriced/lens-saved-model", context="Answer questions", enabled=enabled
|
||||
),
|
||||
)
|
||||
await endpoints.repository().create(original)
|
||||
try:
|
||||
settings: Final = original.settings.model_copy(update={"context": "Use cited sources", "enabled": False})
|
||||
edited: Final = await endpoints.update_lens(original.id, settings, admin)
|
||||
assert edited.settings == settings
|
||||
assert edited.revision == original.revision + 1
|
||||
assert (await endpoints.read_lens(original.id, admin)).settings == settings
|
||||
for operation in (
|
||||
endpoints.run_lens(original.id, RunRequest(), admin),
|
||||
endpoints.update_lens(original.id, settings.model_copy(update={"enabled": True}), admin),
|
||||
endpoints.update_lens(original.id, settings.model_copy(update={"model": "unpriced/other-model"}), admin),
|
||||
):
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await operation
|
||||
assert error.value.status_code == 400
|
||||
assert "Pricing is not configured" in error.value.detail
|
||||
with pytest.raises(HTTPException) as invalid_selection:
|
||||
await endpoints.update_lens(original.id, settings.model_copy(update={"execution_ids": ("invalid",)}), admin)
|
||||
assert invalid_selection.value.status_code == 422
|
||||
assert (await endpoints.read_lens(original.id, admin)).settings == settings
|
||||
finally:
|
||||
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', original.id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_route_requires_a_worker_with_matching_model_access(lens_database: PrismaClient) -> None:
|
||||
admin: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, team_id="lens-test-team-a")
|
||||
name: Final = f"Team route regression {uuid4()}"
|
||||
settings: Final = LensSettings(name=name, model="private/analysis", context="Answer questions", enabled=False)
|
||||
lens: Final = await endpoints.create_lens(settings, admin)
|
||||
key_a: Final = hashlib.sha256(uuid4().bytes).hexdigest()
|
||||
key_b: Final = hashlib.sha256(uuid4().bytes).hexdigest()
|
||||
await lens_database.db.litellm_verificationtoken.create(
|
||||
data={"token": key_a, "team_id": "lens-test-team-a", "models": ["private/*"]}
|
||||
)
|
||||
await lens_database.db.litellm_verificationtoken.create(
|
||||
data={"token": key_b, "team_id": "lens-test-team-b", "models": ["private/*"]}
|
||||
)
|
||||
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
|
||||
for operation in (
|
||||
endpoints.create_lens(settings, admin),
|
||||
endpoints.run_lens(lens.id, RunRequest(), admin),
|
||||
):
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await operation
|
||||
assert error.value.status_code == 400
|
||||
assert "worker" in error.value.detail
|
||||
edited: Final = await endpoints.update_lens(
|
||||
lens.id, settings.model_copy(update={"context": "Use sources"}), admin
|
||||
)
|
||||
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))
|
||||
assert claim is not None and claim.job.worker_id == right_team.worker.id
|
||||
finally:
|
||||
await lens_database.db.execute_raw(
|
||||
"""DELETE FROM "LiteLLM_LensRun" WHERE lens_id IN
|
||||
(SELECT id FROM "LiteLLM_Lens" WHERE data->'settings'->>'name'=$1)""",
|
||||
name,
|
||||
)
|
||||
await lens_database.db.execute_raw("DELETE FROM \"LiteLLM_Lens\" WHERE data->'settings'->>'name'=$1", name)
|
||||
await lens_database.db.execute_raw(
|
||||
"DELETE FROM \"LiteLLM_LensWorker\" WHERE data->>'analysis_key_id' IN ($1, $2)", key_a, key_b
|
||||
)
|
||||
await lens_database.db.execute_raw(
|
||||
'DELETE FROM "LiteLLM_VerificationToken" WHERE token IN ($1, $2)', key_a, key_b
|
||||
)
|
||||
|
||||
|
||||
@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)
|
||||
|
|
@ -186,14 +343,19 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
|
|||
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, storage=None) == finished
|
||||
assert (
|
||||
await endpoints.result(lens.id, claimed.job.id, Result(coverage=Coverage()), worker, storage=None)
|
||||
== 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
|
||||
)
|
||||
edited: Final = await endpoints.update_lens(lens.id, settings.model_copy(update={"interval_minutes": 7}), admin)
|
||||
assert edited.revision == lens.revision + 1
|
||||
with pytest.raises(HTTPException) as unavailable_worker:
|
||||
await endpoints.run_lens(lens.id, RunRequest(lookback_hours=3), admin)
|
||||
assert unavailable_worker.value.status_code == 400
|
||||
await endpoints.set_worker_billing(worker.id, endpoints.WorkerBilling(analysis_key_id=key_id), admin)
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -3,8 +3,105 @@ from typing import Final
|
|||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.lens.endpoints import list_agents, user_scope
|
||||
from litellm import Router
|
||||
from litellm.proxy.lens.endpoints import list_agents, user_scope, validate_model, worker_supports_model
|
||||
from litellm.proxy.lens.models import LensSettings
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def analysis_router(monkeypatch: pytest.MonkeyPatch) -> Router:
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setattr(litellm, "model_cost", {**litellm.model_cost})
|
||||
router: Final = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "openai/*",
|
||||
"litellm_params": {
|
||||
"model": "openai/*",
|
||||
"api_key": "test-key",
|
||||
"input_cost_per_token": 0.001,
|
||||
"output_cost_per_token": 0.002,
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "analysis",
|
||||
"litellm_params": {
|
||||
"model": "openai/test-analysis",
|
||||
"api_key": "test-key",
|
||||
"input_cost_per_token": 0.001,
|
||||
"output_cost_per_token": 0.002,
|
||||
},
|
||||
},
|
||||
{"model_name": "unpriced/*", "litellm_params": {"model": "openai/*", "api_key": "test-key"}},
|
||||
],
|
||||
model_group_alias={"analysis-alias": "analysis"},
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
return router
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ("openai/test-analysis", "analysis", "analysis-alias"))
|
||||
@pytest.mark.asyncio
|
||||
async def test_analysis_accepts_models_served_by_configured_routes(analysis_router: Router, model: str) -> None:
|
||||
settings: Final = LensSettings(name="Research", model=model, context="Answer using cited sources")
|
||||
auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
assert analysis_router.get_model_list(model_name=model)
|
||||
await validate_model(settings, auth)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ("unconfigured", "anthropic/test-analysis"))
|
||||
@pytest.mark.asyncio
|
||||
async def test_analysis_rejects_models_without_a_configured_route(analysis_router: Router, model: str) -> None:
|
||||
settings: Final = LensSettings(name="Research", model=model, context="Answer using cited sources")
|
||||
auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
assert not analysis_router.get_model_list(model_name=model)
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await validate_model(settings, auth)
|
||||
assert error.value.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_analysis_route_resolution_preserves_key_model_restrictions(analysis_router: Router) -> None:
|
||||
settings: Final = LensSettings(name="Research", model="openai/test-analysis", context="Answer using cited sources")
|
||||
auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, models=["analysis"])
|
||||
assert analysis_router.get_model_list(model_name=settings.model)
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await validate_model(settings, auth)
|
||||
assert error.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_analysis_rejects_unpriced_wildcard_before_creating_a_run(analysis_router: Router) -> None:
|
||||
settings: Final = LensSettings(name="Research", model="unpriced/lens-unpriced-test", context="Answer questions")
|
||||
auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
assert analysis_router.get_model_list(model_name=settings.model)
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await validate_model(settings, auth)
|
||||
assert error.value.status_code == 400
|
||||
assert "Pricing is not configured" in error.value.detail
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model,allowed", (("openai/test-analysis", "openai/*"), ("analysis-alias", "analysis")))
|
||||
@pytest.mark.asyncio
|
||||
async def test_analysis_key_accepts_wildcard_and_alias_access(
|
||||
analysis_router: Router, model: str, allowed: str
|
||||
) -> None:
|
||||
auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, models=[allowed])
|
||||
assert analysis_router.get_model_list(model_name=model)
|
||||
await validate_model(LensSettings(name="Research", model=model, context="Answer questions"), auth)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revoked,key_id", ((True, "a" * 64), (False, None)))
|
||||
@pytest.mark.asyncio
|
||||
async def test_worker_without_active_billing_cannot_take_work(revoked: bool, key_id: str | None) -> None:
|
||||
from tests.unit.proxy.lens.test_state import worker
|
||||
|
||||
inactive: Final = worker().model_copy(update={"revoked": revoked, "analysis_key_id": key_id})
|
||||
settings: Final = LensSettings(name="Research", model="analysis", context="Answer questions")
|
||||
assert not await worker_supports_model(inactive, settings)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("role", (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY))
|
||||
|
|
|
|||
|
|
@ -1,11 +1,47 @@
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm.proxy.lens.inference import Deployment, DeploymentParams, completion_charge, quote
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
def test_missing_optional_price_tiers_use_base_rates(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(litellm, "model_cost", {**litellm.model_cost})
|
||||
litellm.register_model(
|
||||
model_cost={
|
||||
"openai/lens-base-rate-test": {
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
"input_cost_per_token": 0.001,
|
||||
"output_cost_per_token": 0.002,
|
||||
"input_cost_per_token_above_200k_tokens": None,
|
||||
"output_cost_per_token_above_200k_tokens": None,
|
||||
"input_cost_per_token_above_128k_tokens": None,
|
||||
"output_cost_per_token_above_128k_tokens": None,
|
||||
}
|
||||
}
|
||||
)
|
||||
deployment: Final = Deployment(litellm_params=DeploymentParams(model="openai/lens-base-rate-test"))
|
||||
explicit: Final = Deployment(
|
||||
litellm_params=DeploymentParams(
|
||||
model="openai/lens-base-rate-test", input_cost_per_token=0.001, output_cost_per_token=0.002
|
||||
)
|
||||
)
|
||||
assert quote((deployment,), "Answer the question") == quote((explicit,), "Answer the question")
|
||||
|
||||
|
||||
def test_unpriced_model_requires_explicit_rates() -> None:
|
||||
deployment: Final = Deployment(litellm_params=DeploymentParams(model="openai/lens-unpriced-test"))
|
||||
with pytest.raises(HTTPException) as error:
|
||||
quote((deployment,), "Answer the question")
|
||||
assert error.value.status_code == 400
|
||||
assert "input_cost_per_token" in error.value.detail
|
||||
assert "output_cost_per_token" in error.value.detail
|
||||
|
||||
|
||||
def test_custom_priced_model_charges_reported_tokens() -> None:
|
||||
deployment: Final = Deployment(
|
||||
litellm_params=DeploymentParams(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue