mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(router): add LLM-based classifier option to complexity router (#32169)
* feat(router): add LLM-based classifier option to complexity router Adds classifier_type: "heuristic" | "llm" to complexity_router_config. When set to "llm", the router calls a configured model (e.g. a small model like haiku) via structured output to pick the complexity tier, falling back to the existing regex/keyword scorer on any error, empty response, or unparseable output. * feat(ui): add classifier_type option to complexity router UI, fix edit flow Adds an "Advanced: Classification Method" section to ComplexityRouterConfig with a heuristic/LLM toggle, revealing a classifier model picker and timeout when LLM is selected. Also fixes the auto router edit modal, which never rendered the complexity router UI at all (it only handled the semantic router), and the "Edit Auto Router" button visibility check, which was gated on auto_router_config and never matched complexity router deployments. * fix(router): attribute classifier calls to caller, raise default timeout Forwards the original request's litellm_metadata into the classifier's acompletion call. Without it, the proxy's cost-tracking gate sees no user_api_key/team_id/user_id and silently drops spend logging and budget accounting for every classifier call, letting an authenticated user rack up unaccounted provider spend via repeated requests. Also raises the default classifier timeout from 400ms to 3000ms (400ms undershoots real LLM latency and would silently degrade to the heuristic scorer on most requests) and corrects the module/class docstrings, which still claimed zero external API calls after the llm classifier path was added. * fix(ci): resolve ruff strict-budget and frontend-lint failures - Use PEP 585 generics (dict/tuple/list) in the new aclassify/_classify_with_llm signatures instead of typing.Dict/Tuple/List, and suppress BLE001 on the intentionally broad except in aclassify's fallback path with a reason. - Fix prettier formatting in ComplexityRouterConfig.tsx. - Regenerate eslint-metrics.json (was stale after the classifier UI changes). * fix(ci): regenerate stale eslint-metrics.json * fix(router): strip parent budget reservation from classifier metadata The classifier's internal acompletion call previously forwarded the parent request's full litellm_metadata, including its budget reservation (user_api_key_budget_reservation / user_api_key_auth). That reservation belongs to the routed completion the classifier is deciding on, not to the classifier call itself, so it's now stripped while key/team attribution fields are still forwarded for spend logging.
This commit is contained in:
parent
dacf1cfb26
commit
109193f26a
10 changed files with 647 additions and 98 deletions
|
|
@ -4,16 +4,21 @@ Complexity-based Auto Router
|
|||
A rule-based routing strategy that uses weighted scoring across multiple dimensions
|
||||
to classify requests by complexity and route them to appropriate models.
|
||||
|
||||
No external API calls - all scoring is local and <1ms.
|
||||
By default, scoring is local (regex/keyword-based) with no external API calls and <1ms
|
||||
latency. Optionally, classifier_type="llm" routes classification through a configured
|
||||
model instead, trading that latency/cost guarantee for potentially better accuracy.
|
||||
|
||||
Inspired by ClawRouter: https://github.com/BlockRunAI/ClawRouter
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
from .config import (
|
||||
DEFAULT_CODE_KEYWORDS,
|
||||
|
|
@ -32,6 +37,24 @@ else:
|
|||
PreRoutingHookResponse = Any
|
||||
|
||||
|
||||
class TierClassification(BaseModel):
|
||||
"""Structured response schema for the LLM-based complexity classifier."""
|
||||
|
||||
tier: Literal["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"]
|
||||
|
||||
|
||||
_CLASSIFICATION_PROMPT_TEMPLATE = """Classify the complexity of the following user request into exactly one tier.
|
||||
|
||||
Tiers:
|
||||
- SIMPLE: factual lookups, greetings, short direct questions with no reasoning or code involved.
|
||||
- MEDIUM: everyday requests needing some explanation or minor code/technical content.
|
||||
- COMPLEX: requests involving non-trivial code, architecture, or multi-step technical work.
|
||||
- REASONING: requests explicitly requiring step-by-step reasoning, analysis, or weighing tradeoffs.
|
||||
|
||||
{system_context}Request:
|
||||
{prompt}"""
|
||||
|
||||
|
||||
def _append_custom_keywords(base_keywords: list[str], custom_keywords: Optional[list[str]]) -> list[str]:
|
||||
if not custom_keywords:
|
||||
return base_keywords
|
||||
|
|
@ -40,6 +63,20 @@ def _append_custom_keywords(base_keywords: list[str], custom_keywords: Optional[
|
|||
return [*base_keywords, *deduped_custom.values()]
|
||||
|
||||
|
||||
# Metadata keys that carry the parent request's budget reservation. These must not
|
||||
# reach the classifier's internal acompletion call: the reservation belongs to the
|
||||
# routed completion that the classifier is deciding on, not to the classifier call
|
||||
# itself, and forwarding it would let the classifier's cost-tracking reconcile
|
||||
# against a reservation it isn't responsible for.
|
||||
_BUDGET_RESERVATION_METADATA_KEYS = frozenset({"user_api_key_budget_reservation", "user_api_key_auth"})
|
||||
|
||||
|
||||
def _classifier_call_metadata(metadata: Optional[dict[str, Any]]) -> Optional[dict[str, Any]]:
|
||||
if not metadata:
|
||||
return metadata
|
||||
return {k: v for k, v in metadata.items() if k not in _BUDGET_RESERVATION_METADATA_KEYS}
|
||||
|
||||
|
||||
class DimensionScore:
|
||||
"""Represents a score for a single dimension with optional signal."""
|
||||
|
||||
|
|
@ -53,10 +90,10 @@ class DimensionScore:
|
|||
|
||||
class ComplexityRouter(CustomLogger):
|
||||
"""
|
||||
Rule-based complexity router that classifies requests and routes to appropriate models.
|
||||
Complexity router that classifies requests and routes to appropriate models.
|
||||
|
||||
Handles requests in <1ms with zero external API calls by using weighted scoring
|
||||
across multiple dimensions:
|
||||
By default, handles requests in <1ms with zero external API calls, using weighted
|
||||
scoring across multiple dimensions:
|
||||
- Token count (short=simple, long=complex)
|
||||
- Code presence (code keywords → complex)
|
||||
- Reasoning markers ("step by step", "think through" → reasoning tier)
|
||||
|
|
@ -297,6 +334,63 @@ class ComplexityRouter(CustomLogger):
|
|||
|
||||
return tier, weighted_score, signals
|
||||
|
||||
async def aclassify(
|
||||
self,
|
||||
prompt: str,
|
||||
system_prompt: Optional[str] = None,
|
||||
request_kwargs: Optional[dict[str, Any]] = None,
|
||||
) -> tuple[ComplexityTier, float, list[str]]:
|
||||
"""
|
||||
Classify a prompt by complexity, using the LLM classifier when configured.
|
||||
|
||||
Falls back to the local heuristic scorer if classifier_type is "heuristic",
|
||||
or if the LLM call fails, times out, or returns an unparseable response.
|
||||
"""
|
||||
if self.config.classifier_type != "llm" or self.config.classifier_llm_config is None:
|
||||
return self.classify(prompt, system_prompt)
|
||||
|
||||
try:
|
||||
tier = await self._classify_with_llm(prompt, system_prompt, request_kwargs)
|
||||
return tier, 1.0, [f"llm-classifier:{tier.value}"]
|
||||
except Exception as e: # noqa: BLE001 -- external LLM call can fail in many distinct ways (timeout, provider error, validation, parse error); any failure must fall back to the heuristic scorer
|
||||
verbose_router_logger.warning(
|
||||
f"ComplexityRouter: LLM classifier failed ({e}), falling back to heuristic scoring"
|
||||
)
|
||||
return self.classify(prompt, system_prompt)
|
||||
|
||||
async def _classify_with_llm(
|
||||
self,
|
||||
prompt: str,
|
||||
system_prompt: Optional[str] = None,
|
||||
request_kwargs: Optional[dict[str, Any]] = None,
|
||||
) -> ComplexityTier:
|
||||
"""Call the configured classifier model and parse its structured tier response."""
|
||||
llm_config = self.config.classifier_llm_config
|
||||
if llm_config is None:
|
||||
raise ValueError("classifier_llm_config is not set")
|
||||
|
||||
system_context = f"Context: {system_prompt}\n\n" if system_prompt else ""
|
||||
classification_prompt = _CLASSIFICATION_PROMPT_TEMPLATE.format(system_context=system_context, prompt=prompt)
|
||||
|
||||
# Forward the original request's metadata so the classifier call's spend is
|
||||
# attributed to the calling key/team instead of being dropped. Excludes the
|
||||
# parent request's budget reservation, which the routed completion (not this
|
||||
# internal classifier call) is responsible for reconciling.
|
||||
metadata = _classifier_call_metadata((request_kwargs or {}).get("litellm_metadata"))
|
||||
|
||||
response: ModelResponse = await self.litellm_router_instance.acompletion(
|
||||
model=llm_config.model,
|
||||
messages=[{"role": "user", "content": classification_prompt}],
|
||||
response_format=TierClassification,
|
||||
timeout=llm_config.timeout_ms / 1000,
|
||||
metadata=metadata,
|
||||
)
|
||||
content = response.choices[0].message.content
|
||||
if not content:
|
||||
raise ValueError("LLM classifier returned empty content")
|
||||
result = TierClassification.model_validate_json(content)
|
||||
return ComplexityTier[result.tier]
|
||||
|
||||
def get_model_for_tier(self, tier: ComplexityTier) -> str:
|
||||
"""
|
||||
Get the model name for a given complexity tier.
|
||||
|
|
@ -445,7 +539,7 @@ class ComplexityRouter(CustomLogger):
|
|||
messages=messages if has_original_messages else None,
|
||||
)
|
||||
|
||||
tier, score, signals = self.classify(user_message, system_prompt)
|
||||
tier, score, signals = await self.aclassify(user_message, system_prompt, request_kwargs)
|
||||
routed_model = self.get_model_for_tier(tier)
|
||||
|
||||
verbose_router_logger.info(
|
||||
|
|
|
|||
|
|
@ -6,9 +6,9 @@ All values are configurable via proxy config.yaml.
|
|||
"""
|
||||
|
||||
from enum import Enum
|
||||
from typing import Dict, List, Optional
|
||||
from typing import Dict, List, Literal, Optional
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
class ComplexityTier(str, Enum):
|
||||
|
|
@ -197,6 +197,18 @@ DEFAULT_TIER_MODELS: Dict[str, str] = {
|
|||
}
|
||||
|
||||
|
||||
class ClassifierLLMConfig(BaseModel):
|
||||
"""Configuration for the LLM-based complexity classifier."""
|
||||
|
||||
model: str = Field(
|
||||
description="Model name (from the router's model_list) to call for classification",
|
||||
)
|
||||
timeout_ms: int = Field(
|
||||
default=3000,
|
||||
description="Timeout budget for the classification call, in milliseconds",
|
||||
)
|
||||
|
||||
|
||||
class ComplexityRouterConfig(BaseModel):
|
||||
"""Configuration for the ComplexityRouter."""
|
||||
|
||||
|
|
@ -257,8 +269,24 @@ class ComplexityRouterConfig(BaseModel):
|
|||
description="Default model to use if tier cannot be determined",
|
||||
)
|
||||
|
||||
# Classifier strategy
|
||||
classifier_type: Literal["heuristic", "llm"] = Field(
|
||||
default="heuristic",
|
||||
description="Classification strategy: local regex/keyword scoring, or an LLM call",
|
||||
)
|
||||
classifier_llm_config: Optional[ClassifierLLMConfig] = Field(
|
||||
default=None,
|
||||
description="Configuration for the LLM classifier; required when classifier_type is 'llm'",
|
||||
)
|
||||
|
||||
model_config = ConfigDict(extra="allow") # Allow additional fields
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_llm_classifier_config(self) -> "ComplexityRouterConfig":
|
||||
if self.classifier_type == "llm" and self.classifier_llm_config is None:
|
||||
raise ValueError("classifier_llm_config is required when classifier_type is 'llm'")
|
||||
return self
|
||||
|
||||
|
||||
# Combined default config
|
||||
DEFAULT_COMPLEXITY_CONFIG = ComplexityRouterConfig()
|
||||
|
|
|
|||
|
|
@ -6,10 +6,11 @@ Tests the rule-based complexity scoring and tier assignment logic.
|
|||
|
||||
import os
|
||||
import sys
|
||||
from typing import Dict, List
|
||||
from unittest.mock import MagicMock, patch
|
||||
from typing import Dict
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../..")
|
||||
|
|
@ -743,7 +744,7 @@ class TestKeywordFalsePositives:
|
|||
# Should NOT detect code presence from 'api' in 'capital'
|
||||
assert not any(
|
||||
"code" in s.lower() for s in signals
|
||||
), f"False positive: got code signal from 'capital'"
|
||||
), "False positive: got code signal from 'capital'"
|
||||
# Should be SIMPLE (definition question)
|
||||
assert tier == ComplexityTier.SIMPLE
|
||||
|
||||
|
|
@ -754,7 +755,7 @@ class TestKeywordFalsePositives:
|
|||
# Should NOT detect code presence from 'git' in 'digital'
|
||||
assert not any(
|
||||
"code" in s.lower() for s in signals
|
||||
), f"False positive: got code signal from 'digital'"
|
||||
), "False positive: got code signal from 'digital'"
|
||||
|
||||
def test_try_not_in_entry(self, complexity_router):
|
||||
"""'try' should not match in 'entry'."""
|
||||
|
|
@ -770,7 +771,7 @@ class TestKeywordFalsePositives:
|
|||
tier, score, signals = complexity_router.classify(prompt)
|
||||
assert not any(
|
||||
"code" in s.lower() for s in signals
|
||||
), f"False positive: got code signal from 'terrorism'"
|
||||
), "False positive: got code signal from 'terrorism'"
|
||||
|
||||
def test_class_not_in_classical(self, complexity_router):
|
||||
"""'class' should not match in 'classical'."""
|
||||
|
|
@ -778,7 +779,7 @@ class TestKeywordFalsePositives:
|
|||
tier, score, signals = complexity_router.classify(prompt)
|
||||
assert not any(
|
||||
"code" in s.lower() for s in signals
|
||||
), f"False positive: got code signal from 'classical'"
|
||||
), "False positive: got code signal from 'classical'"
|
||||
|
||||
def test_merge_not_in_emerged(self, complexity_router):
|
||||
"""'merge' should not match in 'emerged'."""
|
||||
|
|
@ -786,7 +787,7 @@ class TestKeywordFalsePositives:
|
|||
tier, score, signals = complexity_router.classify(prompt)
|
||||
assert not any(
|
||||
"code" in s.lower() for s in signals
|
||||
), f"False positive: got code signal from 'emerged'"
|
||||
), "False positive: got code signal from 'emerged'"
|
||||
|
||||
def test_actual_api_keyword_detected(self, complexity_router):
|
||||
"""Actual 'api' usage should be detected."""
|
||||
|
|
@ -1131,3 +1132,180 @@ class TestExtractUserMessageAndSystemPrompt:
|
|||
)
|
||||
assert user_msg is None
|
||||
assert sys_prompt is None
|
||||
|
||||
|
||||
def _llm_response(content: str):
|
||||
"""Build a fake acompletion response with the given message content."""
|
||||
response = MagicMock()
|
||||
response.choices = [MagicMock()]
|
||||
response.choices[0].message.content = content
|
||||
return response
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def llm_classifier_config() -> Dict:
|
||||
"""Config with an LLM-based classifier wired to a 'haiku-classifier' model."""
|
||||
return {
|
||||
"tiers": {
|
||||
"SIMPLE": "gpt-4o-mini",
|
||||
"MEDIUM": "gpt-4o",
|
||||
"COMPLEX": "claude-sonnet-4-20250514",
|
||||
"REASONING": "o1-preview",
|
||||
},
|
||||
"classifier_type": "llm",
|
||||
"classifier_llm_config": {"model": "haiku-classifier", "timeout_ms": 400},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def llm_complexity_router(mock_router_instance, llm_classifier_config):
|
||||
"""ComplexityRouter configured to classify via an LLM call."""
|
||||
return ComplexityRouter(
|
||||
model_name="test-complexity-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config=llm_classifier_config,
|
||||
)
|
||||
|
||||
|
||||
class TestLLMClassifierConfig:
|
||||
"""Test config validation for the LLM classifier option."""
|
||||
|
||||
def test_llm_classifier_type_requires_config(self):
|
||||
"""classifier_type='llm' without classifier_llm_config must raise."""
|
||||
with pytest.raises(ValidationError):
|
||||
ComplexityRouterConfig(classifier_type="llm")
|
||||
|
||||
def test_heuristic_classifier_type_needs_no_llm_config(self):
|
||||
"""classifier_type='heuristic' (the default) needs no classifier_llm_config."""
|
||||
config = ComplexityRouterConfig()
|
||||
assert config.classifier_type == "heuristic"
|
||||
assert config.classifier_llm_config is None
|
||||
|
||||
|
||||
class TestLLMClassifier:
|
||||
"""Test the LLM-based classifier path (aclassify) and its fallback behavior."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aclassify_heuristic_skips_llm_call(self, complexity_router, mock_router_instance):
|
||||
"""When classifier_type is 'heuristic' (default), aclassify must not call the LLM."""
|
||||
mock_router_instance.acompletion = AsyncMock()
|
||||
tier, score, signals = await complexity_router.aclassify("Hello!")
|
||||
mock_router_instance.acompletion.assert_not_called()
|
||||
assert tier == ComplexityTier.SIMPLE
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aclassify_llm_success_routes_by_llm_verdict(
|
||||
self, llm_complexity_router, mock_router_instance
|
||||
):
|
||||
"""A well-formed structured LLM response should decide the tier directly.
|
||||
|
||||
Uses a prompt that heuristic scoring alone would classify as SIMPLE, to prove
|
||||
the LLM verdict -- not the heuristic scorer -- is what decided the tier.
|
||||
"""
|
||||
mock_router_instance.acompletion = AsyncMock(
|
||||
return_value=_llm_response('{"tier": "COMPLEX"}')
|
||||
)
|
||||
tier, score, signals = await llm_complexity_router.aclassify("hi")
|
||||
assert tier == ComplexityTier.COMPLEX
|
||||
assert "llm-classifier:COMPLEX" in signals
|
||||
mock_router_instance.acompletion.assert_awaited_once()
|
||||
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
|
||||
assert call_kwargs["model"] == "haiku-classifier"
|
||||
assert call_kwargs["timeout"] == 0.4
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aclassify_forwards_request_metadata_for_spend_tracking(
|
||||
self, llm_complexity_router, mock_router_instance
|
||||
):
|
||||
"""The classifier call must carry the original request's metadata.
|
||||
|
||||
Without this, the proxy's cost-tracking gate (_should_track_cost_callback)
|
||||
sees no user_api_key/team_id/user_id and silently drops all spend logging
|
||||
and budget accounting for the classifier call.
|
||||
"""
|
||||
mock_router_instance.acompletion = AsyncMock(
|
||||
return_value=_llm_response('{"tier": "SIMPLE"}')
|
||||
)
|
||||
request_metadata = {"user_api_key": "sk-abc", "user_api_key_team_id": "team-1"}
|
||||
await llm_complexity_router.aclassify(
|
||||
"hi", request_kwargs={"litellm_metadata": request_metadata}
|
||||
)
|
||||
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
|
||||
assert call_kwargs["metadata"] == request_metadata
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aclassify_strips_budget_reservation_from_classifier_metadata(
|
||||
self, llm_complexity_router, mock_router_instance
|
||||
):
|
||||
"""The classifier call must not receive the parent request's budget reservation.
|
||||
|
||||
The reservation belongs to the routed completion the classifier is deciding
|
||||
on, not to this internal classifier call. Forwarding it would let the
|
||||
classifier's own cost-tracking reconcile against a reservation it has no
|
||||
business touching, so it must be stripped while the rest of the attribution
|
||||
metadata (key/team) is preserved.
|
||||
"""
|
||||
mock_router_instance.acompletion = AsyncMock(
|
||||
return_value=_llm_response('{"tier": "SIMPLE"}')
|
||||
)
|
||||
request_metadata = {
|
||||
"user_api_key": "sk-abc",
|
||||
"user_api_key_team_id": "team-1",
|
||||
"user_api_key_budget_reservation": {"reserved_cost": 1.0},
|
||||
"user_api_key_auth": {"budget_reservation": {"reserved_cost": 1.0}},
|
||||
}
|
||||
await llm_complexity_router.aclassify(
|
||||
"hi", request_kwargs={"litellm_metadata": request_metadata}
|
||||
)
|
||||
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
|
||||
assert call_kwargs["metadata"] == {
|
||||
"user_api_key": "sk-abc",
|
||||
"user_api_key_team_id": "team-1",
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aclassify_falls_back_to_heuristic_on_llm_exception(
|
||||
self, llm_complexity_router, mock_router_instance
|
||||
):
|
||||
"""A timeout/error from the classifier model must fall back to heuristic scoring."""
|
||||
mock_router_instance.acompletion = AsyncMock(side_effect=TimeoutError("classifier timed out"))
|
||||
tier, score, signals = await llm_complexity_router.aclassify("Hello!")
|
||||
assert tier == llm_complexity_router.classify("Hello!")[0]
|
||||
assert tier == ComplexityTier.SIMPLE
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aclassify_falls_back_to_heuristic_on_unparseable_response(
|
||||
self, llm_complexity_router, mock_router_instance
|
||||
):
|
||||
"""Non-JSON or schema-violating output must fall back to heuristic scoring, not raise."""
|
||||
mock_router_instance.acompletion = AsyncMock(return_value=_llm_response("not json"))
|
||||
tier, score, signals = await llm_complexity_router.aclassify("Hello!")
|
||||
assert tier == ComplexityTier.SIMPLE
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aclassify_falls_back_to_heuristic_on_empty_content(
|
||||
self, llm_complexity_router, mock_router_instance
|
||||
):
|
||||
"""Empty/None message content (e.g. provider quirk) must fall back, not raise."""
|
||||
mock_router_instance.acompletion = AsyncMock(return_value=_llm_response(None))
|
||||
tier, score, signals = await llm_complexity_router.aclassify("Hello!")
|
||||
assert tier == ComplexityTier.SIMPLE
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_routing_hook_uses_llm_classifier_end_to_end(
|
||||
self, llm_complexity_router, mock_router_instance
|
||||
):
|
||||
"""The full pre-routing hook should route using the LLM classifier's verdict."""
|
||||
mock_router_instance.acompletion = AsyncMock(
|
||||
return_value=_llm_response('{"tier": "REASONING"}')
|
||||
)
|
||||
request_metadata = {"user_api_key": "sk-abc", "user_api_key_team_id": "team-1"}
|
||||
result = await llm_complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={"litellm_metadata": request_metadata},
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
)
|
||||
assert result is not None
|
||||
assert result.model == "o1-preview" # REASONING tier model
|
||||
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
|
||||
assert call_kwargs["metadata"] == request_metadata
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
{
|
||||
"@typescript-eslint/no-explicit-any": 1977,
|
||||
"complexity": 129,
|
||||
"@typescript-eslint/no-explicit-any": 1978,
|
||||
"complexity": 130,
|
||||
"local/no-large-inline-object-arg": 509,
|
||||
"local/no-long-condition-chain": 233,
|
||||
"local/no-long-condition-chain": 234,
|
||||
"max-depth": 59,
|
||||
"no-console": 16
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import { renderWithProviders, screen, within } from "../../../tests/test-utils";
|
||||
import { fireEvent, renderWithProviders, screen, within } from "../../../tests/test-utils";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { vi } from "vitest";
|
||||
import ComplexityRouterConfig from "./ComplexityRouterConfig";
|
||||
import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
|
||||
|
||||
const mockModelInfo = [
|
||||
{ model_group: "gpt-4" },
|
||||
|
|
@ -9,21 +9,24 @@ const mockModelInfo = [
|
|||
{ model_group: "claude-3-opus" },
|
||||
] as any[];
|
||||
|
||||
const defaultTiers = {
|
||||
SIMPLE: "gpt-3.5-turbo",
|
||||
MEDIUM: "gpt-3.5-turbo",
|
||||
COMPLEX: "gpt-4",
|
||||
REASONING: "claude-3-opus",
|
||||
const defaultValue: ComplexityRouterConfigValue = {
|
||||
tiers: {
|
||||
SIMPLE: "gpt-3.5-turbo",
|
||||
MEDIUM: "gpt-3.5-turbo",
|
||||
COMPLEX: "gpt-4",
|
||||
REASONING: "claude-3-opus",
|
||||
},
|
||||
classifier_type: "heuristic",
|
||||
};
|
||||
|
||||
describe("ComplexityRouterConfig", () => {
|
||||
it("should render", () => {
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultTiers} onChange={vi.fn()} />);
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
|
||||
expect(screen.getByText("Complexity Tier Configuration")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display all four tier labels", () => {
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultTiers} onChange={vi.fn()} />);
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
|
||||
expect(screen.getByText("Simple Tier")).toBeInTheDocument();
|
||||
expect(screen.getByText("Medium Tier")).toBeInTheDocument();
|
||||
expect(screen.getByText("Complex Tier")).toBeInTheDocument();
|
||||
|
|
@ -31,7 +34,7 @@ describe("ComplexityRouterConfig", () => {
|
|||
});
|
||||
|
||||
it("should show example queries for each tier", () => {
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultTiers} onChange={vi.fn()} />);
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
|
||||
expect(screen.getByText(/Hello!/)).toBeInTheDocument();
|
||||
expect(screen.getByText(/Explain how REST APIs work/)).toBeInTheDocument();
|
||||
expect(screen.getByText(/Design a microservices architecture/)).toBeInTheDocument();
|
||||
|
|
@ -39,20 +42,56 @@ describe("ComplexityRouterConfig", () => {
|
|||
});
|
||||
|
||||
it("should display the how classification works section", () => {
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultTiers} onChange={vi.fn()} />);
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
|
||||
expect(screen.getByText("How Classification Works")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show score thresholds in the classification section", () => {
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultTiers} onChange={vi.fn()} />);
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
|
||||
expect(screen.getByText(/Score < 0.15/)).toBeInTheDocument();
|
||||
expect(screen.getByText(/Score 0.15 - 0.35/)).toBeInTheDocument();
|
||||
expect(screen.getByText(/Score 0.35 - 0.60/)).toBeInTheDocument();
|
||||
expect(screen.getByText(/Score > 0.60/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should default to heuristic and hide classifier model/timeout fields", () => {
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
|
||||
expect(screen.getByText("Advanced: Classification Method")).toBeInTheDocument();
|
||||
expect(screen.queryByText("Classifier Model")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should reveal classifier model and timeout fields when llm is selected", () => {
|
||||
const onChange = vi.fn();
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={onChange} />);
|
||||
|
||||
// Collapse panel content isn't rendered until first expanded.
|
||||
fireEvent.click(screen.getByText("Advanced: Classification Method"));
|
||||
fireEvent.click(screen.getByText("LLM Classifier"));
|
||||
|
||||
expect(onChange).toHaveBeenCalledWith({
|
||||
...defaultValue,
|
||||
classifier_type: "llm",
|
||||
classifier_llm_config: { model: "", timeout_ms: 3000 },
|
||||
});
|
||||
});
|
||||
|
||||
it("should show classifier fields and use the configured values when classifier_type is llm", () => {
|
||||
const llmValue: ComplexityRouterConfigValue = {
|
||||
...defaultValue,
|
||||
classifier_type: "llm",
|
||||
classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 750 },
|
||||
};
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={llmValue} onChange={vi.fn()} />);
|
||||
|
||||
fireEvent.click(screen.getByText("Advanced: Classification Method"));
|
||||
|
||||
expect(screen.getByText("Classifier Model")).toBeInTheDocument();
|
||||
expect(screen.getByText("Timeout (ms)")).toBeInTheDocument();
|
||||
expect(screen.getByDisplayValue("750")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render the custom technical keywords field", () => {
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultTiers} onChange={vi.fn()} />);
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
|
||||
expect(screen.getByText("Custom Technical Keywords")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
|
|
@ -60,7 +99,7 @@ describe("ComplexityRouterConfig", () => {
|
|||
renderWithProviders(
|
||||
<ComplexityRouterConfig
|
||||
modelInfo={mockModelInfo}
|
||||
value={defaultTiers}
|
||||
value={defaultValue}
|
||||
onChange={vi.fn()}
|
||||
customTechnicalKeywords={["udp", "kafka"]}
|
||||
onCustomTechnicalKeywordsChange={vi.fn()}
|
||||
|
|
@ -76,7 +115,7 @@ describe("ComplexityRouterConfig", () => {
|
|||
renderWithProviders(
|
||||
<ComplexityRouterConfig
|
||||
modelInfo={mockModelInfo}
|
||||
value={defaultTiers}
|
||||
value={defaultValue}
|
||||
onChange={vi.fn()}
|
||||
customTechnicalKeywords={[]}
|
||||
onCustomTechnicalKeywordsChange={onCustomTechnicalKeywordsChange}
|
||||
|
|
|
|||
|
|
@ -1,21 +1,36 @@
|
|||
import { InfoCircleOutlined } from "@ant-design/icons";
|
||||
import { Select as AntdSelect, Card, Divider, Space, Tooltip, Typography } from "antd";
|
||||
import { Select as AntdSelect, Card, Collapse, Divider, InputNumber, Radio, Space, Tooltip, Typography } from "antd";
|
||||
import React from "react";
|
||||
import { ModelGroup } from "@/components/llm_calls/fetch_models";
|
||||
|
||||
const { Text } = Typography;
|
||||
|
||||
interface ComplexityTiers {
|
||||
export const DEFAULT_CLASSIFIER_TIMEOUT_MS = 3000;
|
||||
|
||||
export interface ComplexityTiers {
|
||||
SIMPLE: string;
|
||||
MEDIUM: string;
|
||||
COMPLEX: string;
|
||||
REASONING: string;
|
||||
}
|
||||
|
||||
export interface ClassifierLLMConfig {
|
||||
model: string;
|
||||
timeout_ms: number;
|
||||
}
|
||||
|
||||
export type ClassifierType = "heuristic" | "llm";
|
||||
|
||||
export interface ComplexityRouterConfigValue {
|
||||
tiers: ComplexityTiers;
|
||||
classifier_type: ClassifierType;
|
||||
classifier_llm_config?: ClassifierLLMConfig;
|
||||
}
|
||||
|
||||
interface ComplexityRouterConfigProps {
|
||||
modelInfo: ModelGroup[];
|
||||
value: ComplexityTiers;
|
||||
onChange: (tiers: ComplexityTiers) => void;
|
||||
value: ComplexityRouterConfigValue;
|
||||
onChange: (value: ComplexityRouterConfigValue) => void;
|
||||
customTechnicalKeywords?: string[];
|
||||
onCustomTechnicalKeywordsChange?: (keywords: string[]) => void;
|
||||
}
|
||||
|
|
@ -59,7 +74,38 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
|
|||
const handleTierChange = (tier: keyof ComplexityTiers, model: string) => {
|
||||
onChange({
|
||||
...value,
|
||||
[tier]: model,
|
||||
tiers: { ...value.tiers, [tier]: model },
|
||||
});
|
||||
};
|
||||
|
||||
const handleClassifierTypeChange = (classifierType: ClassifierType) => {
|
||||
onChange({
|
||||
...value,
|
||||
classifier_type: classifierType,
|
||||
classifier_llm_config:
|
||||
classifierType === "llm"
|
||||
? value.classifier_llm_config ?? { model: "", timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS }
|
||||
: undefined,
|
||||
});
|
||||
};
|
||||
|
||||
const handleClassifierModelChange = (model: string) => {
|
||||
onChange({
|
||||
...value,
|
||||
classifier_llm_config: {
|
||||
model,
|
||||
timeout_ms: value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
|
||||
},
|
||||
});
|
||||
};
|
||||
|
||||
const handleClassifierTimeoutChange = (timeoutMs: number | null) => {
|
||||
onChange({
|
||||
...value,
|
||||
classifier_llm_config: {
|
||||
model: value.classifier_llm_config?.model ?? "",
|
||||
timeout_ms: timeoutMs ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
|
||||
},
|
||||
});
|
||||
};
|
||||
|
||||
|
|
@ -98,7 +144,7 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
|
|||
Examples: {tierInfo.examples}
|
||||
</Text>
|
||||
<AntdSelect
|
||||
value={value[tier]}
|
||||
value={value.tiers[tier]}
|
||||
onChange={(model) => handleTierChange(tier, model)}
|
||||
placeholder={`Select model for ${tierInfo.label.toLowerCase()} queries`}
|
||||
showSearch
|
||||
|
|
@ -113,6 +159,76 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
|
|||
|
||||
<Divider />
|
||||
|
||||
<Collapse
|
||||
ghost
|
||||
style={{ background: "#f9fafb", borderRadius: 8, border: "1px solid #e5e7eb" }}
|
||||
items={[
|
||||
{
|
||||
key: "classifier",
|
||||
label: (
|
||||
<Text strong style={{ color: "#374151" }}>
|
||||
Advanced: Classification Method
|
||||
</Text>
|
||||
),
|
||||
children: (
|
||||
<>
|
||||
<Radio.Group
|
||||
value={value.classifier_type}
|
||||
onChange={(e) => handleClassifierTypeChange(e.target.value)}
|
||||
className="w-full"
|
||||
>
|
||||
<Space direction="vertical" className="w-full">
|
||||
<Radio value="heuristic">
|
||||
<Text strong>Heuristic</Text>{" "}
|
||||
<Text type="secondary">(default) — rule-based scoring, no API calls, <1ms latency</Text>
|
||||
</Radio>
|
||||
<Radio value="llm">
|
||||
<Text strong>LLM Classifier</Text>{" "}
|
||||
<Text type="secondary">— use a model to decide the tier (e.g. a small/fast model)</Text>
|
||||
</Radio>
|
||||
</Space>
|
||||
</Radio.Group>
|
||||
|
||||
{value.classifier_type === "llm" && (
|
||||
<div className="mt-4 space-y-3">
|
||||
<div>
|
||||
<Text strong style={{ display: "block", marginBottom: 4 }}>
|
||||
Classifier Model
|
||||
</Text>
|
||||
<AntdSelect
|
||||
value={value.classifier_llm_config?.model || undefined}
|
||||
onChange={handleClassifierModelChange}
|
||||
placeholder="Select the model that will classify request complexity"
|
||||
showSearch
|
||||
style={{ width: "100%" }}
|
||||
options={modelOptions}
|
||||
/>
|
||||
</div>
|
||||
<div>
|
||||
<Text strong style={{ display: "block", marginBottom: 4 }}>
|
||||
Timeout (ms)
|
||||
</Text>
|
||||
<InputNumber
|
||||
value={value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS}
|
||||
onChange={handleClassifierTimeoutChange}
|
||||
min={1}
|
||||
style={{ width: "100%" }}
|
||||
/>
|
||||
<Text type="secondary" style={{ fontSize: 12 }}>
|
||||
Falls back to the heuristic scorer if the classifier call errors, times out, or returns an
|
||||
unparseable response.
|
||||
</Text>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</>
|
||||
),
|
||||
},
|
||||
]}
|
||||
/>
|
||||
|
||||
<Divider />
|
||||
|
||||
<Card>
|
||||
<div className="flex items-center gap-2 mb-2">
|
||||
<Text strong style={{ fontSize: 16 }}>
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ import { all_admin_roles } from "@/utils/roles";
|
|||
import { handleAddAutoRouterSubmit } from "./handle_add_auto_router_submit";
|
||||
import { fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models";
|
||||
import RouterConfigBuilder from "./RouterConfigBuilder";
|
||||
import ComplexityRouterConfig from "./ComplexityRouterConfig";
|
||||
import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
|
||||
import NotificationManager from "../molecules/notifications_manager";
|
||||
import { ThunderboltOutlined, BranchesOutlined } from "@ant-design/icons";
|
||||
|
||||
|
|
@ -21,13 +21,6 @@ interface AddAutoRouterTabProps {
|
|||
|
||||
type RouterType = "complexity" | "semantic";
|
||||
|
||||
interface ComplexityTiers {
|
||||
SIMPLE: string;
|
||||
MEDIUM: string;
|
||||
COMPLEX: string;
|
||||
REASONING: string;
|
||||
}
|
||||
|
||||
const { Title, Link } = Typography;
|
||||
|
||||
const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, accessToken, userRole }) => {
|
||||
|
|
@ -48,11 +41,9 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
|
|||
const [routerConfig, setRouterConfig] = useState<any>(null);
|
||||
|
||||
// Complexity router config (new)
|
||||
const [complexityTiers, setComplexityTiers] = useState<ComplexityTiers>({
|
||||
SIMPLE: "",
|
||||
MEDIUM: "",
|
||||
COMPLEX: "",
|
||||
REASONING: "",
|
||||
const [complexityRouterConfig, setComplexityRouterConfig] = useState<ComplexityRouterConfigValue>({
|
||||
tiers: { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" },
|
||||
classifier_type: "heuristic",
|
||||
});
|
||||
|
||||
const [customTechnicalKeywords, setCustomTechnicalKeywords] = useState<string[]>([]);
|
||||
|
|
@ -99,15 +90,20 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
|
|||
// Validation differs based on router type
|
||||
if (routerType === "complexity") {
|
||||
// Complexity Router validation
|
||||
const filledTiers = Object.values(complexityTiers).filter(Boolean);
|
||||
const { tiers, classifier_type, classifier_llm_config } = complexityRouterConfig;
|
||||
const filledTiers = Object.values(tiers).filter(Boolean);
|
||||
if (filledTiers.length === 0) {
|
||||
NotificationManager.fromBackend("Please select at least one model for a complexity tier");
|
||||
return;
|
||||
}
|
||||
|
||||
if (classifier_type === "llm" && !classifier_llm_config?.model) {
|
||||
NotificationManager.fromBackend("Please select a classifier model, or switch back to Heuristic");
|
||||
return;
|
||||
}
|
||||
|
||||
// For complexity router, use the first non-empty tier as default
|
||||
const defaultModel =
|
||||
complexityTiers.MEDIUM || complexityTiers.SIMPLE || complexityTiers.COMPLEX || complexityTiers.REASONING;
|
||||
const defaultModel = tiers.MEDIUM || tiers.SIMPLE || tiers.COMPLEX || tiers.REASONING;
|
||||
|
||||
// Set form values for complexity router
|
||||
form.setFieldsValue({
|
||||
|
|
@ -128,7 +124,9 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
|
|||
// Use special model prefix for complexity router
|
||||
model_type: "complexity_router",
|
||||
complexity_router_config: {
|
||||
tiers: complexityTiers,
|
||||
tiers,
|
||||
classifier_type,
|
||||
...(classifier_type === "llm" ? { classifier_llm_config } : {}),
|
||||
...(customTechnicalKeywords.length > 0 && { custom_technical_keywords: customTechnicalKeywords }),
|
||||
},
|
||||
model_access_group: currentFormValues.model_access_group,
|
||||
|
|
@ -279,9 +277,9 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
|
|||
<div className="w-full mb-4">
|
||||
<ComplexityRouterConfig
|
||||
modelInfo={modelInfo}
|
||||
value={complexityTiers}
|
||||
onChange={(tiers) => {
|
||||
setComplexityTiers(tiers);
|
||||
value={complexityRouterConfig}
|
||||
onChange={(config) => {
|
||||
setComplexityRouterConfig(config);
|
||||
}}
|
||||
customTechnicalKeywords={customTechnicalKeywords}
|
||||
onCustomTechnicalKeywordsChange={setCustomTechnicalKeywords}
|
||||
|
|
|
|||
|
|
@ -4,8 +4,13 @@ import { Text, TextInput } from "@tremor/react";
|
|||
import { modelAvailableCall, modelPatchUpdateCall } from "../networking";
|
||||
import { fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models";
|
||||
import RouterConfigBuilder from "../add_model/RouterConfigBuilder";
|
||||
import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "../add_model/ComplexityRouterConfig";
|
||||
import NotificationsManager from "../molecules/notifications_manager";
|
||||
|
||||
const isComplexityRouterModel = (modelData: any): boolean =>
|
||||
modelData?.litellm_params?.model?.startsWith("auto_router/complexity_router") ||
|
||||
modelData?.litellm_params?.complexity_router_config != null;
|
||||
|
||||
interface EditAutoRouterModalProps {
|
||||
isVisible: boolean;
|
||||
onCancel: () => void;
|
||||
|
|
@ -30,6 +35,11 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
|
|||
const [showCustomDefaultModel, setShowCustomDefaultModel] = useState<boolean>(false);
|
||||
const [showCustomEmbeddingModel, setShowCustomEmbeddingModel] = useState<boolean>(false);
|
||||
const [routerConfig, setRouterConfig] = useState<any>(null);
|
||||
const [complexityRouterConfig, setComplexityRouterConfig] = useState<ComplexityRouterConfigValue>({
|
||||
tiers: { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" },
|
||||
classifier_type: "heuristic",
|
||||
});
|
||||
const isComplexityRouter = isComplexityRouterModel(modelData);
|
||||
|
||||
useEffect(() => {
|
||||
if (isVisible && modelData) {
|
||||
|
|
@ -66,6 +76,31 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
|
|||
|
||||
const initializeForm = () => {
|
||||
try {
|
||||
if (isComplexityRouterModel(modelData)) {
|
||||
// Parse the complexity_router_config if it exists and is a string
|
||||
let parsedConfig = modelData.litellm_params?.complexity_router_config || {};
|
||||
if (typeof parsedConfig === "string") {
|
||||
parsedConfig = JSON.parse(parsedConfig);
|
||||
}
|
||||
|
||||
setComplexityRouterConfig({
|
||||
tiers: {
|
||||
SIMPLE: parsedConfig.tiers?.SIMPLE || "",
|
||||
MEDIUM: parsedConfig.tiers?.MEDIUM || "",
|
||||
COMPLEX: parsedConfig.tiers?.COMPLEX || "",
|
||||
REASONING: parsedConfig.tiers?.REASONING || "",
|
||||
},
|
||||
classifier_type: parsedConfig.classifier_type || "heuristic",
|
||||
classifier_llm_config: parsedConfig.classifier_llm_config,
|
||||
});
|
||||
|
||||
form.setFieldsValue({
|
||||
auto_router_name: modelData.model_name,
|
||||
model_access_group: modelData.model_info?.access_groups || [],
|
||||
});
|
||||
return;
|
||||
}
|
||||
|
||||
// Parse the auto_router_config if it exists and is a string
|
||||
let parsedConfig = null;
|
||||
if (modelData.litellm_params?.auto_router_config) {
|
||||
|
|
@ -101,6 +136,49 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
|
|||
setLoading(true);
|
||||
const values = await form.validateFields();
|
||||
|
||||
if (isComplexityRouter) {
|
||||
const { tiers, classifier_type, classifier_llm_config } = complexityRouterConfig;
|
||||
if (Object.values(tiers).filter(Boolean).length === 0) {
|
||||
NotificationsManager.fromBackend("Please select at least one model for a complexity tier");
|
||||
return;
|
||||
}
|
||||
if (classifier_type === "llm" && !classifier_llm_config?.model) {
|
||||
NotificationsManager.fromBackend("Please select a classifier model, or switch back to Heuristic");
|
||||
return;
|
||||
}
|
||||
|
||||
const defaultModel = tiers.MEDIUM || tiers.SIMPLE || tiers.COMPLEX || tiers.REASONING;
|
||||
const updatedLitellmParams = {
|
||||
...modelData.litellm_params,
|
||||
complexity_router_config: {
|
||||
tiers,
|
||||
classifier_type,
|
||||
...(classifier_type === "llm" ? { classifier_llm_config } : {}),
|
||||
},
|
||||
complexity_router_default_model: defaultModel,
|
||||
};
|
||||
const updatedModelInfo = {
|
||||
...modelData.model_info,
|
||||
access_groups: values.model_access_group || [],
|
||||
};
|
||||
|
||||
await modelPatchUpdateCall(
|
||||
accessToken,
|
||||
{ model_name: values.auto_router_name, litellm_params: updatedLitellmParams, model_info: updatedModelInfo },
|
||||
modelData.model_info.id,
|
||||
);
|
||||
|
||||
NotificationsManager.success("Auto router configuration updated successfully");
|
||||
onSuccess({
|
||||
...modelData,
|
||||
model_name: values.auto_router_name,
|
||||
litellm_params: updatedLitellmParams,
|
||||
model_info: updatedModelInfo,
|
||||
});
|
||||
onCancel();
|
||||
return;
|
||||
}
|
||||
|
||||
// Prepare the updated litellm_params
|
||||
const updatedLitellmParams = {
|
||||
...modelData.litellm_params,
|
||||
|
|
@ -177,45 +255,60 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
|
|||
<TextInput placeholder="e.g., auto_router_1, smart_routing" />
|
||||
</Form.Item>
|
||||
|
||||
{/* Router Configuration Builder */}
|
||||
<div className="w-full">
|
||||
<RouterConfigBuilder
|
||||
modelInfo={modelInfo}
|
||||
value={routerConfig}
|
||||
onChange={(config) => {
|
||||
setRouterConfig(config);
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
{isComplexityRouter ? (
|
||||
/* Complexity Router Configuration */
|
||||
<div className="w-full">
|
||||
<ComplexityRouterConfig
|
||||
modelInfo={modelInfo}
|
||||
value={complexityRouterConfig}
|
||||
onChange={(config) => {
|
||||
setComplexityRouterConfig(config);
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
) : (
|
||||
<>
|
||||
{/* Router Configuration Builder */}
|
||||
<div className="w-full">
|
||||
<RouterConfigBuilder
|
||||
modelInfo={modelInfo}
|
||||
value={routerConfig}
|
||||
onChange={(config) => {
|
||||
setRouterConfig(config);
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* Default Model */}
|
||||
<Form.Item
|
||||
label="Default Model"
|
||||
name="auto_router_default_model"
|
||||
rules={[{ required: true, message: "Default model is required" }]}
|
||||
>
|
||||
<AntdSelect
|
||||
placeholder="Select a default model"
|
||||
onChange={(value) => {
|
||||
setShowCustomDefaultModel(value === "custom");
|
||||
}}
|
||||
options={[...modelOptions, { value: "custom", label: "Enter custom model name" }]}
|
||||
showSearch={true}
|
||||
/>
|
||||
</Form.Item>
|
||||
{/* Default Model */}
|
||||
<Form.Item
|
||||
label="Default Model"
|
||||
name="auto_router_default_model"
|
||||
rules={[{ required: true, message: "Default model is required" }]}
|
||||
>
|
||||
<AntdSelect
|
||||
placeholder="Select a default model"
|
||||
onChange={(value) => {
|
||||
setShowCustomDefaultModel(value === "custom");
|
||||
}}
|
||||
options={[...modelOptions, { value: "custom", label: "Enter custom model name" }]}
|
||||
showSearch={true}
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
{/* Embedding Model */}
|
||||
<Form.Item label="Embedding Model" name="auto_router_embedding_model">
|
||||
<AntdSelect
|
||||
placeholder="Select an embedding model (optional)"
|
||||
onChange={(value) => {
|
||||
setShowCustomEmbeddingModel(value === "custom");
|
||||
}}
|
||||
options={[...modelOptions, { value: "custom", label: "Enter custom model name" }]}
|
||||
showSearch={true}
|
||||
allowClear
|
||||
/>
|
||||
</Form.Item>
|
||||
{/* Embedding Model */}
|
||||
<Form.Item label="Embedding Model" name="auto_router_embedding_model">
|
||||
<AntdSelect
|
||||
placeholder="Select an embedding model (optional)"
|
||||
onChange={(value) => {
|
||||
setShowCustomEmbeddingModel(value === "custom");
|
||||
}}
|
||||
options={[...modelOptions, { value: "custom", label: "Enter custom model name" }]}
|
||||
showSearch={true}
|
||||
allowClear
|
||||
/>
|
||||
</Form.Item>
|
||||
</>
|
||||
)}
|
||||
|
||||
{/* Model Access Groups - Admin only */}
|
||||
{userRole === "Admin" && (
|
||||
|
|
|
|||
|
|
@ -124,7 +124,10 @@ export default function ModelInfoView({
|
|||
const canEditModel =
|
||||
(userRole === "Admin" || modelData?.model_info?.created_by === userID) && modelData?.model_info?.db_model;
|
||||
const isAdmin = userRole === "Admin";
|
||||
const isAutoRouter = modelData?.litellm_params?.auto_router_config != null;
|
||||
const isAutoRouter =
|
||||
modelData?.litellm_params?.auto_router_config != null ||
|
||||
modelData?.litellm_params?.complexity_router_config != null ||
|
||||
modelData?.litellm_params?.model?.startsWith("auto_router/complexity_router");
|
||||
|
||||
const usingExistingCredential =
|
||||
modelData?.litellm_params?.litellm_credential_name != null &&
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
Loading…
Add table
Reference in a new issue