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
app/cli/main.py
ADDED
|
@@ -0,0 +1,287 @@
|
|
|
1
|
+
import sys
|
|
2
|
+
import os
|
|
3
|
+
import asyncio
|
|
4
|
+
import httpx
|
|
5
|
+
import typer
|
|
6
|
+
from rich.console import Console
|
|
7
|
+
from rich.table import Table
|
|
8
|
+
from rich.panel import Panel
|
|
9
|
+
from rich.text import Text
|
|
10
|
+
|
|
11
|
+
# Set UTF-8 encoding support for Windows terminals
|
|
12
|
+
if sys.platform.startswith("win"):
|
|
13
|
+
try:
|
|
14
|
+
sys.stdout.reconfigure(encoding="utf-8")
|
|
15
|
+
sys.stderr.reconfigure(encoding="utf-8")
|
|
16
|
+
except Exception:
|
|
17
|
+
pass
|
|
18
|
+
|
|
19
|
+
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")))
|
|
20
|
+
|
|
21
|
+
from app.config.settings import get_settings
|
|
22
|
+
from app.analyzer.heuristics import analyze_request_heuristics
|
|
23
|
+
from app.router.engine import route_request
|
|
24
|
+
from app.storage.database import init_db, AsyncSessionLocal
|
|
25
|
+
from app.storage.models import ModelRecord, RoutingPolicyRecord
|
|
26
|
+
from app.api.routes import db_model_to_meta
|
|
27
|
+
from app.fallback.handler import execute_with_fallback
|
|
28
|
+
from sqlalchemy import select
|
|
29
|
+
|
|
30
|
+
app = typer.Typer(help="Model Router: Intelligent LLM Request Routing Platform")
|
|
31
|
+
console = Console(safe_box=True)
|
|
32
|
+
settings = get_settings()
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
@app.command()
|
|
36
|
+
def doctor():
|
|
37
|
+
"""Run environment, connectivity, and dependency health checks."""
|
|
38
|
+
console.print(Panel.fit("[bold cyan]Model Router Diagnostics (modelrouter doctor)[/bold cyan]"))
|
|
39
|
+
|
|
40
|
+
# 1. Python Check
|
|
41
|
+
py_ver = f"{sys.version_info.major}.{sys.version_info.minor}.{sys.version_info.micro}"
|
|
42
|
+
console.print(f" [green][OK][/green] Python: [bold]{py_ver}[/bold]")
|
|
43
|
+
|
|
44
|
+
# 2. Database Check
|
|
45
|
+
try:
|
|
46
|
+
asyncio.run(init_db())
|
|
47
|
+
console.print(" [green][OK][/green] SQLite Database: [bold]Initialized & Ready[/bold]")
|
|
48
|
+
except Exception as exc:
|
|
49
|
+
console.print(f" [red][FAIL][/red] Database Error: {exc}")
|
|
50
|
+
|
|
51
|
+
# 3. Ollama Connectivity Check
|
|
52
|
+
ollama_ok = False
|
|
53
|
+
try:
|
|
54
|
+
r = httpx.get(f"{settings.OLLAMA_BASE_URL}/api/tags", timeout=1.5)
|
|
55
|
+
if r.status_code == 200:
|
|
56
|
+
models = [m.get("name") for m in r.json().get("models", [])]
|
|
57
|
+
console.print(f" [green][OK][/green] Ollama (Local): [bold]Connected[/bold] ({len(models)} local models found)")
|
|
58
|
+
ollama_ok = True
|
|
59
|
+
except Exception:
|
|
60
|
+
pass
|
|
61
|
+
if not ollama_ok:
|
|
62
|
+
console.print(f" [yellow][!][/yellow] Ollama (Local): [dim]Not running on {settings.OLLAMA_BASE_URL} (Optional, local fallback available)[/dim]")
|
|
63
|
+
|
|
64
|
+
# 4. Mock Engine
|
|
65
|
+
console.print(" [green][OK][/green] Mock Engine: [bold]Ready (Zero API keys required)[/bold]")
|
|
66
|
+
|
|
67
|
+
# 5. External Providers
|
|
68
|
+
if settings.OPENAI_API_KEY:
|
|
69
|
+
console.print(" [green][OK][/green] OpenAI Provider: [bold]API Key Configured in .env[/bold]")
|
|
70
|
+
else:
|
|
71
|
+
console.print(" [yellow][!][/yellow] OpenAI Provider: [dim]Not configured (Optional)[/dim]")
|
|
72
|
+
|
|
73
|
+
if settings.ANTHROPIC_API_KEY:
|
|
74
|
+
console.print(" [green][OK][/green] Anthropic Provider: [bold]API Key Configured in .env[/bold]")
|
|
75
|
+
else:
|
|
76
|
+
console.print(" [yellow][!][/yellow] Anthropic Provider: [dim]Not configured (Optional)[/dim]")
|
|
77
|
+
|
|
78
|
+
if settings.GEMINI_API_KEY:
|
|
79
|
+
console.print(" [green][OK][/green] Gemini Provider: [bold]API Key Configured in .env[/bold]")
|
|
80
|
+
else:
|
|
81
|
+
console.print(" [yellow][!][/yellow] Gemini Provider: [dim]Not configured (Optional)[/dim]")
|
|
82
|
+
|
|
83
|
+
console.print("\n[bold green]Ready for local-first intelligent routing![/bold green]\n")
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
@app.command()
|
|
87
|
+
def models():
|
|
88
|
+
"""List all registered candidate models and their capability scores."""
|
|
89
|
+
async def _list():
|
|
90
|
+
await init_db()
|
|
91
|
+
async with AsyncSessionLocal() as db:
|
|
92
|
+
res = await db.execute(select(ModelRecord).order_by(ModelRecord.tier, ModelRecord.name))
|
|
93
|
+
records = res.scalars().all()
|
|
94
|
+
|
|
95
|
+
table = Table(title="Model Router Registry", border_style="bright_blue")
|
|
96
|
+
table.add_column("Model ID", style="cyan", no_wrap=True)
|
|
97
|
+
table.add_column("Tier", style="magenta")
|
|
98
|
+
table.add_column("Provider", style="green")
|
|
99
|
+
table.add_column("Context", style="yellow")
|
|
100
|
+
table.add_column("Quality", style="blue")
|
|
101
|
+
table.add_column("Speed", style="cyan")
|
|
102
|
+
table.add_column("Cost / 1K In/Out", style="white")
|
|
103
|
+
|
|
104
|
+
for m in records:
|
|
105
|
+
cost_str = f"${m.cost_per_input_token*1000:.4f} / ${m.cost_per_output_token*1000:.4f}" if m.cost_per_input_token > 0 else "Free ($0)"
|
|
106
|
+
table.add_row(
|
|
107
|
+
m.id,
|
|
108
|
+
m.tier,
|
|
109
|
+
m.provider,
|
|
110
|
+
f"{m.context_window:,}",
|
|
111
|
+
f"{m.quality_score*100:.0f}%",
|
|
112
|
+
f"{m.speed_score*100:.0f}%",
|
|
113
|
+
cost_str,
|
|
114
|
+
)
|
|
115
|
+
console.print(table)
|
|
116
|
+
|
|
117
|
+
asyncio.run(_list())
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
@app.command()
|
|
121
|
+
def route(
|
|
122
|
+
prompt: str = typer.Argument(..., help="Prompt text to route"),
|
|
123
|
+
policy: str = typer.Option("balanced", "--policy", "-p", help="Routing policy to use"),
|
|
124
|
+
):
|
|
125
|
+
"""Analyze prompt and explain why the optimal model was selected (Dry Run)."""
|
|
126
|
+
async def _route():
|
|
127
|
+
await init_db()
|
|
128
|
+
analysis = analyze_request_heuristics(prompt)
|
|
129
|
+
|
|
130
|
+
async with AsyncSessionLocal() as db:
|
|
131
|
+
res_m = await db.execute(select(ModelRecord).filter_by(is_active=True))
|
|
132
|
+
models = [db_model_to_meta(m) for m in res_m.scalars().all()]
|
|
133
|
+
|
|
134
|
+
res_p = await db.execute(select(RoutingPolicyRecord).filter_by(id=policy))
|
|
135
|
+
pol = res_p.scalar_one_or_none()
|
|
136
|
+
weights = {
|
|
137
|
+
"quality_weight": pol.quality_weight if pol else 0.35,
|
|
138
|
+
"cost_weight": pol.cost_weight if pol else 0.25,
|
|
139
|
+
"speed_weight": pol.speed_weight if pol else 0.20,
|
|
140
|
+
"capability_weight": pol.capability_weight if pol else 0.15,
|
|
141
|
+
"reliability_weight": pol.reliability_weight if pol else 0.05,
|
|
142
|
+
}
|
|
143
|
+
|
|
144
|
+
decision = route_request(
|
|
145
|
+
analysis=analysis,
|
|
146
|
+
available_models=models,
|
|
147
|
+
policy_weights=weights,
|
|
148
|
+
policy_name=policy,
|
|
149
|
+
)
|
|
150
|
+
|
|
151
|
+
console.print(Panel.fit(
|
|
152
|
+
f"[bold white]Prompt:[/bold white] \"{prompt}\"\n"
|
|
153
|
+
f"[bold cyan]Task Type:[/bold cyan] {analysis.task_type.value} | [bold cyan]Complexity:[/bold cyan] {analysis.complexity_label.value} ({analysis.complexity:.2f})\n"
|
|
154
|
+
f"[bold green]Selected Model:[/bold green] [bold yellow]{decision.selected_model_name}[/bold yellow] ({decision.selected_model})\n"
|
|
155
|
+
f"[bold magenta]Confidence:[/bold magenta] {decision.confidence * 100:.0f}%\n"
|
|
156
|
+
f"[bold blue]Estimated Latency:[/bold blue] {decision.estimated_latency_ms:.0f}ms | [bold blue]Estimated Cost:[/bold blue] ${decision.estimated_cost_usd:.6f}",
|
|
157
|
+
title="ROUTING DECISION",
|
|
158
|
+
border_style="bright_blue",
|
|
159
|
+
))
|
|
160
|
+
|
|
161
|
+
console.print("[bold cyan]Why this model?[/bold cyan]")
|
|
162
|
+
for r in decision.reasons:
|
|
163
|
+
console.print(f" [green]+[/green] {r}")
|
|
164
|
+
|
|
165
|
+
if decision.rejected_candidates:
|
|
166
|
+
console.print("\n[bold red]Rejected Candidates:[/bold red]")
|
|
167
|
+
for cid, reason in decision.rejected_candidates.items():
|
|
168
|
+
console.print(f" [dim]- {cid}: {reason}[/dim]")
|
|
169
|
+
|
|
170
|
+
asyncio.run(_route())
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
@app.command()
|
|
174
|
+
def run(
|
|
175
|
+
prompt: str = typer.Argument(..., help="Prompt to route and execute"),
|
|
176
|
+
policy: str = typer.Option("balanced", "--policy", "-p", help="Routing policy to use"),
|
|
177
|
+
):
|
|
178
|
+
"""Route request and execute inference through the chosen provider."""
|
|
179
|
+
async def _run():
|
|
180
|
+
await init_db()
|
|
181
|
+
analysis = analyze_request_heuristics(prompt)
|
|
182
|
+
|
|
183
|
+
async with AsyncSessionLocal() as db:
|
|
184
|
+
res_m = await db.execute(select(ModelRecord).filter_by(is_active=True))
|
|
185
|
+
models = [db_model_to_meta(m) for m in res_m.scalars().all()]
|
|
186
|
+
|
|
187
|
+
decision = route_request(
|
|
188
|
+
analysis=analysis,
|
|
189
|
+
available_models=models,
|
|
190
|
+
policy_weights={"quality_weight": 0.35, "cost_weight": 0.25, "speed_weight": 0.20, "capability_weight": 0.15, "reliability_weight": 0.05},
|
|
191
|
+
policy_name=policy,
|
|
192
|
+
)
|
|
193
|
+
|
|
194
|
+
console.print(f"[bold cyan]Routing to:[/bold cyan] {decision.selected_model} via provider: {decision.provider}...")
|
|
195
|
+
|
|
196
|
+
resp, fallback_used, orig_m, fb_reason = await execute_with_fallback(
|
|
197
|
+
prompt=prompt,
|
|
198
|
+
selected_model_id=decision.selected_model,
|
|
199
|
+
selected_provider_id=decision.provider,
|
|
200
|
+
all_models=models,
|
|
201
|
+
)
|
|
202
|
+
|
|
203
|
+
console.print(Panel(
|
|
204
|
+
resp.content,
|
|
205
|
+
title=f"RESPONSE from {resp.model} ({'FALLBACK: ' + orig_m if fallback_used else 'PRIMARY'})",
|
|
206
|
+
subtitle=f"Latency: {resp.provider_latency_ms:.0f}ms | Tokens: {resp.total_tokens}",
|
|
207
|
+
border_style="green" if not fallback_used else "yellow",
|
|
208
|
+
))
|
|
209
|
+
|
|
210
|
+
asyncio.run(_run())
|
|
211
|
+
|
|
212
|
+
|
|
213
|
+
@app.command()
|
|
214
|
+
def traffic(limit: int = 10):
|
|
215
|
+
"""View recent routed traffic log."""
|
|
216
|
+
async def _traffic():
|
|
217
|
+
await init_db()
|
|
218
|
+
from app.storage.models import RequestRecord
|
|
219
|
+
from sqlalchemy import desc
|
|
220
|
+
async with AsyncSessionLocal() as db:
|
|
221
|
+
res = await db.execute(select(RequestRecord).order_by(desc(RequestRecord.timestamp)).limit(limit))
|
|
222
|
+
records = res.scalars().all()
|
|
223
|
+
|
|
224
|
+
table = Table(title="Recent Routed Traffic Feed", border_style="bright_blue")
|
|
225
|
+
table.add_column("Time", style="dim")
|
|
226
|
+
table.add_column("Task", style="cyan")
|
|
227
|
+
table.add_column("Selected Model", style="green")
|
|
228
|
+
table.add_column("Latency", style="yellow")
|
|
229
|
+
table.add_column("Cost", style="magenta")
|
|
230
|
+
table.add_column("Prompt Sample", style="white")
|
|
231
|
+
|
|
232
|
+
for r in records:
|
|
233
|
+
t_str = r.timestamp.strftime("%H:%M:%S") if r.timestamp else "N/A"
|
|
234
|
+
table.add_row(
|
|
235
|
+
t_str,
|
|
236
|
+
r.task_type,
|
|
237
|
+
r.selected_model,
|
|
238
|
+
f"{r.total_latency_ms:.0f}ms",
|
|
239
|
+
f"${r.estimated_cost:.5f}",
|
|
240
|
+
r.prompt[:40] + ("..." if len(r.prompt) > 40 else ""),
|
|
241
|
+
)
|
|
242
|
+
console.print(table)
|
|
243
|
+
|
|
244
|
+
asyncio.run(_traffic())
|
|
245
|
+
|
|
246
|
+
|
|
247
|
+
@app.command()
|
|
248
|
+
def analytics():
|
|
249
|
+
"""Display system-wide routing performance and cost savings analytics."""
|
|
250
|
+
async def _analytics():
|
|
251
|
+
await init_db()
|
|
252
|
+
from app.analytics.service import get_system_analytics
|
|
253
|
+
async with AsyncSessionLocal() as db:
|
|
254
|
+
data = await get_system_analytics(db)
|
|
255
|
+
console.print(Panel.fit(
|
|
256
|
+
f"[bold cyan]Total Routed Requests:[/bold cyan] {data['total_requests']}\n"
|
|
257
|
+
f"[bold green]Average Latency:[/bold green] {data['avg_latency_ms']}ms (Routing overhead: {data['avg_routing_latency_ms']}ms)\n"
|
|
258
|
+
f"[bold yellow]Total Cost:[/bold yellow] ${data['total_cost_usd']:.6f}\n"
|
|
259
|
+
f"[bold magenta]Cost Saved vs Baseline:[/bold magenta] ${data['savings']['cost_saved_usd']:.6f} ({data['savings']['savings_percentage']}%)\n"
|
|
260
|
+
f"[bold blue]Fallback Rate:[/bold blue] {data['fallback_rate_percent']}%\n",
|
|
261
|
+
title="MODEL ROUTER ANALYTICS OVERVIEW",
|
|
262
|
+
border_style="bright_blue",
|
|
263
|
+
))
|
|
264
|
+
|
|
265
|
+
asyncio.run(_analytics())
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
@app.command()
|
|
269
|
+
def ui(
|
|
270
|
+
host: str = typer.Option("127.0.0.1", "--host", "-h", help="Host address to bind the web server"),
|
|
271
|
+
port: int = typer.Option(8000, "--port", "-p", help="Port to run the Model Router Control Room on"),
|
|
272
|
+
):
|
|
273
|
+
"""Launch the Model Router API Gateway & AI Traffic Control Room UI."""
|
|
274
|
+
import uvicorn
|
|
275
|
+
console.print(Panel.fit(
|
|
276
|
+
f"[bold cyan]Starting Model Router AI Traffic Control Room[/bold cyan]\n"
|
|
277
|
+
f"Web Dashboard: [bold green]http://{host}:{port}[/bold green]\n"
|
|
278
|
+
f"API Documentation: [bold green]http://{host}:{port}/docs[/bold green]",
|
|
279
|
+
title="MODEL ROUTER UI LAUNCHER",
|
|
280
|
+
border_style="cyan"
|
|
281
|
+
))
|
|
282
|
+
uvicorn.run("main:app", host=host, port=port, reload=False, app_dir=os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")))
|
|
283
|
+
|
|
284
|
+
|
|
285
|
+
if __name__ == "__main__":
|
|
286
|
+
app()
|
|
287
|
+
|
app/config/settings.py
ADDED
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
from functools import lru_cache
|
|
2
|
+
from typing import Optional
|
|
3
|
+
from pydantic_settings import BaseSettings, SettingsConfigDict
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class Settings(BaseSettings):
|
|
7
|
+
APP_ENV: str = "development"
|
|
8
|
+
DEBUG: bool = True
|
|
9
|
+
PORT: int = 8000
|
|
10
|
+
HOST: str = "0.0.0.0"
|
|
11
|
+
|
|
12
|
+
DATABASE_URL: str = "sqlite+aiosqlite:///./model_router.db"
|
|
13
|
+
ROUTER_ANALYZER: str = "rules" # "rules" or "llm"
|
|
14
|
+
DEFAULT_ROUTING_POLICY: str = "balanced"
|
|
15
|
+
BASELINE_MODEL_ID: str = "mock-power"
|
|
16
|
+
|
|
17
|
+
DEFAULT_PROVIDER: str = "mock"
|
|
18
|
+
OLLAMA_BASE_URL: str = "http://localhost:11434"
|
|
19
|
+
|
|
20
|
+
# External Provider Keys (optional)
|
|
21
|
+
OPENAI_API_KEY: Optional[str] = None
|
|
22
|
+
ANTHROPIC_API_KEY: Optional[str] = None
|
|
23
|
+
GEMINI_API_KEY: Optional[str] = None
|
|
24
|
+
|
|
25
|
+
# Budgets
|
|
26
|
+
DAILY_BUDGET: float = 10.00
|
|
27
|
+
MONTHLY_BUDGET: float = 100.00
|
|
28
|
+
PER_REQUEST_BUDGET: float = 1.00
|
|
29
|
+
|
|
30
|
+
# Execution / Reliability
|
|
31
|
+
MAX_RETRIES: int = 2
|
|
32
|
+
PROVIDER_TIMEOUT_SECONDS: float = 30.0
|
|
33
|
+
|
|
34
|
+
model_config = SettingsConfigDict(
|
|
35
|
+
env_file=".env",
|
|
36
|
+
env_file_encoding="utf-8",
|
|
37
|
+
extra="ignore",
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
@lru_cache()
|
|
42
|
+
def get_settings() -> Settings:
|
|
43
|
+
return Settings()
|
|
@@ -0,0 +1,85 @@
|
|
|
1
|
+
import uuid
|
|
2
|
+
import datetime
|
|
3
|
+
from typing import Dict, Any, List
|
|
4
|
+
from sqlalchemy import select, func
|
|
5
|
+
from sqlalchemy.ext.asyncio import AsyncSession
|
|
6
|
+
from app.storage.models import ExperimentRecord, ExperimentRunRecord
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
async def log_experiment_run(
|
|
10
|
+
session: AsyncSession,
|
|
11
|
+
experiment_id: str,
|
|
12
|
+
request_id: str,
|
|
13
|
+
policy: str,
|
|
14
|
+
model: str,
|
|
15
|
+
cost: float,
|
|
16
|
+
latency_ms: float,
|
|
17
|
+
):
|
|
18
|
+
run = ExperimentRunRecord(
|
|
19
|
+
id=str(uuid.uuid4()),
|
|
20
|
+
experiment_id=experiment_id,
|
|
21
|
+
request_id=request_id,
|
|
22
|
+
assigned_policy=policy,
|
|
23
|
+
selected_model=model,
|
|
24
|
+
cost=cost,
|
|
25
|
+
latency_ms=latency_ms,
|
|
26
|
+
)
|
|
27
|
+
session.add(run)
|
|
28
|
+
|
|
29
|
+
# Increment sample count
|
|
30
|
+
res = await session.execute(select(ExperimentRecord).filter_by(id=experiment_id))
|
|
31
|
+
exp = res.scalar_one_or_none()
|
|
32
|
+
if exp:
|
|
33
|
+
exp.sample_count += 1
|
|
34
|
+
await session.commit()
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
async def get_experiment_summary(
|
|
38
|
+
session: AsyncSession,
|
|
39
|
+
experiment_id: str,
|
|
40
|
+
) -> Dict[str, Any]:
|
|
41
|
+
res = await session.execute(select(ExperimentRecord).filter_by(id=experiment_id))
|
|
42
|
+
exp = res.scalar_one_or_none()
|
|
43
|
+
if not exp:
|
|
44
|
+
return {}
|
|
45
|
+
|
|
46
|
+
# Calculate metrics for Policy A vs Policy B
|
|
47
|
+
runs_a = await session.execute(
|
|
48
|
+
select(
|
|
49
|
+
func.count(ExperimentRunRecord.id),
|
|
50
|
+
func.avg(ExperimentRunRecord.cost),
|
|
51
|
+
func.avg(ExperimentRunRecord.latency_ms),
|
|
52
|
+
).filter_by(experiment_id=experiment_id, assigned_policy=exp.policy_a)
|
|
53
|
+
)
|
|
54
|
+
count_a, cost_a, lat_a = runs_a.one()
|
|
55
|
+
|
|
56
|
+
runs_b = await session.execute(
|
|
57
|
+
select(
|
|
58
|
+
func.count(ExperimentRunRecord.id),
|
|
59
|
+
func.avg(ExperimentRunRecord.cost),
|
|
60
|
+
func.avg(ExperimentRunRecord.latency_ms),
|
|
61
|
+
).filter_by(experiment_id=experiment_id, assigned_policy=exp.policy_b)
|
|
62
|
+
)
|
|
63
|
+
count_b, cost_b, lat_b = runs_b.one()
|
|
64
|
+
|
|
65
|
+
return {
|
|
66
|
+
"id": exp.id,
|
|
67
|
+
"name": exp.name,
|
|
68
|
+
"description": exp.description,
|
|
69
|
+
"policy_a": exp.policy_a,
|
|
70
|
+
"policy_b": exp.policy_b,
|
|
71
|
+
"status": exp.status,
|
|
72
|
+
"total_samples": exp.sample_count,
|
|
73
|
+
"results": {
|
|
74
|
+
exp.policy_a: {
|
|
75
|
+
"sample_count": count_a or 0,
|
|
76
|
+
"avg_cost_usd": round(cost_a or 0.0, 6),
|
|
77
|
+
"avg_latency_ms": round(lat_a or 0.0, 1),
|
|
78
|
+
},
|
|
79
|
+
exp.policy_b: {
|
|
80
|
+
"sample_count": count_b or 0,
|
|
81
|
+
"avg_cost_usd": round(cost_b or 0.0, 6),
|
|
82
|
+
"avg_latency_ms": round(lat_b or 0.0, 1),
|
|
83
|
+
},
|
|
84
|
+
},
|
|
85
|
+
}
|
app/fallback/handler.py
ADDED
|
@@ -0,0 +1,105 @@
|
|
|
1
|
+
import time
|
|
2
|
+
import asyncio
|
|
3
|
+
from typing import List, Optional, Tuple, Dict, Any
|
|
4
|
+
from app.models.schemas import ProviderResponse, ModelMetadata
|
|
5
|
+
from app.providers.registry import provider_registry
|
|
6
|
+
from app.config.settings import get_settings
|
|
7
|
+
|
|
8
|
+
settings = get_settings()
|
|
9
|
+
|
|
10
|
+
RETRYABLE_ERRORS = [
|
|
11
|
+
"timeout",
|
|
12
|
+
"connection error",
|
|
13
|
+
"rate limit",
|
|
14
|
+
"429",
|
|
15
|
+
"503",
|
|
16
|
+
"service unavailable",
|
|
17
|
+
"econnreset",
|
|
18
|
+
]
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def is_retryable(error_str: Optional[str]) -> bool:
|
|
22
|
+
if not error_str:
|
|
23
|
+
return False
|
|
24
|
+
low = error_str.lower()
|
|
25
|
+
return any(keyword in low for keyword in RETRYABLE_ERRORS)
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
async def execute_with_fallback(
|
|
29
|
+
prompt: str,
|
|
30
|
+
selected_model_id: str,
|
|
31
|
+
selected_provider_id: str,
|
|
32
|
+
all_models: List[ModelMetadata],
|
|
33
|
+
system_prompt: Optional[str] = None,
|
|
34
|
+
temperature: float = 0.7,
|
|
35
|
+
max_retries: int = 1,
|
|
36
|
+
) -> Tuple[ProviderResponse, bool, Optional[str], Optional[str]]:
|
|
37
|
+
"""
|
|
38
|
+
Executes model generation with configurable retries and fallback hierarchy.
|
|
39
|
+
Returns: (ProviderResponse, fallback_used, original_model, fallback_reason)
|
|
40
|
+
"""
|
|
41
|
+
current_model_id = selected_model_id
|
|
42
|
+
current_provider_id = selected_provider_id
|
|
43
|
+
fallback_used = False
|
|
44
|
+
original_model = selected_model_id
|
|
45
|
+
fallback_reason = None
|
|
46
|
+
|
|
47
|
+
# Step 1: Attempt primary model with retries
|
|
48
|
+
provider = provider_registry.get_provider(current_provider_id)
|
|
49
|
+
if not provider:
|
|
50
|
+
provider = provider_registry.get_provider("mock")
|
|
51
|
+
current_provider_id = "mock"
|
|
52
|
+
current_model_id = "mock-balanced"
|
|
53
|
+
fallback_used = True
|
|
54
|
+
fallback_reason = f"Provider '{selected_provider_id}' is not loaded. Fell back to Mock."
|
|
55
|
+
|
|
56
|
+
for attempt in range(max_retries + 1):
|
|
57
|
+
resp = await provider.generate(
|
|
58
|
+
prompt=prompt,
|
|
59
|
+
model_id=current_model_id,
|
|
60
|
+
system_prompt=system_prompt,
|
|
61
|
+
temperature=temperature,
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
if not resp.error and resp.finish_reason != "error":
|
|
65
|
+
return resp, fallback_used, original_model if fallback_used else None, fallback_reason
|
|
66
|
+
|
|
67
|
+
# If non-retryable error or last attempt, prepare for fallback
|
|
68
|
+
if not is_retryable(resp.error) or attempt == max_retries:
|
|
69
|
+
fallback_reason = resp.error or "Unknown primary model execution error."
|
|
70
|
+
break
|
|
71
|
+
|
|
72
|
+
await asyncio.sleep(0.5 * (attempt + 1))
|
|
73
|
+
|
|
74
|
+
# Step 2: Determine Fallback Candidate
|
|
75
|
+
fallback_used = True
|
|
76
|
+
# Priority: Mock model of same or balanced tier, or active local Ollama
|
|
77
|
+
candidate_fallbacks = [
|
|
78
|
+
m for m in all_models
|
|
79
|
+
if m.id != selected_model_id and m.is_active and (m.provider == "mock" or m.type == "LOCAL")
|
|
80
|
+
]
|
|
81
|
+
if not candidate_fallbacks:
|
|
82
|
+
candidate_fallbacks = [m for m in all_models if m.id != selected_model_id and m.is_active]
|
|
83
|
+
|
|
84
|
+
fallback_target = candidate_fallbacks[0] if candidate_fallbacks else None
|
|
85
|
+
|
|
86
|
+
if fallback_target:
|
|
87
|
+
fb_provider = provider_registry.get_provider(fallback_target.provider) or provider_registry.get_provider("mock")
|
|
88
|
+
fb_resp = await fb_provider.generate(
|
|
89
|
+
prompt=prompt,
|
|
90
|
+
model_id=fallback_target.id,
|
|
91
|
+
system_prompt=system_prompt,
|
|
92
|
+
temperature=temperature,
|
|
93
|
+
)
|
|
94
|
+
if not fb_resp.error:
|
|
95
|
+
return fb_resp, True, original_model, f"Primary model failed: {fallback_reason} -> Fallback to {fallback_target.id}"
|
|
96
|
+
|
|
97
|
+
# Step 3: Emergency Hard Mock fallback
|
|
98
|
+
emergency_mock = provider_registry.get_provider("mock")
|
|
99
|
+
emergency_resp = await emergency_mock.generate(
|
|
100
|
+
prompt=prompt,
|
|
101
|
+
model_id="mock-balanced",
|
|
102
|
+
system_prompt=system_prompt,
|
|
103
|
+
temperature=temperature,
|
|
104
|
+
)
|
|
105
|
+
return emergency_resp, True, original_model, f"Emergency fallback triggered due to: {fallback_reason}"
|
app/models/schemas.py
ADDED
|
@@ -0,0 +1,127 @@
|
|
|
1
|
+
from enum import Enum
|
|
2
|
+
from typing import List, Optional, Dict, Any
|
|
3
|
+
from pydantic import BaseModel, Field
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class TaskType(str, Enum):
|
|
7
|
+
GENERAL_QA = "GENERAL_QA"
|
|
8
|
+
CODING = "CODING"
|
|
9
|
+
DEBUGGING = "DEBUGGING"
|
|
10
|
+
REASONING = "REASONING"
|
|
11
|
+
SUMMARIZATION = "SUMMARIZATION"
|
|
12
|
+
EXTRACTION = "EXTRACTION"
|
|
13
|
+
WRITING = "WRITING"
|
|
14
|
+
TRANSLATION = "TRANSLATION"
|
|
15
|
+
ANALYSIS = "ANALYSIS"
|
|
16
|
+
MATH = "MATH"
|
|
17
|
+
LONG_CONTEXT = "LONG_CONTEXT"
|
|
18
|
+
CREATIVE = "CREATIVE"
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class PriorityLevel(str, Enum):
|
|
22
|
+
LOW = "LOW"
|
|
23
|
+
MEDIUM = "MEDIUM"
|
|
24
|
+
HIGH = "HIGH"
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
class ModelTier(str, Enum):
|
|
28
|
+
FAST = "FAST"
|
|
29
|
+
BALANCED = "BALANCED"
|
|
30
|
+
POWER = "POWER"
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class ProviderType(str, Enum):
|
|
34
|
+
MOCK = "mock"
|
|
35
|
+
OLLAMA = "ollama"
|
|
36
|
+
OPENAI = "openai"
|
|
37
|
+
ANTHROPIC = "anthropic"
|
|
38
|
+
GEMINI = "gemini"
|
|
39
|
+
CUSTOM = "custom"
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class RequestAnalysis(BaseModel):
|
|
43
|
+
task_type: TaskType = Field(default=TaskType.GENERAL_QA)
|
|
44
|
+
complexity: float = Field(default=0.5, ge=0.0, le=1.0)
|
|
45
|
+
complexity_label: PriorityLevel = Field(default=PriorityLevel.MEDIUM)
|
|
46
|
+
reasoning_required: bool = Field(default=False)
|
|
47
|
+
coding_required: bool = Field(default=False)
|
|
48
|
+
vision_required: bool = Field(default=False)
|
|
49
|
+
tools_required: bool = Field(default=False)
|
|
50
|
+
context_size: int = Field(default=0) # in estimated tokens / characters
|
|
51
|
+
latency_priority: PriorityLevel = Field(default=PriorityLevel.MEDIUM)
|
|
52
|
+
cost_sensitivity: PriorityLevel = Field(default=PriorityLevel.MEDIUM)
|
|
53
|
+
quality_requirement: PriorityLevel = Field(default=PriorityLevel.MEDIUM)
|
|
54
|
+
keywords_detected: List[str] = Field(default_factory=list)
|
|
55
|
+
analyzer_used: str = Field(default="rules")
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
class ModelMetadata(BaseModel):
|
|
59
|
+
id: str
|
|
60
|
+
name: str
|
|
61
|
+
provider: str
|
|
62
|
+
type: str = "LOCAL" # LOCAL, CLOUD, MOCK
|
|
63
|
+
tier: ModelTier = ModelTier.BALANCED
|
|
64
|
+
context_window: int = 8192
|
|
65
|
+
|
|
66
|
+
supports_coding: bool = False
|
|
67
|
+
supports_reasoning: bool = False
|
|
68
|
+
supports_vision: bool = False
|
|
69
|
+
supports_tools: bool = False
|
|
70
|
+
|
|
71
|
+
quality_score: float = 0.75 # 0.0 - 1.0
|
|
72
|
+
speed_score: float = 0.75 # 0.0 - 1.0
|
|
73
|
+
reliability_score: float = 0.95
|
|
74
|
+
|
|
75
|
+
cost_per_input_token: float = 0.0
|
|
76
|
+
cost_per_output_token: float = 0.0
|
|
77
|
+
availability: str = "AVAILABLE"
|
|
78
|
+
is_active: bool = True
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
class CandidateScore(BaseModel):
|
|
82
|
+
model_id: str
|
|
83
|
+
model_name: str
|
|
84
|
+
provider: str
|
|
85
|
+
tier: str
|
|
86
|
+
overall_score: float # 0 - 100
|
|
87
|
+
quality_component: float
|
|
88
|
+
cost_component: float
|
|
89
|
+
speed_component: float
|
|
90
|
+
capability_component: float
|
|
91
|
+
reliability_component: float
|
|
92
|
+
estimated_latency_ms: float
|
|
93
|
+
estimated_cost_usd: float
|
|
94
|
+
eligible: bool = True
|
|
95
|
+
rejection_reason: Optional[str] = None
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
class RoutingDecision(BaseModel):
|
|
99
|
+
decision_id: str
|
|
100
|
+
request_id: str
|
|
101
|
+
selected_model: str
|
|
102
|
+
selected_model_name: str
|
|
103
|
+
provider: str
|
|
104
|
+
tier: str
|
|
105
|
+
confidence: float
|
|
106
|
+
policy_used: str
|
|
107
|
+
reasons: List[str]
|
|
108
|
+
candidate_scores: List[CandidateScore]
|
|
109
|
+
rejected_candidates: Dict[str, str] = Field(default_factory=dict)
|
|
110
|
+
estimated_cost_usd: float
|
|
111
|
+
estimated_latency_ms: float
|
|
112
|
+
rule_applied: Optional[str] = None
|
|
113
|
+
timestamp: str
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
class ProviderResponse(BaseModel):
|
|
117
|
+
content: str
|
|
118
|
+
finish_reason: str = "stop"
|
|
119
|
+
input_tokens: int = 0
|
|
120
|
+
output_tokens: int = 0
|
|
121
|
+
total_tokens: int = 0
|
|
122
|
+
provider_latency_ms: float = 0.0
|
|
123
|
+
time_to_first_token_ms: Optional[float] = None
|
|
124
|
+
is_mock: bool = False
|
|
125
|
+
model: str = ""
|
|
126
|
+
provider: str = ""
|
|
127
|
+
error: Optional[str] = None
|
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import logging
|
|
3
|
+
from typing import Dict, Any, Optional
|
|
4
|
+
import datetime
|
|
5
|
+
|
|
6
|
+
logger = logging.getLogger("model_router.observability")
|
|
7
|
+
logger.setLevel(logging.INFO)
|
|
8
|
+
if not logger.handlers:
|
|
9
|
+
handler = logging.StreamHandler()
|
|
10
|
+
formatter = logging.Formatter('{"time":"%(asctime)s", "level":"%(levelname)s", "event":%(message)s}')
|
|
11
|
+
handler.setFormatter(formatter)
|
|
12
|
+
logger.addHandler(handler)
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def log_router_event(
|
|
16
|
+
event_name: str,
|
|
17
|
+
request_id: str,
|
|
18
|
+
model: Optional[str] = None,
|
|
19
|
+
provider: Optional[str] = None,
|
|
20
|
+
duration_ms: Optional[float] = None,
|
|
21
|
+
metadata: Optional[Dict[str, Any]] = None,
|
|
22
|
+
):
|
|
23
|
+
"""
|
|
24
|
+
Structured event logger for model router pipeline.
|
|
25
|
+
Redacts sensitive credentials automatically.
|
|
26
|
+
"""
|
|
27
|
+
payload = {
|
|
28
|
+
"event": event_name,
|
|
29
|
+
"request_id": request_id,
|
|
30
|
+
"timestamp": datetime.datetime.utcnow().isoformat(),
|
|
31
|
+
"model": model,
|
|
32
|
+
"provider": provider,
|
|
33
|
+
"duration_ms": duration_ms,
|
|
34
|
+
}
|
|
35
|
+
if metadata:
|
|
36
|
+
# Sanitize any key that looks like an API token
|
|
37
|
+
safe_meta = {
|
|
38
|
+
k: ("[REDACTED]" if "key" in k.lower() or "secret" in k.lower() or "token" in k.lower() and "count" not in k.lower() else v)
|
|
39
|
+
for k, v in metadata.items()
|
|
40
|
+
}
|
|
41
|
+
payload["metadata"] = safe_meta
|
|
42
|
+
|
|
43
|
+
logger.info(json.dumps(payload))
|