77 lines
2.2 KiB
Python
77 lines
2.2 KiB
Python
"""Model to provider router."""
|
|
from typing import Optional, Tuple
|
|
from ..database import db_connection
|
|
from ..models import Provider
|
|
from ..config import MODEL_PROVIDER_MAP
|
|
|
|
|
|
class ModelRouter:
|
|
"""Routes model names to providers."""
|
|
|
|
@staticmethod
|
|
def get_provider_name(model: str) -> Optional[str]:
|
|
"""Determine provider name from model name.
|
|
|
|
Args:
|
|
model: The model name (e.g., 'gpt-4o', 'claude-3-opus')
|
|
|
|
Returns:
|
|
Provider name or None if unknown
|
|
"""
|
|
model_lower = model.lower()
|
|
for prefix, provider in MODEL_PROVIDER_MAP.items():
|
|
if model_lower.startswith(prefix):
|
|
return provider
|
|
return None
|
|
|
|
@staticmethod
|
|
async def get_provider(provider_name: str) -> Optional[Provider]:
|
|
"""Get provider configuration by name.
|
|
|
|
Args:
|
|
provider_name: The provider name (e.g., 'openai', 'anthropic')
|
|
|
|
Returns:
|
|
Provider if found and active, None otherwise
|
|
"""
|
|
async with db_connection() as db:
|
|
cursor = await db.execute(
|
|
"""SELECT id, name, base_url, api_type, is_active
|
|
FROM providers WHERE name = ? AND is_active = TRUE""",
|
|
(provider_name,)
|
|
)
|
|
row = await cursor.fetchone()
|
|
if row:
|
|
return Provider(
|
|
id=row["id"],
|
|
name=row["name"],
|
|
base_url=row["base_url"],
|
|
api_type=row["api_type"],
|
|
is_active=row["is_active"]
|
|
)
|
|
return None
|
|
|
|
@staticmethod
|
|
async def route_model(model: str) -> Tuple[Optional[Provider], Optional[str]]:
|
|
"""Route a model to its provider.
|
|
|
|
Args:
|
|
model: The model name
|
|
|
|
Returns:
|
|
Tuple of (Provider, error_message)
|
|
"""
|
|
provider_name = ModelRouter.get_provider_name(model)
|
|
if not provider_name:
|
|
return None, f"Unsupported model: {model}"
|
|
|
|
provider = await ModelRouter.get_provider(provider_name)
|
|
if not provider:
|
|
return None, f"Provider not found: {provider_name}"
|
|
|
|
return provider, None
|
|
|
|
|
|
# Global instance
|
|
model_router = ModelRouter()
|