diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 316b9f35..29cbfcd0 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -6,7 +6,7 @@ from .core.logging import get_logger if TYPE_CHECKING: from .payment.models import Model - from .upstream import UpstreamProvider + from .upstream import BaseUpstreamProvider logger = get_logger(__name__) @@ -59,7 +59,7 @@ def calculate_model_cost_score(model: "Model") -> float: return total_cost -def get_provider_penalty(provider: "UpstreamProvider") -> float: +def get_provider_penalty(provider: "BaseUpstreamProvider") -> float: """Calculate a penalty multiplier for certain providers. This allows applying policy-based adjustments beyond pure cost. @@ -86,9 +86,9 @@ def get_provider_penalty(provider: "UpstreamProvider") -> float: def should_prefer_model( candidate_model: "Model", - candidate_provider: "UpstreamProvider", + candidate_provider: "BaseUpstreamProvider", current_model: "Model", - current_provider: "UpstreamProvider", + current_provider: "BaseUpstreamProvider", alias: str, ) -> bool: """Determine if candidate model should replace current model for an alias. @@ -166,10 +166,10 @@ def should_prefer_model( def create_model_mappings( - upstreams: list["UpstreamProvider"], + upstreams: list["BaseUpstreamProvider"], overrides_by_id: dict[str, tuple], disabled_model_ids: set[str], -) -> tuple[dict[str, "Model"], dict[str, "UpstreamProvider"], dict[str, "Model"]]: +) -> tuple[dict[str, "Model"], dict[str, "BaseUpstreamProvider"], dict[str, "Model"]]: """Create optimal model mappings based on cost and provider preferences. This is the main entry point for the algorithm. It processes all upstream providers @@ -196,12 +196,12 @@ def create_model_mappings( from .upstream import resolve_model_alias model_instances: dict[str, "Model"] = {} - provider_map: dict[str, "UpstreamProvider"] = {} + provider_map: dict[str, "BaseUpstreamProvider"] = {} unique_models: dict[str, "Model"] = {} # Separate OpenRouter from other providers - openrouter: "UpstreamProvider" | None = None - other_upstreams: list["UpstreamProvider"] = [] + openrouter: "BaseUpstreamProvider" | None = None + other_upstreams: list["BaseUpstreamProvider"] = [] for upstream in upstreams: base_url = getattr(upstream, "base_url", "") @@ -215,7 +215,7 @@ def create_model_mappings( return model_id.split("/", 1)[1] if "/" in model_id else model_id def _maybe_set_alias( - alias: str, model: "Model", provider: "UpstreamProvider" + alias: str, model: "Model", provider: "BaseUpstreamProvider" ) -> None: """Set alias to model/provider if not set or if new model is preferred.""" existing_model = model_instances.get(alias) @@ -233,7 +233,7 @@ def create_model_mappings( provider_map[alias] = provider def process_provider_models( - upstream: "UpstreamProvider", is_openrouter: bool = False + upstream: "BaseUpstreamProvider", is_openrouter: bool = False ) -> None: """Process all models from a given provider.""" upstream_prefix = getattr(upstream, "upstream_name", None) diff --git a/routstr/proxy.py b/routstr/proxy.py index 47b98c13..1e6b376f 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -23,14 +23,14 @@ from .payment.helpers import ( get_max_cost_for_model, ) from .payment.models import Model -from .upstream import UpstreamProvider, init_upstreams +from .upstream import BaseUpstreamProvider, init_upstreams logger = get_logger(__name__) proxy_router = APIRouter() -_upstreams: list[UpstreamProvider] = [] +_upstreams: list[BaseUpstreamProvider] = [] _model_instances: dict[str, Model] = {} # All aliases -> Model -_provider_map: dict[str, UpstreamProvider] = {} # All aliases -> Provider +_provider_map: dict[str, BaseUpstreamProvider] = {} # All aliases -> Provider _unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates) @@ -53,7 +53,7 @@ async def reinitialize_upstreams() -> None: await refresh_model_maps() -def get_upstreams() -> list[UpstreamProvider]: +def get_upstreams() -> list[BaseUpstreamProvider]: """Get the initialized upstream providers. Returns: @@ -67,7 +67,7 @@ def get_model_instance(model_id: str) -> Model | None: return _model_instances.get(model_id) -def get_provider_for_model(model_id: str) -> UpstreamProvider | None: +def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None: """Get UpstreamProvider for model ID from global cache.""" return _provider_map.get(model_id) diff --git a/routstr/upstream/__init__.py b/routstr/upstream/__init__.py new file mode 100644 index 00000000..6aa23e84 --- /dev/null +++ b/routstr/upstream/__init__.py @@ -0,0 +1,35 @@ +from .anthropic import AnthropicUpstreamProvider +from .azure import AzureUpstreamProvider +from .base import BaseUpstreamProvider +from .generic import GenericUpstreamProvider +from .helpers import ( + _instantiate_provider, + _seed_providers_from_settings, + get_all_models_with_overrides, + get_model_with_override, + init_upstreams, + refresh_upstreams_models_periodically, + resolve_model_alias, +) +from .ollama import OllamaUpstreamProvider +from .openai import OpenAIUpstreamProvider +from .openrouter import OpenRouterUpstreamProvider + +__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", +] diff --git a/routstr/upstream/anthropic.py b/routstr/upstream/anthropic.py new file mode 100644 index 00000000..68f61f4f --- /dev/null +++ b/routstr/upstream/anthropic.py @@ -0,0 +1,43 @@ +from ..payment.models import Model, async_fetch_openrouter_models +from .base import BaseUpstreamProvider + + +class AnthropicUpstreamProvider(BaseUpstreamProvider): + """Upstream provider specifically configured for Anthropic API.""" + + def __init__(self, api_key: str, provider_fee: float = 1.01): + self.upstream_name = "anthropic" + super().__init__( + base_url="https://api.anthropic.com/v1", + api_key=api_key, + provider_fee=provider_fee, + ) + + 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/"): + model_id = model_id[len("anthropic/") :] + fixed_transforms = { + "claude-haiku-4.5": "claude-haiku-4-5-20251001", + "claude-sonnet-4.5": "claude-sonnet-4-5-20250929", + "claude-opus-4.1": "claude-opus-4-1-20250805", + "claude-opus-4": "claude-opus-4-20250514", + "claude-sonnet-4": "claude-sonnet-4-20250514", + "claude-3.5-haiku": "claude-3-5-haiku-20241022", + "claude-3-haiku": "claude-3-haiku-20240307", + "claude-haiku-4-5": "claude-haiku-4-5-20251001", + "claude-sonnet-4-5": "claude-sonnet-4-5-20250929", + "claude-opus-4-1": "claude-opus-4-1-20250805", + "claude-3-5-haiku": "claude-3-5-haiku-20241022", + } + if model_id in fixed_transforms: + model_id = fixed_transforms[model_id] + return model_id + + async def fetch_models(self) -> list[Model]: + """Fetch Anthropic models from OpenRouter API filtered by anthropic source.""" + models_data = await async_fetch_openrouter_models(source_filter="anthropic") + models = [Model(**model) for model in models_data] # type: ignore + for model in models: + model.alias_ids = [self.transform_model_name(model.id)] + return models diff --git a/routstr/upstream/azure.py b/routstr/upstream/azure.py new file mode 100644 index 00000000..324d303a --- /dev/null +++ b/routstr/upstream/azure.py @@ -0,0 +1,46 @@ +from typing import Mapping + +from .base import BaseUpstreamProvider + + +class AzureUpstreamProvider(BaseUpstreamProvider): + """Upstream provider specifically configured for Azure OpenAI Service.""" + + def __init__( + self, + base_url: str, + api_key: str, + api_version: str, + provider_fee: float = 1.01, + ): + """Initialize Azure provider with API key and version. + + Args: + base_url: Azure OpenAI endpoint base URL + api_key: Azure OpenAI API key for authentication + api_version: Azure OpenAI API version (e.g., "2024-02-15-preview") + provider_fee: Provider fee multiplier (default 1.01 for 1% fee) + """ + super().__init__( + base_url=base_url, + api_key=api_key, + provider_fee=provider_fee, + ) + self.api_version = api_version + + def prepare_params( + self, path: str, query_params: Mapping[str, str] | None + ) -> Mapping[str, str]: + """Prepare query parameters for Azure OpenAI, adding API version. + + Args: + path: Request path + query_params: Original query parameters from the client + + Returns: + Query parameters dict with Azure API version added for chat completions + """ + params = dict(query_params or {}) + if path.endswith("chat/completions"): + params["api-version"] = self.api_version + return params diff --git a/routstr/upstreams/upstream.py b/routstr/upstream/base.py similarity index 92% rename from routstr/upstreams/upstream.py rename to routstr/upstream/base.py index f2591d83..2007996e 100644 --- a/routstr/upstreams/upstream.py +++ b/routstr/upstream/base.py @@ -25,7 +25,6 @@ from ..payment.models import ( Pricing, _calculate_usd_max_costs, _update_model_sats_pricing, - async_fetch_openrouter_models, ) from ..payment.price import sats_usd_price from ..wallet import recieve_token, send_token @@ -33,7 +32,7 @@ from ..wallet import recieve_token, send_token logger = get_logger(__name__) -class UpstreamProvider: +class BaseUpstreamProvider: """Provider for forwarding requests to an upstream AI service API.""" base_url: str @@ -1704,131 +1703,3 @@ class UpstreamProvider: Model object or None if not found """ return self._models_by_id.get(model_id) - - -class OpenAIUpstreamProvider(UpstreamProvider): - """Upstream provider specifically configured for OpenAI API.""" - - 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, - ) - - def transform_model_name(self, model_id: str) -> str: - """Strip 'openai/' prefix for OpenAI API compatibility.""" - return model_id.removeprefix("openai/") - - async def fetch_models(self) -> list[Model]: - """Fetch OpenAI models from OpenRouter API filtered by openai source.""" - models_data = await async_fetch_openrouter_models(source_filter="openai") - return [Model(**model) for model in models_data] # type: ignore - - -class AnthropicUpstreamProvider(UpstreamProvider): - """Upstream provider specifically configured for Anthropic API.""" - - def __init__(self, api_key: str, provider_fee: float = 1.01): - self.upstream_name = "anthropic" - super().__init__( - base_url="https://api.anthropic.com/v1", - api_key=api_key, - provider_fee=provider_fee, - ) - - 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/"): - model_id = model_id[len("anthropic/") :] - fixed_transforms = { - "claude-haiku-4.5": "claude-haiku-4-5-20251001", - "claude-sonnet-4.5": "claude-sonnet-4-5-20250929", - "claude-opus-4.1": "claude-opus-4-1-20250805", - "claude-opus-4": "claude-opus-4-20250514", - "claude-sonnet-4": "claude-sonnet-4-20250514", - "claude-3.5-haiku": "claude-3-5-haiku-20241022", - "claude-3-haiku": "claude-3-haiku-20240307", - "claude-haiku-4-5": "claude-haiku-4-5-20251001", - "claude-sonnet-4-5": "claude-sonnet-4-5-20250929", - "claude-opus-4-1": "claude-opus-4-1-20250805", - "claude-3-5-haiku": "claude-3-5-haiku-20241022", - } - if model_id in fixed_transforms: - model_id = fixed_transforms[model_id] - return model_id - - async def fetch_models(self) -> list[Model]: - """Fetch Anthropic models from OpenRouter API filtered by anthropic source.""" - models_data = await async_fetch_openrouter_models(source_filter="anthropic") - models = [Model(**model) for model in models_data] # type: ignore - for model in models: - model.alias_ids = [self.transform_model_name(model.id)] - return models - - -class AzureUpstreamProvider(UpstreamProvider): - """Upstream provider specifically configured for Azure OpenAI Service.""" - - def __init__( - self, - base_url: str, - api_key: str, - api_version: str, - provider_fee: float = 1.01, - ): - """Initialize Azure provider with API key and version. - - Args: - base_url: Azure OpenAI endpoint base URL - api_key: Azure OpenAI API key for authentication - api_version: Azure OpenAI API version (e.g., "2024-02-15-preview") - provider_fee: Provider fee multiplier (default 1.01 for 1% fee) - """ - super().__init__( - base_url=base_url, - api_key=api_key, - provider_fee=provider_fee, - ) - self.api_version = api_version - - def prepare_params( - self, path: str, query_params: Mapping[str, str] | None - ) -> Mapping[str, str]: - """Prepare query parameters for Azure OpenAI, adding API version. - - Args: - path: Request path - query_params: Original query parameters from the client - - Returns: - Query parameters dict with Azure API version added for chat completions - """ - params = dict(query_params or {}) - if path.endswith("chat/completions"): - params["api-version"] = self.api_version - return params - - -class OpenRouterUpstreamProvider(UpstreamProvider): - """Upstream provider specifically configured for OpenRouter API.""" - - def __init__(self, api_key: str, provider_fee: float = 1.06): - """Initialize OpenRouter provider with API key. - - Args: - 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, - ) - - async def fetch_models(self) -> list[Model]: - """Fetch all OpenRouter models.""" - models_data = await async_fetch_openrouter_models() - return [Model(**model) for model in models_data] # type: ignore diff --git a/routstr/upstream/fireworks.py b/routstr/upstream/fireworks.py new file mode 100644 index 00000000..b68a9f91 --- /dev/null +++ b/routstr/upstream/fireworks.py @@ -0,0 +1,24 @@ +from ..payment.models import Model, async_fetch_openrouter_models +from .base import BaseUpstreamProvider + + +class FireworksUpstreamProvider(BaseUpstreamProvider): + """Upstream provider specifically configured for Fireworks.ai API.""" + + upstream_name = "fireworks" + 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 + ) + + def transform_model_name(self, model_id: str) -> str: + """Strip 'fireworks/' prefix for Fireworks API compatibility.""" + return model_id.removeprefix("fireworks/") + + async def fetch_models(self) -> list[Model]: + """Fetch Fireworks models from OpenRouter API filtered by fireworks source.""" + models_data = await async_fetch_openrouter_models(source_filter="fireworks") + return [Model(**model) for model in models_data] # type: ignore diff --git a/routstr/upstreams/generic.py b/routstr/upstream/generic.py similarity index 98% rename from routstr/upstreams/generic.py rename to routstr/upstream/generic.py index 3fc72609..403b1bd4 100644 --- a/routstr/upstreams/generic.py +++ b/routstr/upstream/generic.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import httpx -from .upstream import UpstreamProvider +from .base import BaseUpstreamProvider if TYPE_CHECKING: from ..payment.models import Model @@ -14,7 +14,7 @@ from ..core.logging import get_logger logger = get_logger(__name__) -class GenericUpstreamProvider(UpstreamProvider): +class GenericUpstreamProvider(BaseUpstreamProvider): """Generic upstream provider that can fetch models from any OpenAI-compatible API.""" def __init__( diff --git a/routstr/upstream/groq.py b/routstr/upstream/groq.py new file mode 100644 index 00000000..e0a6efef --- /dev/null +++ b/routstr/upstream/groq.py @@ -0,0 +1,24 @@ +from ..payment.models import Model, async_fetch_openrouter_models +from .base import BaseUpstreamProvider + + +class GroqUpstreamProvider(BaseUpstreamProvider): + """Upstream provider specifically configured for Groq API.""" + + upstream_name = "groq" + 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 + ) + + def transform_model_name(self, model_id: str) -> str: + """Strip 'groq/' prefix for Groq API compatibility.""" + return model_id.removeprefix("groq/") + + async def fetch_models(self) -> list[Model]: + """Fetch Groq models from OpenRouter API filtered by groq source.""" + models_data = await async_fetch_openrouter_models(source_filter="groq") + return [Model(**model) for model in models_data] # type: ignore diff --git a/routstr/upstream.py b/routstr/upstream/helpers.py similarity index 95% rename from routstr/upstream.py rename to routstr/upstream/helpers.py index 4d22930b..79ea4160 100644 --- a/routstr/upstream.py +++ b/routstr/upstream/helpers.py @@ -5,21 +5,21 @@ import re from typing import TYPE_CHECKING if TYPE_CHECKING: - from .core.settings import Settings + from ..core.settings import Settings -from .core import get_logger -from .core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session -from .payment.models import Model -from .upstreams import ( +from ..core import get_logger +from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session +from ..payment.models import Model +from . import ( AnthropicUpstreamProvider, AzureUpstreamProvider, + BaseUpstreamProvider, + GenericUpstreamProvider, OllamaUpstreamProvider, OpenAIUpstreamProvider, OpenRouterUpstreamProvider, - UpstreamProvider, ) -from .upstreams.generic import GenericUpstreamProvider logger = get_logger(__name__) @@ -75,7 +75,7 @@ def resolve_model_alias( async def get_all_models_with_overrides( - upstreams: list[UpstreamProvider], + upstreams: list[BaseUpstreamProvider], ) -> list[Model]: """Get all models from all providers with database overrides applied. @@ -90,7 +90,7 @@ async def get_all_models_with_overrides( """ from sqlmodel import select - from .payment.models import _row_to_model + from ..payment.models import _row_to_model async with create_session() as session: result = await session.exec(select(ModelRow).where(ModelRow.enabled)) @@ -126,7 +126,7 @@ async def get_all_models_with_overrides( async def refresh_upstreams_models_periodically( - upstreams: list[UpstreamProvider], + upstreams: list[BaseUpstreamProvider], ) -> None: """Background task to periodically refresh models cache for all providers. @@ -136,7 +136,7 @@ async def refresh_upstreams_models_periodically( import asyncio import random - from .core.settings import settings + from ..core.settings import settings interval = getattr(settings, "models_refresh_interval_seconds", 0) if not interval or interval <= 0: @@ -168,7 +168,7 @@ async def refresh_upstreams_models_periodically( break -async def init_upstreams() -> list[UpstreamProvider]: +async def init_upstreams() -> list[BaseUpstreamProvider]: """Initialize upstream providers from database. Seeds database with providers from settings if empty, then loads and instantiates @@ -176,7 +176,7 @@ async def init_upstreams() -> list[UpstreamProvider]: """ from sqlmodel import select - from .core.settings import settings + from ..core.settings import settings async with create_session() as session: result = await session.exec(select(UpstreamProviderRow)) @@ -191,7 +191,7 @@ async def init_upstreams() -> list[UpstreamProvider]: result = await session.exec(select(UpstreamProviderRow)) existing_providers = result.all() - upstreams: list[UpstreamProvider] = [] + upstreams: list[BaseUpstreamProvider] = [] for provider_row in existing_providers: if not provider_row.enabled: logger.debug(f"Skipping disabled provider: {provider_row.base_url}") @@ -222,7 +222,7 @@ async def _seed_providers_from_settings( """ from sqlmodel import select - from .core.settings import settings + from ..core.settings import settings providers_to_add: list[UpstreamProviderRow] = [] seeded_base_urls: set[str] = set() @@ -371,7 +371,9 @@ async def _seed_providers_from_settings( ) -def _instantiate_provider(provider_row: UpstreamProviderRow) -> UpstreamProvider | None: +def _instantiate_provider( + provider_row: UpstreamProviderRow, +) -> BaseUpstreamProvider | None: """Instantiate an UpstreamProvider from a database row. Args: @@ -418,7 +420,7 @@ def _instantiate_provider(provider_row: UpstreamProviderRow) -> UpstreamProvider provider_row.provider_type, ) elif provider_row.provider_type == "custom": - return UpstreamProvider( + return BaseUpstreamProvider( provider_row.base_url, provider_row.api_key, provider_row.provider_fee ) else: diff --git a/routstr/upstreams/ollama.py b/routstr/upstream/ollama.py similarity index 99% rename from routstr/upstreams/ollama.py rename to routstr/upstream/ollama.py index 117b95a7..c3b0fc9e 100644 --- a/routstr/upstreams/ollama.py +++ b/routstr/upstream/ollama.py @@ -6,7 +6,7 @@ import httpx from fastapi import Request from fastapi.responses import Response, StreamingResponse -from .upstream import UpstreamProvider +from .base import BaseUpstreamProvider if TYPE_CHECKING: from ..core.db import ApiKey, AsyncSession @@ -17,7 +17,7 @@ from ..core.logging import get_logger logger = get_logger(__name__) -class OllamaUpstreamProvider(UpstreamProvider): +class OllamaUpstreamProvider(BaseUpstreamProvider): """Upstream provider specifically configured for Ollama API.""" def __init__( diff --git a/routstr/upstream/openai.py b/routstr/upstream/openai.py new file mode 100644 index 00000000..3a0482d9 --- /dev/null +++ b/routstr/upstream/openai.py @@ -0,0 +1,23 @@ +from ..payment.models import Model, async_fetch_openrouter_models +from .base import BaseUpstreamProvider + + +class OpenAIUpstreamProvider(BaseUpstreamProvider): + """Upstream provider specifically configured for OpenAI API.""" + + 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, + ) + + def transform_model_name(self, model_id: str) -> str: + """Strip 'openai/' prefix for OpenAI API compatibility.""" + return model_id.removeprefix("openai/") + + async def fetch_models(self) -> list[Model]: + """Fetch OpenAI models from OpenRouter API filtered by openai source.""" + models_data = await async_fetch_openrouter_models(source_filter="openai") + return [Model(**model) for model in models_data] # type: ignore diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py new file mode 100644 index 00000000..956f4e15 --- /dev/null +++ b/routstr/upstream/openrouter.py @@ -0,0 +1,25 @@ +from ..payment.models import Model, async_fetch_openrouter_models +from .base import BaseUpstreamProvider + + +class OpenRouterUpstreamProvider(BaseUpstreamProvider): + """Upstream provider specifically configured for OpenRouter API.""" + + def __init__(self, api_key: str, provider_fee: float = 1.06): + """Initialize OpenRouter provider with API key. + + Args: + 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, + ) + + async def fetch_models(self) -> list[Model]: + """Fetch all OpenRouter models.""" + models_data = await async_fetch_openrouter_models() + return [Model(**model) for model in models_data] # type: ignore diff --git a/routstr/upstream/perplexity.py b/routstr/upstream/perplexity.py new file mode 100644 index 00000000..775a0229 --- /dev/null +++ b/routstr/upstream/perplexity.py @@ -0,0 +1,26 @@ +from ..payment.models import Model, async_fetch_openrouter_models +from .base import BaseUpstreamProvider + + +class PerplexityUpstreamProvider(BaseUpstreamProvider): + """Upstream provider specifically configured for OpenAI API.""" + + upstream_name = "perplexity" + base_url = "https://api.perplexity.ai/" # without v1 + 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, + api_key=api_key, + provider_fee=provider_fee, + ) + + def transform_model_name(self, model_id: str) -> str: + """Strip 'perplexity/' prefix for Perplexity API compatibility.""" + return model_id.removeprefix("perplexity/") + + async def fetch_models(self) -> list[Model]: + """Fetch Perplexity models from OpenRouter API filtered by perplexity source.""" + models_data = await async_fetch_openrouter_models(source_filter="perplexity") + return [Model(**model) for model in models_data] # type: ignore diff --git a/routstr/upstream/xai.py b/routstr/upstream/xai.py new file mode 100644 index 00000000..059b0dc3 --- /dev/null +++ b/routstr/upstream/xai.py @@ -0,0 +1,24 @@ +from ..payment.models import Model, async_fetch_openrouter_models +from .base import BaseUpstreamProvider + + +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" + + 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 + ) + + def transform_model_name(self, model_id: str) -> str: + """Strip 'xai/' prefix for XAI API compatibility.""" + return model_id.removeprefix("xai/") + + async def fetch_models(self) -> list[Model]: + """Fetch XAI models from OpenRouter API filtered by xai source.""" + models_data = await async_fetch_openrouter_models(source_filter="xai") + return [Model(**model) for model in models_data] # type: ignore diff --git a/routstr/upstreams/__init__.py b/routstr/upstreams/__init__.py deleted file mode 100644 index 397c0828..00000000 --- a/routstr/upstreams/__init__.py +++ /dev/null @@ -1,17 +0,0 @@ -from .ollama import OllamaUpstreamProvider -from .upstream import ( - AnthropicUpstreamProvider, - AzureUpstreamProvider, - OpenAIUpstreamProvider, - OpenRouterUpstreamProvider, - UpstreamProvider, -) - -__all__ = [ - "OllamaUpstreamProvider", - "UpstreamProvider", - "AnthropicUpstreamProvider", - "AzureUpstreamProvider", - "OpenAIUpstreamProvider", - "OpenRouterUpstreamProvider", -]