mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
Merge branch 'upstream-refactor' into v0.2.0
This commit is contained in:
+11
-11
@@ -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
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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",
|
||||||
|
]
|
||||||
@@ -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
|
||||||
@@ -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
|
|
||||||
@@ -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__(
|
||||||
@@ -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__(
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -1,17 +0,0 @@
|
|||||||
from .ollama import OllamaUpstreamProvider
|
|
||||||
from .upstream import (
|
|
||||||
AnthropicUpstreamProvider,
|
|
||||||
AzureUpstreamProvider,
|
|
||||||
OpenAIUpstreamProvider,
|
|
||||||
OpenRouterUpstreamProvider,
|
|
||||||
UpstreamProvider,
|
|
||||||
)
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"OllamaUpstreamProvider",
|
|
||||||
"UpstreamProvider",
|
|
||||||
"AnthropicUpstreamProvider",
|
|
||||||
"AzureUpstreamProvider",
|
|
||||||
"OpenAIUpstreamProvider",
|
|
||||||
"OpenRouterUpstreamProvider",
|
|
||||||
]
|
|
||||||
Reference in New Issue
Block a user