mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +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 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
@@ -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)
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
|
||||||
|
|||||||
Reference in New Issue
Block a user