Add upstream providers management and model integration

- Created migration scripts to establish the `upstream_providers` table and integrate it with the `models` table.
- Enhanced the proxy functionality to support dynamic upstream provider initialization and model resolution.
- Implemented API endpoints for managing upstream providers, including CRUD operations and model fetching.
- Refactored pricing calculations to accommodate upstream provider overrides and ensure accurate cost estimation.
- Updated the admin interface to allow for easy management of upstream providers and their associated models.
This commit is contained in:
Shroominic
2025-10-13 17:17:54 +08:00
parent 4b36fb8d6f
commit 61a0559f8e
10 changed files with 2148 additions and 351 deletions
@@ -0,0 +1,45 @@
"""create upstream_providers table
Revision ID: d1e2f3a4b5c6
Revises: c0ffee123456
Create Date: 2025-10-09 00:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "d1e2f3a4b5c6"
down_revision = "c0ffee123456"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
if "upstream_providers" not in inspector.get_table_names():
op.create_table(
"upstream_providers",
sa.Column(
"id", sa.Integer(), primary_key=True, nullable=False, autoincrement=True
),
sa.Column("provider_type", sa.String(), nullable=False),
sa.Column("base_url", sa.String(), nullable=False, unique=True),
sa.Column("api_key", sa.String(), nullable=False),
sa.Column("api_version", sa.String(), nullable=True),
sa.Column("enabled", sa.Boolean(), nullable=False, default=True),
)
op.create_index(
"ix_upstream_providers_base_url",
"upstream_providers",
["base_url"],
unique=True,
)
def downgrade() -> None:
op.drop_index("ix_upstream_providers_base_url", "upstream_providers")
op.drop_table("upstream_providers")
@@ -0,0 +1,53 @@
"""add upstream_provider and enabled to models
Revision ID: e1f2a3b4c5d6
Revises: d1e2f3a4b5c6
Create Date: 2025-10-13 00:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "e1f2a3b4c5d6"
down_revision = "d1e2f3a4b5c6"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.drop_table("models")
op.create_table(
"models",
sa.Column("id", sa.String(), primary_key=True, nullable=False),
sa.Column("name", sa.String(), nullable=False),
sa.Column("created", sa.Integer(), nullable=False),
sa.Column("description", sa.Text(), nullable=False),
sa.Column("context_length", sa.Integer(), nullable=False),
sa.Column("architecture", sa.Text(), nullable=False),
sa.Column("pricing", sa.Text(), nullable=False),
sa.Column("sats_pricing", sa.Text(), nullable=True),
sa.Column("per_request_limits", sa.Text(), nullable=True),
sa.Column("top_provider", sa.Text(), nullable=True),
sa.Column("enabled", sa.Boolean(), nullable=False, server_default="1"),
sa.Column("upstream_provider_id", sa.Integer(), nullable=True),
sa.ForeignKeyConstraint(["upstream_provider_id"], ["upstream_providers.id"]),
)
def downgrade() -> None:
op.drop_table("models")
op.create_table(
"models",
sa.Column("id", sa.String(), primary_key=True, nullable=False),
sa.Column("name", sa.String(), nullable=False),
sa.Column("created", sa.Integer(), nullable=False),
sa.Column("description", sa.Text(), nullable=False),
sa.Column("context_length", sa.Integer(), nullable=False),
sa.Column("architecture", sa.Text(), nullable=False),
sa.Column("pricing", sa.Text(), nullable=False),
sa.Column("sats_pricing", sa.Text(), nullable=True),
sa.Column("per_request_limits", sa.Text(), nullable=True),
sa.Column("top_provider", sa.Text(), nullable=True),
)
+935 -32
View File
File diff suppressed because it is too large Load Diff
+21 -1
View File
@@ -5,7 +5,7 @@ from typing import AsyncGenerator
from alembic import command from alembic import command
from alembic.config import Config from alembic.config import Config
from sqlalchemy.ext.asyncio.engine import create_async_engine from sqlalchemy.ext.asyncio.engine import create_async_engine
from sqlmodel import Field, SQLModel, func, select from sqlmodel import Field, Relationship, SQLModel, func, select
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
from .logging import get_logger from .logging import get_logger
@@ -64,6 +64,26 @@ class ModelRow(SQLModel, table=True): # type: ignore
sats_pricing: str | None = Field(default=None) sats_pricing: str | None = Field(default=None)
per_request_limits: str | None = Field(default=None) per_request_limits: str | None = Field(default=None)
top_provider: str | None = Field(default=None) top_provider: str | None = Field(default=None)
enabled: bool = Field(default=True, description="Whether this model is enabled")
upstream_provider_id: int | None = Field(
default=None, foreign_key="upstream_providers.id"
)
upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models")
class UpstreamProviderRow(SQLModel, table=True): # type: ignore
__tablename__ = "upstream_providers"
id: int | None = Field(default=None, primary_key=True)
provider_type: str = Field(
description="Provider type: generic, openai, azure, openrouter"
)
base_url: str = Field(unique=True, description="Base URL of the upstream API")
api_key: str = Field(description="API key for the upstream provider")
api_version: str | None = Field(
default=None, description="API version for Azure OpenAI"
)
enabled: bool = Field(default=True, description="Whether this provider is enabled")
models: list["ModelRow"] = Relationship(back_populates="upstream_provider")
async def balances_for_mint_and_unit( async def balances_for_mint_and_unit(
+16 -5
View File
@@ -11,12 +11,10 @@ from ..balance import balance_router, deprecated_wallet_router
from ..discovery import providers_cache_refresher, providers_router from ..discovery import providers_cache_refresher, providers_router
from ..nip91 import announce_provider from ..nip91 import announce_provider
from ..payment.models import ( from ..payment.models import (
ensure_models_bootstrapped,
models_router, models_router,
refresh_models_periodically,
update_sats_pricing, update_sats_pricing,
) )
from ..proxy import proxy_router from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically
from ..wallet import periodic_payout from ..wallet import periodic_payout
from .admin import admin_router from .admin import admin_router
from .db import create_session, init_db, run_migrations from .db import create_session, init_db, run_migrations
@@ -42,6 +40,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
nip91_task = None nip91_task = None
providers_task = None providers_task = None
models_refresh_task = None models_refresh_task = None
model_maps_refresh_task = None
try: try:
# Run database migrations on startup # Run database migrations on startup
@@ -65,10 +64,18 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
except Exception: except Exception:
pass pass
await ensure_models_bootstrapped() # await ensure_models_bootstrapped()
await initialize_upstreams()
from ..proxy import get_upstreams
from ..upstream import refresh_upstreams_models_periodically
pricing_task = asyncio.create_task(update_sats_pricing()) pricing_task = asyncio.create_task(update_sats_pricing())
if global_settings.models_refresh_interval_seconds > 0: if global_settings.models_refresh_interval_seconds > 0:
models_refresh_task = asyncio.create_task(refresh_models_periodically()) models_refresh_task = asyncio.create_task(
refresh_upstreams_models_periodically(get_upstreams())
)
model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically())
payout_task = asyncio.create_task(periodic_payout()) payout_task = asyncio.create_task(periodic_payout())
nip91_task = asyncio.create_task(announce_provider()) nip91_task = asyncio.create_task(announce_provider())
providers_task = asyncio.create_task(providers_cache_refresher()) providers_task = asyncio.create_task(providers_cache_refresher())
@@ -94,6 +101,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
providers_task.cancel() providers_task.cancel()
if models_refresh_task is not None: if models_refresh_task is not None:
models_refresh_task.cancel() models_refresh_task.cancel()
if model_maps_refresh_task is not None:
model_maps_refresh_task.cancel()
try: try:
tasks_to_wait = [] tasks_to_wait = []
@@ -107,6 +116,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
tasks_to_wait.append(providers_task) tasks_to_wait.append(providers_task)
if models_refresh_task is not None: if models_refresh_task is not None:
tasks_to_wait.append(models_refresh_task) tasks_to_wait.append(models_refresh_task)
if model_maps_refresh_task is not None:
tasks_to_wait.append(model_maps_refresh_task)
if tasks_to_wait: if tasks_to_wait:
await asyncio.gather(*tasks_to_wait, return_exceptions=True) await asyncio.gather(*tasks_to_wait, return_exceptions=True)
+12 -16
View File
@@ -1,12 +1,8 @@
import json
import math import math
from pydantic.v1 import BaseModel from pydantic.v1 import BaseModel
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from ..core import get_logger from ..core import get_logger
from ..core.db import ModelRow
from ..core.settings import settings from ..core.settings import settings
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -29,7 +25,7 @@ class CostDataError(BaseModel):
async def calculate_cost( async def calculate_cost(
response_data: dict, max_cost: int, session: AsyncSession | None = None response_data: dict, max_cost: int, session: object | None = None
) -> CostData | MaxCostData | CostDataError: ) -> CostData | MaxCostData | CostDataError:
""" """
Calculate the cost of an API request based on token usage. Calculate the cost of an API request based on token usage.
@@ -74,18 +70,20 @@ async def calculate_cost(
float(settings.fixed_per_1k_output_tokens) * 1000.0 float(settings.fixed_per_1k_output_tokens) * 1000.0
) )
if not settings.fixed_pricing and session is not None: if not settings.fixed_pricing:
response_model = response_data.get("model", "") response_model = response_data.get("model", "")
logger.debug( logger.debug(
"Using model-based pricing", "Using model-based pricing",
extra={"model": response_model}, extra={"model": response_model},
) )
result = await session.exec(select(ModelRow.id)) # type: ignore from ..proxy import get_upstreams
available_ids = [ from ..upstream import get_model_with_override
row[0] if isinstance(row, tuple) else row for row in result.all()
] upstreams = get_upstreams()
if response_model not in available_ids: model_obj = await get_model_with_override(response_model, upstreams)
if not model_obj:
logger.error( logger.error(
"Invalid model in response", "Invalid model in response",
extra={"response_model": response_model}, extra={"response_model": response_model},
@@ -95,8 +93,7 @@ async def calculate_cost(
code="model_not_found", code="model_not_found",
) )
row = await session.get(ModelRow, response_model) if not model_obj.sats_pricing:
if row is None or not row.sats_pricing:
logger.error( logger.error(
"Model pricing not defined", "Model pricing not defined",
extra={"model": response_model, "model_id": response_model}, extra={"model": response_model, "model_id": response_model},
@@ -106,9 +103,8 @@ async def calculate_cost(
) )
try: try:
sats_pricing = json.loads(row.sats_pricing) mspp = float(model_obj.sats_pricing.prompt)
mspp = float(sats_pricing.get("prompt", 0)) mspc = float(model_obj.sats_pricing.completion)
mspc = float(sats_pricing.get("completion", 0))
except Exception: except Exception:
return CostDataError(message="Invalid pricing data", code="pricing_invalid") return CostDataError(message="Invalid pricing data", code="pricing_invalid")
+35 -34
View File
@@ -1,13 +1,12 @@
import json import json
import math import math
from typing import Any
from fastapi import HTTPException, Response from fastapi import HTTPException, Response
from fastapi.requests import Request from fastapi.requests import Request
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
from ..core import get_logger from ..core import get_logger
from ..core.db import ModelRow
from ..core.settings import settings from ..core.settings import settings
from ..wallet import deserialize_token_from_string from ..wallet import deserialize_token_from_string
from .models import Pricing from .models import Pricing
@@ -84,19 +83,19 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N
async def get_max_cost_for_model( async def get_max_cost_for_model(
model: str, session: AsyncSession | None = None model: str,
session: AsyncSession | None = None,
model_obj: Any | None = None,
) -> int: ) -> int:
"""Get the maximum cost for a specific model.""" """Get the maximum cost for a specific model from providers with overrides."""
logger.debug( logger.debug(
"Getting max cost for model", "Getting max cost for model",
extra={ extra={
"model": model, "model": model,
"fixed_pricing": settings.fixed_pricing, "fixed_pricing": settings.fixed_pricing,
"has_models": True,
}, },
) )
# Fixed pricing: always use fixed_cost_per_request
if settings.fixed_pricing: if settings.fixed_pricing:
default_cost_msats = settings.fixed_cost_per_request * 1000 default_cost_msats = settings.fixed_cost_per_request * 1000
logger.debug( logger.debug(
@@ -105,43 +104,42 @@ async def get_max_cost_for_model(
) )
return max(settings.min_request_msat, default_cost_msats) return max(settings.min_request_msat, default_cost_msats)
if session is None: if not model_obj:
# Without a DB session, we can't resolve model pricing; fall back to fixed cost from ..proxy import get_upstreams
fallback_msats = settings.fixed_cost_per_request * 1000 from ..upstream import get_model_with_override
logger.warning(
"No DB session provided for model pricing; using fixed cost",
extra={"requested_model": model, "using_default_cost": fallback_msats},
)
return max(settings.min_request_msat, fallback_msats)
result = await session.exec(select(ModelRow.id)) # type: ignore upstreams = get_upstreams()
available_ids = [row[0] if isinstance(row, tuple) else row for row in result.all()] model_obj = await get_model_with_override(model, upstreams)
if model not in available_ids:
# If no models or unknown model, fall back to fixed cost if provided, else minimal default if not model_obj:
fallback_msats = settings.fixed_cost_per_request * 1000 fallback_msats = settings.fixed_cost_per_request * 1000
logger.warning( logger.warning(
"Model not found in available models", "Model not found in providers or overrides",
extra={ extra={
"requested_model": model, "requested_model": model,
"available_models": available_ids,
"using_default_cost": fallback_msats, "using_default_cost": fallback_msats,
}, },
) )
return max(settings.min_request_msat, fallback_msats) return max(settings.min_request_msat, fallback_msats)
row = await session.get(ModelRow, model) if model_obj.sats_pricing:
if row and row.sats_pricing:
try: try:
sats = Pricing(**json.loads(row.sats_pricing)) # type: ignore max_cost = (
max_cost = sats.max_cost * 1000 * (1 - settings.tolerance_percentage / 100) model_obj.sats_pricing.max_cost
* 1000
* (1 - settings.tolerance_percentage / 100)
)
logger.debug( logger.debug(
"Found model-specific max cost", "Found model-specific max cost",
extra={"model": model, "max_cost_msats": max_cost}, extra={"model": model, "max_cost_msats": max_cost},
) )
calculated_msats = int(max_cost) calculated_msats = int(max_cost)
return max(settings.min_request_msat, calculated_msats) return max(settings.min_request_msat, calculated_msats)
except Exception: except Exception as e:
pass logger.error(
"Error calculating max cost from model pricing",
extra={"model": model, "error": str(e)},
)
logger.warning( logger.warning(
"Model pricing not found, using fixed cost", "Model pricing not found, using fixed cost",
@@ -220,16 +218,19 @@ def estimate_tokens(messages: list) -> int:
async def get_model_cost_info( async def get_model_cost_info(
model_id: str, session: AsyncSession | None = None model_id: str, session: AsyncSession | None = None
) -> Pricing | None: ) -> Pricing | None:
"""Get model pricing info from providers with database overrides."""
if not model_id or model_id == "unknown": if not model_id or model_id == "unknown":
return None return None
if session is None:
return None from ..proxy import get_upstreams
row = await session.get(ModelRow, model_id) from ..upstream import get_model_with_override
if row and row.sats_pricing:
try: upstreams = get_upstreams()
return Pricing(**json.loads(row.sats_pricing)) # type: ignore model_obj = await get_model_with_override(model_id, upstreams)
except Exception:
return None if model_obj and model_obj.sats_pricing:
return model_obj.sats_pricing
return None return None
+304 -115
View File
@@ -4,6 +4,7 @@ import random
from pathlib import Path from pathlib import Path
from urllib.request import urlopen from urllib.request import urlopen
import httpx
from fastapi import APIRouter, Depends from fastapi import APIRouter, Depends
from pydantic.v1 import BaseModel from pydantic.v1 import BaseModel
from sqlmodel import select from sqlmodel import select
@@ -56,6 +57,12 @@ class Model(BaseModel):
sats_pricing: Pricing | None = None sats_pricing: Pricing | None = None
per_request_limits: dict | None = None per_request_limits: dict | None = None
top_provider: TopProvider | None = None top_provider: TopProvider | None = None
enabled: bool = True
upstream_provider_id: int | None = None
canonical_slug: str | None = None
def __hash__(self) -> int:
return hash(self.id)
def fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: def fetch_openrouter_models(source_filter: str | None = None) -> list[dict]:
@@ -97,6 +104,47 @@ def fetch_openrouter_models(source_filter: str | None = None) -> list[dict]:
return [] return []
async def async_fetch_openrouter_models(source_filter: str | None = None) -> list[dict]:
"""Asynchronously fetch model information from OpenRouter API."""
base_url = "https://openrouter.ai/api/v1"
try:
async with httpx.AsyncClient() as client:
response = await client.get(f"{base_url}/models", timeout=30)
response.raise_for_status()
data = response.json()
models_data: list[dict] = []
for model in data.get("data", []):
model_id = model.get("id", "")
if source_filter:
source_prefix = f"{source_filter}/"
if not model_id.startswith(source_prefix):
continue
model = dict(model)
model["id"] = model_id[len(source_prefix) :]
model_id = model["id"]
if (
"(free)" in model.get("name", "")
or model_id == "openrouter/auto"
or model_id == "google/gemini-2.5-pro-exp-03-25"
or model_id == "opengvlab/internvl3-78b"
or model_id == "openrouter/sonoma-dusk-alpha"
or model_id == "openrouter/sonoma-sky-alpha"
):
continue
models_data.append(model)
return models_data
except Exception as e:
logger.error(f"Error (async) fetching models from OpenRouter API: {e}")
return []
def is_openrouter_upstream() -> bool: def is_openrouter_upstream() -> bool:
try: try:
base = (settings.upstream_base_url or "").strip().rstrip("/") base = (settings.upstream_base_url or "").strip().rstrip("/")
@@ -188,10 +236,13 @@ def _row_to_model(row: ModelRow) -> Model:
sats_pricing=Pricing.parse_obj(sats_pricing) if sats_pricing else None, sats_pricing=Pricing.parse_obj(sats_pricing) if sats_pricing else None,
per_request_limits=per_request_limits, per_request_limits=per_request_limits,
top_provider=TopProvider.parse_obj(top_provider) if top_provider else None, top_provider=TopProvider.parse_obj(top_provider) if top_provider else None,
enabled=row.enabled,
upstream_provider_id=row.upstream_provider_id,
canonical_slug=getattr(row, "canonical_slug", None),
) )
def _model_to_row_payload(model: Model) -> dict[str, str | int | None]: def _model_to_row_payload(model: Model) -> dict[str, str | int | bool | None]:
return { return {
"id": model.id, "id": model.id,
"name": model.name, "name": model.name,
@@ -209,18 +260,28 @@ def _model_to_row_payload(model: Model) -> dict[str, str | int | None]:
"top_provider": json.dumps(model.top_provider.dict()) "top_provider": json.dumps(model.top_provider.dict())
if model.top_provider is not None if model.top_provider is not None
else None, else None,
"enabled": model.enabled,
"upstream_provider_id": model.upstream_provider_id,
} }
async def list_models(session: AsyncSession | None = None) -> list[Model]: async def list_models(
session: AsyncSession | None = None,
upstream_id: int | None = None,
include_disabled: bool = False,
) -> list[Model]:
from sqlmodel import select
query = select(ModelRow)
if upstream_id is not None:
query = query.where(ModelRow.upstream_provider_id == upstream_id)
if not include_disabled:
query = query.where(ModelRow.enabled)
if session is not None: if session is not None:
result = await session.exec(select(ModelRow)) # type: ignore return [_row_to_model(r) for r in (await session.exec(query)).all()] # type: ignore
rows = result.all()
return [_row_to_model(r) for r in rows]
async with create_session() as s: async with create_session() as s:
result = await s.exec(select(ModelRow)) # type: ignore return [_row_to_model(r) for r in (await s.exec(query)).all()] # type: ignore
rows = result.all()
return [_row_to_model(r) for r in rows]
async def get_model_by_id( async def get_model_by_id(
@@ -228,10 +289,101 @@ async def get_model_by_id(
) -> Model | None: ) -> Model | None:
if session is not None: if session is not None:
row = await session.get(ModelRow, model_id) row = await session.get(ModelRow, model_id)
return _row_to_model(row) if row else None return _row_to_model(row) if row and row.enabled else None
async with create_session() as s: async with create_session() as s:
row = await s.get(ModelRow, model_id) row = await s.get(ModelRow, model_id)
return _row_to_model(row) if row else None return _row_to_model(row) if row and row.enabled else None
def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model:
"""Update a model's sats_pricing based on USD pricing and exchange rate.
Args:
model: Model object to update
sats_to_usd: Current sats to USD exchange rate
Returns:
Updated Model object with new sats_pricing
"""
try:
sats = Pricing.parse_obj(
{k: v / sats_to_usd for k, v in model.pricing.dict().items()}
)
min_req_msat = max(1, int(getattr(settings, "min_request_msat", 1)))
min_req_sats = float(min_req_msat) / 1000.0
if sats.request <= 0.0:
sats.request = min_req_sats
mspp = sats.prompt
mspc = sats.completion
if model.top_provider and (
model.top_provider.context_length
or model.top_provider.max_completion_tokens
):
if (cl := model.top_provider.context_length) and (
mct := model.top_provider.max_completion_tokens
):
max_prompt_cost = (cl - mct) * mspp
max_completion_cost = mct * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif cl := model.top_provider.context_length:
max_prompt_cost = cl * 0.8 * mspp
max_completion_cost = cl * 0.2 * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif mct := model.top_provider.max_completion_tokens:
max_prompt_cost = mct * 4 * mspp
max_completion_cost = mct * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif model.context_length:
max_prompt_cost = mspp * model.context_length * 0.8
max_completion_cost = mspc * model.context_length * 0.2
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
else:
p = mspp * 1_000_000
c = mspc * 32_000
r = sats.request * 100_000
i = sats.image * 100
w = sats.web_search * 1000
ir = sats.internal_reasoning * 100
sats.max_prompt_cost = p
sats.max_completion_cost = c
sats.max_cost = p + c + r + i + w + ir
if (sats.max_cost or 0.0) < min_req_sats:
sats.max_cost = min_req_sats
return Model(
id=model.id,
name=model.name,
created=model.created,
description=model.description,
context_length=model.context_length,
architecture=model.architecture,
pricing=model.pricing,
sats_pricing=sats,
per_request_limits=model.per_request_limits,
top_provider=model.top_provider,
)
except Exception as e:
logger.error(
"Failed to update sats pricing for model",
extra={
"model_id": model.id,
"error": str(e),
"error_type": type(e).__name__,
},
)
return model
async def ensure_models_bootstrapped() -> None: async def ensure_models_bootstrapped() -> None:
@@ -285,113 +437,134 @@ async def ensure_models_bootstrapped() -> None:
await s.commit() await s.commit()
async def update_sats_pricing() -> None: async def _update_sats_pricing_once() -> None:
while True: """Update sats pricing once for all provider models and database overrides."""
try: from ..proxy import get_upstreams
sats_to_usd = await sats_usd_ask_price()
upstreams = get_upstreams()
updated_count = 0
for upstream in upstreams:
updated_models = [
_update_model_sats_pricing(m, sats_to_usd)
for m in upstream.get_cached_models()
]
upstream._models_cache = updated_models
upstream._models_by_id = {m.id: m for m in updated_models}
updated_count += len(updated_models)
async with create_session() as s:
result = await s.exec(
select(ModelRow).where(ModelRow.upstream_provider_id.isnot(None)) # type: ignore
) # type: ignore
rows = result.all()
changed = 0
for row in rows:
try: try:
if not settings.enable_pricing_refresh: pricing = Pricing.parse_obj(json.loads(row.pricing))
return top_provider = (
except Exception: TopProvider.parse_obj(json.loads(row.top_provider))
pass if row.top_provider
sats_to_usd = await sats_usd_ask_price() else None
async with create_session() as s: )
result = await s.exec(select(ModelRow)) # type: ignore sats = Pricing.parse_obj(
rows = result.all() {k: v / sats_to_usd for k, v in pricing.dict().items()}
changed = 0 )
for row in rows: min_req_msat = max(1, int(getattr(settings, "min_request_msat", 1)))
try: min_req_sats = float(min_req_msat) / 1000.0
pricing = Pricing.parse_obj(json.loads(row.pricing)) if sats.request <= 0.0:
top_provider = ( sats.request = min_req_sats
TopProvider.parse_obj(json.loads(row.top_provider)) mspp = sats.prompt
if row.top_provider mspc = sats.completion
else None if top_provider and (
) top_provider.context_length or top_provider.max_completion_tokens
sats = Pricing.parse_obj( ):
{k: v / sats_to_usd for k, v in pricing.dict().items()} if (cl := top_provider.context_length) and (
) mct := top_provider.max_completion_tokens
# Enforce minimum per-request charge floor in sats ):
try: max_prompt_cost = (cl - mct) * mspp
min_req_msat = max( max_completion_cost = mct * mspc
1, int(getattr(settings, "min_request_msat", 1)) sats.max_prompt_cost = max_prompt_cost
) sats.max_completion_cost = max_completion_cost
except Exception: sats.max_cost = max_prompt_cost + max_completion_cost
min_req_msat = 1 elif cl := top_provider.context_length:
min_req_sats = float(min_req_msat) / 1000.0 max_prompt_cost = cl * 0.8 * mspp
if sats.request <= 0.0: max_completion_cost = cl * 0.2 * mspc
sats.request = min_req_sats sats.max_prompt_cost = max_prompt_cost
mspp = sats.prompt sats.max_completion_cost = max_completion_cost
mspc = sats.completion sats.max_cost = max_prompt_cost + max_completion_cost
if top_provider and ( elif mct := top_provider.max_completion_tokens:
top_provider.context_length max_prompt_cost = mct * 4 * mspp
or top_provider.max_completion_tokens max_completion_cost = mct * mspc
): sats.max_prompt_cost = max_prompt_cost
if (cl := top_provider.context_length) and ( sats.max_completion_cost = max_completion_cost
mct := top_provider.max_completion_tokens sats.max_cost = max_prompt_cost + max_completion_cost
): else:
max_prompt_cost = (cl - mct) * mspp max_prompt_cost = 1_000_000 * mspp
max_completion_cost = mct * mspc max_completion_cost = 32_000 * mspc
sats.max_prompt_cost = max_prompt_cost sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost sats.max_cost = max_prompt_cost + max_completion_cost
elif cl := top_provider.context_length: elif row.context_length:
max_prompt_cost = cl * 0.8 * mspp max_prompt_cost = mspp * row.context_length * 0.8
max_completion_cost = cl * 0.2 * mspc max_completion_cost = mspc * row.context_length * 0.2
sats.max_prompt_cost = max_prompt_cost sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost sats.max_cost = max_prompt_cost + max_completion_cost
elif mct := top_provider.max_completion_tokens: else:
max_prompt_cost = mct * 4 * mspp p = mspp * 1_000_000
max_completion_cost = mct * mspc c = mspc * 32_000
sats.max_prompt_cost = max_prompt_cost r = sats.request * 100_000
sats.max_completion_cost = max_completion_cost i = sats.image * 100
sats.max_cost = max_prompt_cost + max_completion_cost w = sats.web_search * 1000
else: ir = sats.internal_reasoning * 100
max_prompt_cost = 1_000_000 * mspp sats.max_prompt_cost = p
max_completion_cost = 32_000 * mspc sats.max_completion_cost = c
sats.max_prompt_cost = max_prompt_cost sats.max_cost = p + c + r + i + w + ir
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif row.context_length:
max_prompt_cost = mspp * row.context_length * 0.8
max_completion_cost = mspc * row.context_length * 0.2
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
else:
p = mspp * 1_000_000
c = mspc * 32_000
r = sats.request * 100_000
i = sats.image * 100
w = sats.web_search * 1000
ir = sats.internal_reasoning * 100
sats.max_prompt_cost = p
sats.max_completion_cost = c
sats.max_cost = p + c + r + i + w + ir
# Ensure overall minimum per-request total cost floor if (sats.max_cost or 0.0) < min_req_sats:
if (sats.max_cost or 0.0) < min_req_sats: sats.max_cost = min_req_sats
sats.max_cost = min_req_sats
new_json = json.dumps(sats.dict()) new_json = json.dumps(sats.dict())
if row.sats_pricing != new_json: if row.sats_pricing != new_json:
row.sats_pricing = new_json row.sats_pricing = new_json
s.add(row) s.add(row)
changed += 1 changed += 1
except Exception as per_row_error: except Exception as per_row_error:
logger.error( logger.error(
"Failed to update pricing for model", "Failed to update pricing for model",
extra={ extra={
"model_id": row.id, "model_id": row.id,
"error": str(per_row_error), "error": str(per_row_error),
"error_type": type(per_row_error).__name__, "error_type": type(per_row_error).__name__,
}, },
) )
if changed: if changed:
await s.commit() await s.commit()
except asyncio.CancelledError:
break if updated_count > 0 or changed > 0:
except Exception as e: logger.info(
logger.error(f"Error updating sats pricing: {e}") "Updated sats pricing",
extra={
"provider_models_updated": updated_count,
"database_overrides_updated": changed,
},
)
async def update_sats_pricing() -> None:
"""Periodically update sats pricing for all provider models and database overrides."""
try:
if not settings.enable_pricing_refresh:
return
except Exception:
pass
await _update_sats_pricing_once()
while True:
try: try:
interval = getattr(settings, "pricing_refresh_interval_seconds", 120) interval = getattr(settings, "pricing_refresh_interval_seconds", 120)
jitter = max(0.0, float(interval) * 0.1) jitter = max(0.0, float(interval) * 0.1)
@@ -399,6 +572,19 @@ async def update_sats_pricing() -> None:
except asyncio.CancelledError: except asyncio.CancelledError:
break break
try:
try:
if not settings.enable_pricing_refresh:
return
except Exception:
pass
await _update_sats_pricing_once()
except asyncio.CancelledError:
break
except Exception as e:
logger.error(f"Error updating sats pricing: {e}")
async def refresh_models_periodically() -> None: async def refresh_models_periodically() -> None:
"""Background task: periodically fetch OpenRouter models and insert new ones. """Background task: periodically fetch OpenRouter models and insert new ones.
@@ -473,5 +659,8 @@ async def refresh_models_periodically() -> None:
@models_router.get("/v1/models") @models_router.get("/v1/models")
@models_router.get("/models", include_in_schema=False) @models_router.get("/models", include_in_schema=False)
async def models(session: AsyncSession = Depends(get_session)) -> dict: async def models(session: AsyncSession = Depends(get_session)) -> dict:
items = await list_models(session) """Get all available models from all providers with database overrides applied."""
from ..proxy import get_unique_models
items = get_unique_models()
return {"data": items} return {"data": items}
+212 -92
View File
@@ -1,33 +1,173 @@
import json import json
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Request from fastapi import APIRouter, Depends, HTTPException, Request
from fastapi.responses import Response, StreamingResponse from fastapi.responses import Response, StreamingResponse
from sqlmodel import select
from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key
from .core import get_logger from .core import get_logger
from .core.db import ApiKey, AsyncSession, get_session from .core.db import ApiKey, AsyncSession, ModelRow, create_session, get_session
from .core.settings import settings
from .payment.helpers import ( from .payment.helpers import (
calculate_discounted_max_cost, calculate_discounted_max_cost,
check_token_balance, check_token_balance,
create_error_response, create_error_response,
get_max_cost_for_model, get_max_cost_for_model,
) )
from .upstream import init_upstreams from .payment.models import Model, _row_to_model
from .upstream import UpstreamProvider, init_upstreams, resolve_model_alias
logger = get_logger(__name__) logger = get_logger(__name__)
proxy_router = APIRouter() proxy_router = APIRouter()
upstreams = init_upstreams(settings.upstream_base_url, settings.upstream_api_key) _upstreams: list[UpstreamProvider] = []
upstream = upstreams[0] _model_instances: dict[str, Model] = {} # All aliases -> Model
_provider_map: dict[str, UpstreamProvider] = {} # All aliases -> Provider
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
async def initialize_upstreams() -> None:
"""Initialize upstream providers from database during application startup."""
global _upstreams
_upstreams = await init_upstreams()
logger.info(f"Initialized {len(_upstreams)} upstream providers")
await refresh_model_maps()
async def reinitialize_upstreams() -> None:
"""Re-initialize upstream providers from database (called after admin changes)."""
global _upstreams
_upstreams = await init_upstreams()
logger.info(
"Re-initialized upstream providers from admin action",
extra={"provider_count": len(_upstreams)},
)
await refresh_model_maps()
def get_upstreams() -> list[UpstreamProvider]:
"""Get the initialized upstream providers.
Returns:
List of upstream provider instances
"""
return _upstreams
def get_model_instance(model_id: str) -> Model | None:
"""Get Model instance by ID from global cache."""
return _model_instances.get(model_id)
def get_provider_for_model(model_id: str) -> UpstreamProvider | None:
"""Get UpstreamProvider for model ID from global cache."""
return _provider_map.get(model_id)
def get_unique_models() -> list[Model]:
"""Get list of unique models (no duplicates from aliases)."""
return list(_unique_models.values())
async def refresh_model_maps() -> None:
"""Refresh global model and provider maps in-place."""
global _model_instances, _provider_map, _unique_models
model_instances: dict[str, Model] = {}
provider_map: dict[str, UpstreamProvider] = {}
unique_models: dict[str, Model] = {}
openrouter: UpstreamProvider | None = None
other_upstreams: list[UpstreamProvider] = []
async with create_session() as session:
result = await session.exec(select(ModelRow).where(ModelRow.enabled))
override_rows = result.all()
overrides_by_id = {
row.id: row for row in override_rows if row.upstream_provider_id is not None
}
for upstream in _upstreams:
if upstream.base_url == "https://openrouter.ai/api/v1":
openrouter = upstream
else:
other_upstreams.append(upstream)
def get_base_model_id(model_id: str) -> str:
"""Get base model ID by removing provider prefix."""
return model_id.split("/", 1)[1] if "/" in model_id else model_id
if openrouter:
for model in openrouter.get_cached_models():
if model.enabled:
model_to_use = (
_row_to_model(overrides_by_id[model.id])
if model.id in overrides_by_id
else model
)
base_id = get_base_model_id(model_to_use.id)
if base_id not in unique_models:
unique_models[base_id] = model_to_use
for alias in resolve_model_alias(model.id, model_to_use.canonical_slug):
model_instances[alias] = model_to_use
provider_map[alias] = openrouter
for upstream in other_upstreams:
upstream_prefix = getattr(upstream, "upstream_name", None)
for model in upstream.get_cached_models():
if model.enabled:
model_to_use = (
_row_to_model(overrides_by_id[model.id])
if model.id in overrides_by_id
else model
)
base_id = get_base_model_id(model_to_use.id)
unique_models[base_id] = model_to_use
aliases = resolve_model_alias(model.id, model_to_use.canonical_slug)
if upstream_prefix and "/" not in model.id:
prefixed_id = f"{upstream_prefix}/{model.id}"
if prefixed_id not in aliases:
aliases.append(prefixed_id)
for alias in aliases:
model_instances[alias] = model_to_use
provider_map[alias] = upstream
_model_instances = model_instances
_provider_map = provider_map
_unique_models = unique_models
logger.debug(
"Refreshed model maps",
extra={
"unique_model_count": len(_unique_models),
"total_alias_count": len(_model_instances),
},
)
async def refresh_model_maps_periodically() -> None:
"""Background task to refresh model maps every minute."""
import asyncio
while True:
try:
await asyncio.sleep(60)
await refresh_model_maps()
except asyncio.CancelledError:
break
except Exception as e:
logger.error(
"Error refreshing model maps",
extra={"error": str(e), "error_type": type(e).__name__},
)
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) @proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
async def proxy( async def proxy(
request: Request, path: str, session: AsyncSession = Depends(get_session) request: Request, path: str, session: AsyncSession = Depends(get_session)
) -> Response | StreamingResponse: ) -> Response | StreamingResponse:
"""Main proxy endpoint handler."""
request_body = await request.body()
headers = dict(request.headers) headers = dict(request.headers)
if "x-cashu" not in headers and "authorization" not in headers.keys(): if "x-cashu" not in headers and "authorization" not in headers.keys():
@@ -35,7 +175,7 @@ async def proxy(
"unauthorized", "Unauthorized", 401, request=request "unauthorized", "Unauthorized", 401, request=request
) )
logger.info( logger.info( # TODO: move to middleware, async
"Received proxy request", "Received proxy request",
extra={ extra={
"method": request.method, "method": request.method,
@@ -45,76 +185,47 @@ async def proxy(
}, },
) )
# Parse JSON body if present, handle empty/invalid JSON request_body = await request.body()
request_body_dict = {} request_body_dict = parse_request_body_json(request_body, path)
if request_body:
try:
request_body_dict = json.loads(request_body)
logger.debug(
"Request body parsed",
extra={
"path": path,
"body_keys": list(request_body_dict.keys()),
"model": request_body_dict.get("model", "not_specified"),
},
)
except json.JSONDecodeError as e:
logger.error(
"Invalid JSON in request body",
extra={
"error": str(e),
"path": path,
"body_preview": request_body[:200].decode(errors="ignore")
if request_body
else "empty",
},
)
return Response(
content=json.dumps(
{"error": {"type": "invalid_request_error", "code": "invalid_json"}}
),
status_code=400,
media_type="application/json",
)
model = request_body_dict.get("model", "unknown") model_id = request_body_dict.get("model", "unknown")
_max_cost_for_model = await get_max_cost_for_model(model=model, session=session)
model_obj = get_model_instance(model_id)
if not model_obj:
return create_error_response(
"invalid_model", f"Model '{model_id}' not found", 400, request=request
)
upstream = get_provider_for_model(model_id)
if not upstream:
return create_error_response(
"invalid_model",
f"No provider found for model '{model_id}'",
400,
request=request,
)
_max_cost_for_model = await get_max_cost_for_model(
model=model_id, session=session, model_obj=model_obj
)
max_cost_for_model = await calculate_discounted_max_cost( max_cost_for_model = await calculate_discounted_max_cost(
_max_cost_for_model, request_body_dict, session _max_cost_for_model, request_body_dict, session
) )
check_token_balance(headers, request_body_dict, max_cost_for_model) check_token_balance(headers, request_body_dict, max_cost_for_model)
# Handle authentication
if x_cashu := headers.get("x-cashu", None): if x_cashu := headers.get("x-cashu", None):
logger.info(
"Processing X-Cashu payment",
extra={
"path": path,
"token_preview": x_cashu[:20] + "..." if len(x_cashu) > 20 else x_cashu,
},
)
return await upstream.handle_x_cashu(request, x_cashu, path, max_cost_for_model) return await upstream.handle_x_cashu(request, x_cashu, path, max_cost_for_model)
elif auth := headers.get("authorization", None): elif auth := headers.get("authorization", None):
logger.debug(
"Processing bearer token authentication",
extra={
"path": path,
"token_preview": auth[:20] + "..." if len(auth) > 20 else auth,
},
)
key = await get_bearer_token_key(headers, path, session, auth) key = await get_bearer_token_key(headers, path, session, auth)
else: else:
if request.method not in ["GET"]: if request.method not in ["GET"]:
logger.warning( raise HTTPException(
"Unauthorized request - no authentication provided",
extra={"method": request.method, "path": path},
)
return Response(
content=json.dumps({"detail": "Unauthorized"}),
status_code=401, status_code=401,
media_type="application/json", detail={
"error": {"type": "invalid_request_error", "code": "unauthorized"}
},
) )
logger.debug("Processing unauthenticated GET request", extra={"path": path}) logger.debug("Processing unauthenticated GET request", extra={"path": path})
@@ -124,38 +235,13 @@ async def proxy(
# Only pay for request if we have request body data (for completions endpoints) # Only pay for request if we have request body data (for completions endpoints)
if request_body_dict: if request_body_dict:
logger.info(
"Processing payment for request",
extra={
"path": path,
"key_hash": key.hashed_key[:8] + "...",
"key_balance_before": key.balance,
"model": request_body_dict.get("model", "unknown"),
},
)
try: try:
await pay_for_request(key, max_cost_for_model, session) await pay_for_request(key, max_cost_for_model, session)
logger.info( except Exception:
"Payment processed successfully", raise HTTPException(
extra={ status_code=402,
"path": path, detail={"error": {"type": "payment_error", "code": "payment_error"}},
"key_hash": key.hashed_key[:8] + "...",
"key_balance_after": key.balance,
"model": request_body_dict.get("model", "unknown"),
},
) )
except Exception as e:
logger.error(
"Payment processing failed",
extra={
"error": str(e),
"error_type": type(e).__name__,
"path": path,
"key_hash": key.hashed_key[:8] + "...",
},
)
raise
# Prepare headers for upstream # Prepare headers for upstream
headers = upstream.prepare_headers(dict(request.headers)) headers = upstream.prepare_headers(dict(request.headers))
@@ -270,3 +356,37 @@ async def get_bearer_token_key(
}, },
) )
raise raise
def parse_request_body_json(request_body: bytes, path: str) -> dict[str, Any]:
request_body_dict = {}
if request_body:
try:
request_body_dict = json.loads(request_body)
logger.debug(
"Request body parsed",
extra={
"path": path,
"body_keys": list(request_body_dict.keys()),
"model": request_body_dict.get("model", "not_specified"),
},
)
except json.JSONDecodeError as e:
logger.error(
"Invalid JSON in request body",
extra={
"error": str(e),
"path": path,
"body_preview": request_body[:200].decode(errors="ignore")
if request_body
else "empty",
},
)
raise HTTPException(
status_code=400,
detail={
"error": {"type": "invalid_request_error", "code": "invalid_json"}
},
)
return request_body_dict
+515 -56
View File
@@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Mapping
import httpx import httpx
if TYPE_CHECKING: if TYPE_CHECKING:
from .core.settings import Settings
from .payment.cost_caculation import CostData, MaxCostData from .payment.cost_caculation import CostData, MaxCostData
from fastapi import BackgroundTasks, HTTPException, Request from fastapi import BackgroundTasks, HTTPException, Request
@@ -16,50 +17,415 @@ from fastapi.responses import Response, StreamingResponse
from .auth import adjust_payment_for_tokens from .auth import adjust_payment_for_tokens
from .core import get_logger from .core import get_logger
from .core.db import ApiKey, AsyncSession, create_session from .core.db import ApiKey, AsyncSession, ModelRow, UpstreamProviderRow, create_session
from .payment.helpers import create_error_response from .payment.helpers import create_error_response
from .payment.models import Model from .payment.models import Model, async_fetch_openrouter_models
logger = get_logger(__name__) logger = get_logger(__name__)
def init_upstreams( def resolve_model_alias(model_id: str, canonical_slug: str | None = None) -> list[str]:
base_url: str, api_key: str, api_version: str | None = None """Resolve model ID to all possible aliases.
) -> list[UpstreamProvider]:
"""Initialize upstream providers based on settings. Returns list of aliases including canonical slug and variations without provider prefix.
Args: Args:
base_url: Base URL of the upstream API endpoint model_id: Model identifier (e.g., "gpt-5-mini" or "openai/gpt-5-mini")
api_key: API key for authenticating with the upstream service canonical_slug: Optional canonical slug from provider (e.g., "openai/gpt-5-pro-2025-10-06")
api_version: API version for Azure OpenAI
Returns:
List of possible model ID aliases
""" """
aliases = [model_id]
base_model = model_id
if "/" in model_id:
without_prefix = model_id.split("/", 1)[1]
aliases.append(without_prefix)
base_model = without_prefix
date_pattern = re.compile(r"-\d{4}-\d{2}-\d{2}$")
if date_pattern.search(base_model):
base_without_date = date_pattern.sub("", base_model)
if base_without_date not in aliases:
aliases.append(base_without_date)
if "/" in model_id:
prefix = model_id.split("/", 1)[0]
prefixed_without_date = f"{prefix}/{base_without_date}"
if prefixed_without_date not in aliases:
aliases.append(prefixed_without_date)
if canonical_slug and canonical_slug not in aliases:
aliases.append(canonical_slug)
if "/" in canonical_slug:
canonical_without_prefix = canonical_slug.split("/", 1)[1]
if canonical_without_prefix not in aliases:
aliases.append(canonical_without_prefix)
if date_pattern.search(canonical_without_prefix):
canonical_base = date_pattern.sub("", canonical_without_prefix)
if canonical_base not in aliases:
aliases.append(canonical_base)
return aliases
async def get_all_models_with_overrides(
upstreams: list[UpstreamProvider],
) -> list[Model]:
"""Get all models from all providers with database overrides applied.
Models in the database with upstream_provider_id set are treated as overrides
that replace the provider's model with the same ID.
Args:
upstreams: List of upstream provider instances
Returns:
List of Model objects with overrides applied
"""
from sqlmodel import select
from .payment.models import _row_to_model
async with create_session() as session:
result = await session.exec(select(ModelRow).where(ModelRow.enabled))
override_rows = result.all()
overrides_by_id = {
row.id: row for row in override_rows if row.upstream_provider_id is not None
}
all_models: dict[str, Model] = {}
for upstream in upstreams:
for model in upstream.get_cached_models():
if model.id in overrides_by_id:
all_models[model.id] = _row_to_model(overrides_by_id[model.id])
elif model.enabled:
all_models[model.id] = model
return list(all_models.values())
async def get_model_with_override(
model_id: str,
upstreams: list[UpstreamProvider],
) -> Model | None:
"""Get a specific model from providers with database override applied.
Resolves model aliases automatically (e.g., both "gpt-5-mini" and "openai/gpt-5-mini").
Args:
model_id: Model identifier (with or without provider prefix)
upstreams: List of upstream provider instances
Returns:
Model object or None if not found
"""
from sqlmodel import select
from .payment.models import _row_to_model
aliases = resolve_model_alias(model_id)
async with create_session() as session:
for alias in aliases:
result = await session.exec(
select(ModelRow).where(
ModelRow.id == alias,
ModelRow.upstream_provider_id.isnot(None), # type: ignore
ModelRow.enabled,
)
)
override_row = result.first()
if override_row:
return _row_to_model(override_row)
for alias in aliases:
for upstream in upstreams:
model = upstream.get_cached_model_by_id(alias)
if model and model.enabled:
return model
return None
async def refresh_upstreams_models_periodically(
upstreams: list[UpstreamProvider],
) -> None:
"""Background task to periodically refresh models cache for all providers.
Args:
upstreams: List of upstream provider instances
"""
import asyncio
import random
from .core.settings import settings from .core.settings import settings
upstreams: list[UpstreamProvider] = [] interval = getattr(settings, "models_refresh_interval_seconds", 0)
if settings.chat_completions_api_version: if not interval or interval <= 0:
upstreams.append( logger.info("Provider models refresh disabled (interval <= 0)")
AzureUpstreamProvider( return
settings.upstream_base_url,
settings.upstream_api_key, while True:
settings.chat_completions_api_version, try:
for upstream in upstreams:
try:
await upstream.refresh_models_cache()
except Exception as e:
logger.error(
f"Error refreshing models for {upstream.upstream_name or upstream.base_url}",
extra={"error": str(e), "error_type": type(e).__name__},
)
except asyncio.CancelledError:
break
except Exception as e:
logger.error(
"Error in provider models refresh loop",
extra={"error": str(e), "error_type": type(e).__name__},
) )
try:
jitter = max(0.0, float(interval) * 0.1)
await asyncio.sleep(interval + random.uniform(0, jitter))
except asyncio.CancelledError:
break
import os
async def init_upstreams() -> list[UpstreamProvider]:
"""Initialize upstream providers from database.
Seeds database with providers from settings if empty, then loads and instantiates
provider instances from database records, and refreshes their models cache.
"""
from sqlmodel import select
from .core.settings import settings
async with create_session() as session:
result = await session.exec(select(UpstreamProviderRow))
existing_providers = result.all()
if not existing_providers:
logger.info(
"No upstream providers found in database, seeding from settings"
)
await _seed_providers_from_settings(session, settings)
await session.commit()
result = await session.exec(select(UpstreamProviderRow))
existing_providers = result.all()
upstreams: list[UpstreamProvider] = []
for provider_row in existing_providers:
if not provider_row.enabled:
logger.debug(f"Skipping disabled provider: {provider_row.base_url}")
continue
provider = _instantiate_provider(provider_row)
if provider:
await provider.refresh_models_cache()
upstreams.append(provider)
logger.info(
f"Initialized {provider_row.provider_type} provider",
extra={
"base_url": provider_row.base_url,
"models_cached": len(provider.get_cached_models()),
},
)
return upstreams
async def _seed_providers_from_settings(
session: AsyncSession, settings: "Settings"
) -> None:
"""Seed database with upstream providers from environment variables.
Args:
session: Database session
"""
from sqlmodel import select
from .core.settings import settings
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)
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)
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,
)
)
seeded_base_urls.add(base_url)
if settings.chat_completions_api_version and settings.upstream_base_url:
base_url = settings.upstream_base_url
if base_url not in seeded_base_urls:
result = await session.exec(
select(UpstreamProviderRow).where(
UpstreamProviderRow.base_url == base_url
)
)
if not result.first():
providers_to_add.append(
UpstreamProviderRow(
provider_type="azure",
base_url=base_url,
api_key=settings.upstream_api_key,
api_version=settings.chat_completions_api_version,
enabled=True,
)
)
seeded_base_urls.add(base_url)
if settings.upstream_base_url and settings.upstream_api_key:
base_url = settings.upstream_base_url
if base_url not in seeded_base_urls:
result = await session.exec(
select(UpstreamProviderRow).where(
UpstreamProviderRow.base_url == base_url
)
)
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 "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="generic",
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",
extra={"base_url": provider.base_url},
) )
if "api.openai.com" in settings.upstream_base_url.lower():
upstreams.append(OpenAIUpstreamProvider(settings.upstream_api_key))
elif "openrouter.ai/api/v1" in settings.upstream_base_url.lower():
upstreams.append(OpenRouterUpstreamProvider(settings.upstream_api_key))
else:
upstreams.append(
UpstreamProvider(settings.upstream_base_url, settings.upstream_api_key)
)
return upstreams def _instantiate_provider(provider_row: UpstreamProviderRow) -> UpstreamProvider | None:
"""Instantiate an UpstreamProvider from a database row.
Args:
provider_row: Database row containing provider configuration
Returns:
Instantiated provider or None if provider type is unknown
"""
try:
if provider_row.provider_type == "openai":
return OpenAIUpstreamProvider(provider_row.api_key)
elif provider_row.provider_type == "azure":
if not provider_row.api_version:
logger.error(
"Azure provider missing api_version",
extra={"base_url": provider_row.base_url},
)
return None
return AzureUpstreamProvider(
provider_row.base_url,
provider_row.api_key,
provider_row.api_version,
)
elif provider_row.provider_type == "openrouter":
return OpenRouterUpstreamProvider(provider_row.api_key)
elif provider_row.provider_type == "generic":
return UpstreamProvider(provider_row.base_url, provider_row.api_key)
else:
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}",
extra={
"provider_type": provider_row.provider_type,
"base_url": provider_row.base_url,
"error": str(e),
},
)
return None
class UpstreamProvider: class UpstreamProvider:
"""Provider for forwarding requests to an upstream AI service API.""" """Provider for forwarding requests to an upstream AI service API."""
base_url: str
api_key: str
upstream_name: str | None = None
_models_cache: list[Model] = []
_models_by_id: dict[str, Model] = {}
def __init__(self, base_url: str, api_key: str): def __init__(self, base_url: str, api_key: str):
"""Initialize the upstream provider. """Initialize the upstream provider.
@@ -69,6 +435,8 @@ class UpstreamProvider:
""" """
self.base_url = base_url self.base_url = base_url
self.api_key = api_key self.api_key = api_key
self._models_cache = []
self._models_by_id = {}
def prepare_headers(self, request_headers: dict) -> dict: def prepare_headers(self, request_headers: dict) -> dict:
"""Prepare headers for upstream request by removing proxy-specific headers and adding authentication. """Prepare headers for upstream request by removing proxy-specific headers and adding authentication.
@@ -136,6 +504,60 @@ class UpstreamProvider:
""" """
return query_params or {} return query_params or {}
def transform_model_name(self, model_id: str) -> str:
"""Transform model ID for this provider's API format.
Base implementation returns model_id unchanged. Override in subclasses for provider-specific transformations.
Args:
model_id: Model identifier (may include provider prefix)
Returns:
Transformed model ID for this provider
"""
return model_id
def prepare_request_body(self, body: bytes | None) -> bytes | None:
"""Transform request body for provider-specific requirements.
Automatically transforms model names in the request body.
Args:
body: Original request body bytes
Returns:
Transformed request body bytes
"""
if not body:
return body
try:
data = json.loads(body)
if isinstance(data, dict) and "model" in data:
original_model = data["model"]
transformed_model = self.transform_model_name(original_model)
if transformed_model != original_model:
data["model"] = transformed_model
logger.debug(
"Transformed model name in request",
extra={
"original": original_model,
"transformed": transformed_model,
"provider": self.upstream_name or self.base_url,
},
)
return json.dumps(data).encode()
except Exception as e:
logger.debug(
"Could not transform request body",
extra={
"error": str(e),
"provider": self.upstream_name or self.base_url,
},
)
return body
def _extract_upstream_error_message( def _extract_upstream_error_message(
self, body_bytes: bytes self, body_bytes: bytes
) -> tuple[str, str | None]: ) -> tuple[str, str | None]:
@@ -560,6 +982,8 @@ class UpstreamProvider:
url = f"{self.base_url}/{path}" url = f"{self.base_url}/{path}"
transformed_body = self.prepare_request_body(request_body)
logger.info( logger.info(
"Forwarding request to upstream", "Forwarding request to upstream",
extra={ extra={
@@ -578,13 +1002,13 @@ class UpstreamProvider:
) )
try: try:
if request_body is not None: if transformed_body is not None:
response = await client.send( response = await client.send(
client.build_request( client.build_request(
request.method, request.method,
url, url,
headers=headers, headers=headers,
content=request_body, content=transformed_body,
params=self.prepare_params(path, request.query_params), params=self.prepare_params(path, request.query_params),
), ),
stream=True, stream=True,
@@ -1326,6 +1750,9 @@ class UpstreamProvider:
url = f"{self.base_url}/{path}" url = f"{self.base_url}/{path}"
request_body = await request.body()
transformed_body = self.prepare_request_body(request_body)
logger.debug( logger.debug(
"Forwarding request to upstream", "Forwarding request to upstream",
extra={ extra={
@@ -1347,7 +1774,7 @@ class UpstreamProvider:
request.method, request.method,
url, url,
headers=headers, headers=headers,
content=request.stream(), content=transformed_body if transformed_body else request_body,
params=self.prepare_params(path, request.query_params), params=self.prepare_params(path, request.query_params),
), ),
stream=True, stream=True,
@@ -1546,12 +1973,66 @@ class UpstreamProvider:
token=x_cashu_token, token=x_cashu_token,
) )
async def fetch_models(self) -> list[Model]:
"""Fetch available models from upstream API and update cache.
Returns:
List of Model objects with pricing
"""
logger.debug(f"Fetching models for {self.upstream_name or self.base_url}")
return []
async def refresh_models_cache(self) -> None:
"""Refresh the in-memory models cache from upstream API."""
try:
models = await self.fetch_models()
self._models_cache = models
self._models_by_id = {m.id: m for m in models}
logger.info(
f"Refreshed models cache for {self.upstream_name 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}",
extra={"error": str(e), "error_type": type(e).__name__},
)
def get_cached_models(self) -> list[Model]:
"""Get cached models for this provider.
Returns:
List of cached Model objects
"""
return self._models_cache
def get_cached_model_by_id(self, model_id: str) -> Model | None:
"""Get a specific cached model by ID.
Args:
model_id: Model identifier
Returns:
Model object or None if not found
"""
return self._models_by_id.get(model_id)
class OpenAIUpstreamProvider(UpstreamProvider): class OpenAIUpstreamProvider(UpstreamProvider):
"""Upstream provider specifically configured for OpenAI API.""" """Upstream provider specifically configured for OpenAI API."""
def __init__(self, api_key: str): def __init__(self, api_key: str):
super().__init__(base_url="https://api.openai.com", api_key=api_key) self.upstream_name = "openai"
super().__init__(base_url="https://api.openai.com/v1", api_key=api_key)
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 AzureUpstreamProvider(UpstreamProvider): class AzureUpstreamProvider(UpstreamProvider):
@@ -1595,32 +2076,10 @@ class OpenRouterUpstreamProvider(UpstreamProvider):
Args: Args:
api_key: OpenRouter API key for authentication api_key: OpenRouter API key for authentication
""" """
self.upstream_name = "openrouter"
super().__init__(base_url="https://openrouter.ai/api/v1", api_key=api_key) super().__init__(base_url="https://openrouter.ai/api/v1", api_key=api_key)
async def fetch_models(self) -> dict: async def fetch_models(self) -> list[Model]:
"""Fetch available models from OpenRouter API. """Fetch all OpenRouter models."""
models_data = await async_fetch_openrouter_models()
Returns: return [Model(**model) for model in models_data] # type: ignore
Raw JSON response containing model data
"""
async with httpx.AsyncClient() as client:
response = await client.get(
"https://openrouter.ai/api/v1/models",
headers={"Authorization": f"Bearer {self.api_key}"},
)
return response.json()
async def models(self) -> list[Model]:
"""Get list of available models from OpenRouter.
Returns:
List of Model objects representing available models
"""
response_data = await self.fetch_models()
models_list: list[Model] = []
for model_data in response_data.get("data", []):
try:
models_list.append(Model(**model_data)) # type: ignore
except Exception:
continue
return models_list