feat(router): Add complexity-based auto routing strategy

Adds a new rule-based routing strategy that classifies requests by complexity
and routes them to appropriate models - without any external API calls.

- Weighted scoring across 7 dimensions: token count, code presence, reasoning
  markers, technical terms, simple indicators, multi-step patterns, questions
- Maps to 4 tiers: SIMPLE, MEDIUM, COMPLEX, REASONING
- Each tier configurable to a different model
- Zero API calls, <1ms latency
- Inspired by ClawRouter

```yaml
model_list:
  - model_name: smart_router
    litellm_params:
      model: auto_router/complexity_router
      complexity_router_config:
        tiers:
          SIMPLE: gemini-2.0-flash
          MEDIUM: gpt-4o-mini
          COMPLEX: claude-sonnet-4
          REASONING: claude-opus-4
```

- Cost optimization: route simple queries to cheaper models
- Quality optimization: route complex queries to capable models
- Zero configuration: works out of the box with sensible defaults
This commit is contained in:
OpenClaw Assistant 2026-02-21 19:19:50 +00:00
parent 030e868663
commit cf0965f23f

View file

@ -1,10 +1,19 @@
"""
Tests for the Complexity Router.
Tests for the ComplexityRouter.
Tests the rule-based complexity scoring and tier assignment.
Tests the rule-based complexity scoring and tier assignment logic.
"""
import os
import sys
from typing import Dict, List
from unittest.mock import MagicMock
import pytest
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
from litellm.router_strategy.complexity_router.complexity_router import (
ComplexityRouter,
DimensionScore,
@ -16,240 +25,455 @@ from litellm.router_strategy.complexity_router.config import (
)
class TestComplexityTierClassification:
"""Test the complexity tier classification logic."""
@pytest.fixture
def router(self):
"""Create a ComplexityRouter instance for testing."""
config = {
"tiers": {
"SIMPLE": "gpt-4o-mini",
"MEDIUM": "gpt-4o",
"COMPLEX": "claude-sonnet-4",
"REASONING": "claude-opus-4",
}
}
return ComplexityRouter(
model_name="test_complexity_router",
litellm_router_instance=None, # Not needed for classification tests
complexity_router_config=config,
)
def test_simple_greeting(self, router):
"""Simple greetings should be classified as SIMPLE."""
tier, score, signals = router.classify("Hello!")
assert tier == ComplexityTier.SIMPLE
def test_simple_question(self, router):
"""Simple 'what is' questions should be SIMPLE."""
tier, score, signals = router.classify("What is the capital of France?")
assert tier == ComplexityTier.SIMPLE
assert any("simple" in s.lower() for s in signals)
def test_code_question(self, router):
"""Code-related questions should trend toward COMPLEX."""
tier, score, signals = router.classify(
"Write a Python function that implements binary search with error handling"
)
assert tier in [ComplexityTier.MEDIUM, ComplexityTier.COMPLEX]
assert any("code" in s.lower() for s in signals)
def test_reasoning_explicit(self, router):
"""Explicit reasoning markers should trigger REASONING tier."""
tier, score, signals = router.classify(
"Let's think step by step about how to solve this problem. "
"Analyze this carefully and show your reasoning."
)
assert tier == ComplexityTier.REASONING
assert any("reasoning" in s.lower() for s in signals)
def test_technical_content(self, router):
"""Technical content should trend toward COMPLEX."""
tier, score, signals = router.classify(
"Explain the architecture of a distributed microservice system "
"with Kubernetes orchestration and gRPC communication."
)
assert tier in [ComplexityTier.COMPLEX, ComplexityTier.REASONING]
def test_multi_step_patterns(self, router):
"""Multi-step patterns should increase complexity score."""
tier, score, signals = router.classify(
"First, analyze the requirements. "
"Then, design the database schema. "
"Finally, implement the API endpoints."
)
assert any("multi-step" in s.lower() for s in signals)
assert tier in [ComplexityTier.MEDIUM, ComplexityTier.COMPLEX]
def test_multiple_questions(self, router):
"""Multiple questions should increase complexity."""
tier, score, signals = router.classify(
"What is the difference between TCP and UDP? "
"When should I use each? "
"What are the performance implications? "
"How do they handle packet loss?"
)
assert any("question" in s.lower() for s in signals)
def test_long_prompt(self, router):
"""Long prompts should trend toward complex."""
long_text = "This is a detailed analysis request. " * 100
tier, score, signals = router.classify(long_text)
assert any("long" in s.lower() for s in signals)
@pytest.fixture
def mock_router_instance():
"""Create a mock LiteLLM Router instance."""
router = MagicMock()
return router
class TestModelSelection:
"""Test that correct models are returned for each tier."""
@pytest.fixture
def router(self):
"""Create a ComplexityRouter with custom tier mapping."""
config = {
"tiers": {
"SIMPLE": "gemini-2.0-flash",
"MEDIUM": "gpt-4o-mini",
"COMPLEX": "claude-sonnet-4",
"REASONING": "claude-opus-4",
}
}
return ComplexityRouter(
model_name="test_router",
litellm_router_instance=None,
complexity_router_config=config,
)
def test_simple_tier_model(self, router):
"""SIMPLE tier should return the configured simple model."""
model = router.get_model_for_tier(ComplexityTier.SIMPLE)
assert model == "gemini-2.0-flash"
def test_medium_tier_model(self, router):
"""MEDIUM tier should return the configured medium model."""
model = router.get_model_for_tier(ComplexityTier.MEDIUM)
assert model == "gpt-4o-mini"
def test_complex_tier_model(self, router):
"""COMPLEX tier should return the configured complex model."""
model = router.get_model_for_tier(ComplexityTier.COMPLEX)
assert model == "claude-sonnet-4"
def test_reasoning_tier_model(self, router):
"""REASONING tier should return the configured reasoning model."""
model = router.get_model_for_tier(ComplexityTier.REASONING)
assert model == "claude-opus-4"
@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,
},
}
class TestConfig:
"""Test configuration handling."""
def test_default_config(self):
"""Default config should have all required fields."""
config = DEFAULT_COMPLEXITY_CONFIG
assert config.tiers is not None
assert config.tier_boundaries is not None
assert config.dimension_weights is not None
def test_custom_tier_boundaries(self):
"""Custom tier boundaries should be respected."""
config = ComplexityRouterConfig(
tiers={"SIMPLE": "model-a", "MEDIUM": "model-b", "COMPLEX": "model-c", "REASONING": "model-d"},
tier_boundaries={
"simple_medium": 0.1,
"medium_complex": 0.3,
"complex_reasoning": 0.5,
}
)
assert config.tier_boundaries["simple_medium"] == 0.1
assert config.tier_boundaries["medium_complex"] == 0.3
assert config.tier_boundaries["complex_reasoning"] == 0.5
def test_custom_dimension_weights(self):
"""Custom dimension weights should be respected."""
config = ComplexityRouterConfig(
tiers={"SIMPLE": "model-a", "MEDIUM": "model-b", "COMPLEX": "model-c", "REASONING": "model-d"},
dimension_weights={
"tokenCount": 0.5,
"codePresence": 0.5,
}
)
assert config.dimension_weights["tokenCount"] == 0.5
assert config.dimension_weights["codePresence"] == 0.5
@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_with_signal(self):
"""DimensionScore should store name, score, and signal."""
ds = DimensionScore("test", 0.5, "test signal")
assert ds.name == "test"
assert ds.score == 0.5
assert ds.signal == "test signal"
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_without_signal(self):
"""DimensionScore should work without a signal."""
ds = DimensionScore("test", 0.5)
assert ds.name == "test"
assert ds.score == 0.5
assert ds.signal is None
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"
assert router.config == DEFAULT_COMPLEXITY_CONFIG
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 TestEdgeCases:
"""Test edge cases and error handling."""
def test_empty_prompt(self):
"""Empty prompts should be handled gracefully."""
router = ComplexityRouter(
model_name="test",
litellm_router_instance=None,
complexity_router_config={
"tiers": {
"SIMPLE": "model-a",
"MEDIUM": "model-b",
"COMPLEX": "model-c",
"REASONING": "model-d",
}
},
)
tier, score, signals = router.classify("")
assert tier == ComplexityTier.SIMPLE # Empty = simple
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_unicode_content(self):
"""Unicode content should be handled correctly."""
router = ComplexityRouter(
model_name="test",
litellm_router_instance=None,
complexity_router_config={
"tiers": {
"SIMPLE": "model-a",
"MEDIUM": "model-b",
"COMPLEX": "model-c",
"REASONING": "model-d",
}
},
)
tier, score, signals = router.classify("こんにちは、世界!🌍")
assert tier is not None
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_with_system_prompt(self):
"""System prompt should be considered in classification."""
router = ComplexityRouter(
model_name="test",
litellm_router_instance=None,
complexity_router_config={
"tiers": {
"SIMPLE": "model-a",
"MEDIUM": "model-b",
"COMPLEX": "model-c",
"REASONING": "model-d",
}
},
)
# System prompt with code context
tier, score, signals = router.classify(
"Hello",
system_prompt="You are a Python programming assistant."
)
# The code keywords in system prompt should influence scoring
assert tier is not None
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}"