mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
feat(complexity-router): add a bundled Nadir classifier plugin
classifier_type: custom already accepts any classifier that answers a tier. This ships one for Nadir's /v1/bucket endpoint, a trained complexity classifier rather than an LLM call, whose simple/medium/complex verdict maps to the SIMPLE/MEDIUM/COMPLEX tiers. A router names the module-level instance by dotted path; a renamed or custom tier set passes its own names as tier_map. Decision-only, so nothing moves off the proxy: the tier's model pool, the provider call, the operator's keys, fallbacks and spend tracking are unchanged, and Nadir sees only the messages it classifies. Every failure is the existing one, since a decline, an exception or a call past classifier_plugin_timeout_ms falls through to classifier_fallback. NADIR_API_KEY attributes decisions to an account and lifts the anonymous rate limit; NADIR_API_BASE points at a self-hosted deployment. The URL builder tolerates a base that already ends in /v1, which is how Nadir's own docs advertise it for OpenAI-compatible clients, so the common paste does not produce /v1/v1/bucket. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
parent
9cdedf81cd
commit
e460907f00
4 changed files with 378 additions and 0 deletions
|
|
@ -258,6 +258,64 @@ Spend logs record `routing_decision.cause: heuristic_v2`, the detected request
|
|||
type, and all four predicted probabilities. Existing `classifier_type: heuristic`
|
||||
configurations keep the original weighted scorer unchanged
|
||||
|
||||
### Custom classifier plugins
|
||||
|
||||
`classifier_type: custom` hands the tier decision to a plugin of your own. The plugin is an
|
||||
object with one `async classify(context) -> str | None` method, named in the config by the
|
||||
dotted path to an instance. It receives a `RoutingContext` (the request messages, both raw
|
||||
and normalized to chat-completions shape, the candidate models, and the request metadata
|
||||
including caller identity) and returns the name of the tier to route to, or `None` to
|
||||
decline. A decline, an exception, or a call that overruns `classifier_plugin_timeout_ms`
|
||||
(default 3000) all fall through to `classifier_fallback`, so a classifier that is down
|
||||
cannot fail a completion.
|
||||
|
||||
Everything downstream of the tier is unchanged: the tier's model pool, the provider call,
|
||||
your provider keys, fallbacks, and spend tracking all stay on this proxy.
|
||||
|
||||
#### Nadir
|
||||
|
||||
`litellm.router_strategy.complexity_router.nadir_classifier.nadir_classifier` is a bundled
|
||||
plugin that asks [Nadir](https://getnadir.com)'s `/v1/bucket` endpoint, a trained
|
||||
complexity classifier rather than an LLM call, which of `simple` / `medium` / `complex` a
|
||||
request is. Those map to the SIMPLE / MEDIUM / COMPLEX tiers:
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: smart-router
|
||||
litellm_params:
|
||||
model: auto_router/complexity_router
|
||||
complexity_router_config:
|
||||
classifier_type: custom
|
||||
classifier_plugin: litellm.router_strategy.complexity_router.nadir_classifier.nadir_classifier
|
||||
classifier_fallback: heuristic
|
||||
tiers:
|
||||
SIMPLE: gpt-4o-mini
|
||||
MEDIUM: gpt-4o
|
||||
COMPLEX: o1-preview
|
||||
```
|
||||
|
||||
`NADIR_API_KEY` attributes decisions to an account and lifts the anonymous rate limit;
|
||||
without it the endpoint still answers, burst-limited per IP and stored nowhere.
|
||||
`NADIR_API_BASE` points at a self-hosted or on-prem Nadir instead of the hosted API.
|
||||
|
||||
A router whose tiers are renamed with `tier_labels`, or defined with `tier_definitions`,
|
||||
constructs the classifier with the matching names instead of using the module instance:
|
||||
|
||||
```python
|
||||
from litellm.router_strategy.complexity_router.nadir_classifier import NadirComplexityClassifier
|
||||
|
||||
classifier = NadirComplexityClassifier(tier_map={"simple": "Cheap", "medium": "Standard", "complex": "Premium"})
|
||||
```
|
||||
|
||||
Two properties to weigh against the local scorers. The classification is a network call, so
|
||||
it costs a round trip per request where `heuristic` and `heuristic_v2` cost under a
|
||||
millisecond (a laptop against the hosted API measured about 160ms warm, and the first call
|
||||
in a fresh process also pays connection setup; anything that overruns
|
||||
`classifier_plugin_timeout_ms` routes on the fallback instead of waiting). And the messages are sent to the configured Nadir host, which is a third party
|
||||
unless that host is yours, the same disclosure `classifier_type: llm` carries when the
|
||||
classifier model is a hosted one. A REASONING tier is never returned, since Nadir grades
|
||||
three buckets; keyword rules and the reasoning override still place requests there.
|
||||
|
||||
### Renaming the tiers
|
||||
|
||||
`tier_labels` puts your own vocabulary on the four tiers:
|
||||
|
|
|
|||
129
litellm/router_strategy/complexity_router/nadir_classifier.py
Normal file
129
litellm/router_strategy/complexity_router/nadir_classifier.py
Normal file
|
|
@ -0,0 +1,129 @@
|
|||
"""
|
||||
Nadir as the Complexity Router's classifier.
|
||||
|
||||
A decision-only integration: the tier comes from a call to Nadir's ``/v1/bucket``
|
||||
endpoint, and everything else stays here. The tier's model pool, the provider call, the
|
||||
operator's own provider keys, fallbacks, spend tracking and the response path are
|
||||
untouched, so Nadir sees the messages it is asked to classify and never the completion.
|
||||
|
||||
``/v1/bucket`` runs a trained complexity classifier (a fine-tuned encoder, not an LLM
|
||||
call) and answers ``simple`` / ``medium`` / ``complex``, which map to the router's
|
||||
default SIMPLE / MEDIUM / COMPLEX tiers. A router whose tiers are renamed with
|
||||
``tier_labels``, or defined with ``tier_definitions``, passes its own names as
|
||||
``tier_map``.
|
||||
|
||||
Configure it in the proxy by pointing ``classifier_plugin`` at the module-level
|
||||
instance::
|
||||
|
||||
model_list:
|
||||
- model_name: smart-router
|
||||
litellm_params:
|
||||
model: auto_router/complexity_router
|
||||
complexity_router_config:
|
||||
classifier_type: custom
|
||||
classifier_plugin: litellm.router_strategy.complexity_router.nadir_classifier.nadir_classifier
|
||||
tiers:
|
||||
SIMPLE: gpt-4o-mini
|
||||
MEDIUM: gpt-4o
|
||||
COMPLEX: o1-preview
|
||||
|
||||
``NADIR_API_KEY`` attributes the decision to an account and lifts the anonymous rate
|
||||
limit; without it the endpoint still answers, burst-limited per IP and stored nowhere.
|
||||
``NADIR_API_BASE`` overrides the host for a self-hosted or on-prem deployment.
|
||||
|
||||
Every failure mode is the router's existing one: this returns None to decline, and a
|
||||
network error, a timeout past ``classifier_plugin_timeout_ms`` or an unknown bucket all
|
||||
hand the request to ``classifier_fallback``. Nothing here can fail a completion.
|
||||
|
||||
Privacy note for operators: the messages are sent to the configured Nadir host, which is
|
||||
a third party unless that host is your own. That is the same disclosure a remote LLM
|
||||
classifier carries, and it is the reason the local heuristic scorers exist.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
from litellm.types.router import RoutingContext
|
||||
|
||||
DEFAULT_NADIR_API_BASE: Final = "https://api.getnadir.com"
|
||||
|
||||
DEFAULT_TIER_MAP: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{"simple": "SIMPLE", "medium": "MEDIUM", "complex": "COMPLEX"}
|
||||
)
|
||||
|
||||
_API_PATH: Final = "/v1/bucket"
|
||||
|
||||
|
||||
def _bucket_url(api_base: str) -> str:
|
||||
"""Join the endpoint path to a host, tolerating a base that already ends in ``/v1``.
|
||||
|
||||
Nadir's own docs advertise ``https://api.getnadir.com/v1`` as the base URL, because the
|
||||
OpenAI-compatible clients that consume it append ``/chat/completions``. An operator who
|
||||
copies that value into ``NADIR_API_BASE`` would otherwise send ``/v1/v1/bucket`` and get a
|
||||
404 that reads like the endpoint does not exist.
|
||||
"""
|
||||
base: Final = api_base.rstrip("/").removesuffix("/v1")
|
||||
return f"{base}{_API_PATH}"
|
||||
|
||||
|
||||
class NadirComplexityClassifier:
|
||||
"""Classifier plugin that asks Nadir's decision API which tier a request belongs to.
|
||||
|
||||
Args:
|
||||
api_base: Nadir host. Defaults to ``NADIR_API_BASE``, then to the hosted API.
|
||||
api_key: Nadir API key. Defaults to ``NADIR_API_KEY``; anonymous when unset.
|
||||
tier_map: Nadir bucket name -> the tier name this router routes on. Defaults to the
|
||||
router's built-in tier names.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
tier_map: Mapping[str, str] | None = None,
|
||||
) -> None:
|
||||
self._api_base: Final = api_base
|
||||
self._api_key: Final = api_key
|
||||
# A plain dict, not a MappingProxyType: the proxy deepcopies a deployment's
|
||||
# litellm_params, this instance travels inside them, and a mappingproxy cannot be
|
||||
# deepcopied. Router construction would fail before the first request.
|
||||
self._tier_map: Final[dict[str, str]] = dict( # mutable-ok: read-only after __init__, deepcopy-safe
|
||||
DEFAULT_TIER_MAP if tier_map is None else tier_map
|
||||
)
|
||||
|
||||
def _headers(self) -> dict[str, str]:
|
||||
api_key: Final = self._api_key or os.getenv("NADIR_API_KEY")
|
||||
return {"X-API-Key": api_key} if api_key else {} # mutable-ok: request headers, handed straight to httpx
|
||||
|
||||
async def classify(self, context: RoutingContext) -> str | None:
|
||||
"""Return the tier Nadir places this request in, or None to decline.
|
||||
|
||||
Declines rather than guesses on anything the endpoint cannot grade: a request with no
|
||||
messages, and a bucket name outside ``tier_map`` (which is what a renamed tier set looks
|
||||
like before ``tier_map`` is configured). Both leave the decision to ``classifier_fallback``,
|
||||
where a local scorer still routes the request.
|
||||
"""
|
||||
messages: Final = context.structured_messages or context.raw_messages
|
||||
if not messages:
|
||||
return None
|
||||
api_base: Final = self._api_base or os.getenv("NADIR_API_BASE") or DEFAULT_NADIR_API_BASE
|
||||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.ComplexityClassifier)
|
||||
body: Final = {"messages": list(messages), "source": "litellm"} # mutable-ok: one request body
|
||||
response: Final = await client.post(url=_bucket_url(api_base), json=body, headers=self._headers())
|
||||
payload: Final = response.json()
|
||||
if not isinstance(payload, Mapping):
|
||||
return None
|
||||
bucket: Final = payload.get("bucket")
|
||||
if not isinstance(bucket, str):
|
||||
return None
|
||||
return self._tier_map.get(bucket.strip().lower())
|
||||
|
||||
|
||||
nadir_classifier: Final = NadirComplexityClassifier()
|
||||
"""Module-level instance, so the proxy config can name it by dotted path."""
|
||||
|
|
@ -37,6 +37,7 @@ class httpxSpecialProvider(str, Enum):
|
|||
PasswordBreachCheck = "password_breach_check"
|
||||
ASGI = "asgi"
|
||||
AgentHarness = "agent_harness"
|
||||
ComplexityClassifier = "complexity_classifier"
|
||||
|
||||
|
||||
VerifyTypes = str | bool | ssl.SSLContext
|
||||
|
|
|
|||
|
|
@ -0,0 +1,190 @@
|
|||
"""Tests for the Nadir classifier plugin (litellm/router_strategy/complexity_router/nadir_classifier.py)."""
|
||||
|
||||
import copy
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.router_strategy.complexity_router import ComplexityRouter
|
||||
from litellm.router_strategy.complexity_router.nadir_classifier import (
|
||||
DEFAULT_NADIR_API_BASE,
|
||||
NadirComplexityClassifier,
|
||||
_bucket_url,
|
||||
nadir_classifier,
|
||||
)
|
||||
from litellm.types.router import ClassifierPlugin, RoutingContext
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, payload):
|
||||
self._payload = payload
|
||||
|
||||
def json(self):
|
||||
return self._payload
|
||||
|
||||
|
||||
class _FakeClient:
|
||||
"""Records the one request the classifier makes, and answers with a canned bucket."""
|
||||
|
||||
def __init__(self, payload=None, error=None):
|
||||
self._payload = payload
|
||||
self._error = error
|
||||
self.calls = []
|
||||
|
||||
async def post(self, url, json=None, headers=None):
|
||||
self.calls.append({"url": url, "json": json, "headers": headers})
|
||||
if self._error is not None:
|
||||
raise self._error
|
||||
return _FakeResponse(self._payload)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_client(monkeypatch):
|
||||
def _install(payload=None, error=None):
|
||||
client = _FakeClient(payload=payload, error=error)
|
||||
monkeypatch.setattr(
|
||||
"litellm.router_strategy.complexity_router.nadir_classifier.get_async_httpx_client",
|
||||
lambda llm_provider: client,
|
||||
)
|
||||
return client
|
||||
|
||||
return _install
|
||||
|
||||
|
||||
def _context(messages=None):
|
||||
resolved = [{"role": "user", "content": "refactor the retry loop"}] if messages is None else messages
|
||||
return RoutingContext(
|
||||
raw_messages=resolved,
|
||||
structured_messages=resolved,
|
||||
candidate_models=["gpt-4o-mini", "gpt-4o"],
|
||||
)
|
||||
|
||||
|
||||
class TestBucketURL:
|
||||
"""The endpoint URL, including the documented double-/v1 footgun."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_base",
|
||||
[
|
||||
"https://api.getnadir.com",
|
||||
"https://api.getnadir.com/",
|
||||
"https://api.getnadir.com/v1",
|
||||
"https://api.getnadir.com/v1/",
|
||||
],
|
||||
)
|
||||
def test_v1_is_never_doubled(self, api_base):
|
||||
"""Nadir's docs advertise the base URL with /v1, so operators paste it with and without."""
|
||||
assert _bucket_url(api_base) == "https://api.getnadir.com/v1/bucket"
|
||||
|
||||
def test_self_hosted_host_with_a_path_prefix_is_preserved(self):
|
||||
assert _bucket_url("https://gateway.internal/nadir") == "https://gateway.internal/nadir/v1/bucket"
|
||||
|
||||
|
||||
class TestClassify:
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("bucket", "expected_tier"), [("simple", "SIMPLE"), ("medium", "MEDIUM"), ("complex", "COMPLEX")]
|
||||
)
|
||||
async def test_each_bucket_maps_to_its_default_tier(self, fake_client, bucket, expected_tier):
|
||||
client = fake_client(payload={"bucket": bucket, "confidence": 0.91})
|
||||
assert await NadirComplexityClassifier().classify(_context()) == expected_tier
|
||||
assert client.calls[0]["url"] == f"{DEFAULT_NADIR_API_BASE}/v1/bucket"
|
||||
assert client.calls[0]["json"]["messages"] == [{"role": "user", "content": "refactor the retry loop"}]
|
||||
assert client.calls[0]["json"]["source"] == "litellm"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tier_map_override_serves_renamed_tiers(self, fake_client):
|
||||
"""tier_labels/tier_definitions rename the tiers; the plugin must answer in the router's names."""
|
||||
fake_client(payload={"bucket": "complex"})
|
||||
classifier = NadirComplexityClassifier(tier_map={"simple": "cheap", "medium": "mid", "complex": "deep"})
|
||||
assert await classifier.classify(_context()) == "deep"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_key_is_sent_when_configured(self, fake_client):
|
||||
client = fake_client(payload={"bucket": "simple"})
|
||||
await NadirComplexityClassifier(api_key="ndr_test").classify(_context())
|
||||
assert client.calls[0]["headers"] == {"X-API-Key": "ndr_test"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_env_supplies_key_and_base(self, fake_client, monkeypatch):
|
||||
client = fake_client(payload={"bucket": "simple"})
|
||||
monkeypatch.setenv("NADIR_API_KEY", "ndr_from_env")
|
||||
monkeypatch.setenv("NADIR_API_BASE", "https://nadir.internal")
|
||||
await NadirComplexityClassifier().classify(_context())
|
||||
assert client.calls[0]["headers"] == {"X-API-Key": "ndr_from_env"}
|
||||
assert client.calls[0]["url"] == "https://nadir.internal/v1/bucket"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anonymous_when_no_key_is_configured(self, fake_client, monkeypatch):
|
||||
client = fake_client(payload={"bucket": "simple"})
|
||||
monkeypatch.delenv("NADIR_API_KEY", raising=False)
|
||||
await NadirComplexityClassifier().classify(_context())
|
||||
assert client.calls[0]["headers"] == {}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_messages_declines_without_a_call(self, fake_client):
|
||||
"""An empty request is a 400 from the endpoint; decline locally instead of spending the trip."""
|
||||
client = fake_client(payload={"bucket": "simple"})
|
||||
assert await NadirComplexityClassifier().classify(_context(messages=[])) is None
|
||||
assert client.calls == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("payload", [{"bucket": "reasoning"}, {"bucket": None}, {}, ["not", "a", "mapping"]])
|
||||
async def test_unusable_verdicts_decline(self, fake_client, payload):
|
||||
"""Anything outside tier_map is declined, so classifier_fallback decides rather than this guessing."""
|
||||
fake_client(payload=payload)
|
||||
assert await NadirComplexityClassifier().classify(_context()) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bucket_name_is_matched_case_insensitively(self, fake_client):
|
||||
fake_client(payload={"bucket": " Complex "})
|
||||
assert await NadirComplexityClassifier().classify(_context()) == "COMPLEX"
|
||||
|
||||
def test_module_instance_satisfies_the_plugin_protocol(self):
|
||||
"""The proxy resolves the dotted path to this instance and interface-checks it at startup."""
|
||||
assert isinstance(nadir_classifier, ClassifierPlugin)
|
||||
|
||||
|
||||
class TestDeploymentConfigSurvivesRouterInit:
|
||||
"""The plugin instance travels inside a deployment's litellm_params, which get deepcopied."""
|
||||
|
||||
@pytest.mark.parametrize("tier_map", [None, {"simple": "Cheap", "medium": "Standard", "complex": "Premium"}])
|
||||
def test_classifier_is_deepcopyable(self, tier_map):
|
||||
"""A MappingProxyType here would raise `cannot pickle 'mappingproxy'` at Router init, i.e. at
|
||||
proxy startup, long before any request reaches the classifier."""
|
||||
classifier = NadirComplexityClassifier(tier_map=tier_map)
|
||||
copied = copy.deepcopy({"complexity_router_config": {"classifier_plugin": classifier}})
|
||||
assert isinstance(copied["complexity_router_config"]["classifier_plugin"], NadirComplexityClassifier)
|
||||
|
||||
|
||||
class TestThroughTheComplexityRouter:
|
||||
"""The plugin as the router actually drives it: verdict in, tier out, failures fall back."""
|
||||
|
||||
def _router(self, mock_router_instance=None, **overrides):
|
||||
return ComplexityRouter(
|
||||
model_name="test-nadir-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o", "COMPLEX": "o1-preview"},
|
||||
"classifier_type": "custom",
|
||||
"classifier_plugin": NadirComplexityClassifier(),
|
||||
**overrides,
|
||||
},
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_verdict_becomes_the_routed_tier(self, fake_client):
|
||||
fake_client(payload={"bucket": "complex"})
|
||||
outcome = await self._router().aclassify(
|
||||
prompt="port the scheduler to the new executor",
|
||||
raw_messages=[{"role": "user", "content": "port the scheduler to the new executor"}],
|
||||
)
|
||||
assert outcome.tier.value == "COMPLEX"
|
||||
assert outcome.cause == "classifier_plugin"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_network_failure_falls_back_to_the_local_scorer(self, fake_client):
|
||||
"""A Nadir outage must never fail a completion: the heuristic scorer still places the request."""
|
||||
fake_client(error=ConnectionError("nadir unreachable"))
|
||||
outcome = await self._router().aclassify(prompt="hi", raw_messages=[{"role": "user", "content": "hi"}])
|
||||
assert outcome.cause != "classifier_plugin"
|
||||
assert outcome.tier is not None
|
||||
Loading…
Add table
Reference in a new issue