diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 29cbfcd0..efd2a566 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -193,7 +193,7 @@ def create_model_mappings( Tuple of (model_instances, provider_map, unique_models) """ from .payment.models import _row_to_model - from .upstream import resolve_model_alias + from .upstream.helpers import resolve_model_alias model_instances: dict[str, "Model"] = {} provider_map: dict[str, "BaseUpstreamProvider"] = {} diff --git a/routstr/core/admin.py b/routstr/core/admin.py index bc44f369..3baad78b 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -2739,45 +2739,9 @@ async def delete_upstream_provider(provider_id: int) -> dict[str, object]: @admin_router.get("/api/provider-types", dependencies=[Depends(require_admin_api)]) async def get_provider_types() -> list[dict[str, object]]: """Get metadata about available provider types including default URLs and whether they're fixed.""" - provider_types = [ - { - "id": "openrouter", - "name": "OpenRouter", - "default_base_url": "https://openrouter.ai/api/v1", - "fixed_base_url": True, - }, - { - "id": "openai", - "name": "OpenAI", - "default_base_url": "https://api.openai.com/v1", - "fixed_base_url": True, - }, - { - "id": "anthropic", - "name": "Anthropic", - "default_base_url": "https://api.anthropic.com/v1", - "fixed_base_url": True, - }, - { - "id": "azure", - "name": "Azure OpenAI", - "default_base_url": "", - "fixed_base_url": False, - }, - { - "id": "ollama", - "name": "Ollama", - "default_base_url": "http://localhost:11434", - "fixed_base_url": False, - }, - { - "id": "generic", - "name": "Generic", - "default_base_url": "", - "fixed_base_url": False, - }, - ] - return provider_types + from ..upstream import upstream_provider_classes + + return [cls.get_provider_metadata() for cls in upstream_provider_classes] @admin_router.get( @@ -2785,7 +2749,7 @@ async def get_provider_types() -> list[dict[str, object]]: dependencies=[Depends(require_admin_api)], ) async def get_provider_models(provider_id: int) -> dict[str, object]: - from ..upstream import _instantiate_provider + from ..upstream.helpers import _instantiate_provider async with create_session() as session: provider = await session.get(UpstreamProviderRow, provider_id) diff --git a/routstr/core/main.py b/routstr/core/main.py index 22bcb5ce..be6e1f5c 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -78,7 +78,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: from ..payment.price import _update_prices from ..proxy import get_upstreams - from ..upstream import refresh_upstreams_models_periodically + from ..upstream.helpers import refresh_upstreams_models_periodically await _update_prices() await initialize_upstreams() diff --git a/routstr/proxy.py b/routstr/proxy.py index 1e6b376f..ce558fc5 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -23,7 +23,8 @@ from .payment.helpers import ( get_max_cost_for_model, ) from .payment.models import Model -from .upstream import BaseUpstreamProvider, init_upstreams +from .upstream import BaseUpstreamProvider +from .upstream.helpers import init_upstreams logger = get_logger(__name__) proxy_router = APIRouter() diff --git a/routstr/upstream/__init__.py b/routstr/upstream/__init__.py index 82d649aa..7c2f4961 100644 --- a/routstr/upstream/__init__.py +++ b/routstr/upstream/__init__.py @@ -1,34 +1,31 @@ from .anthropic import AnthropicUpstreamProvider from .azure import AzureUpstreamProvider from .base import BaseUpstreamProvider +from .fireworks import FireworksUpstreamProvider from .generic import GenericUpstreamProvider -from .helpers import ( - _instantiate_provider, - _seed_providers_from_settings, - get_all_models_with_overrides, - init_upstreams, - refresh_upstreams_models_periodically, - resolve_model_alias, -) +from .groq import GroqUpstreamProvider from .ollama import OllamaUpstreamProvider from .openai import OpenAIUpstreamProvider from .openrouter import OpenRouterUpstreamProvider +from .perplexity import PerplexityUpstreamProvider +from .xai import XAIUpstreamProvider + +upstream_provider_classes: list[type[BaseUpstreamProvider]] = [ + AnthropicUpstreamProvider, + AzureUpstreamProvider, + FireworksUpstreamProvider, + GenericUpstreamProvider, + GroqUpstreamProvider, + OllamaUpstreamProvider, + OpenAIUpstreamProvider, + OpenRouterUpstreamProvider, + PerplexityUpstreamProvider, + XAIUpstreamProvider, +] +"""List of all upstream classes""" __all__ = [ - # upstreams - "AnthropicUpstreamProvider", - "AzureUpstreamProvider", "BaseUpstreamProvider", - "GenericUpstreamProvider", - "OllamaUpstreamProvider", - "OpenAIUpstreamProvider", - "OpenRouterUpstreamProvider", - # helpers - "resolve_model_alias", - "get_all_models_with_overrides", - "get_model_with_override", - "refresh_upstreams_models_periodically", - "init_upstreams", - "_seed_providers_from_settings", - "_instantiate_provider", + *[cls.__name__ for cls in upstream_provider_classes], + "upstream_provider_classes", ] diff --git a/routstr/upstream/anthropic.py b/routstr/upstream/anthropic.py index 68f61f4f..3f228e9c 100644 --- a/routstr/upstream/anthropic.py +++ b/routstr/upstream/anthropic.py @@ -1,18 +1,45 @@ +from typing import TYPE_CHECKING + from ..payment.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider +if TYPE_CHECKING: + from ..core.db import UpstreamProviderRow + class AnthropicUpstreamProvider(BaseUpstreamProvider): """Upstream provider specifically configured for Anthropic API.""" + provider_type = "anthropic" + default_base_url = "https://api.anthropic.com/v1" + platform_url = "https://console.anthropic.com/settings/keys" + def __init__(self, api_key: str, provider_fee: float = 1.01): - self.upstream_name = "anthropic" super().__init__( - base_url="https://api.anthropic.com/v1", + base_url=self.default_base_url, api_key=api_key, provider_fee=provider_fee, ) + @classmethod + def from_db_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "AnthropicUpstreamProvider": + return cls( + api_key=provider_row.api_key, + provider_fee=provider_row.provider_fee, + ) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + return { + "id": cls.provider_type, + "name": "Anthropic", + "default_base_url": cls.default_base_url, + "fixed_base_url": True, + "platform_url": cls.platform_url, + } + def transform_model_name(self, model_id: str) -> str: """Strip 'anthropic/' prefix for Anthropic API compatibility and transform model names.""" if model_id.startswith("anthropic/"): diff --git a/routstr/upstream/azure.py b/routstr/upstream/azure.py index 324d303a..b6240fbd 100644 --- a/routstr/upstream/azure.py +++ b/routstr/upstream/azure.py @@ -1,11 +1,18 @@ -from typing import Mapping +from typing import TYPE_CHECKING, Mapping from .base import BaseUpstreamProvider +if TYPE_CHECKING: + from ..core.db import UpstreamProviderRow + class AzureUpstreamProvider(BaseUpstreamProvider): """Upstream provider specifically configured for Azure OpenAI Service.""" + provider_type = "azure" + default_base_url = None + platform_url = "https://portal.azure.com/" + def __init__( self, base_url: str, @@ -28,6 +35,29 @@ class AzureUpstreamProvider(BaseUpstreamProvider): ) self.api_version = api_version + @classmethod + def from_db_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "AzureUpstreamProvider | None": + if not provider_row.api_version: + return None + return cls( + base_url=provider_row.base_url, + api_key=provider_row.api_key, + api_version=provider_row.api_version, + provider_fee=provider_row.provider_fee, + ) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + return { + "id": cls.provider_type, + "name": "Azure OpenAI", + "default_base_url": "", + "fixed_base_url": False, + "platform_url": cls.platform_url, + } + def prepare_params( self, path: str, query_params: Mapping[str, str] | None ) -> Mapping[str, str]: diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 2007996e..bc47f650 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -4,7 +4,7 @@ import json import re import traceback from collections.abc import AsyncGenerator -from typing import Mapping +from typing import TYPE_CHECKING, Mapping import httpx from fastapi import BackgroundTasks, HTTPException, Request @@ -13,6 +13,10 @@ from fastapi.responses import Response, StreamingResponse from ..auth import adjust_payment_for_tokens from ..core import get_logger from ..core.db import ApiKey, AsyncSession, create_session + +if TYPE_CHECKING: + from ..core.db import UpstreamProviderRow + from ..payment.cost_caculation import ( CostData, CostDataError, @@ -35,9 +39,12 @@ logger = get_logger(__name__) class BaseUpstreamProvider: """Provider for forwarding requests to an upstream AI service API.""" + provider_type: str = "base" + default_base_url: str | None = None + platform_url: str | None = None + base_url: str api_key: str - upstream_name: str | None = None provider_fee: float = 1.05 _models_cache: list[Model] = [] _models_by_id: dict[str, Model] = {} @@ -56,6 +63,39 @@ class BaseUpstreamProvider: self._models_cache = [] self._models_by_id = {} + @classmethod + def from_db_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "BaseUpstreamProvider | None": + """Factory method to instantiate provider from database row. + + Args: + provider_row: Database row containing provider configuration + + Returns: + Instantiated provider or None if instantiation fails + """ + return cls( + base_url=provider_row.base_url, + api_key=provider_row.api_key, + provider_fee=provider_row.provider_fee, + ) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + """Get metadata about this provider type for API responses. + + Returns: + Dict with provider type metadata including id, name, default_base_url, fixed_base_url, platform_url + """ + return { + "id": cls.provider_type, + "name": cls.provider_type.title(), + "default_base_url": cls.default_base_url or "", + "fixed_base_url": bool(cls.default_base_url), + "platform_url": cls.platform_url, + } + def prepare_headers(self, request_headers: dict) -> dict: """Prepare headers for upstream request by removing proxy-specific headers and adding authentication. @@ -162,7 +202,7 @@ class BaseUpstreamProvider: extra={ "original": original_model, "transformed": transformed_model, - "provider": self.upstream_name or self.base_url, + "provider": self.provider_type or self.base_url, }, ) return json.dumps(data).encode() @@ -171,7 +211,7 @@ class BaseUpstreamProvider: "Could not transform request body", extra={ "error": str(e), - "provider": self.upstream_name or self.base_url, + "provider": self.provider_type or self.base_url, }, ) @@ -1657,7 +1697,7 @@ class BaseUpstreamProvider: Returns: List of Model objects with pricing """ - logger.debug(f"Fetching models for {self.upstream_name or self.base_url}") + logger.debug(f"Fetching models for {self.provider_type or self.base_url}") return [] async def refresh_models_cache(self) -> None: @@ -1676,12 +1716,12 @@ class BaseUpstreamProvider: self._models_by_id = {m.id: m for m in self._models_cache} logger.info( - f"Refreshed models cache for {self.upstream_name or self.base_url}", + f"Refreshed models cache for {self.provider_type or self.base_url}", extra={"model_count": len(models)}, ) except Exception as e: logger.error( - f"Failed to refresh models cache for {self.upstream_name or self.base_url}", + f"Failed to refresh models cache for {self.provider_type or self.base_url}", extra={"error": str(e), "error_type": type(e).__name__}, ) diff --git a/routstr/upstream/fireworks.py b/routstr/upstream/fireworks.py index b68a9f91..4f66f10c 100644 --- a/routstr/upstream/fireworks.py +++ b/routstr/upstream/fireworks.py @@ -1,19 +1,43 @@ +from typing import TYPE_CHECKING + from ..payment.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider +if TYPE_CHECKING: + from ..core.db import UpstreamProviderRow + class FireworksUpstreamProvider(BaseUpstreamProvider): """Upstream provider specifically configured for Fireworks.ai API.""" - upstream_name = "fireworks" - base_url = "https://api.fireworks.ai/inference/v1" + provider_type = "fireworks" + default_base_url = "https://api.fireworks.ai/inference/v1" platform_url = "https://app.fireworks.ai/settings/users/api-keys" def __init__(self, api_key: str, provider_fee: float = 1.01): super().__init__( - base_url=self.base_url, api_key=api_key, provider_fee=provider_fee + base_url=self.default_base_url, api_key=api_key, provider_fee=provider_fee ) + @classmethod + def from_db_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "FireworksUpstreamProvider": + return cls( + api_key=provider_row.api_key, + provider_fee=provider_row.provider_fee, + ) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + return { + "id": cls.provider_type, + "name": "Fireworks", + "default_base_url": cls.default_base_url, + "fixed_base_url": True, + "platform_url": cls.platform_url, + } + def transform_model_name(self, model_id: str) -> str: """Strip 'fireworks/' prefix for Fireworks API compatibility.""" return model_id.removeprefix("fireworks/") diff --git a/routstr/upstream/generic.py b/routstr/upstream/generic.py index 403b1bd4..390c8372 100644 --- a/routstr/upstream/generic.py +++ b/routstr/upstream/generic.py @@ -7,6 +7,7 @@ import httpx from .base import BaseUpstreamProvider if TYPE_CHECKING: + from ..core.db import UpstreamProviderRow from ..payment.models import Model from ..core.logging import get_logger @@ -17,6 +18,10 @@ logger = get_logger(__name__) class GenericUpstreamProvider(BaseUpstreamProvider): """Generic upstream provider that can fetch models from any OpenAI-compatible API.""" + provider_type = "generic" + default_base_url = "http://localhost:8888" + platform_url = None + def __init__( self, base_url: str, @@ -39,6 +44,26 @@ class GenericUpstreamProvider(BaseUpstreamProvider): provider_fee=provider_fee, ) + @classmethod + def from_db_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "GenericUpstreamProvider": + return cls( + base_url=provider_row.base_url, + api_key=provider_row.api_key, + provider_fee=provider_row.provider_fee, + ) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + return { + "id": cls.provider_type, + "name": "Generic", + "default_base_url": cls.default_base_url, + "fixed_base_url": False, + "platform_url": cls.platform_url, + } + async def fetch_models(self) -> list[Model]: """Fetch models from upstream API using /models endpoint.""" from ..payment.models import Architecture, Model, Pricing, TopProvider diff --git a/routstr/upstream/groq.py b/routstr/upstream/groq.py index e0a6efef..36aae48e 100644 --- a/routstr/upstream/groq.py +++ b/routstr/upstream/groq.py @@ -1,19 +1,41 @@ +from typing import TYPE_CHECKING + from ..payment.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider +if TYPE_CHECKING: + from ..core.db import UpstreamProviderRow + class GroqUpstreamProvider(BaseUpstreamProvider): """Upstream provider specifically configured for Groq API.""" - upstream_name = "groq" - base_url = "https://api.groq.com/openai/v1" + provider_type = "groq" + default_base_url = "https://api.groq.com/openai/v1" platform_url = "https://console.groq.com/keys" def __init__(self, api_key: str, provider_fee: float = 1.01): super().__init__( - base_url=self.base_url, api_key=api_key, provider_fee=provider_fee + base_url=self.default_base_url, api_key=api_key, provider_fee=provider_fee ) + @classmethod + def from_db_row(cls, provider_row: "UpstreamProviderRow") -> "GroqUpstreamProvider": + return cls( + api_key=provider_row.api_key, + provider_fee=provider_row.provider_fee, + ) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + return { + "id": cls.provider_type, + "name": "Groq", + "default_base_url": cls.default_base_url, + "fixed_base_url": True, + "platform_url": cls.platform_url, + } + def transform_model_name(self, model_id: str) -> str: """Strip 'groq/' prefix for Groq API compatibility.""" return model_id.removeprefix("groq/") diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index 5e446837..3d91550f 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -11,13 +11,7 @@ if TYPE_CHECKING: from ..core import get_logger from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session from ..payment.models import Model -from .anthropic import AnthropicUpstreamProvider -from .azure import AzureUpstreamProvider from .base import BaseUpstreamProvider -from .generic import GenericUpstreamProvider -from .ollama import OllamaUpstreamProvider -from .openai import OpenAIUpstreamProvider -from .openrouter import OpenRouterUpstreamProvider logger = get_logger(__name__) @@ -148,7 +142,7 @@ async def refresh_upstreams_models_periodically( await upstream.refresh_models_cache() except Exception as e: logger.error( - f"Error refreshing models for {upstream.upstream_name or upstream.base_url}", + f"Error refreshing models for {upstream.base_url}", extra={"error": str(e), "error_type": type(e).__name__}, ) except asyncio.CancelledError: @@ -220,61 +214,47 @@ async def _seed_providers_from_settings( """ from sqlmodel import select - from ..core.settings import settings + from . import upstream_provider_classes providers_to_add: list[UpstreamProviderRow] = [] seeded_base_urls: set[str] = set() - openai_api_key = os.environ.get("OPENAI_API_KEY") - if openai_api_key: - base_url = "https://api.openai.com/v1" - result = await session.exec( - select(UpstreamProviderRow).where(UpstreamProviderRow.base_url == base_url) - ) - if not result.first(): - providers_to_add.append( - UpstreamProviderRow( - provider_type="openai", - base_url=base_url, - api_key=openai_api_key, - enabled=True, - ) - ) - seeded_base_urls.add(base_url) + provider_classes_by_type = { + cls.provider_type: cls + for cls in upstream_provider_classes # type: ignore[attr-defined] + } - anthropic_api_key = os.environ.get("ANTHROPIC_API_KEY") - if anthropic_api_key: - base_url = "https://api.anthropic.com/v1" - result = await session.exec( - select(UpstreamProviderRow).where(UpstreamProviderRow.base_url == base_url) - ) - if not result.first(): - providers_to_add.append( - UpstreamProviderRow( - provider_type="anthropic", - base_url=base_url, - api_key=anthropic_api_key, - enabled=True, - ) - ) - seeded_base_urls.add(base_url) + env_mappings: list[tuple[str, str, str | None, str | None]] = [ + ("OPENAI_API_KEY", "openai", None, None), + ("ANTHROPIC_API_KEY", "anthropic", None, None), + ("OPENROUTER_API_KEY", "openrouter", None, None), + ("GROQ_API_KEY", "groq", None, None), + ("PERPLEXITY_API_KEY", "perplexity", None, None), + ("FIREWORKS_API_KEY", "fireworks", None, None), + ("XAI_API_KEY", "xai", None, None), + ] - openrouter_api_key = os.environ.get("OPENROUTER_API_KEY") - if openrouter_api_key: - base_url = "https://openrouter.ai/api/v1" - result = await session.exec( - select(UpstreamProviderRow).where(UpstreamProviderRow.base_url == base_url) - ) - if not result.first(): - providers_to_add.append( - UpstreamProviderRow( - provider_type="openrouter", - base_url=base_url, - api_key=openrouter_api_key, - enabled=True, + for env_key, provider_type, _, _ in env_mappings: + api_key = os.environ.get(env_key) + if api_key and provider_type in provider_classes_by_type: + provider_class = provider_classes_by_type[provider_type] + if provider_class.default_base_url: # type: ignore[attr-defined] + base_url = provider_class.default_base_url # type: ignore[attr-defined] + result = await session.exec( + select(UpstreamProviderRow).where( + UpstreamProviderRow.base_url == base_url + ) ) - ) - seeded_base_urls.add(base_url) + if not result.first(): + providers_to_add.append( + UpstreamProviderRow( + provider_type=provider_type, + base_url=base_url, + api_key=api_key, + enabled=True, + ) + ) + seeded_base_urls.add(base_url) ollama_base_url = os.environ.get("OLLAMA_BASE_URL") if ollama_base_url: @@ -323,48 +303,20 @@ async def _seed_providers_from_settings( ) ) if not result.first(): - if "api.openai.com" in base_url.lower(): - providers_to_add.append( - UpstreamProviderRow( - provider_type="openai", - base_url=base_url, - api_key=settings.upstream_api_key, - enabled=True, - ) - ) - elif "api.anthropic.com" in base_url.lower(): - providers_to_add.append( - UpstreamProviderRow( - provider_type="anthropic", - base_url=base_url, - api_key=settings.upstream_api_key, - enabled=True, - ) - ) - elif "openrouter.ai/api/v1" in base_url.lower(): - providers_to_add.append( - UpstreamProviderRow( - provider_type="openrouter", - base_url=base_url, - api_key=settings.upstream_api_key, - enabled=True, - ) - ) - else: - providers_to_add.append( - UpstreamProviderRow( - provider_type="custom", - base_url=base_url, - api_key=settings.upstream_api_key, - enabled=True, - ) + providers_to_add.append( + UpstreamProviderRow( + provider_type="custom", + base_url=base_url, + api_key=settings.upstream_api_key, + enabled=True, ) + ) seeded_base_urls.add(base_url) for provider in providers_to_add: session.add(provider) logger.info( - f"Seeding {provider.provider_type} provider", + f"Seeding {provider.provider_type} provider", # type: ignore[str-format] extra={"base_url": provider.base_url}, ) @@ -380,53 +332,35 @@ def _instantiate_provider( Returns: Instantiated provider or None if provider type is unknown """ + from . import upstream_provider_classes + try: - if provider_row.provider_type == "openai": - return OpenAIUpstreamProvider( - provider_row.api_key, provider_row.provider_fee - ) - elif provider_row.provider_type == "anthropic": - return AnthropicUpstreamProvider( - provider_row.api_key, provider_row.provider_fee - ) - elif provider_row.provider_type == "azure": - if not provider_row.api_version: + provider_classes_by_type = { + cls.provider_type: cls + for cls in upstream_provider_classes # type: ignore[attr-defined] + } + + provider_class = provider_classes_by_type.get(provider_row.provider_type) + + if provider_class: + provider = provider_class.from_db_row(provider_row) # type: ignore[attr-defined] + if provider is None: logger.error( - "Azure provider missing api_version", + f"Failed to instantiate {provider_row.provider_type} provider", extra={"base_url": provider_row.base_url}, ) - return None - return AzureUpstreamProvider( - provider_row.base_url, - provider_row.api_key, - provider_row.api_version, - provider_row.provider_fee, - ) - elif provider_row.provider_type == "openrouter": - return OpenRouterUpstreamProvider( - provider_row.api_key, provider_row.provider_fee - ) - elif provider_row.provider_type == "ollama": - return OllamaUpstreamProvider( - provider_row.base_url, provider_row.api_key, provider_row.provider_fee - ) - elif provider_row.provider_type == "generic": - return GenericUpstreamProvider( - provider_row.base_url, - provider_row.api_key, - provider_row.provider_fee, - provider_row.provider_type, - ) - elif provider_row.provider_type == "custom": + return provider + + if provider_row.provider_type == "custom": return BaseUpstreamProvider( provider_row.base_url, provider_row.api_key, provider_row.provider_fee ) - else: - logger.error( - f"Unknown provider type: {provider_row.provider_type}", - extra={"base_url": provider_row.base_url}, - ) - return None + + logger.error( + f"Unknown provider type: {provider_row.provider_type}", + extra={"base_url": provider_row.base_url}, + ) + return None except Exception as e: logger.error( f"Failed to instantiate provider: {e}", diff --git a/routstr/upstream/ollama.py b/routstr/upstream/ollama.py index c3b0fc9e..923a6018 100644 --- a/routstr/upstream/ollama.py +++ b/routstr/upstream/ollama.py @@ -9,7 +9,7 @@ from fastapi.responses import Response, StreamingResponse from .base import BaseUpstreamProvider if TYPE_CHECKING: - from ..core.db import ApiKey, AsyncSession + from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow from ..payment.models import Model from ..core.logging import get_logger @@ -20,6 +20,10 @@ logger = get_logger(__name__) class OllamaUpstreamProvider(BaseUpstreamProvider): """Upstream provider specifically configured for Ollama API.""" + provider_type = "ollama" + default_base_url = "http://localhost:11434" + platform_url = "https://ollama.com/" + def __init__( self, base_url: str = "http://localhost:11434", @@ -33,13 +37,32 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): api_key: Optional API key (Ollama typically doesn't require one) provider_fee: Provider fee multiplier (default 1.01 for 1% fee) """ - self.upstream_name = "ollama" super().__init__( base_url=base_url, api_key=api_key, provider_fee=provider_fee, ) + @classmethod + def from_db_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "OllamaUpstreamProvider": + return cls( + base_url=provider_row.base_url, + api_key=provider_row.api_key, + provider_fee=provider_row.provider_fee, + ) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + return { + "id": cls.provider_type, + "name": "Ollama", + "default_base_url": cls.default_base_url, + "fixed_base_url": False, + "platform_url": cls.platform_url, + } + def transform_model_name(self, model_id: str) -> str: """Strip 'ollama/' prefix for Ollama API compatibility.""" return model_id.removeprefix("ollama/") @@ -192,12 +215,12 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): self._models_by_id = {m.id: m for m in self._models_cache} logger.info( - f"Refreshed models cache for {self.upstream_name or self.base_url}", + f"Refreshed models cache for {self.base_url}", extra={"model_count": len(models)}, ) except Exception as e: logger.error( - f"Failed to refresh models cache for {self.upstream_name or self.base_url}", + f"Failed to refresh models cache for {self.base_url}", extra={"error": str(e), "error_type": type(e).__name__}, ) diff --git a/routstr/upstream/openai.py b/routstr/upstream/openai.py index 3a0482d9..11cc4336 100644 --- a/routstr/upstream/openai.py +++ b/routstr/upstream/openai.py @@ -1,18 +1,43 @@ +from typing import TYPE_CHECKING + from ..payment.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider +if TYPE_CHECKING: + from ..core.db import UpstreamProviderRow + class OpenAIUpstreamProvider(BaseUpstreamProvider): """Upstream provider specifically configured for OpenAI API.""" + provider_type = "openai" + default_base_url = "https://api.openai.com/v1" + platform_url = "https://platform.openai.com/api-keys" + def __init__(self, api_key: str, provider_fee: float = 1.01): - self.upstream_name = "openai" super().__init__( - base_url="https://api.openai.com/v1", - api_key=api_key, - provider_fee=provider_fee, + base_url=self.default_base_url, api_key=api_key, provider_fee=provider_fee ) + @classmethod + def from_db_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "OpenAIUpstreamProvider": + return cls( + api_key=provider_row.api_key, + provider_fee=provider_row.provider_fee, + ) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + return { + "id": cls.provider_type, + "name": "OpenAI", + "default_base_url": cls.default_base_url, + "fixed_base_url": True, + "platform_url": cls.platform_url, + } + def transform_model_name(self, model_id: str) -> str: """Strip 'openai/' prefix for OpenAI API compatibility.""" return model_id.removeprefix("openai/") diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index 956f4e15..cc0e2908 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -1,10 +1,19 @@ +from typing import TYPE_CHECKING + from ..payment.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider +if TYPE_CHECKING: + from ..core.db import UpstreamProviderRow + class OpenRouterUpstreamProvider(BaseUpstreamProvider): """Upstream provider specifically configured for OpenRouter API.""" + provider_type = "openrouter" + default_base_url = "https://openrouter.ai/api/v1" + platform_url = "https://openrouter.ai/settings/keys" + def __init__(self, api_key: str, provider_fee: float = 1.06): """Initialize OpenRouter provider with API key. @@ -12,13 +21,29 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): api_key: OpenRouter API key for authentication provider_fee: Provider fee multiplier (default 1.06 for 6% fee) """ - self.upstream_name = "openrouter" super().__init__( - base_url="https://openrouter.ai/api/v1", - api_key=api_key, - provider_fee=provider_fee, + base_url=self.default_base_url, api_key=api_key, provider_fee=provider_fee ) + @classmethod + def from_db_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "OpenRouterUpstreamProvider": + return cls( + api_key=provider_row.api_key, + provider_fee=provider_row.provider_fee, + ) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + return { + "id": cls.provider_type, + "name": "OpenRouter", + "default_base_url": cls.default_base_url, + "fixed_base_url": True, + "platform_url": cls.platform_url, + } + async def fetch_models(self) -> list[Model]: """Fetch all OpenRouter models.""" models_data = await async_fetch_openrouter_models() diff --git a/routstr/upstream/perplexity.py b/routstr/upstream/perplexity.py index 775a0229..b73881d8 100644 --- a/routstr/upstream/perplexity.py +++ b/routstr/upstream/perplexity.py @@ -1,21 +1,45 @@ +from typing import TYPE_CHECKING + from ..payment.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider +if TYPE_CHECKING: + from ..core.db import UpstreamProviderRow + class PerplexityUpstreamProvider(BaseUpstreamProvider): - """Upstream provider specifically configured for OpenAI API.""" + """Upstream provider specifically configured for Perplexity API.""" - upstream_name = "perplexity" - base_url = "https://api.perplexity.ai/" # without v1 + provider_type = "perplexity" + default_base_url = "https://api.perplexity.ai/" platform_url = "https://www.perplexity.ai/account/api/keys" def __init__(self, api_key: str, provider_fee: float = 1.01): super().__init__( - base_url=self.base_url, + base_url=self.default_base_url, api_key=api_key, provider_fee=provider_fee, ) + @classmethod + def from_db_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "PerplexityUpstreamProvider": + return cls( + api_key=provider_row.api_key, + provider_fee=provider_row.provider_fee, + ) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + return { + "id": cls.provider_type, + "name": "Perplexity", + "default_base_url": cls.default_base_url, + "fixed_base_url": True, + "platform_url": cls.platform_url, + } + def transform_model_name(self, model_id: str) -> str: """Strip 'perplexity/' prefix for Perplexity API compatibility.""" return model_id.removeprefix("perplexity/") diff --git a/routstr/upstream/xai.py b/routstr/upstream/xai.py index 059b0dc3..95332f85 100644 --- a/routstr/upstream/xai.py +++ b/routstr/upstream/xai.py @@ -1,19 +1,41 @@ +from typing import TYPE_CHECKING + from ..payment.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider +if TYPE_CHECKING: + from ..core.db import UpstreamProviderRow + class XAIUpstreamProvider(BaseUpstreamProvider): """Upstream provider specifically configured for XAI API.""" - upstream_name = "xai" - base_url = "https://api.x.ai/v1" - platform_url = "https://accounts.x.ai/sign-up" + provider_type = "xai" + default_base_url = "https://api.x.ai/v1" + platform_url = "https://console.x.ai/" def __init__(self, api_key: str, provider_fee: float = 1.01): super().__init__( - base_url=self.base_url, api_key=api_key, provider_fee=provider_fee + base_url=self.default_base_url, api_key=api_key, provider_fee=provider_fee ) + @classmethod + def from_db_row(cls, provider_row: "UpstreamProviderRow") -> "XAIUpstreamProvider": + return cls( + api_key=provider_row.api_key, + provider_fee=provider_row.provider_fee, + ) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + return { + "id": cls.provider_type, + "name": "xAI", + "default_base_url": cls.default_base_url, + "fixed_base_url": True, + "platform_url": cls.platform_url, + } + def transform_model_name(self, model_id: str) -> str: """Strip 'xai/' prefix for XAI API compatibility.""" return model_id.removeprefix("xai/") diff --git a/tests/unit/test_algorithm.py b/tests/unit/test_algorithm.py index 22155e79..d367819f 100644 --- a/tests/unit/test_algorithm.py +++ b/tests/unit/test_algorithm.py @@ -49,7 +49,7 @@ def create_test_model( def create_test_provider(name: str, base_url: str = "http://test.com") -> Mock: """Helper to create a test provider mock.""" provider = Mock() - provider.upstream_name = name + provider.provider_type = name provider.base_url = base_url return provider