From d9d2e17e5d44f565761707c356986eab57523a1f Mon Sep 17 00:00:00 2001 From: Shroominic Date: Thu, 6 Nov 2025 17:23:37 +0800 Subject: [PATCH 1/4] refactor upstream files and classes --- routstr/algorithm.py | 22 ++-- routstr/proxy.py | 10 +- routstr/upstream/__init__.py | 35 ++++++ routstr/upstream/anthropic.py | 23 ++++ routstr/upstream/azure.py | 46 ++++++++ .../upstream.py => upstream/base.py} | 111 +----------------- routstr/{upstreams => upstream}/generic.py | 4 +- routstr/{upstream.py => upstream/helpers.py} | 40 ++++--- routstr/{upstreams => upstream}/ollama.py | 4 +- routstr/upstream/openai.py | 23 ++++ routstr/upstream/openrouter.py | 25 ++++ routstr/upstreams/__init__.py | 17 --- 12 files changed, 194 insertions(+), 166 deletions(-) create mode 100644 routstr/upstream/__init__.py create mode 100644 routstr/upstream/anthropic.py create mode 100644 routstr/upstream/azure.py rename routstr/{upstreams/upstream.py => upstream/base.py} (94%) rename routstr/{upstreams => upstream}/generic.py (98%) rename routstr/{upstream.py => upstream/helpers.py} (94%) rename routstr/{upstreams => upstream}/ollama.py (99%) create mode 100644 routstr/upstream/openai.py create mode 100644 routstr/upstream/openrouter.py delete mode 100644 routstr/upstreams/__init__.py diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 66f25d2b..88e0282d 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..36702919 --- /dev/null +++ b/routstr/upstream/anthropic.py @@ -0,0 +1,23 @@ +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.""" + return model_id.removeprefix("anthropic/") + + 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") + return [Model(**model) for model in models_data] # type: ignore 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 94% rename from routstr/upstreams/upstream.py rename to routstr/upstream/base.py index 9d3b1e35..73560ef6 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 @@ -1702,111 +1701,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.""" - return model_id.removeprefix("anthropic/") - - 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") - return [Model(**model) for model in models_data] # type: ignore - - -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/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.py b/routstr/upstream/helpers.py similarity index 94% rename from routstr/upstream.py rename to routstr/upstream/helpers.py index c1744585..8ae135d4 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__) @@ -70,7 +70,7 @@ def resolve_model_alias(model_id: str, canonical_slug: str | None = None) -> lis 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. @@ -85,7 +85,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)) @@ -122,7 +122,7 @@ async def get_all_models_with_overrides( async def get_model_with_override( model_id: str, - upstreams: list[UpstreamProvider], + upstreams: list[BaseUpstreamProvider], session: AsyncSession, ) -> Model | None: """Get a specific model from providers with database override applied. @@ -138,7 +138,7 @@ async def get_model_with_override( """ from sqlmodel import select - from .payment.models import _row_to_model + from ..payment.models import _row_to_model aliases = resolve_model_alias(model_id) @@ -170,7 +170,7 @@ async def get_model_with_override( async def refresh_upstreams_models_periodically( - upstreams: list[UpstreamProvider], + upstreams: list[BaseUpstreamProvider], ) -> None: """Background task to periodically refresh models cache for all providers. @@ -180,7 +180,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: @@ -212,7 +212,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 @@ -220,7 +220,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)) @@ -235,7 +235,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}") @@ -266,7 +266,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() @@ -415,7 +415,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: @@ -462,7 +464,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/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", -] From c9f458b8ba314b6e12d6e91e07ad9aebc003cbe3 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Thu, 6 Nov 2025 18:27:17 +0800 Subject: [PATCH 2/4] experimentation to get better fetching algorighm --- model_fetching_experiments.py | 58 +++++++++++++++++++++++++++++++++++ 1 file changed, 58 insertions(+) create mode 100644 model_fetching_experiments.py diff --git a/model_fetching_experiments.py b/model_fetching_experiments.py new file mode 100644 index 00000000..a9ab46a7 --- /dev/null +++ b/model_fetching_experiments.py @@ -0,0 +1,58 @@ +import asyncio +import os + +import httpx + + +async def fetch_models(base_url: str, api_key: str | None = None) -> dict: + url = f"{base_url.rstrip('/')}/models" + headers = {"Authorization": f"Bearer {api_key}"} if api_key else None + async with httpx.AsyncClient(timeout=30.0) as client: + response = await client.get(url, headers=headers) + response.raise_for_status() + return response.json() + + +def parse_model_ids(response: dict) -> list[str]: + return [model.get("id") for model in response.get("data", []) if "id" in model] + + +if __name__ == "__main__": + api_key = os.environ["API_KEY"] + base_url = os.environ.get("API_BASE_URL", "https://api.openai.com/v1") + + async def main() -> None: + or_models_response, provider_models_response = await asyncio.gather( + fetch_models("https://openrouter.ai/api/v1"), + fetch_models(base_url, api_key), + ) + provider_model_ids = parse_model_ids(provider_models_response) + or_models = or_models_response.get("data", []) + + found_models = [] + not_found_models = [] + + for model_id in provider_model_ids: + model = next( + ( + model + for model in or_models + if (model.get("id") == model_id) + or (model.get("id").split("/")[-1] == model_id) + or (model.get("canonical_slug") == model_id) + or (model.get("canonical_slug").split("/")[-1] == model_id) + ), + None, + ) + if model: + found_models.append(model) + else: + not_found_models.append(model_id) + print("\nFound models:") + for model in found_models: + print(model.get("id").split("/")[-1], model.get("pricing").get("prompt")) + print("\nNot found models:") + for model in not_found_models: + print(model) + + asyncio.run(main()) From 4b935a6f4d8fafb38dbe8655dc7341f05e35c42a Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sat, 8 Nov 2025 12:30:12 +0800 Subject: [PATCH 3/4] more upstream + model fetchin wip --- model_fetching_exp.ipynb | 174 +++++++++++++++++++++++++++++++++ routstr/upstream/fireworks.py | 24 +++++ routstr/upstream/groq.py | 24 +++++ routstr/upstream/perplexity.py | 26 +++++ routstr/upstream/xai.py | 24 +++++ 5 files changed, 272 insertions(+) create mode 100644 model_fetching_exp.ipynb create mode 100644 routstr/upstream/fireworks.py create mode 100644 routstr/upstream/groq.py create mode 100644 routstr/upstream/perplexity.py create mode 100644 routstr/upstream/xai.py diff --git a/model_fetching_exp.ipynb b/model_fetching_exp.ipynb new file mode 100644 index 00000000..10904175 --- /dev/null +++ b/model_fetching_exp.ipynb @@ -0,0 +1,174 @@ +{ + "cells": [ + { + "cell_type": "code", + "execution_count": 20, + "id": "e9b0df81", + "metadata": {}, + "outputs": [], + "source": [ + "import asyncio\n", + "import os\n", + "\n", + "import httpx\n", + "\n", + "\n", + "async def fetch_models(base_url: str, api_key: str | None = None) -> dict:\n", + " url = f\"{base_url.rstrip('/')}/models\"\n", + " headers = {\"Authorization\": f\"Bearer {api_key}\"} if api_key else None\n", + " async with httpx.AsyncClient(timeout=30.0) as client:\n", + " response = await client.get(url, headers=headers)\n", + " response.raise_for_status()\n", + " return response.json()\n", + "\n", + "\n", + "def parse_model_ids(response: dict) -> list[str]:\n", + " return [model.get(\"id\") for model in response.get(\"data\", []) if \"id\" in model]\n" + ] + }, + { + "cell_type": "markdown", + "id": "7f52e59d", + "metadata": {}, + "source": [ + "Notes:\n", + "\n", + "- reasoning: enabled \n", + "- sometimes different pricing for +128k context\n" + ] + }, + { + "cell_type": "code", + "execution_count": 21, + "id": "dac611eb", + "metadata": {}, + "outputs": [], + "source": [ + "api_key = \"xai-5cOB2MswjBKa9DEDU3gbCo02XVUUAC0nY6NYJ6Df0EYhL6PXoz2PEI9fTi4ilCUnEW9zNPbfzedrMgG9\"\n", + "base_url = \"https://api.x.ai/v1\"\n", + "provider_id: str | None = \"x-ai\"" + ] + }, + { + "cell_type": "code", + "execution_count": 25, + "id": "72968c71", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "openrouter models: 340\n", + "provider models: 9\n", + "filtered models: 7\n" + ] + } + ], + "source": [ + "openrouter_models_response = await fetch_models(\"https://openrouter.ai/api/v1\")\n", + "openrouter_models_list = openrouter_models_response.get(\"data\", [])\n", + "print(\"openrouter models:\", len(openrouter_models_list))\n", + "\n", + "provider_models_response = await fetch_models(base_url, api_key)\n", + "provider_model_ids = parse_model_ids(provider_models_response)\n", + "print(\"provider models:\", len(provider_model_ids))\n", + "\n", + "if provider_id:\n", + " openrouter_models_list = [model for model in openrouter_models_list if model.get(\"id\").startswith(provider_id)]\n", + " print(\"filtered models:\", len(openrouter_models_list))\n" + ] + }, + { + "cell_type": "code", + "execution_count": 28, + "id": "d5ef841d", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "\n", + "Found models:\n", + "grok-3 x-ai/grok-3\n", + "grok-3-mini x-ai/grok-3-mini\n", + "grok-code-fast-1 x-ai/grok-code-fast-1\n", + "\n", + "Not found models:\n", + "grok-2-1212\n", + "grok-2-vision-1212\n", + "grok-4-0709\n", + "grok-4-fast-non-reasoning\n", + "grok-4-fast-reasoning\n", + "grok-2-image-1212\n", + "\n", + "Unmatched models:\n", + "x-ai/grok-4-fast x-ai/grok-4-fast\n", + "x-ai/grok-4 x-ai/grok-4-07-09\n", + "x-ai/grok-3-mini-beta x-ai/grok-3-mini-beta\n", + "x-ai/grok-3-beta x-ai/grok-3-beta\n" + ] + } + ], + "source": [ + "models_map: dict[str, dict | None] = {}\n", + "\n", + "for model_id in provider_model_ids:\n", + " model = next(\n", + " (\n", + " model\n", + " for model in openrouter_models_list\n", + " if (model.get(\"id\") == model_id)\n", + " or (model.get(\"id\").split(\"/\")[-1] == model_id)\n", + " or (model.get(\"canonical_slug\") == model_id)\n", + " or (model.get(\"canonical_slug\").split(\"/\")[-1] == model_id)\n", + " ),\n", + " None,\n", + " )\n", + " if model:\n", + " models_map[model_id] = model\n", + " else:\n", + " models_map[model_id] = None\n", + "\n", + "print(\"\\nFound models:\")\n", + "for model_id, or_model in models_map.items():\n", + " if or_model:\n", + " print(model_id, or_model.get(\"id\"))\n", + "\n", + "print(\"\\nNot found models:\")\n", + "for model_id, or_model in models_map.items():\n", + " if not or_model:\n", + " print(model_id)\n", + "\n", + "print(\"\\nUnmatched models:\")\n", + "for model in openrouter_models_list:\n", + " if model.get(\"id\").split(\"/\")[-1] not in models_map:\n", + " print(model.get(\"id\"), model.get(\"canonical_slug\"))\n", + "\n", + "\n" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": ".venv", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.11.11" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} 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/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/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 From 2c8ba93312b6b167be4e37dad6eeec1fd7070e47 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Tue, 11 Nov 2025 10:24:34 +0800 Subject: [PATCH 4/4] experiments --- model_fetching_exp.ipynb | 2 -- 1 file changed, 2 deletions(-) diff --git a/model_fetching_exp.ipynb b/model_fetching_exp.ipynb index 10904175..9d69ade2 100644 --- a/model_fetching_exp.ipynb +++ b/model_fetching_exp.ipynb @@ -7,8 +7,6 @@ "metadata": {}, "outputs": [], "source": [ - "import asyncio\n", - "import os\n", "\n", "import httpx\n", "\n",