Merge branch 'upstream-refactor' into v0.2.0

This commit is contained in:
Shroominic
2025-11-11 13:09:02 +08:00
16 changed files with 310 additions and 184 deletions
+11 -11
View File
@@ -6,7 +6,7 @@ from .core.logging import get_logger
if TYPE_CHECKING: if TYPE_CHECKING:
from .payment.models import Model from .payment.models import Model
from .upstream import UpstreamProvider from .upstream import BaseUpstreamProvider
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -59,7 +59,7 @@ def calculate_model_cost_score(model: "Model") -> float:
return total_cost return total_cost
def get_provider_penalty(provider: "UpstreamProvider") -> float: def get_provider_penalty(provider: "BaseUpstreamProvider") -> float:
"""Calculate a penalty multiplier for certain providers. """Calculate a penalty multiplier for certain providers.
This allows applying policy-based adjustments beyond pure cost. This allows applying policy-based adjustments beyond pure cost.
@@ -86,9 +86,9 @@ def get_provider_penalty(provider: "UpstreamProvider") -> float:
def should_prefer_model( def should_prefer_model(
candidate_model: "Model", candidate_model: "Model",
candidate_provider: "UpstreamProvider", candidate_provider: "BaseUpstreamProvider",
current_model: "Model", current_model: "Model",
current_provider: "UpstreamProvider", current_provider: "BaseUpstreamProvider",
alias: str, alias: str,
) -> bool: ) -> bool:
"""Determine if candidate model should replace current model for an alias. """Determine if candidate model should replace current model for an alias.
@@ -166,10 +166,10 @@ def should_prefer_model(
def create_model_mappings( def create_model_mappings(
upstreams: list["UpstreamProvider"], upstreams: list["BaseUpstreamProvider"],
overrides_by_id: dict[str, tuple], overrides_by_id: dict[str, tuple],
disabled_model_ids: set[str], 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. """Create optimal model mappings based on cost and provider preferences.
This is the main entry point for the algorithm. It processes all upstream providers 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 from .upstream import resolve_model_alias
model_instances: dict[str, "Model"] = {} model_instances: dict[str, "Model"] = {}
provider_map: dict[str, "UpstreamProvider"] = {} provider_map: dict[str, "BaseUpstreamProvider"] = {}
unique_models: dict[str, "Model"] = {} unique_models: dict[str, "Model"] = {}
# Separate OpenRouter from other providers # Separate OpenRouter from other providers
openrouter: "UpstreamProvider" | None = None openrouter: "BaseUpstreamProvider" | None = None
other_upstreams: list["UpstreamProvider"] = [] other_upstreams: list["BaseUpstreamProvider"] = []
for upstream in upstreams: for upstream in upstreams:
base_url = getattr(upstream, "base_url", "") 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 return model_id.split("/", 1)[1] if "/" in model_id else model_id
def _maybe_set_alias( def _maybe_set_alias(
alias: str, model: "Model", provider: "UpstreamProvider" alias: str, model: "Model", provider: "BaseUpstreamProvider"
) -> None: ) -> None:
"""Set alias to model/provider if not set or if new model is preferred.""" """Set alias to model/provider if not set or if new model is preferred."""
existing_model = model_instances.get(alias) existing_model = model_instances.get(alias)
@@ -233,7 +233,7 @@ def create_model_mappings(
provider_map[alias] = provider provider_map[alias] = provider
def process_provider_models( def process_provider_models(
upstream: "UpstreamProvider", is_openrouter: bool = False upstream: "BaseUpstreamProvider", is_openrouter: bool = False
) -> None: ) -> None:
"""Process all models from a given provider.""" """Process all models from a given provider."""
upstream_prefix = getattr(upstream, "upstream_name", None) upstream_prefix = getattr(upstream, "upstream_name", None)
+5 -5
View File
@@ -23,14 +23,14 @@ from .payment.helpers import (
get_max_cost_for_model, get_max_cost_for_model,
) )
from .payment.models import Model from .payment.models import Model
from .upstream import UpstreamProvider, init_upstreams from .upstream import BaseUpstreamProvider, init_upstreams
logger = get_logger(__name__) logger = get_logger(__name__)
proxy_router = APIRouter() proxy_router = APIRouter()
_upstreams: list[UpstreamProvider] = [] _upstreams: list[BaseUpstreamProvider] = []
_model_instances: dict[str, Model] = {} # All aliases -> Model _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) _unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
@@ -53,7 +53,7 @@ async def reinitialize_upstreams() -> None:
await refresh_model_maps() await refresh_model_maps()
def get_upstreams() -> list[UpstreamProvider]: def get_upstreams() -> list[BaseUpstreamProvider]:
"""Get the initialized upstream providers. """Get the initialized upstream providers.
Returns: Returns:
@@ -67,7 +67,7 @@ def get_model_instance(model_id: str) -> Model | None:
return _model_instances.get(model_id) 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.""" """Get UpstreamProvider for model ID from global cache."""
return _provider_map.get(model_id) return _provider_map.get(model_id)
+35
View File
@@ -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",
]
+43
View File
@@ -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
+46
View File
@@ -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
@@ -25,7 +25,6 @@ from ..payment.models import (
Pricing, Pricing,
_calculate_usd_max_costs, _calculate_usd_max_costs,
_update_model_sats_pricing, _update_model_sats_pricing,
async_fetch_openrouter_models,
) )
from ..payment.price import sats_usd_price from ..payment.price import sats_usd_price
from ..wallet import recieve_token, send_token from ..wallet import recieve_token, send_token
@@ -33,7 +32,7 @@ from ..wallet import recieve_token, send_token
logger = get_logger(__name__) logger = get_logger(__name__)
class UpstreamProvider: class BaseUpstreamProvider:
"""Provider for forwarding requests to an upstream AI service API.""" """Provider for forwarding requests to an upstream AI service API."""
base_url: str base_url: str
@@ -1704,131 +1703,3 @@ class UpstreamProvider:
Model object or None if not found Model object or None if not found
""" """
return self._models_by_id.get(model_id) 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
+24
View File
@@ -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
@@ -4,7 +4,7 @@ from typing import TYPE_CHECKING
import httpx import httpx
from .upstream import UpstreamProvider from .base import BaseUpstreamProvider
if TYPE_CHECKING: if TYPE_CHECKING:
from ..payment.models import Model from ..payment.models import Model
@@ -14,7 +14,7 @@ from ..core.logging import get_logger
logger = get_logger(__name__) logger = get_logger(__name__)
class GenericUpstreamProvider(UpstreamProvider): class GenericUpstreamProvider(BaseUpstreamProvider):
"""Generic upstream provider that can fetch models from any OpenAI-compatible API.""" """Generic upstream provider that can fetch models from any OpenAI-compatible API."""
def __init__( def __init__(
+24
View File
@@ -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
@@ -5,21 +5,21 @@ import re
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
if TYPE_CHECKING: if TYPE_CHECKING:
from .core.settings import Settings from ..core.settings import Settings
from .core import get_logger from ..core import get_logger
from .core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session
from .payment.models import Model from ..payment.models import Model
from .upstreams import ( from . import (
AnthropicUpstreamProvider, AnthropicUpstreamProvider,
AzureUpstreamProvider, AzureUpstreamProvider,
BaseUpstreamProvider,
GenericUpstreamProvider,
OllamaUpstreamProvider, OllamaUpstreamProvider,
OpenAIUpstreamProvider, OpenAIUpstreamProvider,
OpenRouterUpstreamProvider, OpenRouterUpstreamProvider,
UpstreamProvider,
) )
from .upstreams.generic import GenericUpstreamProvider
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -75,7 +75,7 @@ def resolve_model_alias(
async def get_all_models_with_overrides( async def get_all_models_with_overrides(
upstreams: list[UpstreamProvider], upstreams: list[BaseUpstreamProvider],
) -> list[Model]: ) -> list[Model]:
"""Get all models from all providers with database overrides applied. """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 sqlmodel import select
from .payment.models import _row_to_model from ..payment.models import _row_to_model
async with create_session() as session: async with create_session() as session:
result = await session.exec(select(ModelRow).where(ModelRow.enabled)) 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( async def refresh_upstreams_models_periodically(
upstreams: list[UpstreamProvider], upstreams: list[BaseUpstreamProvider],
) -> None: ) -> None:
"""Background task to periodically refresh models cache for all providers. """Background task to periodically refresh models cache for all providers.
@@ -136,7 +136,7 @@ async def refresh_upstreams_models_periodically(
import asyncio import asyncio
import random import random
from .core.settings import settings from ..core.settings import settings
interval = getattr(settings, "models_refresh_interval_seconds", 0) interval = getattr(settings, "models_refresh_interval_seconds", 0)
if not interval or interval <= 0: if not interval or interval <= 0:
@@ -168,7 +168,7 @@ async def refresh_upstreams_models_periodically(
break break
async def init_upstreams() -> list[UpstreamProvider]: async def init_upstreams() -> list[BaseUpstreamProvider]:
"""Initialize upstream providers from database. """Initialize upstream providers from database.
Seeds database with providers from settings if empty, then loads and instantiates 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 sqlmodel import select
from .core.settings import settings from ..core.settings import settings
async with create_session() as session: async with create_session() as session:
result = await session.exec(select(UpstreamProviderRow)) result = await session.exec(select(UpstreamProviderRow))
@@ -191,7 +191,7 @@ async def init_upstreams() -> list[UpstreamProvider]:
result = await session.exec(select(UpstreamProviderRow)) result = await session.exec(select(UpstreamProviderRow))
existing_providers = result.all() existing_providers = result.all()
upstreams: list[UpstreamProvider] = [] upstreams: list[BaseUpstreamProvider] = []
for provider_row in existing_providers: for provider_row in existing_providers:
if not provider_row.enabled: if not provider_row.enabled:
logger.debug(f"Skipping disabled provider: {provider_row.base_url}") 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 sqlmodel import select
from .core.settings import settings from ..core.settings import settings
providers_to_add: list[UpstreamProviderRow] = [] providers_to_add: list[UpstreamProviderRow] = []
seeded_base_urls: set[str] = set() 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. """Instantiate an UpstreamProvider from a database row.
Args: Args:
@@ -418,7 +420,7 @@ def _instantiate_provider(provider_row: UpstreamProviderRow) -> UpstreamProvider
provider_row.provider_type, provider_row.provider_type,
) )
elif provider_row.provider_type == "custom": elif provider_row.provider_type == "custom":
return UpstreamProvider( return BaseUpstreamProvider(
provider_row.base_url, provider_row.api_key, provider_row.provider_fee provider_row.base_url, provider_row.api_key, provider_row.provider_fee
) )
else: else:
@@ -6,7 +6,7 @@ import httpx
from fastapi import Request from fastapi import Request
from fastapi.responses import Response, StreamingResponse from fastapi.responses import Response, StreamingResponse
from .upstream import UpstreamProvider from .base import BaseUpstreamProvider
if TYPE_CHECKING: if TYPE_CHECKING:
from ..core.db import ApiKey, AsyncSession from ..core.db import ApiKey, AsyncSession
@@ -17,7 +17,7 @@ from ..core.logging import get_logger
logger = get_logger(__name__) logger = get_logger(__name__)
class OllamaUpstreamProvider(UpstreamProvider): class OllamaUpstreamProvider(BaseUpstreamProvider):
"""Upstream provider specifically configured for Ollama API.""" """Upstream provider specifically configured for Ollama API."""
def __init__( def __init__(
+23
View File
@@ -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
+25
View File
@@ -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
+26
View File
@@ -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
+24
View File
@@ -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
-17
View File
@@ -1,17 +0,0 @@
from .ollama import OllamaUpstreamProvider
from .upstream import (
AnthropicUpstreamProvider,
AzureUpstreamProvider,
OpenAIUpstreamProvider,
OpenRouterUpstreamProvider,
UpstreamProvider,
)
__all__ = [
"OllamaUpstreamProvider",
"UpstreamProvider",
"AnthropicUpstreamProvider",
"AzureUpstreamProvider",
"OpenAIUpstreamProvider",
"OpenRouterUpstreamProvider",
]