Files
zzrouter/backend/app/services/router.py

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()