mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
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:
@@ -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
File diff suppressed because it is too large
Load Diff
+21
-1
@@ -5,7 +5,7 @@ from typing import AsyncGenerator
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
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 .logging import get_logger
|
||||
@@ -64,6 +64,26 @@ class ModelRow(SQLModel, table=True): # type: ignore
|
||||
sats_pricing: str | None = Field(default=None)
|
||||
per_request_limits: 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(
|
||||
|
||||
+16
-5
@@ -11,12 +11,10 @@ from ..balance import balance_router, deprecated_wallet_router
|
||||
from ..discovery import providers_cache_refresher, providers_router
|
||||
from ..nip91 import announce_provider
|
||||
from ..payment.models import (
|
||||
ensure_models_bootstrapped,
|
||||
models_router,
|
||||
refresh_models_periodically,
|
||||
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 .admin import admin_router
|
||||
from .db import create_session, init_db, run_migrations
|
||||
@@ -42,6 +40,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
nip91_task = None
|
||||
providers_task = None
|
||||
models_refresh_task = None
|
||||
model_maps_refresh_task = None
|
||||
|
||||
try:
|
||||
# Run database migrations on startup
|
||||
@@ -65,10 +64,18 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
except Exception:
|
||||
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())
|
||||
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())
|
||||
nip91_task = asyncio.create_task(announce_provider())
|
||||
providers_task = asyncio.create_task(providers_cache_refresher())
|
||||
@@ -94,6 +101,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
providers_task.cancel()
|
||||
if models_refresh_task is not None:
|
||||
models_refresh_task.cancel()
|
||||
if model_maps_refresh_task is not None:
|
||||
model_maps_refresh_task.cancel()
|
||||
|
||||
try:
|
||||
tasks_to_wait = []
|
||||
@@ -107,6 +116,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
tasks_to_wait.append(providers_task)
|
||||
if models_refresh_task is not None:
|
||||
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:
|
||||
await asyncio.gather(*tasks_to_wait, return_exceptions=True)
|
||||
|
||||
@@ -1,12 +1,8 @@
|
||||
import json
|
||||
import math
|
||||
|
||||
from pydantic.v1 import BaseModel
|
||||
from sqlmodel import select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from ..core import get_logger
|
||||
from ..core.db import ModelRow
|
||||
from ..core.settings import settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
@@ -29,7 +25,7 @@ class CostDataError(BaseModel):
|
||||
|
||||
|
||||
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:
|
||||
"""
|
||||
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
|
||||
)
|
||||
|
||||
if not settings.fixed_pricing and session is not None:
|
||||
if not settings.fixed_pricing:
|
||||
response_model = response_data.get("model", "")
|
||||
logger.debug(
|
||||
"Using model-based pricing",
|
||||
extra={"model": response_model},
|
||||
)
|
||||
|
||||
result = await session.exec(select(ModelRow.id)) # type: ignore
|
||||
available_ids = [
|
||||
row[0] if isinstance(row, tuple) else row for row in result.all()
|
||||
]
|
||||
if response_model not in available_ids:
|
||||
from ..proxy import get_upstreams
|
||||
from ..upstream import get_model_with_override
|
||||
|
||||
upstreams = get_upstreams()
|
||||
model_obj = await get_model_with_override(response_model, upstreams)
|
||||
|
||||
if not model_obj:
|
||||
logger.error(
|
||||
"Invalid model in response",
|
||||
extra={"response_model": response_model},
|
||||
@@ -95,8 +93,7 @@ async def calculate_cost(
|
||||
code="model_not_found",
|
||||
)
|
||||
|
||||
row = await session.get(ModelRow, response_model)
|
||||
if row is None or not row.sats_pricing:
|
||||
if not model_obj.sats_pricing:
|
||||
logger.error(
|
||||
"Model pricing not defined",
|
||||
extra={"model": response_model, "model_id": response_model},
|
||||
@@ -106,9 +103,8 @@ async def calculate_cost(
|
||||
)
|
||||
|
||||
try:
|
||||
sats_pricing = json.loads(row.sats_pricing)
|
||||
mspp = float(sats_pricing.get("prompt", 0))
|
||||
mspc = float(sats_pricing.get("completion", 0))
|
||||
mspp = float(model_obj.sats_pricing.prompt)
|
||||
mspc = float(model_obj.sats_pricing.completion)
|
||||
except Exception:
|
||||
return CostDataError(message="Invalid pricing data", code="pricing_invalid")
|
||||
|
||||
|
||||
+35
-34
@@ -1,13 +1,12 @@
|
||||
import json
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException, Response
|
||||
from fastapi.requests import Request
|
||||
from sqlmodel import select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from ..core import get_logger
|
||||
from ..core.db import ModelRow
|
||||
from ..core.settings import settings
|
||||
from ..wallet import deserialize_token_from_string
|
||||
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(
|
||||
model: str, session: AsyncSession | None = None
|
||||
model: str,
|
||||
session: AsyncSession | None = None,
|
||||
model_obj: Any | None = None,
|
||||
) -> int:
|
||||
"""Get the maximum cost for a specific model."""
|
||||
"""Get the maximum cost for a specific model from providers with overrides."""
|
||||
logger.debug(
|
||||
"Getting max cost for model",
|
||||
extra={
|
||||
"model": model,
|
||||
"fixed_pricing": settings.fixed_pricing,
|
||||
"has_models": True,
|
||||
},
|
||||
)
|
||||
|
||||
# Fixed pricing: always use fixed_cost_per_request
|
||||
if settings.fixed_pricing:
|
||||
default_cost_msats = settings.fixed_cost_per_request * 1000
|
||||
logger.debug(
|
||||
@@ -105,43 +104,42 @@ async def get_max_cost_for_model(
|
||||
)
|
||||
return max(settings.min_request_msat, default_cost_msats)
|
||||
|
||||
if session is None:
|
||||
# Without a DB session, we can't resolve model pricing; fall back to fixed cost
|
||||
fallback_msats = settings.fixed_cost_per_request * 1000
|
||||
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)
|
||||
if not model_obj:
|
||||
from ..proxy import get_upstreams
|
||||
from ..upstream import get_model_with_override
|
||||
|
||||
result = await session.exec(select(ModelRow.id)) # type: ignore
|
||||
available_ids = [row[0] if isinstance(row, tuple) else row for row in result.all()]
|
||||
if model not in available_ids:
|
||||
# If no models or unknown model, fall back to fixed cost if provided, else minimal default
|
||||
upstreams = get_upstreams()
|
||||
model_obj = await get_model_with_override(model, upstreams)
|
||||
|
||||
if not model_obj:
|
||||
fallback_msats = settings.fixed_cost_per_request * 1000
|
||||
logger.warning(
|
||||
"Model not found in available models",
|
||||
"Model not found in providers or overrides",
|
||||
extra={
|
||||
"requested_model": model,
|
||||
"available_models": available_ids,
|
||||
"using_default_cost": fallback_msats,
|
||||
},
|
||||
)
|
||||
return max(settings.min_request_msat, fallback_msats)
|
||||
|
||||
row = await session.get(ModelRow, model)
|
||||
if row and row.sats_pricing:
|
||||
if model_obj.sats_pricing:
|
||||
try:
|
||||
sats = Pricing(**json.loads(row.sats_pricing)) # type: ignore
|
||||
max_cost = sats.max_cost * 1000 * (1 - settings.tolerance_percentage / 100)
|
||||
max_cost = (
|
||||
model_obj.sats_pricing.max_cost
|
||||
* 1000
|
||||
* (1 - settings.tolerance_percentage / 100)
|
||||
)
|
||||
logger.debug(
|
||||
"Found model-specific max cost",
|
||||
extra={"model": model, "max_cost_msats": max_cost},
|
||||
)
|
||||
calculated_msats = int(max_cost)
|
||||
return max(settings.min_request_msat, calculated_msats)
|
||||
except Exception:
|
||||
pass
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Error calculating max cost from model pricing",
|
||||
extra={"model": model, "error": str(e)},
|
||||
)
|
||||
|
||||
logger.warning(
|
||||
"Model pricing not found, using fixed cost",
|
||||
@@ -220,16 +218,19 @@ def estimate_tokens(messages: list) -> int:
|
||||
async def get_model_cost_info(
|
||||
model_id: str, session: AsyncSession | None = None
|
||||
) -> Pricing | None:
|
||||
"""Get model pricing info from providers with database overrides."""
|
||||
if not model_id or model_id == "unknown":
|
||||
return None
|
||||
if session is None:
|
||||
return None
|
||||
row = await session.get(ModelRow, model_id)
|
||||
if row and row.sats_pricing:
|
||||
try:
|
||||
return Pricing(**json.loads(row.sats_pricing)) # type: ignore
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
from ..proxy import get_upstreams
|
||||
from ..upstream import get_model_with_override
|
||||
|
||||
upstreams = get_upstreams()
|
||||
model_obj = await get_model_with_override(model_id, upstreams)
|
||||
|
||||
if model_obj and model_obj.sats_pricing:
|
||||
return model_obj.sats_pricing
|
||||
|
||||
return None
|
||||
|
||||
|
||||
|
||||
+304
-115
@@ -4,6 +4,7 @@ import random
|
||||
from pathlib import Path
|
||||
from urllib.request import urlopen
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends
|
||||
from pydantic.v1 import BaseModel
|
||||
from sqlmodel import select
|
||||
@@ -56,6 +57,12 @@ class Model(BaseModel):
|
||||
sats_pricing: Pricing | None = None
|
||||
per_request_limits: dict | 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]:
|
||||
@@ -97,6 +104,47 @@ def fetch_openrouter_models(source_filter: str | None = None) -> list[dict]:
|
||||
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:
|
||||
try:
|
||||
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,
|
||||
per_request_limits=per_request_limits,
|
||||
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 {
|
||||
"id": model.id,
|
||||
"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())
|
||||
if model.top_provider is not 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:
|
||||
result = await session.exec(select(ModelRow)) # type: ignore
|
||||
rows = result.all()
|
||||
return [_row_to_model(r) for r in rows]
|
||||
return [_row_to_model(r) for r in (await session.exec(query)).all()] # type: ignore
|
||||
async with create_session() as s:
|
||||
result = await s.exec(select(ModelRow)) # type: ignore
|
||||
rows = result.all()
|
||||
return [_row_to_model(r) for r in rows]
|
||||
return [_row_to_model(r) for r in (await s.exec(query)).all()] # type: ignore
|
||||
|
||||
|
||||
async def get_model_by_id(
|
||||
@@ -228,10 +289,101 @@ async def get_model_by_id(
|
||||
) -> Model | None:
|
||||
if session is not None:
|
||||
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:
|
||||
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:
|
||||
@@ -285,113 +437,134 @@ async def ensure_models_bootstrapped() -> None:
|
||||
await s.commit()
|
||||
|
||||
|
||||
async def update_sats_pricing() -> None:
|
||||
while True:
|
||||
try:
|
||||
async def _update_sats_pricing_once() -> None:
|
||||
"""Update sats pricing once for all provider models and database overrides."""
|
||||
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:
|
||||
if not settings.enable_pricing_refresh:
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
sats_to_usd = await sats_usd_ask_price()
|
||||
async with create_session() as s:
|
||||
result = await s.exec(select(ModelRow)) # type: ignore
|
||||
rows = result.all()
|
||||
changed = 0
|
||||
for row in rows:
|
||||
try:
|
||||
pricing = Pricing.parse_obj(json.loads(row.pricing))
|
||||
top_provider = (
|
||||
TopProvider.parse_obj(json.loads(row.top_provider))
|
||||
if row.top_provider
|
||||
else None
|
||||
)
|
||||
sats = Pricing.parse_obj(
|
||||
{k: v / sats_to_usd for k, v in pricing.dict().items()}
|
||||
)
|
||||
# Enforce minimum per-request charge floor in sats
|
||||
try:
|
||||
min_req_msat = max(
|
||||
1, int(getattr(settings, "min_request_msat", 1))
|
||||
)
|
||||
except Exception:
|
||||
min_req_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 top_provider and (
|
||||
top_provider.context_length
|
||||
or top_provider.max_completion_tokens
|
||||
):
|
||||
if (cl := top_provider.context_length) and (
|
||||
mct := 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 := 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 := 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
|
||||
else:
|
||||
max_prompt_cost = 1_000_000 * mspp
|
||||
max_completion_cost = 32_000 * 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 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
|
||||
pricing = Pricing.parse_obj(json.loads(row.pricing))
|
||||
top_provider = (
|
||||
TopProvider.parse_obj(json.loads(row.top_provider))
|
||||
if row.top_provider
|
||||
else None
|
||||
)
|
||||
sats = Pricing.parse_obj(
|
||||
{k: v / sats_to_usd for k, v in 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 top_provider and (
|
||||
top_provider.context_length or top_provider.max_completion_tokens
|
||||
):
|
||||
if (cl := top_provider.context_length) and (
|
||||
mct := 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 := 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 := 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
|
||||
else:
|
||||
max_prompt_cost = 1_000_000 * mspp
|
||||
max_completion_cost = 32_000 * 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 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:
|
||||
sats.max_cost = min_req_sats
|
||||
if (sats.max_cost or 0.0) < min_req_sats:
|
||||
sats.max_cost = min_req_sats
|
||||
|
||||
new_json = json.dumps(sats.dict())
|
||||
if row.sats_pricing != new_json:
|
||||
row.sats_pricing = new_json
|
||||
s.add(row)
|
||||
changed += 1
|
||||
except Exception as per_row_error:
|
||||
logger.error(
|
||||
"Failed to update pricing for model",
|
||||
extra={
|
||||
"model_id": row.id,
|
||||
"error": str(per_row_error),
|
||||
"error_type": type(per_row_error).__name__,
|
||||
},
|
||||
)
|
||||
if changed:
|
||||
await s.commit()
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating sats pricing: {e}")
|
||||
new_json = json.dumps(sats.dict())
|
||||
if row.sats_pricing != new_json:
|
||||
row.sats_pricing = new_json
|
||||
s.add(row)
|
||||
changed += 1
|
||||
except Exception as per_row_error:
|
||||
logger.error(
|
||||
"Failed to update pricing for model",
|
||||
extra={
|
||||
"model_id": row.id,
|
||||
"error": str(per_row_error),
|
||||
"error_type": type(per_row_error).__name__,
|
||||
},
|
||||
)
|
||||
if changed:
|
||||
await s.commit()
|
||||
|
||||
if updated_count > 0 or changed > 0:
|
||||
logger.info(
|
||||
"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:
|
||||
interval = getattr(settings, "pricing_refresh_interval_seconds", 120)
|
||||
jitter = max(0.0, float(interval) * 0.1)
|
||||
@@ -399,6 +572,19 @@ async def update_sats_pricing() -> None:
|
||||
except asyncio.CancelledError:
|
||||
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:
|
||||
"""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("/models", include_in_schema=False)
|
||||
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}
|
||||
|
||||
+212
-92
@@ -1,33 +1,173 @@
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from fastapi.responses import Response, StreamingResponse
|
||||
from sqlmodel import select
|
||||
|
||||
from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key
|
||||
from .core import get_logger
|
||||
from .core.db import ApiKey, AsyncSession, get_session
|
||||
from .core.settings import settings
|
||||
from .core.db import ApiKey, AsyncSession, ModelRow, create_session, get_session
|
||||
from .payment.helpers import (
|
||||
calculate_discounted_max_cost,
|
||||
check_token_balance,
|
||||
create_error_response,
|
||||
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__)
|
||||
proxy_router = APIRouter()
|
||||
|
||||
upstreams = init_upstreams(settings.upstream_base_url, settings.upstream_api_key)
|
||||
upstream = upstreams[0]
|
||||
_upstreams: list[UpstreamProvider] = []
|
||||
_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)
|
||||
async def proxy(
|
||||
request: Request, path: str, session: AsyncSession = Depends(get_session)
|
||||
) -> Response | StreamingResponse:
|
||||
"""Main proxy endpoint handler."""
|
||||
request_body = await request.body()
|
||||
headers = dict(request.headers)
|
||||
|
||||
if "x-cashu" not in headers and "authorization" not in headers.keys():
|
||||
@@ -35,7 +175,7 @@ async def proxy(
|
||||
"unauthorized", "Unauthorized", 401, request=request
|
||||
)
|
||||
|
||||
logger.info(
|
||||
logger.info( # TODO: move to middleware, async
|
||||
"Received proxy request",
|
||||
extra={
|
||||
"method": request.method,
|
||||
@@ -45,76 +185,47 @@ async def proxy(
|
||||
},
|
||||
)
|
||||
|
||||
# Parse JSON body if present, handle empty/invalid JSON
|
||||
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",
|
||||
},
|
||||
)
|
||||
return Response(
|
||||
content=json.dumps(
|
||||
{"error": {"type": "invalid_request_error", "code": "invalid_json"}}
|
||||
),
|
||||
status_code=400,
|
||||
media_type="application/json",
|
||||
)
|
||||
request_body = await request.body()
|
||||
request_body_dict = parse_request_body_json(request_body, path)
|
||||
|
||||
model = request_body_dict.get("model", "unknown")
|
||||
_max_cost_for_model = await get_max_cost_for_model(model=model, session=session)
|
||||
model_id = request_body_dict.get("model", "unknown")
|
||||
|
||||
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, request_body_dict, session
|
||||
)
|
||||
check_token_balance(headers, request_body_dict, max_cost_for_model)
|
||||
|
||||
# Handle authentication
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
else:
|
||||
if request.method not in ["GET"]:
|
||||
logger.warning(
|
||||
"Unauthorized request - no authentication provided",
|
||||
extra={"method": request.method, "path": path},
|
||||
)
|
||||
return Response(
|
||||
content=json.dumps({"detail": "Unauthorized"}),
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
media_type="application/json",
|
||||
detail={
|
||||
"error": {"type": "invalid_request_error", "code": "unauthorized"}
|
||||
},
|
||||
)
|
||||
|
||||
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)
|
||||
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:
|
||||
await pay_for_request(key, max_cost_for_model, session)
|
||||
logger.info(
|
||||
"Payment processed successfully",
|
||||
extra={
|
||||
"path": path,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"key_balance_after": key.balance,
|
||||
"model": request_body_dict.get("model", "unknown"),
|
||||
},
|
||||
except Exception:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={"error": {"type": "payment_error", "code": "payment_error"}},
|
||||
)
|
||||
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
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
@@ -270,3 +356,37 @@ async def get_bearer_token_key(
|
||||
},
|
||||
)
|
||||
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
@@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Mapping
|
||||
import httpx
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .core.settings import Settings
|
||||
from .payment.cost_caculation import CostData, MaxCostData
|
||||
|
||||
from fastapi import BackgroundTasks, HTTPException, Request
|
||||
@@ -16,50 +17,415 @@ 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
|
||||
from .core.db import ApiKey, AsyncSession, ModelRow, UpstreamProviderRow, create_session
|
||||
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__)
|
||||
|
||||
|
||||
def init_upstreams(
|
||||
base_url: str, api_key: str, api_version: str | None = None
|
||||
) -> list[UpstreamProvider]:
|
||||
"""Initialize upstream providers based on settings.
|
||||
def resolve_model_alias(model_id: str, canonical_slug: str | None = None) -> list[str]:
|
||||
"""Resolve model ID to all possible aliases.
|
||||
|
||||
Returns list of aliases including canonical slug and variations without provider prefix.
|
||||
|
||||
Args:
|
||||
base_url: Base URL of the upstream API endpoint
|
||||
api_key: API key for authenticating with the upstream service
|
||||
api_version: API version for Azure OpenAI
|
||||
model_id: Model identifier (e.g., "gpt-5-mini" or "openai/gpt-5-mini")
|
||||
canonical_slug: Optional canonical slug from provider (e.g., "openai/gpt-5-pro-2025-10-06")
|
||||
|
||||
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
|
||||
|
||||
upstreams: list[UpstreamProvider] = []
|
||||
if settings.chat_completions_api_version:
|
||||
upstreams.append(
|
||||
AzureUpstreamProvider(
|
||||
settings.upstream_base_url,
|
||||
settings.upstream_api_key,
|
||||
settings.chat_completions_api_version,
|
||||
interval = getattr(settings, "models_refresh_interval_seconds", 0)
|
||||
if not interval or interval <= 0:
|
||||
logger.info("Provider models refresh disabled (interval <= 0)")
|
||||
return
|
||||
|
||||
while True:
|
||||
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:
|
||||
"""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):
|
||||
"""Initialize the upstream provider.
|
||||
|
||||
@@ -69,6 +435,8 @@ class UpstreamProvider:
|
||||
"""
|
||||
self.base_url = base_url
|
||||
self.api_key = api_key
|
||||
self._models_cache = []
|
||||
self._models_by_id = {}
|
||||
|
||||
def prepare_headers(self, request_headers: dict) -> dict:
|
||||
"""Prepare headers for upstream request by removing proxy-specific headers and adding authentication.
|
||||
@@ -136,6 +504,60 @@ class UpstreamProvider:
|
||||
"""
|
||||
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(
|
||||
self, body_bytes: bytes
|
||||
) -> tuple[str, str | None]:
|
||||
@@ -560,6 +982,8 @@ class UpstreamProvider:
|
||||
|
||||
url = f"{self.base_url}/{path}"
|
||||
|
||||
transformed_body = self.prepare_request_body(request_body)
|
||||
|
||||
logger.info(
|
||||
"Forwarding request to upstream",
|
||||
extra={
|
||||
@@ -578,13 +1002,13 @@ class UpstreamProvider:
|
||||
)
|
||||
|
||||
try:
|
||||
if request_body is not None:
|
||||
if transformed_body is not None:
|
||||
response = await client.send(
|
||||
client.build_request(
|
||||
request.method,
|
||||
url,
|
||||
headers=headers,
|
||||
content=request_body,
|
||||
content=transformed_body,
|
||||
params=self.prepare_params(path, request.query_params),
|
||||
),
|
||||
stream=True,
|
||||
@@ -1326,6 +1750,9 @@ class UpstreamProvider:
|
||||
|
||||
url = f"{self.base_url}/{path}"
|
||||
|
||||
request_body = await request.body()
|
||||
transformed_body = self.prepare_request_body(request_body)
|
||||
|
||||
logger.debug(
|
||||
"Forwarding request to upstream",
|
||||
extra={
|
||||
@@ -1347,7 +1774,7 @@ class UpstreamProvider:
|
||||
request.method,
|
||||
url,
|
||||
headers=headers,
|
||||
content=request.stream(),
|
||||
content=transformed_body if transformed_body else request_body,
|
||||
params=self.prepare_params(path, request.query_params),
|
||||
),
|
||||
stream=True,
|
||||
@@ -1546,12 +1973,66 @@ class UpstreamProvider:
|
||||
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):
|
||||
"""Upstream provider specifically configured for OpenAI API."""
|
||||
|
||||
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):
|
||||
@@ -1595,32 +2076,10 @@ class OpenRouterUpstreamProvider(UpstreamProvider):
|
||||
Args:
|
||||
api_key: OpenRouter API key for authentication
|
||||
"""
|
||||
self.upstream_name = "openrouter"
|
||||
super().__init__(base_url="https://openrouter.ai/api/v1", api_key=api_key)
|
||||
|
||||
async def fetch_models(self) -> dict:
|
||||
"""Fetch available models from OpenRouter API.
|
||||
|
||||
Returns:
|
||||
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
|
||||
async def fetch_models(self) -> list[Model]:
|
||||
"""Fetch all OpenRouter models."""
|
||||
models_data = await async_fetch_openrouter_models()
|
||||
return [Model(**model) for model in models_data] # type: ignore
|
||||
|
||||
Reference in New Issue
Block a user