mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(auto-router): preserve Nimble provider during member edits
This commit is contained in:
parent
91dc4414f9
commit
337e5bde07
7 changed files with 157 additions and 41 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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"),)
|
||||
|
|
|
|||
|
|
@ -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() }),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue