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:
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)
+5 -5
View File
@@ -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)
+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,
_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
+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
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__(
+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
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:
@@ -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__(
+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",
]