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:
moe-berri 2026-10-02 14:20:59 -07:00 • committed by GitHub
parent 0238ec9721
commit 9fb327e8c6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 428 additions and 30 deletions

View file

@ -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;

View file

@ -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)

View file

@ -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:

View file

@ -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(

View file

@ -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)

View file

@ -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))

View file

@ -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(