model-router-cli 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- app/analytics/service.py +119 -0
- app/analyzer/analyzer.py +67 -0
- app/analyzer/heuristics.py +192 -0
- app/api/routes.py +589 -0
- app/budgets/manager.py +39 -0
- app/cli/main.py +287 -0
- app/config/settings.py +43 -0
- app/experiments/service.py +85 -0
- app/fallback/handler.py +105 -0
- app/models/schemas.py +127 -0
- app/observability/events.py +43 -0
- app/providers/base.py +46 -0
- app/providers/external_providers.py +321 -0
- app/providers/mock_provider.py +108 -0
- app/providers/ollama_provider.py +141 -0
- app/providers/registry.py +35 -0
- app/router/engine.py +150 -0
- app/router/rules_engine.py +73 -0
- app/router/scoring.py +154 -0
- app/static/assets/index-CQFztymk.js +63 -0
- app/static/assets/index-DWa3sE4Y.css +2 -0
- app/static/favicon.png +0 -0
- app/static/favicon.svg +1 -0
- app/static/icons.svg +24 -0
- app/static/index.html +17 -0
- app/static/logo.png +0 -0
- app/storage/database.py +366 -0
- app/storage/models.py +202 -0
- model_router_cli-1.0.0.dist-info/METADATA +343 -0
- model_router_cli-1.0.0.dist-info/RECORD +38 -0
- model_router_cli-1.0.0.dist-info/WHEEL +5 -0
- model_router_cli-1.0.0.dist-info/entry_points.txt +2 -0
- model_router_cli-1.0.0.dist-info/licenses/LICENSE +22 -0
- model_router_cli-1.0.0.dist-info/top_level.txt +2 -0
- tests/test_analyzer.py +41 -0
- tests/test_e2e.py +127 -0
- tests/test_providers.py +21 -0
- tests/test_router.py +78 -0
tests/test_e2e.py
ADDED
|
@@ -0,0 +1,127 @@
|
|
|
1
|
+
import pytest
|
|
2
|
+
from httpx import AsyncClient, ASGITransport
|
|
3
|
+
from main import app
|
|
4
|
+
from app.storage.database import init_db
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
@pytest.fixture(autouse=True)
|
|
8
|
+
async def setup_db():
|
|
9
|
+
await init_db()
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
@pytest.mark.asyncio
|
|
13
|
+
async def test_api_health():
|
|
14
|
+
transport = ASGITransport(app=app)
|
|
15
|
+
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
16
|
+
res = await client.get("/api/health")
|
|
17
|
+
assert res.status_code == 200
|
|
18
|
+
data = res.json()
|
|
19
|
+
assert data["status"] == "healthy"
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
@pytest.mark.asyncio
|
|
23
|
+
async def test_api_route_only():
|
|
24
|
+
transport = ASGITransport(app=app)
|
|
25
|
+
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
26
|
+
payload = {"prompt": "Write a Python script to benchmark models", "policy": "balanced"}
|
|
27
|
+
res = await client.post("/api/route", json=payload)
|
|
28
|
+
assert res.status_code == 200
|
|
29
|
+
data = res.json()
|
|
30
|
+
assert "decision" in data
|
|
31
|
+
assert "analysis" in data
|
|
32
|
+
assert data["analysis"]["task_type"] == "CODING"
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
@pytest.mark.asyncio
|
|
36
|
+
async def test_api_generate_end_to_end():
|
|
37
|
+
transport = ASGITransport(app=app)
|
|
38
|
+
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
39
|
+
payload = {"prompt": "Debug this memory leak in python subprocess", "policy": "balanced"}
|
|
40
|
+
res = await client.post("/api/generate", json=payload)
|
|
41
|
+
assert res.status_code == 200
|
|
42
|
+
data = res.json()
|
|
43
|
+
assert "response" in data
|
|
44
|
+
assert "metrics" in data
|
|
45
|
+
assert len(data["response"]["content"]) > 0
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
@pytest.mark.asyncio
|
|
49
|
+
async def test_api_models_list():
|
|
50
|
+
transport = ASGITransport(app=app)
|
|
51
|
+
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
52
|
+
res = await client.get("/api/models")
|
|
53
|
+
assert res.status_code == 200
|
|
54
|
+
models = res.json()
|
|
55
|
+
assert len(models) >= 3
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
@pytest.mark.asyncio
|
|
59
|
+
async def test_api_analytics():
|
|
60
|
+
transport = ASGITransport(app=app)
|
|
61
|
+
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
62
|
+
res = await client.get("/api/analytics")
|
|
63
|
+
assert res.status_code == 200
|
|
64
|
+
data = res.json()
|
|
65
|
+
assert "total_requests" in data
|
|
66
|
+
assert "savings" in data
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
@pytest.mark.asyncio
|
|
70
|
+
async def test_api_traffic_export_json_contains_audit_fields():
|
|
71
|
+
transport = ASGITransport(app=app)
|
|
72
|
+
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
73
|
+
generate_res = await client.post(
|
|
74
|
+
"/api/generate",
|
|
75
|
+
json={"prompt": "Write a Python function for a traffic export"},
|
|
76
|
+
)
|
|
77
|
+
assert generate_res.status_code == 200
|
|
78
|
+
|
|
79
|
+
export_res = await client.get("/api/traffic/export?format=json")
|
|
80
|
+
|
|
81
|
+
assert export_res.status_code == 200
|
|
82
|
+
assert export_res.headers["content-type"].startswith("application/json")
|
|
83
|
+
assert "attachment" in export_res.headers["content-disposition"]
|
|
84
|
+
records = export_res.json()
|
|
85
|
+
assert records
|
|
86
|
+
assert {
|
|
87
|
+
"timestamp",
|
|
88
|
+
"request_id",
|
|
89
|
+
"prompt_preview",
|
|
90
|
+
"task_type",
|
|
91
|
+
"complexity",
|
|
92
|
+
"selected_model",
|
|
93
|
+
"input_tokens",
|
|
94
|
+
"output_tokens",
|
|
95
|
+
"total_tokens",
|
|
96
|
+
"cost_saved",
|
|
97
|
+
"total_latency_ms",
|
|
98
|
+
}.issubset(records[0])
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
@pytest.mark.asyncio
|
|
102
|
+
async def test_api_traffic_export_csv_returns_downloadable_rows():
|
|
103
|
+
transport = ASGITransport(app=app)
|
|
104
|
+
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
105
|
+
generate_res = await client.post(
|
|
106
|
+
"/api/generate",
|
|
107
|
+
json={"prompt": "Summarize this traffic record"},
|
|
108
|
+
)
|
|
109
|
+
assert generate_res.status_code == 200
|
|
110
|
+
|
|
111
|
+
export_res = await client.get("/api/traffic/export?format=csv")
|
|
112
|
+
|
|
113
|
+
assert export_res.status_code == 200
|
|
114
|
+
assert export_res.headers["content-type"].startswith("text/csv")
|
|
115
|
+
assert "traffic-export.csv" in export_res.headers["content-disposition"]
|
|
116
|
+
lines = export_res.text.splitlines()
|
|
117
|
+
assert lines[0].startswith("timestamp,request_id,prompt_preview,")
|
|
118
|
+
assert len(lines) >= 2
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
@pytest.mark.asyncio
|
|
122
|
+
async def test_api_traffic_export_rejects_unknown_format():
|
|
123
|
+
transport = ASGITransport(app=app)
|
|
124
|
+
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
125
|
+
export_res = await client.get("/api/traffic/export?format=xml")
|
|
126
|
+
|
|
127
|
+
assert export_res.status_code == 422
|
tests/test_providers.py
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
import pytest
|
|
2
|
+
from app.providers.mock_provider import MockProvider
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
@pytest.mark.asyncio
|
|
6
|
+
async def test_mock_provider_generate():
|
|
7
|
+
provider = MockProvider()
|
|
8
|
+
resp = await provider.generate("Write a quick Python sort function", model_id="mock-fast")
|
|
9
|
+
assert resp.is_mock is True
|
|
10
|
+
assert "DEMO/MOCK" in resp.content
|
|
11
|
+
assert resp.input_tokens > 0
|
|
12
|
+
assert resp.output_tokens > 0
|
|
13
|
+
assert resp.provider_latency_ms > 0
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
@pytest.mark.asyncio
|
|
17
|
+
async def test_mock_provider_health():
|
|
18
|
+
provider = MockProvider()
|
|
19
|
+
health = await provider.check_health()
|
|
20
|
+
assert health["status"] == "CONNECTED"
|
|
21
|
+
assert "mock-fast" in health["models_available"]
|
tests/test_router.py
ADDED
|
@@ -0,0 +1,78 @@
|
|
|
1
|
+
import pytest
|
|
2
|
+
from app.models.schemas import ModelMetadata, ModelTier, RequestAnalysis, TaskType, PriorityLevel
|
|
3
|
+
from app.router.scoring import filter_candidate_models, compute_candidate_score
|
|
4
|
+
from app.router.engine import route_request
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
@pytest.fixture
|
|
8
|
+
def sample_models():
|
|
9
|
+
return [
|
|
10
|
+
ModelMetadata(
|
|
11
|
+
id="mock-fast",
|
|
12
|
+
name="Mock Fast",
|
|
13
|
+
provider="mock",
|
|
14
|
+
tier=ModelTier.FAST,
|
|
15
|
+
context_window=8192,
|
|
16
|
+
supports_coding=True,
|
|
17
|
+
supports_reasoning=False,
|
|
18
|
+
quality_score=0.75,
|
|
19
|
+
speed_score=0.98,
|
|
20
|
+
cost_per_input_token=0.0,
|
|
21
|
+
cost_per_output_token=0.0,
|
|
22
|
+
),
|
|
23
|
+
ModelMetadata(
|
|
24
|
+
id="mock-power",
|
|
25
|
+
name="Mock Power",
|
|
26
|
+
provider="mock",
|
|
27
|
+
tier=ModelTier.POWER,
|
|
28
|
+
context_window=128000,
|
|
29
|
+
supports_coding=True,
|
|
30
|
+
supports_reasoning=True,
|
|
31
|
+
quality_score=0.98,
|
|
32
|
+
speed_score=0.60,
|
|
33
|
+
cost_per_input_token=0.000005,
|
|
34
|
+
cost_per_output_token=0.000015,
|
|
35
|
+
),
|
|
36
|
+
]
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def test_candidate_filtering_context_overflow(sample_models):
|
|
40
|
+
large_analysis = RequestAnalysis(
|
|
41
|
+
task_type=TaskType.LONG_CONTEXT,
|
|
42
|
+
complexity=0.7,
|
|
43
|
+
context_size=15000, # exceeds mock-fast 8192
|
|
44
|
+
)
|
|
45
|
+
eligible, rejected = filter_candidate_models(sample_models, large_analysis)
|
|
46
|
+
assert len(eligible) == 1
|
|
47
|
+
assert eligible[0].id == "mock-power"
|
|
48
|
+
assert "mock-fast" in rejected
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def test_scoring_weights(sample_models):
|
|
52
|
+
analysis = RequestAnalysis(
|
|
53
|
+
task_type=TaskType.CODING,
|
|
54
|
+
complexity=0.3,
|
|
55
|
+
latency_priority=PriorityLevel.HIGH,
|
|
56
|
+
cost_sensitivity=PriorityLevel.HIGH,
|
|
57
|
+
)
|
|
58
|
+
weights = {"quality_weight": 0.1, "cost_weight": 0.5, "speed_weight": 0.4, "capability_weight": 0.0, "reliability_weight": 0.0}
|
|
59
|
+
score_fast = compute_candidate_score(sample_models[0], analysis, weights)
|
|
60
|
+
score_power = compute_candidate_score(sample_models[1], analysis, weights)
|
|
61
|
+
assert score_fast.overall_score > score_power.overall_score
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def test_routing_decision_explainability(sample_models):
|
|
65
|
+
analysis = RequestAnalysis(
|
|
66
|
+
task_type=TaskType.DEBUGGING,
|
|
67
|
+
complexity=0.85,
|
|
68
|
+
complexity_label=PriorityLevel.HIGH,
|
|
69
|
+
reasoning_required=True,
|
|
70
|
+
coding_required=True,
|
|
71
|
+
quality_requirement=PriorityLevel.HIGH,
|
|
72
|
+
)
|
|
73
|
+
weights = {"quality_weight": 0.6, "cost_weight": 0.1, "speed_weight": 0.1, "capability_weight": 0.1, "reliability_weight": 0.1}
|
|
74
|
+
decision = route_request(analysis, sample_models, weights, policy_name="highest_quality")
|
|
75
|
+
assert decision.selected_model == "mock-power"
|
|
76
|
+
assert decision.confidence > 0.70
|
|
77
|
+
assert len(decision.reasons) > 0
|
|
78
|
+
assert any("DEBUGGING" in r for r in decision.reasons)
|