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/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
+ }
@@ -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))