mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
auto populate providers from available classes
This commit is contained in:
@@ -193,7 +193,7 @@ def create_model_mappings(
|
||||
Tuple of (model_instances, provider_map, unique_models)
|
||||
"""
|
||||
from .payment.models import _row_to_model
|
||||
from .upstream import resolve_model_alias
|
||||
from .upstream.helpers import resolve_model_alias
|
||||
|
||||
model_instances: dict[str, "Model"] = {}
|
||||
provider_map: dict[str, "BaseUpstreamProvider"] = {}
|
||||
|
||||
+4
-40
@@ -2739,45 +2739,9 @@ async def delete_upstream_provider(provider_id: int) -> dict[str, object]:
|
||||
@admin_router.get("/api/provider-types", dependencies=[Depends(require_admin_api)])
|
||||
async def get_provider_types() -> list[dict[str, object]]:
|
||||
"""Get metadata about available provider types including default URLs and whether they're fixed."""
|
||||
provider_types = [
|
||||
{
|
||||
"id": "openrouter",
|
||||
"name": "OpenRouter",
|
||||
"default_base_url": "https://openrouter.ai/api/v1",
|
||||
"fixed_base_url": True,
|
||||
},
|
||||
{
|
||||
"id": "openai",
|
||||
"name": "OpenAI",
|
||||
"default_base_url": "https://api.openai.com/v1",
|
||||
"fixed_base_url": True,
|
||||
},
|
||||
{
|
||||
"id": "anthropic",
|
||||
"name": "Anthropic",
|
||||
"default_base_url": "https://api.anthropic.com/v1",
|
||||
"fixed_base_url": True,
|
||||
},
|
||||
{
|
||||
"id": "azure",
|
||||
"name": "Azure OpenAI",
|
||||
"default_base_url": "",
|
||||
"fixed_base_url": False,
|
||||
},
|
||||
{
|
||||
"id": "ollama",
|
||||
"name": "Ollama",
|
||||
"default_base_url": "http://localhost:11434",
|
||||
"fixed_base_url": False,
|
||||
},
|
||||
{
|
||||
"id": "generic",
|
||||
"name": "Generic",
|
||||
"default_base_url": "",
|
||||
"fixed_base_url": False,
|
||||
},
|
||||
]
|
||||
return provider_types
|
||||
from ..upstream import upstream_provider_classes
|
||||
|
||||
return [cls.get_provider_metadata() for cls in upstream_provider_classes]
|
||||
|
||||
|
||||
@admin_router.get(
|
||||
@@ -2785,7 +2749,7 @@ async def get_provider_types() -> list[dict[str, object]]:
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def get_provider_models(provider_id: int) -> dict[str, object]:
|
||||
from ..upstream import _instantiate_provider
|
||||
from ..upstream.helpers import _instantiate_provider
|
||||
|
||||
async with create_session() as session:
|
||||
provider = await session.get(UpstreamProviderRow, provider_id)
|
||||
|
||||
@@ -78,7 +78,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
|
||||
from ..payment.price import _update_prices
|
||||
from ..proxy import get_upstreams
|
||||
from ..upstream import refresh_upstreams_models_periodically
|
||||
from ..upstream.helpers import refresh_upstreams_models_periodically
|
||||
|
||||
await _update_prices()
|
||||
await initialize_upstreams()
|
||||
|
||||
+2
-1
@@ -23,7 +23,8 @@ from .payment.helpers import (
|
||||
get_max_cost_for_model,
|
||||
)
|
||||
from .payment.models import Model
|
||||
from .upstream import BaseUpstreamProvider, init_upstreams
|
||||
from .upstream import BaseUpstreamProvider
|
||||
from .upstream.helpers import init_upstreams
|
||||
|
||||
logger = get_logger(__name__)
|
||||
proxy_router = APIRouter()
|
||||
|
||||
@@ -1,34 +1,31 @@
|
||||
from .anthropic import AnthropicUpstreamProvider
|
||||
from .azure import AzureUpstreamProvider
|
||||
from .base import BaseUpstreamProvider
|
||||
from .fireworks import FireworksUpstreamProvider
|
||||
from .generic import GenericUpstreamProvider
|
||||
from .helpers import (
|
||||
_instantiate_provider,
|
||||
_seed_providers_from_settings,
|
||||
get_all_models_with_overrides,
|
||||
init_upstreams,
|
||||
refresh_upstreams_models_periodically,
|
||||
resolve_model_alias,
|
||||
)
|
||||
from .groq import GroqUpstreamProvider
|
||||
from .ollama import OllamaUpstreamProvider
|
||||
from .openai import OpenAIUpstreamProvider
|
||||
from .openrouter import OpenRouterUpstreamProvider
|
||||
from .perplexity import PerplexityUpstreamProvider
|
||||
from .xai import XAIUpstreamProvider
|
||||
|
||||
upstream_provider_classes: list[type[BaseUpstreamProvider]] = [
|
||||
AnthropicUpstreamProvider,
|
||||
AzureUpstreamProvider,
|
||||
FireworksUpstreamProvider,
|
||||
GenericUpstreamProvider,
|
||||
GroqUpstreamProvider,
|
||||
OllamaUpstreamProvider,
|
||||
OpenAIUpstreamProvider,
|
||||
OpenRouterUpstreamProvider,
|
||||
PerplexityUpstreamProvider,
|
||||
XAIUpstreamProvider,
|
||||
]
|
||||
"""List of all upstream classes"""
|
||||
|
||||
__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",
|
||||
*[cls.__name__ for cls in upstream_provider_classes],
|
||||
"upstream_provider_classes",
|
||||
]
|
||||
|
||||
@@ -1,18 +1,45 @@
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ..payment.models import Model, async_fetch_openrouter_models
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
|
||||
|
||||
class AnthropicUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Upstream provider specifically configured for Anthropic API."""
|
||||
|
||||
provider_type = "anthropic"
|
||||
default_base_url = "https://api.anthropic.com/v1"
|
||||
platform_url = "https://console.anthropic.com/settings/keys"
|
||||
|
||||
def __init__(self, api_key: str, provider_fee: float = 1.01):
|
||||
self.upstream_name = "anthropic"
|
||||
super().__init__(
|
||||
base_url="https://api.anthropic.com/v1",
|
||||
base_url=self.default_base_url,
|
||||
api_key=api_key,
|
||||
provider_fee=provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_db_row(
|
||||
cls, provider_row: "UpstreamProviderRow"
|
||||
) -> "AnthropicUpstreamProvider":
|
||||
return cls(
|
||||
api_key=provider_row.api_key,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": "Anthropic",
|
||||
"default_base_url": cls.default_base_url,
|
||||
"fixed_base_url": True,
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
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/"):
|
||||
|
||||
@@ -1,11 +1,18 @@
|
||||
from typing import Mapping
|
||||
from typing import TYPE_CHECKING, Mapping
|
||||
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
|
||||
|
||||
class AzureUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Upstream provider specifically configured for Azure OpenAI Service."""
|
||||
|
||||
provider_type = "azure"
|
||||
default_base_url = None
|
||||
platform_url = "https://portal.azure.com/"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
@@ -28,6 +35,29 @@ class AzureUpstreamProvider(BaseUpstreamProvider):
|
||||
)
|
||||
self.api_version = api_version
|
||||
|
||||
@classmethod
|
||||
def from_db_row(
|
||||
cls, provider_row: "UpstreamProviderRow"
|
||||
) -> "AzureUpstreamProvider | None":
|
||||
if not provider_row.api_version:
|
||||
return None
|
||||
return cls(
|
||||
base_url=provider_row.base_url,
|
||||
api_key=provider_row.api_key,
|
||||
api_version=provider_row.api_version,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": "Azure OpenAI",
|
||||
"default_base_url": "",
|
||||
"fixed_base_url": False,
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
def prepare_params(
|
||||
self, path: str, query_params: Mapping[str, str] | None
|
||||
) -> Mapping[str, str]:
|
||||
|
||||
@@ -4,7 +4,7 @@ import json
|
||||
import re
|
||||
import traceback
|
||||
from collections.abc import AsyncGenerator
|
||||
from typing import Mapping
|
||||
from typing import TYPE_CHECKING, Mapping
|
||||
|
||||
import httpx
|
||||
from fastapi import BackgroundTasks, HTTPException, Request
|
||||
@@ -13,6 +13,10 @@ from fastapi.responses import Response, StreamingResponse
|
||||
from ..auth import adjust_payment_for_tokens
|
||||
from ..core import get_logger
|
||||
from ..core.db import ApiKey, AsyncSession, create_session
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
|
||||
from ..payment.cost_caculation import (
|
||||
CostData,
|
||||
CostDataError,
|
||||
@@ -35,9 +39,12 @@ logger = get_logger(__name__)
|
||||
class BaseUpstreamProvider:
|
||||
"""Provider for forwarding requests to an upstream AI service API."""
|
||||
|
||||
provider_type: str = "base"
|
||||
default_base_url: str | None = None
|
||||
platform_url: str | None = None
|
||||
|
||||
base_url: str
|
||||
api_key: str
|
||||
upstream_name: str | None = None
|
||||
provider_fee: float = 1.05
|
||||
_models_cache: list[Model] = []
|
||||
_models_by_id: dict[str, Model] = {}
|
||||
@@ -56,6 +63,39 @@ class BaseUpstreamProvider:
|
||||
self._models_cache = []
|
||||
self._models_by_id = {}
|
||||
|
||||
@classmethod
|
||||
def from_db_row(
|
||||
cls, provider_row: "UpstreamProviderRow"
|
||||
) -> "BaseUpstreamProvider | None":
|
||||
"""Factory method to instantiate provider from database row.
|
||||
|
||||
Args:
|
||||
provider_row: Database row containing provider configuration
|
||||
|
||||
Returns:
|
||||
Instantiated provider or None if instantiation fails
|
||||
"""
|
||||
return cls(
|
||||
base_url=provider_row.base_url,
|
||||
api_key=provider_row.api_key,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
"""Get metadata about this provider type for API responses.
|
||||
|
||||
Returns:
|
||||
Dict with provider type metadata including id, name, default_base_url, fixed_base_url, platform_url
|
||||
"""
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": cls.provider_type.title(),
|
||||
"default_base_url": cls.default_base_url or "",
|
||||
"fixed_base_url": bool(cls.default_base_url),
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
def prepare_headers(self, request_headers: dict) -> dict:
|
||||
"""Prepare headers for upstream request by removing proxy-specific headers and adding authentication.
|
||||
|
||||
@@ -162,7 +202,7 @@ class BaseUpstreamProvider:
|
||||
extra={
|
||||
"original": original_model,
|
||||
"transformed": transformed_model,
|
||||
"provider": self.upstream_name or self.base_url,
|
||||
"provider": self.provider_type or self.base_url,
|
||||
},
|
||||
)
|
||||
return json.dumps(data).encode()
|
||||
@@ -171,7 +211,7 @@ class BaseUpstreamProvider:
|
||||
"Could not transform request body",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"provider": self.upstream_name or self.base_url,
|
||||
"provider": self.provider_type or self.base_url,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -1657,7 +1697,7 @@ class BaseUpstreamProvider:
|
||||
Returns:
|
||||
List of Model objects with pricing
|
||||
"""
|
||||
logger.debug(f"Fetching models for {self.upstream_name or self.base_url}")
|
||||
logger.debug(f"Fetching models for {self.provider_type or self.base_url}")
|
||||
return []
|
||||
|
||||
async def refresh_models_cache(self) -> None:
|
||||
@@ -1676,12 +1716,12 @@ class BaseUpstreamProvider:
|
||||
|
||||
self._models_by_id = {m.id: m for m in self._models_cache}
|
||||
logger.info(
|
||||
f"Refreshed models cache for {self.upstream_name or self.base_url}",
|
||||
f"Refreshed models cache for {self.provider_type or self.base_url}",
|
||||
extra={"model_count": len(models)},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Failed to refresh models cache for {self.upstream_name or self.base_url}",
|
||||
f"Failed to refresh models cache for {self.provider_type or self.base_url}",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
|
||||
|
||||
@@ -1,19 +1,43 @@
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ..payment.models import Model, async_fetch_openrouter_models
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
|
||||
|
||||
class FireworksUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Upstream provider specifically configured for Fireworks.ai API."""
|
||||
|
||||
upstream_name = "fireworks"
|
||||
base_url = "https://api.fireworks.ai/inference/v1"
|
||||
provider_type = "fireworks"
|
||||
default_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
|
||||
base_url=self.default_base_url, api_key=api_key, provider_fee=provider_fee
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_db_row(
|
||||
cls, provider_row: "UpstreamProviderRow"
|
||||
) -> "FireworksUpstreamProvider":
|
||||
return cls(
|
||||
api_key=provider_row.api_key,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": "Fireworks",
|
||||
"default_base_url": cls.default_base_url,
|
||||
"fixed_base_url": True,
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
"""Strip 'fireworks/' prefix for Fireworks API compatibility."""
|
||||
return model_id.removeprefix("fireworks/")
|
||||
|
||||
@@ -7,6 +7,7 @@ import httpx
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
from ..payment.models import Model
|
||||
|
||||
from ..core.logging import get_logger
|
||||
@@ -17,6 +18,10 @@ logger = get_logger(__name__)
|
||||
class GenericUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Generic upstream provider that can fetch models from any OpenAI-compatible API."""
|
||||
|
||||
provider_type = "generic"
|
||||
default_base_url = "http://localhost:8888"
|
||||
platform_url = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
@@ -39,6 +44,26 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
|
||||
provider_fee=provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_db_row(
|
||||
cls, provider_row: "UpstreamProviderRow"
|
||||
) -> "GenericUpstreamProvider":
|
||||
return cls(
|
||||
base_url=provider_row.base_url,
|
||||
api_key=provider_row.api_key,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": "Generic",
|
||||
"default_base_url": cls.default_base_url,
|
||||
"fixed_base_url": False,
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
async def fetch_models(self) -> list[Model]:
|
||||
"""Fetch models from upstream API using /models endpoint."""
|
||||
from ..payment.models import Architecture, Model, Pricing, TopProvider
|
||||
|
||||
@@ -1,19 +1,41 @@
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ..payment.models import Model, async_fetch_openrouter_models
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
|
||||
|
||||
class GroqUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Upstream provider specifically configured for Groq API."""
|
||||
|
||||
upstream_name = "groq"
|
||||
base_url = "https://api.groq.com/openai/v1"
|
||||
provider_type = "groq"
|
||||
default_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
|
||||
base_url=self.default_base_url, api_key=api_key, provider_fee=provider_fee
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_db_row(cls, provider_row: "UpstreamProviderRow") -> "GroqUpstreamProvider":
|
||||
return cls(
|
||||
api_key=provider_row.api_key,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": "Groq",
|
||||
"default_base_url": cls.default_base_url,
|
||||
"fixed_base_url": True,
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
"""Strip 'groq/' prefix for Groq API compatibility."""
|
||||
return model_id.removeprefix("groq/")
|
||||
|
||||
+65
-131
@@ -11,13 +11,7 @@ if TYPE_CHECKING:
|
||||
from ..core import get_logger
|
||||
from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session
|
||||
from ..payment.models import Model
|
||||
from .anthropic import AnthropicUpstreamProvider
|
||||
from .azure import AzureUpstreamProvider
|
||||
from .base import BaseUpstreamProvider
|
||||
from .generic import GenericUpstreamProvider
|
||||
from .ollama import OllamaUpstreamProvider
|
||||
from .openai import OpenAIUpstreamProvider
|
||||
from .openrouter import OpenRouterUpstreamProvider
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
@@ -148,7 +142,7 @@ async def refresh_upstreams_models_periodically(
|
||||
await upstream.refresh_models_cache()
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error refreshing models for {upstream.upstream_name or upstream.base_url}",
|
||||
f"Error refreshing models for {upstream.base_url}",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
except asyncio.CancelledError:
|
||||
@@ -220,61 +214,47 @@ async def _seed_providers_from_settings(
|
||||
"""
|
||||
from sqlmodel import select
|
||||
|
||||
from ..core.settings import settings
|
||||
from . import upstream_provider_classes
|
||||
|
||||
providers_to_add: list[UpstreamProviderRow] = []
|
||||
seeded_base_urls: set[str] = set()
|
||||
|
||||
openai_api_key = os.environ.get("OPENAI_API_KEY")
|
||||
if openai_api_key:
|
||||
base_url = "https://api.openai.com/v1"
|
||||
result = await session.exec(
|
||||
select(UpstreamProviderRow).where(UpstreamProviderRow.base_url == base_url)
|
||||
)
|
||||
if not result.first():
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
provider_type="openai",
|
||||
base_url=base_url,
|
||||
api_key=openai_api_key,
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
seeded_base_urls.add(base_url)
|
||||
provider_classes_by_type = {
|
||||
cls.provider_type: cls
|
||||
for cls in upstream_provider_classes # type: ignore[attr-defined]
|
||||
}
|
||||
|
||||
anthropic_api_key = os.environ.get("ANTHROPIC_API_KEY")
|
||||
if anthropic_api_key:
|
||||
base_url = "https://api.anthropic.com/v1"
|
||||
result = await session.exec(
|
||||
select(UpstreamProviderRow).where(UpstreamProviderRow.base_url == base_url)
|
||||
)
|
||||
if not result.first():
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
provider_type="anthropic",
|
||||
base_url=base_url,
|
||||
api_key=anthropic_api_key,
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
seeded_base_urls.add(base_url)
|
||||
env_mappings: list[tuple[str, str, str | None, str | None]] = [
|
||||
("OPENAI_API_KEY", "openai", None, None),
|
||||
("ANTHROPIC_API_KEY", "anthropic", None, None),
|
||||
("OPENROUTER_API_KEY", "openrouter", None, None),
|
||||
("GROQ_API_KEY", "groq", None, None),
|
||||
("PERPLEXITY_API_KEY", "perplexity", None, None),
|
||||
("FIREWORKS_API_KEY", "fireworks", None, None),
|
||||
("XAI_API_KEY", "xai", None, None),
|
||||
]
|
||||
|
||||
openrouter_api_key = os.environ.get("OPENROUTER_API_KEY")
|
||||
if openrouter_api_key:
|
||||
base_url = "https://openrouter.ai/api/v1"
|
||||
result = await session.exec(
|
||||
select(UpstreamProviderRow).where(UpstreamProviderRow.base_url == base_url)
|
||||
)
|
||||
if not result.first():
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
provider_type="openrouter",
|
||||
base_url=base_url,
|
||||
api_key=openrouter_api_key,
|
||||
enabled=True,
|
||||
for env_key, provider_type, _, _ in env_mappings:
|
||||
api_key = os.environ.get(env_key)
|
||||
if api_key and provider_type in provider_classes_by_type:
|
||||
provider_class = provider_classes_by_type[provider_type]
|
||||
if provider_class.default_base_url: # type: ignore[attr-defined]
|
||||
base_url = provider_class.default_base_url # type: ignore[attr-defined]
|
||||
result = await session.exec(
|
||||
select(UpstreamProviderRow).where(
|
||||
UpstreamProviderRow.base_url == base_url
|
||||
)
|
||||
)
|
||||
)
|
||||
seeded_base_urls.add(base_url)
|
||||
if not result.first():
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
provider_type=provider_type,
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
seeded_base_urls.add(base_url)
|
||||
|
||||
ollama_base_url = os.environ.get("OLLAMA_BASE_URL")
|
||||
if ollama_base_url:
|
||||
@@ -323,48 +303,20 @@ async def _seed_providers_from_settings(
|
||||
)
|
||||
)
|
||||
if not result.first():
|
||||
if "api.openai.com" in base_url.lower():
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
provider_type="openai",
|
||||
base_url=base_url,
|
||||
api_key=settings.upstream_api_key,
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
elif "api.anthropic.com" in base_url.lower():
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
provider_type="anthropic",
|
||||
base_url=base_url,
|
||||
api_key=settings.upstream_api_key,
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
elif "openrouter.ai/api/v1" in base_url.lower():
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
provider_type="openrouter",
|
||||
base_url=base_url,
|
||||
api_key=settings.upstream_api_key,
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
else:
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
provider_type="custom",
|
||||
base_url=base_url,
|
||||
api_key=settings.upstream_api_key,
|
||||
enabled=True,
|
||||
)
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
provider_type="custom",
|
||||
base_url=base_url,
|
||||
api_key=settings.upstream_api_key,
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
seeded_base_urls.add(base_url)
|
||||
|
||||
for provider in providers_to_add:
|
||||
session.add(provider)
|
||||
logger.info(
|
||||
f"Seeding {provider.provider_type} provider",
|
||||
f"Seeding {provider.provider_type} provider", # type: ignore[str-format]
|
||||
extra={"base_url": provider.base_url},
|
||||
)
|
||||
|
||||
@@ -380,53 +332,35 @@ def _instantiate_provider(
|
||||
Returns:
|
||||
Instantiated provider or None if provider type is unknown
|
||||
"""
|
||||
from . import upstream_provider_classes
|
||||
|
||||
try:
|
||||
if provider_row.provider_type == "openai":
|
||||
return OpenAIUpstreamProvider(
|
||||
provider_row.api_key, provider_row.provider_fee
|
||||
)
|
||||
elif provider_row.provider_type == "anthropic":
|
||||
return AnthropicUpstreamProvider(
|
||||
provider_row.api_key, provider_row.provider_fee
|
||||
)
|
||||
elif provider_row.provider_type == "azure":
|
||||
if not provider_row.api_version:
|
||||
provider_classes_by_type = {
|
||||
cls.provider_type: cls
|
||||
for cls in upstream_provider_classes # type: ignore[attr-defined]
|
||||
}
|
||||
|
||||
provider_class = provider_classes_by_type.get(provider_row.provider_type)
|
||||
|
||||
if provider_class:
|
||||
provider = provider_class.from_db_row(provider_row) # type: ignore[attr-defined]
|
||||
if provider is None:
|
||||
logger.error(
|
||||
"Azure provider missing api_version",
|
||||
f"Failed to instantiate {provider_row.provider_type} provider",
|
||||
extra={"base_url": provider_row.base_url},
|
||||
)
|
||||
return None
|
||||
return AzureUpstreamProvider(
|
||||
provider_row.base_url,
|
||||
provider_row.api_key,
|
||||
provider_row.api_version,
|
||||
provider_row.provider_fee,
|
||||
)
|
||||
elif provider_row.provider_type == "openrouter":
|
||||
return OpenRouterUpstreamProvider(
|
||||
provider_row.api_key, provider_row.provider_fee
|
||||
)
|
||||
elif provider_row.provider_type == "ollama":
|
||||
return OllamaUpstreamProvider(
|
||||
provider_row.base_url, provider_row.api_key, provider_row.provider_fee
|
||||
)
|
||||
elif provider_row.provider_type == "generic":
|
||||
return GenericUpstreamProvider(
|
||||
provider_row.base_url,
|
||||
provider_row.api_key,
|
||||
provider_row.provider_fee,
|
||||
provider_row.provider_type,
|
||||
)
|
||||
elif provider_row.provider_type == "custom":
|
||||
return provider
|
||||
|
||||
if provider_row.provider_type == "custom":
|
||||
return BaseUpstreamProvider(
|
||||
provider_row.base_url, provider_row.api_key, provider_row.provider_fee
|
||||
)
|
||||
else:
|
||||
logger.error(
|
||||
f"Unknown provider type: {provider_row.provider_type}",
|
||||
extra={"base_url": provider_row.base_url},
|
||||
)
|
||||
return None
|
||||
|
||||
logger.error(
|
||||
f"Unknown provider type: {provider_row.provider_type}",
|
||||
extra={"base_url": provider_row.base_url},
|
||||
)
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Failed to instantiate provider: {e}",
|
||||
|
||||
@@ -9,7 +9,7 @@ from fastapi.responses import Response, StreamingResponse
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import ApiKey, AsyncSession
|
||||
from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow
|
||||
from ..payment.models import Model
|
||||
|
||||
from ..core.logging import get_logger
|
||||
@@ -20,6 +20,10 @@ logger = get_logger(__name__)
|
||||
class OllamaUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Upstream provider specifically configured for Ollama API."""
|
||||
|
||||
provider_type = "ollama"
|
||||
default_base_url = "http://localhost:11434"
|
||||
platform_url = "https://ollama.com/"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str = "http://localhost:11434",
|
||||
@@ -33,13 +37,32 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
||||
api_key: Optional API key (Ollama typically doesn't require one)
|
||||
provider_fee: Provider fee multiplier (default 1.01 for 1% fee)
|
||||
"""
|
||||
self.upstream_name = "ollama"
|
||||
super().__init__(
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
provider_fee=provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_db_row(
|
||||
cls, provider_row: "UpstreamProviderRow"
|
||||
) -> "OllamaUpstreamProvider":
|
||||
return cls(
|
||||
base_url=provider_row.base_url,
|
||||
api_key=provider_row.api_key,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": "Ollama",
|
||||
"default_base_url": cls.default_base_url,
|
||||
"fixed_base_url": False,
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
"""Strip 'ollama/' prefix for Ollama API compatibility."""
|
||||
return model_id.removeprefix("ollama/")
|
||||
@@ -192,12 +215,12 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
||||
|
||||
self._models_by_id = {m.id: m for m in self._models_cache}
|
||||
logger.info(
|
||||
f"Refreshed models cache for {self.upstream_name or self.base_url}",
|
||||
f"Refreshed models cache for {self.base_url}",
|
||||
extra={"model_count": len(models)},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Failed to refresh models cache for {self.upstream_name or self.base_url}",
|
||||
f"Failed to refresh models cache for {self.base_url}",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
|
||||
|
||||
@@ -1,18 +1,43 @@
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ..payment.models import Model, async_fetch_openrouter_models
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
|
||||
|
||||
class OpenAIUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Upstream provider specifically configured for OpenAI API."""
|
||||
|
||||
provider_type = "openai"
|
||||
default_base_url = "https://api.openai.com/v1"
|
||||
platform_url = "https://platform.openai.com/api-keys"
|
||||
|
||||
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,
|
||||
base_url=self.default_base_url, api_key=api_key, provider_fee=provider_fee
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_db_row(
|
||||
cls, provider_row: "UpstreamProviderRow"
|
||||
) -> "OpenAIUpstreamProvider":
|
||||
return cls(
|
||||
api_key=provider_row.api_key,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": "OpenAI",
|
||||
"default_base_url": cls.default_base_url,
|
||||
"fixed_base_url": True,
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
"""Strip 'openai/' prefix for OpenAI API compatibility."""
|
||||
return model_id.removeprefix("openai/")
|
||||
|
||||
@@ -1,10 +1,19 @@
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ..payment.models import Model, async_fetch_openrouter_models
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
|
||||
|
||||
class OpenRouterUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Upstream provider specifically configured for OpenRouter API."""
|
||||
|
||||
provider_type = "openrouter"
|
||||
default_base_url = "https://openrouter.ai/api/v1"
|
||||
platform_url = "https://openrouter.ai/settings/keys"
|
||||
|
||||
def __init__(self, api_key: str, provider_fee: float = 1.06):
|
||||
"""Initialize OpenRouter provider with API key.
|
||||
|
||||
@@ -12,13 +21,29 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider):
|
||||
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,
|
||||
base_url=self.default_base_url, api_key=api_key, provider_fee=provider_fee
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_db_row(
|
||||
cls, provider_row: "UpstreamProviderRow"
|
||||
) -> "OpenRouterUpstreamProvider":
|
||||
return cls(
|
||||
api_key=provider_row.api_key,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": "OpenRouter",
|
||||
"default_base_url": cls.default_base_url,
|
||||
"fixed_base_url": True,
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
async def fetch_models(self) -> list[Model]:
|
||||
"""Fetch all OpenRouter models."""
|
||||
models_data = await async_fetch_openrouter_models()
|
||||
|
||||
@@ -1,21 +1,45 @@
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ..payment.models import Model, async_fetch_openrouter_models
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
|
||||
|
||||
class PerplexityUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Upstream provider specifically configured for OpenAI API."""
|
||||
"""Upstream provider specifically configured for Perplexity API."""
|
||||
|
||||
upstream_name = "perplexity"
|
||||
base_url = "https://api.perplexity.ai/" # without v1
|
||||
provider_type = "perplexity"
|
||||
default_base_url = "https://api.perplexity.ai/"
|
||||
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,
|
||||
base_url=self.default_base_url,
|
||||
api_key=api_key,
|
||||
provider_fee=provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_db_row(
|
||||
cls, provider_row: "UpstreamProviderRow"
|
||||
) -> "PerplexityUpstreamProvider":
|
||||
return cls(
|
||||
api_key=provider_row.api_key,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": "Perplexity",
|
||||
"default_base_url": cls.default_base_url,
|
||||
"fixed_base_url": True,
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
"""Strip 'perplexity/' prefix for Perplexity API compatibility."""
|
||||
return model_id.removeprefix("perplexity/")
|
||||
|
||||
+26
-4
@@ -1,19 +1,41 @@
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ..payment.models import Model, async_fetch_openrouter_models
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
|
||||
|
||||
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"
|
||||
provider_type = "xai"
|
||||
default_base_url = "https://api.x.ai/v1"
|
||||
platform_url = "https://console.x.ai/"
|
||||
|
||||
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
|
||||
base_url=self.default_base_url, api_key=api_key, provider_fee=provider_fee
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_db_row(cls, provider_row: "UpstreamProviderRow") -> "XAIUpstreamProvider":
|
||||
return cls(
|
||||
api_key=provider_row.api_key,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": "xAI",
|
||||
"default_base_url": cls.default_base_url,
|
||||
"fixed_base_url": True,
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
"""Strip 'xai/' prefix for XAI API compatibility."""
|
||||
return model_id.removeprefix("xai/")
|
||||
|
||||
@@ -49,7 +49,7 @@ def create_test_model(
|
||||
def create_test_provider(name: str, base_url: str = "http://test.com") -> Mock:
|
||||
"""Helper to create a test provider mock."""
|
||||
provider = Mock()
|
||||
provider.upstream_name = name
|
||||
provider.provider_type = name
|
||||
provider.base_url = base_url
|
||||
return provider
|
||||
|
||||
|
||||
Reference in New Issue
Block a user