diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index fea4a342a23..8124fda4d80 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -2149,14 +2149,21 @@ class ModelManagementAuthChecks: raise HTTPException( status_code=400, detail="An auto-router configuration and model catalog are required." ) + incoming: Final = incoming_model_params if incoming_model_params is not None else model_params return await authorize_member_auto_router_write( - incoming=incoming_model_params if incoming_model_params is not None else model_params, + incoming=incoming, existing=model_params if member_operation == "update" else None, user_api_key_dict=user_api_key_dict, team=team_obj, premium_user=premium_user, prisma_client=prisma_client, llm_router=llm_router, + effective_config=TypeAdapter(Mapping[str, object] | None).validate_python( + _effective_complexity_router_config( + incoming.litellm_params, + model_params.litellm_params if member_operation == "update" else None, + ) + ), ) return ModelManagementAuthChecks.can_user_make_team_model_call( diff --git a/litellm/proxy/management_helpers/auto_router_permissions.py b/litellm/proxy/management_helpers/auto_router_permissions.py index 777c05a982a..111ffa7b3d7 100644 --- a/litellm/proxy/management_helpers/auto_router_permissions.py +++ b/litellm/proxy/management_helpers/auto_router_permissions.py @@ -289,6 +289,7 @@ async def authorize_member_auto_router_write( premium_user: bool, prisma_client: DatabaseClient, llm_router: Router, + effective_config: Mapping[str, object] | None = None, ) -> MemberAutoRouterWrite: authorize_member_auto_router_team(user_api_key_dict=user_api_key_dict, team=team, premium_user=premium_user) stored: Final = StoredAutoRouterIdentity.model_validate(existing.model_dump()) if existing is not None else None @@ -346,7 +347,7 @@ async def authorize_member_auto_router_write( ) if raw_config is None: raise HTTPException(status_code=400, detail="A complexity_router_config is required.") - config: Final = validate_member_auto_router_config(raw_config) + config: Final = validate_member_auto_router_config(effective_config if effective_config is not None else raw_config) stored_default: Final = existing.litellm_params.complexity_router_default_model if existing is not None else None default_model: Final = ( params.complexity_router_default_model diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 1e774577163..00e38977f8f 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -55,7 +55,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( from litellm.litellm_core_utils.sensitive_data_masker import mask_credentials_in_payload from litellm.llms.anthropic.common_utils import is_claude_code_user_agent from litellm.llms.base_llm.base_utils import type_to_response_format_param -from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client from litellm.router_strategy.adaptive_router.classifier import classify_prompt from litellm.router_strategy.complexity_router.context_compaction import compaction_pending from litellm.router_strategy.complexity_router.tier_predictor import ( @@ -1308,7 +1308,9 @@ class ComplexityRouter(CustomLogger): """ @staticmethod - def _build_jev_client(config: JevClassifierConfig) -> JevClassifierClient: + def _build_jev_client( + config: JevClassifierConfig, http_client: AsyncHTTPHandler | None = None + ) -> JevClassifierClient: env_prefix: Final = config.provider.upper() api_key: Final = config.api_key or ( get_secret_str(f"{env_prefix}_API_KEY") if config.api_base is None else None @@ -1327,7 +1329,9 @@ class ComplexityRouter(CustomLogger): return HttpJevClassifierClient( api_key=api_key or None, api_base=api_base, - http_client=get_async_httpx_client(httpxSpecialProvider.PassThroughEndpoint), + http_client=http_client + if http_client is not None + else get_async_httpx_client(httpxSpecialProvider.PassThroughEndpoint), provider=config.provider, ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 9ad9af957d1..447062af10f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -9,6 +9,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi.testclient import TestClient +from pydantic import TypeAdapter from litellm._uuid import uuid from litellm.models.credentials import CredentialItem @@ -7613,7 +7614,7 @@ class TestTeamMemberAutoRouterWrites: ], ) async def test_dashboard_save_keeps_the_stored_jev_provider_with_its_own_base( - self, supplied: dict[str, str], carried: dict[str, str] + self, supplied: Mapping[str, str | None], carried: Mapping[str, str] ) -> None: from litellm.proxy.management_endpoints.model_management_endpoints import patch_model @@ -7646,10 +7647,67 @@ class TestTeamMemberAutoRouterWrites: ) with self._environment(database, row): await patch_model(row.model_id, request, UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)) - written: Final = database.db.litellm_proxymodeltable.update.await_args.kwargs["data"] - saved: Final = json.loads(written["litellm_params"])["complexity_router_config"]["jev_classifier_config"] + written: Final = TypeAdapter(str).validate_python( + database.db.litellm_proxymodeltable.update.await_args.kwargs["data"]["litellm_params"] + ) + params: Final = LiteLLM_Params.model_validate_json(written) + config: Final = TypeAdapter(Mapping[str, object]).validate_python(params.complexity_router_config) + saved: Final = TypeAdapter(Mapping[str, object]).validate_python(config["jev_classifier_config"]) assert saved == {**carried, **incoming} + @pytest.mark.asyncio + @pytest.mark.parametrize("endpoint", ["patch", "legacy"]) + @pytest.mark.parametrize("can_use_nimble", [True, False]) + async def test_member_partial_update_authorizes_the_persisted_classifier_provider( + self, endpoint: str, can_use_nimble: bool + ) -> None: + from fastapi import HTTPException + + evaluation: Final = "bespoke_nimble/nimble-latest" if can_use_nimble else "typesafe/jev-latest" + models: Final = ["allowed", evaluation] + team: Final = self._team().model_copy(update={"models": models}) + row: Final = self._row().model_copy( + update={ + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "classifier_type": "jev", + "tiers": {"SIMPLE": "allowed"}, + "jev_classifier_config": {"provider": "bespoke_nimble", "model": "nimble-latest"}, + }, + }, + } + ) + database: Final = self._database(team, row) + request: Final = updateDeployment( + litellm_params=updateLiteLLMParams( + complexity_router_config={ + "classifier_type": "jev", + "tiers": {"SIMPLE": "allowed"}, + "jev_classifier_config": {"timeout_ms": 900}, + } + ), + model_info=ModelInfo(id=row.model_id, team_id=team.team_id), + ) + actor: Final = UserAPIKeyAuth(user_id="owner", user_role=LitellmUserRoles.INTERNAL_USER, models=models) + with self._environment(database, row): + operation: Final = patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor) + if not can_use_nimble: + with pytest.raises((HTTPException, ProxyException)) as denied: + await operation + assert "bespoke_nimble/nimble-latest" in str(denied.value) + database.transaction.litellm_proxymodeltable.update.assert_not_awaited() + return + await operation + written: Final = TypeAdapter(str).validate_python( + database.transaction.litellm_proxymodeltable.update.await_args.kwargs["data"]["litellm_params"] + ) + params: Final = LiteLLM_Params.model_validate_json(written) + config: Final = TypeAdapter(Mapping[str, object]).validate_python(params.complexity_router_config) + assert config["jev_classifier_config"] == { + "provider": "bespoke_nimble", "model": "nimble-latest", "timeout_ms": 900 + } + @pytest.mark.asyncio @pytest.mark.parametrize("endpoint", ["patch", "legacy"]) @pytest.mark.parametrize("access", ["owner", "peer", "limited-key"]) diff --git a/tests/unit/router_strategy/complexity_router/test_jev_classifier.py b/tests/unit/router_strategy/complexity_router/test_jev_classifier.py index d3bb1933440..a2c5b6a1d1d 100644 --- a/tests/unit/router_strategy/complexity_router/test_jev_classifier.py +++ b/tests/unit/router_strategy/complexity_router/test_jev_classifier.py @@ -8,6 +8,7 @@ from unittest.mock import create_autospec import httpx import pytest +from pydantic import TypeAdapter import litellm from litellm._logging import verbose_router_logger @@ -579,7 +580,8 @@ async def test_http_jev_classifier_client_posts_to_system_one() -> None: ({"provider": "bespoke_nimble", "model": "nimble-latest"}, {"TYPESAFE_API_BASE": "http://ts.test"}, None), ], ) -def test_jev_provider_credentials_never_leave_their_own_environment_pair( +@pytest.mark.asyncio +async def test_jev_provider_credentials_never_leave_their_own_environment_pair( monkeypatch: pytest.MonkeyPatch, config: Mapping[str, str], environment: Mapping[str, str], @@ -589,18 +591,28 @@ def test_jev_provider_credentials_never_leave_their_own_environment_pair( monkeypatch.delenv(name, raising=False) for name, value in environment.items(): monkeypatch.setenv(name, value) - captured: Final[list[Mapping[str, object]]] = [] - monkeypatch.setattr( - "litellm.router_strategy.complexity_router.complexity_router.HttpJevClassifierClient", - lambda **kwargs: captured.append(kwargs), - ) validated: Final = JevClassifierConfig.model_validate(config) if expected is None: with pytest.raises(ValueError, match="BESPOKE_NIMBLE_API_BASE is required"): ComplexityRouter._build_jev_client(validated) return - ComplexityRouter._build_jev_client(validated) - assert (captured[0]["api_base"], captured[0]["api_key"], captured[0]["provider"]) == (*expected, validated.provider) + base, key = expected + request: Final = build_jev_request("hello", None, validated.model, "Choose a tier", {"SIMPLE": "small talk"}) + answer: Final = JevSystemOneResponse(answers={"tier": _answer()}) + + def respond(sent: httpx.Request) -> httpx.Response: + assert str(sent.url) == f"{base}/v1/systemone" + assert sent.headers.get("authorization") == (f"Bearer {key}" if key else None) + assert json.loads(sent.content) == request.model_dump(mode="json") + return httpx.Response(200, json=answer.model_dump(mode="json")) + + handler: Final = AsyncHTTPHandler(transport=httpx.MockTransport(respond)) + try: + client: Final = ComplexityRouter._build_jev_client(validated, http_client=handler) + assert await client.evaluate(request, 1.0) == answer + await GLOBAL_LOGGING_WORKER.flush() + finally: + await handler.client.aclose() def test_non_typesafe_provider_requires_an_explicit_model() -> None: @@ -612,25 +624,29 @@ def test_non_typesafe_provider_requires_an_explicit_model() -> None: async def test_keyless_bespoke_nimble_classifies_and_reports_under_its_own_provider( monkeypatch: pytest.MonkeyPatch, ) -> None: - logged: Final[list[Mapping[str, object]]] = [] - class _Recorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.events: tuple[tuple[str, str], ...] = () + async def async_log_success_event( self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime ) -> None: if str(kwargs.get("model", "")).endswith("/nimble-accounting"): - logged.append(kwargs) + event: Final = TypeAdapter(tuple[str, str]).validate_python( + (kwargs["model"], kwargs["custom_llm_provider"]) + ) + self.events = (*self.events, event) - monkeypatch.setattr(litellm, "_async_success_callback", [_Recorder()]) + recorder: Final = _Recorder() + monkeypatch.setattr(litellm, "_async_success_callback", [recorder]) monkeypatch.setitem( litellm.model_cost, "bespoke_nimble/nimble-accounting", {"input_cost_per_token": 0.001, "output_cost_per_token": 0.0}, ) - seen: Final[list[httpx.Request]] = [] - def respond(request: httpx.Request) -> httpx.Response: - seen.append(request) + assert "authorization" not in request.headers return httpx.Response( 200, json={ @@ -642,31 +658,30 @@ async def test_keyless_bespoke_nimble_classifies_and_reports_under_its_own_provi }, ) - handler: Final = AsyncHTTPHandler() - handler.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + handler: Final = AsyncHTTPHandler(transport=httpx.MockTransport(respond)) + config: Final = JevClassifierConfig(provider="bespoke_nimble", model="nimble-accounting", api_base="http://nimble.test") router: Final = ComplexityRouter( "nimble-router", litellm.Router(model_list=[]), { "classifier_type": "jev", - "jev_classifier_config": { - "provider": "bespoke_nimble", - "model": "nimble-accounting", - "api_base": "http://n", - }, + "jev_classifier_config": config.model_dump(), "tiers": {"SIMPLE": "cheap"}, }, - jev_client=HttpJevClassifierClient(None, "http://nimble.test", handler, provider="bespoke_nimble"), + jev_client=ComplexityRouter._build_jev_client(config, http_client=handler), derive_savings_baseline=False, ) - outcome: Final = await router.aclassify("hello") - await GLOBAL_LOGGING_WORKER.flush() - await handler.client.aclose() + try: + outcome: Final = await router.async_pre_routing_hook( + model="nimble-router", request_kwargs={}, messages=[{"role": "user", "content": "hello"}] + ) + await GLOBAL_LOGGING_WORKER.flush() + finally: + await handler.client.aclose() - assert "authorization" not in seen[0].headers - assert outcome.cause == "jev_classifier" - assert outcome.jev_verdict is not None - assert (outcome.jev_verdict.provider, outcome.jev_verdict.cost) == ("bespoke_nimble", pytest.approx(0.003)) - assert [(event["model"], event["custom_llm_provider"]) for event in logged] == [ - ("bespoke_nimble/nimble-accounting", "bespoke_nimble") - ] + assert outcome is not None and outcome.model == "cheap" + assert outcome.routing_decision is not None + assert outcome.routing_decision["cause"] == "jev_classifier" + assert outcome.routing_decision["classifier_model"] == "bespoke_nimble/nimble-accounting" + assert outcome.routing_decision["classifier_cost"] == pytest.approx(0.003) + assert recorder.events == (("bespoke_nimble/nimble-accounting", "bespoke_nimble"),) diff --git a/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts b/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts index 478c763351c..ec8af913346 100644 --- a/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts @@ -1,6 +1,7 @@ import { z } from "zod"; const jevClassifierConfigFields = { + provider: z.enum(["typesafe", "bespoke_nimble"]).optional(), model: z.string().trim().min(1).default("jev-latest"), timeout_ms: z.number().int().positive().default(3000), instructions: z @@ -20,6 +21,7 @@ export const defaultJevClassifierConfig = (): JevClassifierConfig => jevClassifi export const normalizeJevClassifierConfig = ( config: JevClassifierConfig = defaultJevClassifierConfig(), ): JevClassifierConfig => ({ + ...(config.provider !== undefined && { provider: config.provider }), model: config.model.trim(), timeout_ms: config.timeout_ms, ...(config.instructions?.trim() && { instructions: config.instructions.trim() }), diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts index 4b8e298191d..5c034a5017c 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts @@ -48,6 +48,35 @@ const hydratedState: KeywordMatchingState = { }; describe("buildUpdatedComplexityRouterConfig keyword matching", () => { + it("preserves Nimble through dashboard validation and saves without exposing transport credentials", () => { + const stored = { + classifier_type: "jev" as const, + tiers: FORM_VALUE.tiers, + jev_classifier_config: { + provider: "bespoke_nimble" as const, + model: "nimble-latest", + timeout_ms: 6100, + api_key: "masked-key", + api_base: "https://nimble.example.com", + }, + }; + const hydrated = hydrateComplexityRouterConfig(stored, undefined); + expect(hydrated.jev_classifier_config).toEqual({ + provider: "bespoke_nimble", + model: "nimble-latest", + timeout_ms: 6100, + }); + const edited = { + ...hydrated, + jev_classifier_config: { ...hydrated.jev_classifier_config!, timeout_ms: 900 }, + }; + expect(buildUpdatedComplexityRouterConfig(stored, edited).jev_classifier_config).toEqual({ + provider: "bespoke_nimble", + model: "nimble-latest", + timeout_ms: 900, + }); + }); + it.each([false, true])("omits masked JEV credentials from dashboard saves, edited: %s", (edited) => { const stored = { classifier_type: "jev" as const,