From 9fb327e8c6d0d6f24664922d3337117bc5dc9511 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Fri, 2 Oct 2026 14:20:59 -0700 Subject: [PATCH] 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 --- .../migration.sql | 3 + litellm/proxy/lens/endpoints.py | 91 +++++++-- litellm/proxy/lens/inference.py | 23 ++- litellm/proxy/lens/repository.py | 32 +++- tests/proxy_behavior/lens/test_lifecycle.py | 174 +++++++++++++++++- tests/unit/proxy/lens/test_endpoints.py | 99 +++++++++- tests/unit/proxy/lens/test_inference.py | 36 ++++ 7 files changed, 428 insertions(+), 30 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20261002220000_lens_worker_scope_index/migration.sql diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261002220000_lens_worker_scope_index/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261002220000_lens_worker_scope_index/migration.sql new file mode 100644 index 00000000000..124e5713994 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261002220000_lens_worker_scope_index/migration.sql @@ -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; diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index 349ecb9a353..dc0ae8c985c 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -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) diff --git a/litellm/proxy/lens/inference.py b/litellm/proxy/lens/inference.py index 8b306932931..b8a9d7754ae 100644 --- a/litellm/proxy/lens/inference.py +++ b/litellm/proxy/lens/inference.py @@ -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: diff --git a/litellm/proxy/lens/repository.py b/litellm/proxy/lens/repository.py index 4aa840e181b..6e1e2da112a 100644 --- a/litellm/proxy/lens/repository.py +++ b/litellm/proxy/lens/repository.py @@ -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( diff --git a/tests/proxy_behavior/lens/test_lifecycle.py b/tests/proxy_behavior/lens/test_lifecycle.py index 849b2186a62..77e1421675a 100644 --- a/tests/proxy_behavior/lens/test_lifecycle.py +++ b/tests/proxy_behavior/lens/test_lifecycle.py @@ -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) diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index ca19277da08..a1441b34ffa 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -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)) diff --git a/tests/unit/proxy/lens/test_inference.py b/tests/unit/proxy/lens/test_inference.py index 2243759b773..3b69d624a7e 100644 --- a/tests/unit/proxy/lens/test_inference.py +++ b/tests/unit/proxy/lens/test_inference.py @@ -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(