mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
* feat(router): add auto_router/quality_router for quality-tier routing (#25987) * feat(router): add auto_router/quality_router for quality-tier routing Adds a new auto-router type that routes a request to a model at a target quality tier. The quality tier is inferred by re-using the existing ComplexityRouter's classification, then mapped through an admin-configured complexity_to_quality table. Each candidate model declares its own quality_tier in model_info.litellm_routing_preferences. Resolution strategy: exact tier match, else round up to the next higher tier, else fall back to default_model. Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com> * feat(quality_router): add capability-based filtering Each deployment can declare a `capabilities: List[str]` field in `model_info.litellm_routing_preferences` (e.g. ["vision", "function_calling"]). Requests can pass `litellm_capabilities` in `request_kwargs` to require specific capabilities — the router will only route to deployments whose declared capabilities are a superset. Resolution still walks tier (exact → round up), but at each tier filters by capability before picking. Falls back to default_model only when it also satisfies the required capabilities; otherwise raises rather than silently routing to a model that lacks a required capability. Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com> * feat(quality_router): expose routing decision in response headers For transparency, expose the QualityRouter's routing decision in the proxy response headers: x-litellm-quality-router-model → picked model_name (e.g. "haiku-vision") x-litellm-quality-router-tier → resolved quality tier (e.g. "1") x-litellm-quality-router-complexity → ComplexityTier name (e.g. "SIMPLE") Mechanism: the pre-routing hook stashes the decision in request_kwargs["metadata"]["quality_router_decision"]. After the call returns, Router.set_response_headers lifts the decision into response._hidden_params["additional_headers"] alongside the existing x-litellm-model-group / x-litellm-model-id headers. Existing metadata keys (trace_id, user_id, etc.) are preserved. Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com> * feat(quality_router): replace capabilities with keyword override Drops the capability-based filtering in favor of a keyword-based override for v0: - RoutingPreferences.keywords: List[str] (replaces capabilities) — each deployment can declare substring keywords. - If any declared keyword (case-insensitive) appears in the user message, the router short-circuits the complexity-classification flow and routes to the matching deployment. - Tiebreaker for overlapping keyword matches: quality_tier DESC, then cheapest model_info.input_cost_per_token ASC. Unpriced models lose ties to priced ones. Decision metadata + headers now expose the override: x-litellm-quality-router-via → "keyword" | "quality_tier" x-litellm-quality-router-keyword → matched keyword (only on keyword route) x-litellm-quality-router-complexity → complexity tier (only on tier route) Removes: - request_kwargs["litellm_capabilities"] reading - _model_capabilities, _model_supports_capabilities, _first_capable_model_at_tier, capability filter in _resolve_model_for_quality_tier Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com> * feat(quality_router): add explicit `order` to RoutingPreferences Adds an explicit priority field to RoutingPreferences for resolving collisions deterministically: RoutingPreferences.order: Optional[int] # lower wins; unset = +inf Used as the PRIMARY tiebreaker in two places: 1. Keyword overlap: when multiple deployments declare the same matching keyword, sort by (order ASC, quality_tier DESC, input_cost_per_token ASC, model_name ASC). Explicit always beats implicit. 2. Tier resolution: when multiple deployments share a quality tier, `_resolve_model_for_quality_tier` picks the one with the lowest order. The tier list is now sorted at index-build time. This lets admins make routing decisions explicit when the natural quality-and-price ordering would pick the wrong model. Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com> * feat(quality_router): reorder tiebreak to (quality, order, price) Changes the tiebreak ordering so quality_tier always wins first, then explicit `order` is used to break ties within the same tier, then price breaks the rest: 1. quality_tier DESC ← best model wins first 2. order ASC ← explicit priority within a tier 3. input_cost_per_token ASC 4. model_name ASC Previously `order` was the primary key — that meant a tier-2 model with `order=1` would beat a tier-3 model with no `order`, which is the wrong default. Now `order` only resolves collisions among same-tier candidates. Tier resolution (within a single tier) keeps the same key minus quality: (order ASC, cost ASC, name). Test renames + flips: - test_explicit_order_overrides_quality_tier → test_quality_wins_over_explicit_order - new: test_order_breaks_tie_within_same_quality_tier Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com> * fix(quality_router): resolve Greptile review feedback Addresses four P1 findings from PR review plus test coverage: 1. set_model_list missing quality_routers reset - Hot-reloading the Router would leave stale QualityRouter instances pointing at the old model_list. `set_model_list` now clears `self.quality_routers` alongside the other indices. 2. Round-down fallback before default_model - `_resolve_model_for_quality_tier` now rounds DOWN to the closest lower tier after round-up fails, before falling back to `default_model`. Degrades gracefully rather than jumping straight off-tier. 3. RoutingPreferences validation bypass - `_build_tier_index` now instantiates `RoutingPreferences(**prefs)` so invalid shapes (e.g. non-int quality_tier) raise a clear ValueError instead of silently succeeding. 4. Config-ordering dependency - `_tier_to_models` is now built lazily on first access. Previously, eager construction in `__init__` meant a QualityRouter deployment had to appear AFTER all its referenced models in config.yaml, because `Router._create_deployment` populates `model_list` incrementally. Any `available_models` defined after the router entry would silently be reported as missing. Also adds 6 new tests covering each fix: - test_invalid_quality_tier_type_raises_clear_error - test_router_can_be_instantiated_before_its_targets_exist - test_set_model_list_clears_quality_routers_registry - test_rounds_down_when_no_higher_tier_exists - test_rounds_down_prefers_closest_lower_tier - test_prefers_round_up_over_round_down Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com> * style: apply black 24.10.0 formatting to pre-existing offenders Unblocks the LiteLLM Linting check for this PR — these 12 files are already failing `black --check` on main (the lint workflow only runs on PRs, so main drifts). No behavior changes; formatting-only. Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com> * Update litellm/router.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --------- Co-authored-by: Claude Opus 4 (1M context) <noreply@anthropic.com> Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * Support /v1/responses in complexity router (#26137) * feat(proxy): add --reload flag for uvicorn hot reload (dev only) Opt-in CLI flag, off by default, no env var. Only affects the uvicorn run path; gunicorn/hypercorn paths and prod (which doesn't pass the flag) are unaffected. * Feature/add audio support for scaleway (#26110) * feat(scaleway): add SCALEWAY to LlmProviders enum * feat(scaleway): add audio transcription config and dispatch wiring Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> * test(scaleway): add behavior tests for audio transcription config Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> * chore(scaleway): advertise audio_transcriptions in endpoint-support JSON * docs(scaleway): document audio transcription support * fix(scaleway): address PR review — plain-text response_format + missing-key fail-fast Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> * test(scaleway): cover new response paths, drop gettysburg.wav coupling Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> --------- Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com> * Prompt Compression - add it to the proxy (#25729) * refactor: new agentic loop event hook simplifies how to create logic for tool based multi llm calls * fix: compress - make it work on anthropic input as well * fix(compress.py): working prompt compression for claude code ensures claude code messages can run through proxy easily * docs: add agentic loop hook guide * docs: add agentic_loop_hook to sidebar * fix: fix multiple arguments error * fix: fix tool call loop for compression on streaming /v1/messages * fix: fix linting errors * fix: fix ci/cd errors * feat(litellm_pre_call_utils.py): use claude code session for litellm session id allows claude code logs to be stitched together, making it easy to know they were all part of the same conversation * fix: suppress incorrect mypy warning rE: module * revert: drop PR's changes to litellm/proxy/_experimental/out/ Restores the 34 HTML files under _experimental/out/ to their pre-PR paths (X/index.html -> X.html). All renames are R100 (content unchanged); no other files are touched. * fix: address greptile review comments on PR #25729 - Skip ``kwargs["tools"] = []`` injection when compression is a no-op — Anthropic Messages rejects empty tool arrays on requests that did not originally declare tools. - Move agentic-loop safety guards (fingerprint cycle / max depth) out of the per-callback try/except so they propagate instead of being swallowed by the generic exception handler. Extracted _check_agentic_loop_safety. - Gate generic ``x-<vendor>-session-id`` capture behind the LITELLM_CAPTURE_VENDOR_SESSION_HEADERS env var (off by default) to preserve backwards compatibility; explicit x-litellm-* headers are unaffected. - Fix monkeypatch target in pre-call-hook test to patch the actual module-level binding (litellm.integrations.compression_interception.handler.compress). - Add regression tests for empty-tools skip and opt-in session capture. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * revert: drop LITELLM_CAPTURE_VENDOR_SESSION_HEADERS flag Generic x-<vendor>-session-id header capture is a new feature and only runs *after* the explicit x-litellm-trace-id / x-litellm-session-id checks, so it does not change behavior for any existing caller that was already using the LiteLLM headers — no backwards-incompatibility to gate. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * refactor(compress): replace input_type with CallTypes call_type Drop the bespoke ``CompressionInputType`` literal and use the existing ``litellm.types.utils.CallTypes`` enum instead. ``litellm.compress()`` now takes ``call_type: Union[CallTypes, str]`` (default ``CallTypes.completion``) — no new concept to learn, and the enum is already the way the rest of the codebase talks about request shapes. Supported values: ``completion`` / ``acompletion`` (OpenAI chat-completions shape) and ``anthropic_messages`` (Anthropic structured content blocks). Updated: compress(), the compression_interception handler, tests, docs, and the two eval scripts. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com> * Support /v1/responses in complexity router Adds cross-format support to the complexity router via the guardrail translation handler dispatch. Adds get_structured_messages to base translation plus OpenAI chat, Responses, and Anthropic handlers. Auto-router helper _extract_text_from_messages handles tool-call and multimodal messages. Widens async_pre_routing_hook messages type to Dict[str, Any]. Fixes https://github.com/BerriAI/litellm/issues/25134 * chore: apply black formatting * fix: fallback to trying each handler when route inference fails --------- Co-authored-by: Ryan Crabbe <ryan@berri.ai> Co-authored-by: nhyy244 <106547304+nhyy244@users.noreply.github.com> Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com> * test: cover _is_quality_router_deployment and init_quality_router_deployment * fix: reset auto_routers on set_model_list to prevent hot-reload ValueError * style: apply black formatting to websearch_interception and agentic_streaming_iterator --------- Co-authored-by: yuneng-jiang <yuneng@berri.ai> Co-authored-by: Claude Opus 4 (1M context) <noreply@anthropic.com> Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> Co-authored-by: Ryan Crabbe <ryan@berri.ai> Co-authored-by: nhyy244 <106547304+nhyy244@users.noreply.github.com>
1049 lines
41 KiB
Python
1049 lines
41 KiB
Python
"""
|
|
Tests for the ComplexityRouter.
|
|
|
|
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
|
|
|
|
import pytest
|
|
|
|
sys.path.insert(
|
|
0, os.path.abspath("../../..")
|
|
) # Adds the parent directory to the system path
|
|
|
|
from litellm import Router
|
|
from litellm.router_strategy.complexity_router.complexity_router import (
|
|
ComplexityRouter,
|
|
DimensionScore,
|
|
)
|
|
from litellm.router_strategy.complexity_router.config import (
|
|
DEFAULT_COMPLEXITY_CONFIG,
|
|
ComplexityRouterConfig,
|
|
ComplexityTier,
|
|
)
|
|
|
|
|
|
@pytest.fixture
|
|
def mock_router_instance():
|
|
"""Create a mock LiteLLM Router instance."""
|
|
router = MagicMock()
|
|
return router
|
|
|
|
|
|
@pytest.fixture
|
|
def basic_config() -> Dict:
|
|
"""Basic configuration with tier mappings."""
|
|
return {
|
|
"tiers": {
|
|
"SIMPLE": "gpt-4o-mini",
|
|
"MEDIUM": "gpt-4o",
|
|
"COMPLEX": "claude-sonnet-4-20250514",
|
|
"REASONING": "o1-preview",
|
|
},
|
|
"tier_boundaries": {
|
|
"simple_medium": 0.25,
|
|
"medium_complex": 0.50,
|
|
"complex_reasoning": 0.75,
|
|
},
|
|
}
|
|
|
|
|
|
@pytest.fixture
|
|
def complexity_router(mock_router_instance, basic_config):
|
|
"""Create a ComplexityRouter instance with basic config."""
|
|
return ComplexityRouter(
|
|
model_name="test-complexity-router",
|
|
litellm_router_instance=mock_router_instance,
|
|
complexity_router_config=basic_config,
|
|
)
|
|
|
|
|
|
class TestDimensionScore:
|
|
"""Test the DimensionScore class."""
|
|
|
|
def test_dimension_score_creation(self):
|
|
"""Test creating a DimensionScore."""
|
|
score = DimensionScore("tokenCount", 0.5, "short (25 tokens)")
|
|
assert score.name == "tokenCount"
|
|
assert score.score == 0.5
|
|
assert score.signal == "short (25 tokens)"
|
|
|
|
def test_dimension_score_no_signal(self):
|
|
"""Test creating a DimensionScore without signal."""
|
|
score = DimensionScore("tokenCount", 0)
|
|
assert score.name == "tokenCount"
|
|
assert score.score == 0
|
|
assert score.signal is None
|
|
|
|
|
|
class TestComplexityRouterInit:
|
|
"""Test ComplexityRouter initialization."""
|
|
|
|
def test_init_with_config(self, mock_router_instance, basic_config):
|
|
"""Test initialization with configuration."""
|
|
router = ComplexityRouter(
|
|
model_name="test-router",
|
|
litellm_router_instance=mock_router_instance,
|
|
complexity_router_config=basic_config,
|
|
)
|
|
assert router.model_name == "test-router"
|
|
assert router.config.tiers["SIMPLE"] == "gpt-4o-mini"
|
|
assert router.config.tiers["REASONING"] == "o1-preview"
|
|
|
|
def test_init_without_config(self, mock_router_instance):
|
|
"""Test initialization without configuration uses defaults."""
|
|
router = ComplexityRouter(
|
|
model_name="test-router",
|
|
litellm_router_instance=mock_router_instance,
|
|
)
|
|
assert router.model_name == "test-router"
|
|
# Should have equivalent default values but NOT be the same instance
|
|
assert router.config.tiers == DEFAULT_COMPLEXITY_CONFIG.tiers
|
|
assert router.config is not DEFAULT_COMPLEXITY_CONFIG # Not a singleton
|
|
|
|
def test_init_with_default_model(self, mock_router_instance, basic_config):
|
|
"""Test initialization with default_model override."""
|
|
router = ComplexityRouter(
|
|
model_name="test-router",
|
|
litellm_router_instance=mock_router_instance,
|
|
complexity_router_config=basic_config,
|
|
default_model="fallback-model",
|
|
)
|
|
assert router.config.default_model == "fallback-model"
|
|
|
|
|
|
class TestTokenScoring:
|
|
"""Test token count scoring."""
|
|
|
|
def test_short_prompt_negative_score(self, complexity_router):
|
|
"""Short prompts should get negative scores (simple indicator)."""
|
|
tier, score, signals = complexity_router.classify("What is Python?")
|
|
# Should be classified as SIMPLE due to short length and simple indicator
|
|
assert tier == ComplexityTier.SIMPLE
|
|
assert any("short" in s.lower() for s in signals) or any(
|
|
"simple" in s.lower() for s in signals
|
|
)
|
|
|
|
def test_long_prompt_positive_score(self, complexity_router):
|
|
"""Long prompts should get positive scores (complex indicator)."""
|
|
# Create a long prompt (~600 tokens)
|
|
long_prompt = "Explain the following concept in detail: " + " ".join(
|
|
["distributed systems architecture and microservices patterns"] * 50
|
|
)
|
|
tier, score, signals = complexity_router.classify(long_prompt)
|
|
# Should have positive score and detect long token count or technical terms
|
|
assert score > 0, f"Expected positive score for long prompt, got {score}"
|
|
assert any("long" in s.lower() for s in signals) or any(
|
|
"technical" in s.lower() for s in signals
|
|
)
|
|
|
|
|
|
class TestCodePresenceScoring:
|
|
"""Test code-related keyword scoring."""
|
|
|
|
def test_code_keywords_increase_complexity(self, complexity_router):
|
|
"""Code keywords should increase complexity score."""
|
|
prompt = "Write a Python function that implements a binary search algorithm with async support"
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
# Should detect code presence
|
|
assert any("code" in s.lower() for s in signals)
|
|
# Score should be positive (code keywords add to complexity)
|
|
assert score > -0.5 # Not heavily negative
|
|
|
|
def test_multiple_code_keywords(self, complexity_router):
|
|
"""Multiple code keywords should strongly increase complexity."""
|
|
prompt = (
|
|
"Debug this Python function that uses async/await with try/catch "
|
|
"for API endpoint error handling in the database query"
|
|
)
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
assert any("code" in s.lower() for s in signals)
|
|
|
|
|
|
class TestReasoningMarkerScoring:
|
|
"""Test reasoning marker detection."""
|
|
|
|
def test_single_reasoning_marker(self, complexity_router):
|
|
"""Single reasoning marker should increase score."""
|
|
prompt = "Think through this problem step by step and explain your reasoning"
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
assert any("reasoning" in s.lower() for s in signals)
|
|
|
|
def test_multiple_reasoning_markers_override(self, complexity_router):
|
|
"""Multiple reasoning markers should force REASONING tier."""
|
|
prompt = "Let's think step by step. Analyze this carefully and reason through each option. Show your work."
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
# 2+ reasoning markers should force REASONING tier
|
|
assert tier == ComplexityTier.REASONING
|
|
|
|
def test_system_prompt_reasoning_not_counted(self, complexity_router):
|
|
"""Reasoning markers in system prompt should not count for override."""
|
|
user_prompt = "What is 2+2?"
|
|
system_prompt = "Think step by step before answering."
|
|
tier, score, signals = complexity_router.classify(user_prompt, system_prompt)
|
|
# Should still be SIMPLE since user message is simple
|
|
assert tier in [ComplexityTier.SIMPLE, ComplexityTier.MEDIUM]
|
|
|
|
|
|
class TestSimpleIndicatorScoring:
|
|
"""Test simple indicator detection."""
|
|
|
|
def test_simple_greeting(self, complexity_router):
|
|
"""Simple greetings should be classified as SIMPLE."""
|
|
tier, score, signals = complexity_router.classify("Hello, how are you?")
|
|
assert tier == ComplexityTier.SIMPLE
|
|
|
|
def test_definition_questions(self, complexity_router):
|
|
"""Definition questions should be classified as SIMPLE."""
|
|
prompts = [
|
|
"What is machine learning?",
|
|
"Define artificial intelligence",
|
|
"Who is Alan Turing?",
|
|
]
|
|
for prompt in prompts:
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
assert tier == ComplexityTier.SIMPLE, f"Expected SIMPLE for: {prompt}"
|
|
|
|
|
|
class TestMultiStepPatterns:
|
|
"""Test multi-step pattern detection."""
|
|
|
|
def test_first_then_pattern(self, complexity_router):
|
|
"""'First...then' patterns should increase complexity."""
|
|
prompt = (
|
|
"First analyze the data, then create a visualization, then write a report"
|
|
)
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
assert any("multi-step" in s.lower() for s in signals)
|
|
|
|
def test_numbered_steps(self, complexity_router):
|
|
"""Numbered steps should increase complexity."""
|
|
prompt = "1. Set up the environment 2. Install dependencies 3. Run the tests"
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
assert any("multi-step" in s.lower() for s in signals)
|
|
|
|
|
|
class TestQuestionComplexity:
|
|
"""Test question complexity scoring."""
|
|
|
|
def test_multiple_questions(self, complexity_router):
|
|
"""Multiple questions should increase complexity."""
|
|
prompt = "What is the capital? Where is it located? How many people live there? What's the climate like?"
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
assert any("question" in s.lower() for s in signals)
|
|
|
|
|
|
class TestTierAssignment:
|
|
"""Test tier assignment based on scores."""
|
|
|
|
def test_simple_tier(self, complexity_router):
|
|
"""Simple prompts should get SIMPLE tier."""
|
|
tier, score, signals = complexity_router.classify("Hi there!")
|
|
assert tier == ComplexityTier.SIMPLE
|
|
|
|
def test_medium_tier(self, complexity_router):
|
|
"""Moderately complex prompts should get MEDIUM tier."""
|
|
prompt = "Explain how REST APIs work with HTTP methods"
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
assert tier in [ComplexityTier.SIMPLE, ComplexityTier.MEDIUM]
|
|
|
|
def test_complex_tier(self, complexity_router):
|
|
"""Complex prompts should get positive complexity score with technical signals."""
|
|
prompt = (
|
|
"Design a distributed microservice architecture for a high-throughput "
|
|
"real-time data processing pipeline with Kubernetes orchestration, "
|
|
"implementing proper authentication and encryption protocols"
|
|
)
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
# Should detect technical terms
|
|
assert any(
|
|
"technical" in s.lower() for s in signals
|
|
), f"Expected technical signals, got {signals}"
|
|
# Score should be positive due to technical content
|
|
assert score > 0, f"Expected positive score, got {score}"
|
|
|
|
def test_reasoning_tier(self, complexity_router):
|
|
"""Reasoning prompts should get REASONING tier."""
|
|
prompt = (
|
|
"Think step by step and reason through this: Analyze the pros and cons "
|
|
"of different database architectures for our distributed system, "
|
|
"considering performance, scalability, and consistency tradeoffs"
|
|
)
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
assert tier == ComplexityTier.REASONING
|
|
|
|
|
|
class TestModelSelection:
|
|
"""Test model selection based on tier."""
|
|
|
|
def test_get_model_for_simple(self, complexity_router):
|
|
"""Should return correct model for SIMPLE tier."""
|
|
model = complexity_router.get_model_for_tier(ComplexityTier.SIMPLE)
|
|
assert model == "gpt-4o-mini"
|
|
|
|
def test_get_model_for_complex(self, complexity_router):
|
|
"""Should return correct model for COMPLEX tier."""
|
|
model = complexity_router.get_model_for_tier(ComplexityTier.COMPLEX)
|
|
assert model == "claude-sonnet-4-20250514"
|
|
|
|
def test_get_model_for_reasoning(self, complexity_router):
|
|
"""Should return correct model for REASONING tier."""
|
|
model = complexity_router.get_model_for_tier(ComplexityTier.REASONING)
|
|
assert model == "o1-preview"
|
|
|
|
def test_get_model_fallback_to_default(self, mock_router_instance):
|
|
"""Should fallback to default_model if tier not configured."""
|
|
config = {
|
|
"tiers": {}, # Empty tiers
|
|
"default_model": "fallback-model",
|
|
}
|
|
router = ComplexityRouter(
|
|
model_name="test-router",
|
|
litellm_router_instance=mock_router_instance,
|
|
complexity_router_config=config,
|
|
)
|
|
model = router.get_model_for_tier(ComplexityTier.SIMPLE)
|
|
assert model == "fallback-model"
|
|
|
|
|
|
class TestPreRoutingHook:
|
|
"""Test the async_pre_routing_hook method."""
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_simple_message(self, complexity_router):
|
|
"""Test pre-routing hook with a simple message."""
|
|
messages = [{"role": "user", "content": "Hello!"}]
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=messages,
|
|
)
|
|
assert result is not None
|
|
assert result.model == "gpt-4o-mini" # SIMPLE tier model
|
|
assert result.messages == messages
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_complex_message(self, complexity_router):
|
|
"""Test pre-routing hook with a message containing technical content."""
|
|
messages = [
|
|
{
|
|
"role": "user",
|
|
"content": (
|
|
"Design a distributed microservice architecture with Kubernetes "
|
|
"orchestration, implementing proper authentication, encryption, "
|
|
"and database optimization for high throughput. Think step by step "
|
|
"about the performance implications and scalability requirements."
|
|
),
|
|
}
|
|
]
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=messages,
|
|
)
|
|
assert result is not None
|
|
# Should return a valid model from the configured tiers
|
|
assert result.model in [
|
|
"gpt-4o-mini",
|
|
"gpt-4o",
|
|
"claude-sonnet-4-20250514",
|
|
"o1-preview",
|
|
]
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_no_messages(self, complexity_router):
|
|
"""Test pre-routing hook returns None when no messages."""
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=None,
|
|
)
|
|
assert result is None
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_empty_messages(self, complexity_router):
|
|
"""Test pre-routing hook returns None when messages empty."""
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=[],
|
|
)
|
|
assert result is None
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_with_system_prompt(self, complexity_router):
|
|
"""Test pre-routing hook considers system prompt."""
|
|
messages = [
|
|
{"role": "system", "content": "You are a helpful assistant."},
|
|
{"role": "user", "content": "Hello!"},
|
|
]
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=messages,
|
|
)
|
|
assert result is not None
|
|
# Should still be SIMPLE
|
|
assert result.model == "gpt-4o-mini"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_reasoning_message(self, complexity_router):
|
|
"""Test pre-routing hook with reasoning markers."""
|
|
messages = [
|
|
{
|
|
"role": "user",
|
|
"content": "Let's think step by step and reason through this problem carefully.",
|
|
}
|
|
]
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=messages,
|
|
)
|
|
assert result is not None
|
|
assert result.model == "o1-preview" # REASONING tier model
|
|
|
|
|
|
class TestConfigOverrides:
|
|
"""Test configuration override functionality."""
|
|
|
|
def test_custom_tier_boundaries(self, mock_router_instance):
|
|
"""Test custom tier boundaries work correctly."""
|
|
config = {
|
|
"tiers": {
|
|
"SIMPLE": "mini-model",
|
|
"MEDIUM": "medium-model",
|
|
"COMPLEX": "complex-model",
|
|
"REASONING": "reasoning-model",
|
|
},
|
|
"tier_boundaries": {
|
|
"simple_medium": -0.5, # Very low threshold - anything above -0.5 is MEDIUM+
|
|
"medium_complex": -0.3,
|
|
"complex_reasoning": 0.0,
|
|
},
|
|
}
|
|
router = ComplexityRouter(
|
|
model_name="test-router",
|
|
litellm_router_instance=mock_router_instance,
|
|
complexity_router_config=config,
|
|
)
|
|
# With very low thresholds, even neutral prompts should be COMPLEX or higher
|
|
tier, score, signals = router.classify(
|
|
"Explain how HTTP works with REST APIs and distributed systems"
|
|
)
|
|
# With boundaries this low, should be at least MEDIUM (anything above -0.5)
|
|
assert (
|
|
tier != ComplexityTier.SIMPLE
|
|
), f"Expected non-SIMPLE tier, got {tier} with score {score}"
|
|
|
|
def test_custom_token_thresholds(self, mock_router_instance):
|
|
"""Test custom token thresholds work correctly."""
|
|
config = {
|
|
"tiers": {
|
|
"SIMPLE": "mini-model",
|
|
"MEDIUM": "medium-model",
|
|
"COMPLEX": "complex-model",
|
|
"REASONING": "reasoning-model",
|
|
},
|
|
"token_thresholds": {
|
|
"simple": 10, # Very low - prompts with >10 tokens are not "short"
|
|
"complex": 100, # Lower than default - prompts with >100 tokens are "long"
|
|
},
|
|
}
|
|
router = ComplexityRouter(
|
|
model_name="test-router",
|
|
litellm_router_instance=mock_router_instance,
|
|
complexity_router_config=config,
|
|
)
|
|
# A longer prompt (~150 tokens) should be considered "long" with these thresholds
|
|
long_prompt = "This is a test prompt " * 30 # ~120 tokens
|
|
tier, score, signals = router.classify(long_prompt)
|
|
# Should get token length signal indicating "long"
|
|
assert any(
|
|
"long" in s.lower() if s else False for s in signals
|
|
), f"Expected 'long' signal, got {signals}"
|
|
|
|
|
|
class TestAsyncPreRoutingHookEdgeCases:
|
|
"""Test edge cases for async_pre_routing_hook method."""
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_multi_turn_conversation(self, complexity_router):
|
|
"""Test pre-routing hook with multi-turn conversation uses last user message."""
|
|
messages = [
|
|
{"role": "user", "content": "What is Python?"},
|
|
{"role": "assistant", "content": "Python is a programming language."},
|
|
{"role": "user", "content": "Hello!"}, # Last user message - simple
|
|
]
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=messages,
|
|
)
|
|
assert result is not None
|
|
assert result.model == "gpt-4o-mini" # SIMPLE tier based on last message
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_multi_user_messages(self, complexity_router):
|
|
"""Test pre-routing hook uses the last user message for classification."""
|
|
# Multiple user messages - should classify based on the LAST one
|
|
messages = [
|
|
{
|
|
"role": "user",
|
|
"content": "Design a complex distributed system",
|
|
}, # Complex prompt
|
|
{"role": "assistant", "content": "I can help with that."},
|
|
{
|
|
"role": "user",
|
|
"content": "Hello!",
|
|
}, # Simple prompt - this should be used
|
|
]
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=messages,
|
|
)
|
|
assert result is not None
|
|
# Should use the last user message "Hello!" which is SIMPLE
|
|
assert result.model == "gpt-4o-mini"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_no_user_message(self, complexity_router):
|
|
"""Test pre-routing hook falls back to default model when no user message found."""
|
|
messages = [
|
|
{"role": "system", "content": "You are helpful."},
|
|
{"role": "assistant", "content": "Hello!"},
|
|
]
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=messages,
|
|
)
|
|
# Should return default model rather than None (None would cause
|
|
# the complexity_router deployment itself to be selected, crashing)
|
|
assert result is not None
|
|
assert result.model in [
|
|
"gpt-4o-mini",
|
|
"gpt-4o",
|
|
"claude-sonnet-4-20250514",
|
|
"o1-preview",
|
|
]
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_list_content(self, complexity_router):
|
|
"""Test pre-routing hook handles list-format message content (OpenAI multi-part format)."""
|
|
messages = [
|
|
{
|
|
"role": "user",
|
|
"content": [{"type": "text", "text": "Hello, how are you?"}],
|
|
},
|
|
]
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=messages,
|
|
)
|
|
# Should extract text from list content and classify normally
|
|
assert result is not None
|
|
assert result.model == "gpt-4o-mini" # "Hello, how are you?" is SIMPLE
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_list_content_complex(self, complexity_router):
|
|
"""Test pre-routing hook classifies list-format content by complexity."""
|
|
messages = [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{
|
|
"type": "text",
|
|
"text": "Think step by step and reason through this: design a distributed system",
|
|
},
|
|
{
|
|
"type": "image_url",
|
|
"image_url": {"url": "data:image/png;base64,abc"},
|
|
},
|
|
],
|
|
}
|
|
]
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=messages,
|
|
)
|
|
assert result is not None
|
|
assert result.model == "o1-preview" # REASONING tier
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_preserves_messages(self, complexity_router):
|
|
"""Test pre-routing hook preserves original messages in response."""
|
|
messages = [
|
|
{"role": "system", "content": "Be helpful"},
|
|
{"role": "user", "content": "Hello!"},
|
|
]
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=messages,
|
|
)
|
|
assert result is not None
|
|
assert result.messages == messages
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pre_routing_hook_empty_string_content(self, complexity_router):
|
|
"""Test pre-routing hook falls back to default model for empty string content."""
|
|
messages = [
|
|
{"role": "user", "content": ""},
|
|
]
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=messages,
|
|
)
|
|
# Empty string content → no extractable user message → routes to default model
|
|
assert result is not None
|
|
assert result.model in [
|
|
"gpt-4o-mini",
|
|
"gpt-4o",
|
|
"claude-sonnet-4-20250514",
|
|
"o1-preview",
|
|
]
|
|
|
|
|
|
class TestSingletonMutation:
|
|
"""Test that the config singleton is not mutated."""
|
|
|
|
def test_default_config_not_mutated(self, mock_router_instance):
|
|
"""Test that creating routers without config doesn't mutate defaults."""
|
|
from litellm.router_strategy.complexity_router.config import (
|
|
ComplexityRouterConfig,
|
|
)
|
|
|
|
# Get original default
|
|
original_default = ComplexityRouterConfig().default_model
|
|
|
|
# Create router with empty config and custom default_model
|
|
router1 = ComplexityRouter(
|
|
model_name="test-router-1",
|
|
litellm_router_instance=mock_router_instance,
|
|
complexity_router_config=None,
|
|
default_model="custom-fallback",
|
|
)
|
|
|
|
# Create another router without config
|
|
router2 = ComplexityRouter(
|
|
model_name="test-router-2",
|
|
litellm_router_instance=mock_router_instance,
|
|
complexity_router_config=None,
|
|
)
|
|
|
|
# Router2 should have fresh defaults, not router1's custom default_model
|
|
# Create a fresh config to check
|
|
fresh_config = ComplexityRouterConfig()
|
|
assert fresh_config.default_model == original_default
|
|
assert router1.config.default_model == "custom-fallback"
|
|
# Router2's config should be independent
|
|
assert router2.config is not router1.config
|
|
|
|
|
|
class TestKeywordFalsePositives:
|
|
"""Test that keyword matching uses word boundaries to avoid false positives."""
|
|
|
|
def test_api_not_in_capital(self, complexity_router):
|
|
"""'api' should not match in 'capital'."""
|
|
prompt = "What is the capital of France?"
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
# 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'"
|
|
# Should be SIMPLE (definition question)
|
|
assert tier == ComplexityTier.SIMPLE
|
|
|
|
def test_git_not_in_digital(self, complexity_router):
|
|
"""'git' should not match in 'digital'."""
|
|
prompt = "Explain digital marketing strategies"
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
# 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'"
|
|
|
|
def test_try_not_in_entry(self, complexity_router):
|
|
"""'try' should not match in 'entry'."""
|
|
prompt = "What is the entry point for this application?"
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
# 'entry' contains 'try' but should not trigger code detection
|
|
# Note: 'application' might trigger something, but 'try' should not
|
|
pass # Just ensure no crash; false positive check is the main goal
|
|
|
|
def test_error_not_in_terrorism(self, complexity_router):
|
|
"""'error' should not match in 'terrorism'."""
|
|
prompt = "The country is dealing with terrorism"
|
|
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'"
|
|
|
|
def test_class_not_in_classical(self, complexity_router):
|
|
"""'class' should not match in 'classical'."""
|
|
prompt = "I enjoy listening to classical music"
|
|
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'"
|
|
|
|
def test_merge_not_in_emerged(self, complexity_router):
|
|
"""'merge' should not match in 'emerged'."""
|
|
prompt = "A new leader emerged from the crowd"
|
|
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'"
|
|
|
|
def test_actual_api_keyword_detected(self, complexity_router):
|
|
"""Actual 'api' usage should be detected."""
|
|
prompt = "How do I call the REST api endpoint?"
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
# Should detect code presence from actual 'api' usage
|
|
assert any(
|
|
"code" in s.lower() for s in signals
|
|
), f"Expected code signal for 'api', got {signals}"
|
|
|
|
def test_actual_git_keyword_detected(self, complexity_router):
|
|
"""Actual 'git' usage should be detected."""
|
|
prompt = "How do I use git to commit changes?"
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
# Should detect code presence from actual 'git' usage
|
|
assert any(
|
|
"code" in s.lower() for s in signals
|
|
), f"Expected code signal for 'git', got {signals}"
|
|
|
|
|
|
class TestEdgeCases:
|
|
"""Test edge cases and error handling."""
|
|
|
|
def test_empty_prompt(self, complexity_router):
|
|
"""Test handling of empty prompt."""
|
|
tier, score, signals = complexity_router.classify("")
|
|
assert tier == ComplexityTier.SIMPLE
|
|
assert score <= 0
|
|
|
|
def test_very_long_prompt(self, complexity_router):
|
|
"""Test handling of very long prompt."""
|
|
# 10000+ character prompt
|
|
long_prompt = "explain " * 2000
|
|
tier, score, signals = complexity_router.classify(long_prompt)
|
|
# Should have positive score due to length
|
|
assert score > 0, f"Expected positive score for very long prompt, got {score}"
|
|
# Should detect long token count
|
|
assert any(
|
|
"long" in s.lower() for s in signals
|
|
), f"Expected 'long' signal, got {signals}"
|
|
|
|
def test_unicode_prompt(self, complexity_router):
|
|
"""Test handling of unicode characters."""
|
|
prompt = "What is 日本語? Explain émojis 🎉 and symbols ∑∏∫"
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
# Should not crash, should be classified
|
|
assert tier in [ComplexityTier.SIMPLE, ComplexityTier.MEDIUM]
|
|
|
|
def test_multiline_prompt(self, complexity_router):
|
|
"""Test handling of multiline prompts with step patterns."""
|
|
prompt = """
|
|
Step 1: Analyze the problem.
|
|
Step 2: Propose a solution.
|
|
Step 3: Implement it.
|
|
"""
|
|
tier, score, signals = complexity_router.classify(prompt)
|
|
# The "step N" pattern should be detected
|
|
assert any(
|
|
"multi-step" in s.lower() for s in signals
|
|
), f"Expected multi-step signal, got {signals}"
|
|
|
|
|
|
class TestRouterComplexityDeploymentMethods:
|
|
"""Tests for Router._is_complexity_router_deployment and Router.init_complexity_router_deployment."""
|
|
|
|
def test_is_complexity_router_deployment_true(self):
|
|
"""_is_complexity_router_deployment returns True for complexity router models."""
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "gpt-4o-mini",
|
|
"litellm_params": {"model": "openai/gpt-4o-mini"},
|
|
}
|
|
]
|
|
)
|
|
from litellm.types.router import LiteLLM_Params
|
|
|
|
params = LiteLLM_Params(model="auto_router/complexity_router/my-router")
|
|
assert router._is_complexity_router_deployment(params) is True
|
|
|
|
def test_is_complexity_router_deployment_false(self):
|
|
"""_is_complexity_router_deployment returns False for regular models."""
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "gpt-4o-mini",
|
|
"litellm_params": {"model": "openai/gpt-4o-mini"},
|
|
}
|
|
]
|
|
)
|
|
from litellm.types.router import LiteLLM_Params
|
|
|
|
params = LiteLLM_Params(model="openai/gpt-4o-mini")
|
|
assert router._is_complexity_router_deployment(params) is False
|
|
|
|
def test_init_complexity_router_deployment(self):
|
|
"""init_complexity_router_deployment registers a ComplexityRouter."""
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "gpt-4o-mini",
|
|
"litellm_params": {"model": "openai/gpt-4o-mini"},
|
|
}
|
|
]
|
|
)
|
|
from litellm.types.router import Deployment, LiteLLM_Params
|
|
|
|
deployment = Deployment(
|
|
model_name="auto_router/complexity_router/test-router",
|
|
litellm_params=LiteLLM_Params(
|
|
model="auto_router/complexity_router/test-router",
|
|
complexity_router_default_model="gpt-4o-mini",
|
|
complexity_router_config={
|
|
"tiers": {
|
|
"SIMPLE": "gpt-4o-mini",
|
|
"MEDIUM": "gpt-4o",
|
|
"COMPLEX": "claude-sonnet-4-20250514",
|
|
"REASONING": "o1-preview",
|
|
}
|
|
},
|
|
),
|
|
model_info={"id": "test-id"},
|
|
)
|
|
router.init_complexity_router_deployment(deployment)
|
|
assert "auto_router/complexity_router/test-router" in router.complexity_routers
|
|
|
|
|
|
class TestAsyncPreRoutingHookMultiFormat:
|
|
"""Test async_pre_routing_hook with multiple input formats."""
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_should_route_with_chat_completions_messages(self, complexity_router):
|
|
"""Test routing with standard chat completions messages."""
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=[{"role": "user", "content": "What is 2+2?"}],
|
|
)
|
|
assert result is not None
|
|
assert result.model is not None
|
|
assert result.messages is not None
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_should_route_with_responses_api_string_input(
|
|
self, complexity_router
|
|
):
|
|
"""Test routing with Responses API string input via handler dispatch."""
|
|
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
|
OpenAIResponsesHandler,
|
|
)
|
|
from litellm.types.utils import CallTypes
|
|
|
|
mock_mappings = {CallTypes.responses: OpenAIResponsesHandler}
|
|
|
|
with patch(
|
|
"litellm.llms.load_guardrail_translation_mappings",
|
|
return_value=mock_mappings,
|
|
):
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={"input": "What is the capital of France?"},
|
|
messages=None,
|
|
input="What is the capital of France?",
|
|
)
|
|
|
|
assert result is not None
|
|
assert result.model is not None
|
|
# messages should be None since the original request didn't have messages
|
|
assert result.messages is None
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_should_route_with_responses_api_list_input(self, complexity_router):
|
|
"""Test routing with Responses API list input via handler dispatch."""
|
|
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
|
OpenAIResponsesHandler,
|
|
)
|
|
from litellm.types.utils import CallTypes
|
|
|
|
mock_mappings = {CallTypes.responses: OpenAIResponsesHandler}
|
|
|
|
list_input = [
|
|
{"role": "user", "content": "Hello"},
|
|
{"role": "assistant", "content": "Hi there!"},
|
|
{
|
|
"role": "user",
|
|
"content": "Write a Python function to sort a list using merge sort",
|
|
},
|
|
]
|
|
|
|
with patch(
|
|
"litellm.llms.load_guardrail_translation_mappings",
|
|
return_value=mock_mappings,
|
|
):
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={"input": list_input},
|
|
messages=None,
|
|
input=list_input,
|
|
)
|
|
|
|
assert result is not None
|
|
assert result.model is not None
|
|
assert result.messages is None
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_should_use_route_based_inference(self, complexity_router):
|
|
"""Test that route-based call type inference is used when available."""
|
|
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
|
OpenAIResponsesHandler,
|
|
)
|
|
from litellm.types.utils import CallTypes
|
|
|
|
mock_mappings = {CallTypes.responses: OpenAIResponsesHandler}
|
|
|
|
with patch(
|
|
"litellm.llms.load_guardrail_translation_mappings",
|
|
return_value=mock_mappings,
|
|
):
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={
|
|
"input": "Roll 2d4+1",
|
|
"litellm_metadata": {
|
|
"user_api_key_request_route": "/v1/responses",
|
|
},
|
|
},
|
|
messages=None,
|
|
)
|
|
|
|
assert result is not None
|
|
assert result.model is not None
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_should_return_none_when_no_messages_or_input(
|
|
self, complexity_router
|
|
):
|
|
"""Test that None is returned when neither messages nor input is available."""
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={},
|
|
messages=None,
|
|
input=None,
|
|
)
|
|
assert result is None
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_should_prefer_original_messages_over_conversion(
|
|
self, complexity_router
|
|
):
|
|
"""Test that original messages are used when both messages and input are available."""
|
|
messages = [{"role": "user", "content": "What is 2+2?"}]
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={"input": "This should be ignored"},
|
|
messages=messages,
|
|
)
|
|
assert result is not None
|
|
assert result.messages == messages
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_should_include_instructions_in_classification(
|
|
self, complexity_router
|
|
):
|
|
"""Test that Responses API instructions influence classification via system message."""
|
|
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
|
OpenAIResponsesHandler,
|
|
)
|
|
from litellm.types.utils import CallTypes
|
|
|
|
mock_mappings = {CallTypes.responses: OpenAIResponsesHandler}
|
|
|
|
with patch(
|
|
"litellm.llms.load_guardrail_translation_mappings",
|
|
return_value=mock_mappings,
|
|
):
|
|
result = await complexity_router.async_pre_routing_hook(
|
|
model="test-model",
|
|
request_kwargs={
|
|
"input": "Write merge sort",
|
|
"instructions": "You are an expert Python developer. Use advanced algorithms and optimize for performance.",
|
|
},
|
|
messages=None,
|
|
)
|
|
|
|
assert result is not None
|
|
assert result.model is not None
|
|
|
|
|
|
class TestExtractUserMessageAndSystemPrompt:
|
|
"""Test the _extract_user_message_and_system_prompt static method."""
|
|
|
|
def test_should_extract_user_message(self):
|
|
"""Test extraction of the last user message."""
|
|
messages = [
|
|
{"role": "system", "content": "You are helpful."},
|
|
{"role": "user", "content": "Hello"},
|
|
{"role": "assistant", "content": "Hi!"},
|
|
{"role": "user", "content": "How are you?"},
|
|
]
|
|
user_msg, sys_prompt = ComplexityRouter._extract_user_message_and_system_prompt(
|
|
messages
|
|
)
|
|
assert user_msg == "How are you?"
|
|
assert sys_prompt == "You are helpful."
|
|
|
|
def test_should_handle_no_user_message(self):
|
|
"""Test when there is no user message."""
|
|
messages = [
|
|
{"role": "system", "content": "You are helpful."},
|
|
{"role": "assistant", "content": "Hi!"},
|
|
]
|
|
user_msg, sys_prompt = ComplexityRouter._extract_user_message_and_system_prompt(
|
|
messages
|
|
)
|
|
assert user_msg is None
|
|
assert sys_prompt == "You are helpful."
|
|
|
|
def test_should_handle_multipart_content(self):
|
|
"""Test extraction from multipart content messages."""
|
|
messages = [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{"type": "text", "text": "Describe this image"},
|
|
{
|
|
"type": "image_url",
|
|
"image_url": {"url": "https://example.com/img.png"},
|
|
},
|
|
],
|
|
}
|
|
]
|
|
user_msg, sys_prompt = ComplexityRouter._extract_user_message_and_system_prompt(
|
|
messages
|
|
)
|
|
assert user_msg == "Describe this image"
|
|
assert sys_prompt is None
|
|
|
|
def test_should_handle_empty_messages(self):
|
|
"""Test with empty messages list."""
|
|
user_msg, sys_prompt = ComplexityRouter._extract_user_message_and_system_prompt(
|
|
[]
|
|
)
|
|
assert user_msg is None
|
|
assert sys_prompt is None
|