From 61a0559f8e6ef462dace009295872588da64c9cb Mon Sep 17 00:00:00 2001 From: Shroominic Date: Mon, 13 Oct 2025 17:17:54 +0800 Subject: [PATCH] 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. --- ...3a4b5c6_create_upstream_providers_table.py | 45 + ...3b4c5d6_add_upstream_provider_to_models.py | 53 + routstr/core/admin.py | 967 +++++++++++++++++- routstr/core/db.py | 22 +- routstr/core/main.py | 21 +- routstr/payment/cost_caculation.py | 28 +- routstr/payment/helpers.py | 69 +- routstr/payment/models.py | 419 +++++--- routstr/proxy.py | 304 ++++-- routstr/upstream.py | 571 ++++++++++- 10 files changed, 2148 insertions(+), 351 deletions(-) create mode 100644 migrations/versions/d1e2f3a4b5c6_create_upstream_providers_table.py create mode 100644 migrations/versions/e1f2a3b4c5d6_add_upstream_provider_to_models.py diff --git a/migrations/versions/d1e2f3a4b5c6_create_upstream_providers_table.py b/migrations/versions/d1e2f3a4b5c6_create_upstream_providers_table.py new file mode 100644 index 00000000..9f36c39f --- /dev/null +++ b/migrations/versions/d1e2f3a4b5c6_create_upstream_providers_table.py @@ -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") diff --git a/migrations/versions/e1f2a3b4c5d6_add_upstream_provider_to_models.py b/migrations/versions/e1f2a3b4c5d6_add_upstream_provider_to_models.py new file mode 100644 index 00000000..523e8483 --- /dev/null +++ b/migrations/versions/e1f2a3b4c5d6_add_upstream_provider_to_models.py @@ -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), + ) diff --git a/routstr/core/admin.py b/routstr/core/admin.py index c6cc8ee0..bb54ffe7 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -8,7 +8,8 @@ from fastapi.responses import HTMLResponse from pydantic import BaseModel from sqlmodel import select -from ..payment.models import Model, get_model_by_id, list_models +from ..payment.models import Model, _row_to_model, list_models +from ..proxy import refresh_model_maps, reinitialize_upstreams from ..wallet import ( fetch_all_balances, get_proofs_per_mint_and_unit, @@ -16,7 +17,7 @@ from ..wallet import ( send_token, slow_filter_spend_proofs, ) -from .db import ApiKey, ModelRow, create_session +from .db import ApiKey, ModelRow, UpstreamProviderRow, create_session from .logging import get_logger from .settings import SettingsService, settings @@ -613,8 +614,8 @@ async def dashboard(request: Request) -> str: - + + + + + + + + + + + + + + +
IDTypeBase URLStatusActions
Loading…
+ + + + + + + + + + """ + ) + + +@admin_router.get("/upstream-providers", response_class=HTMLResponse) +async def admin_upstream_providers(request: Request) -> str: + if is_admin_authenticated(request): + return upstream_providers_page() + return admin_auth() + + @admin_router.get("/api/models", dependencies=[Depends(require_admin_api)]) async def get_models_admin_api(request: Request) -> list[dict[str, object]]: items = await list_models() return [m.dict() for m in items] # type: ignore +class ModelCreate(BaseModel): + id: str + name: str + description: str + created: int + context_length: int + architecture: dict[str, object] + pricing: dict[str, object] + per_request_limits: dict[str, object] | None = None + top_provider: dict[str, object] | None = None + upstream_provider_id: int | None = None + enabled: bool = True + + @admin_router.post("/api/models", dependencies=[Depends(require_admin_api)]) -async def create_model_admin_api(payload: Model) -> dict[str, object]: +async def create_model_admin_api(payload: ModelCreate) -> dict[str, object]: async with create_session() as session: exists = await session.get(ModelRow, payload.id) if exists: raise HTTPException( status_code=409, detail="Model with this ID already exists" ) - pricing_dict = payload.pricing.dict() - for k in ("max_prompt_cost", "max_completion_cost", "max_cost"): - pricing_dict.pop(k, None) row = ModelRow( id=payload.id, name=payload.name, description=payload.description, created=int(payload.created), context_length=int(payload.context_length), - architecture=json.dumps(payload.architecture.dict()), - pricing=json.dumps(pricing_dict), + architecture=json.dumps(payload.architecture), + pricing=json.dumps(payload.pricing), sats_pricing=None, per_request_limits=( json.dumps(payload.per_request_limits) @@ -1416,16 +2121,17 @@ async def create_model_admin_api(payload: Model) -> dict[str, object]: else None ), top_provider=( - json.dumps(payload.top_provider.dict()) - if payload.top_provider - else None + json.dumps(payload.top_provider) if payload.top_provider else None ), + upstream_provider_id=payload.upstream_provider_id, + enabled=payload.enabled, ) session.add(row) await session.commit() + await session.refresh(row) - created_model = await get_model_by_id(payload.id) - return created_model.dict() if created_model else {"id": payload.id} # type: ignore + await refresh_model_maps() + return _row_to_model(row).dict() # type: ignore @admin_router.post("/api/models/batch", dependencies=[Depends(require_admin_api)]) @@ -1475,6 +2181,8 @@ async def batch_create_models(payload: dict[str, object]) -> dict[str, int]: created += 1 if created: await session.commit() + if created: + await refresh_model_maps() return {"created": created, "skipped": skipped} @@ -1482,16 +2190,33 @@ async def batch_create_models(payload: dict[str, object]) -> dict[str, int]: "/api/models/{model_id:path}", dependencies=[Depends(require_admin_api)] ) async def get_model_admin_api(model_id: str) -> dict[str, object]: - model = await get_model_by_id(model_id) - if not model: - raise HTTPException(status_code=404, detail="Model not found") - return model.dict() # type: ignore + async with create_session() as session: + row = await session.get(ModelRow, model_id) + if not row: + raise HTTPException(status_code=404, detail="Model not found") + return _row_to_model(row).dict() # type: ignore + + +class ModelUpdate(BaseModel): + id: str + name: str + description: str + created: int + context_length: int + architecture: dict[str, object] + pricing: dict[str, object] + per_request_limits: dict[str, object] | None = None + top_provider: dict[str, object] | None = None + upstream_provider_id: int | None = None + enabled: bool = True @admin_router.patch( "/api/models/{model_id:path}", dependencies=[Depends(require_admin_api)] ) -async def update_model_admin_api(model_id: str, payload: Model) -> dict[str, object]: +async def update_model_admin_api( + model_id: str, payload: ModelUpdate +) -> dict[str, object]: if payload.id != model_id: raise HTTPException(status_code=400, detail="Path id does not match payload id") @@ -1504,11 +2229,8 @@ async def update_model_admin_api(model_id: str, payload: Model) -> dict[str, obj row.description = payload.description row.created = int(payload.created) row.context_length = int(payload.context_length) - row.architecture = json.dumps(payload.architecture.dict()) - pricing_dict = payload.pricing.dict() - for k in ("max_prompt_cost", "max_completion_cost", "max_cost"): - pricing_dict.pop(k, None) - row.pricing = json.dumps(pricing_dict) + row.architecture = json.dumps(payload.architecture) + row.pricing = json.dumps(payload.pricing) row.sats_pricing = None row.per_request_limits = ( json.dumps(payload.per_request_limits) @@ -1516,16 +2238,17 @@ async def update_model_admin_api(model_id: str, payload: Model) -> dict[str, obj else None ) row.top_provider = ( - json.dumps(payload.top_provider.dict()) if payload.top_provider else None + json.dumps(payload.top_provider) if payload.top_provider else None ) + row.upstream_provider_id = payload.upstream_provider_id + row.enabled = payload.enabled session.add(row) await session.commit() + await session.refresh(row) - updated = await get_model_by_id(model_id) - if not updated: - raise HTTPException(status_code=404, detail="Model not found after update") - return updated.dict() # type: ignore + await refresh_model_maps() + return _row_to_model(row).dict() # type: ignore @admin_router.delete( @@ -1538,6 +2261,7 @@ async def delete_model_admin_api(model_id: str) -> dict[str, object]: raise HTTPException(status_code=404, detail="Model not found") await session.delete(row) await session.commit() + await refresh_model_maps() return {"ok": True, "deleted_id": model_id} @@ -1549,9 +2273,188 @@ async def delete_all_models_admin_api() -> dict[str, object]: for row in rows: await session.delete(row) # type: ignore await session.commit() + await refresh_model_maps() return {"ok": True, "deleted": "all"} +class UpstreamProviderCreate(BaseModel): + provider_type: str + base_url: str + api_key: str + api_version: str | None = None + enabled: bool = True + + +class UpstreamProviderUpdate(BaseModel): + provider_type: str | None = None + base_url: str | None = None + api_key: str | None = None + api_version: str | None = None + enabled: bool | None = None + + +@admin_router.get("/api/upstream-providers", dependencies=[Depends(require_admin_api)]) +async def get_upstream_providers() -> list[dict[str, object]]: + async with create_session() as session: + result = await session.exec(select(UpstreamProviderRow)) + providers = result.all() + return [ + { + "id": p.id, + "provider_type": p.provider_type, + "base_url": p.base_url, + "api_key": "[REDACTED]" if p.api_key else "", + "api_version": p.api_version, + "enabled": p.enabled, + } + for p in providers + ] + + +@admin_router.post("/api/upstream-providers", dependencies=[Depends(require_admin_api)]) +async def create_upstream_provider( + payload: UpstreamProviderCreate, +) -> dict[str, object]: + async with create_session() as session: + result = await session.exec( + select(UpstreamProviderRow).where( + UpstreamProviderRow.base_url == payload.base_url + ) + ) + if result.first(): + raise HTTPException( + status_code=409, detail="Provider with this base URL already exists" + ) + + provider = UpstreamProviderRow( + provider_type=payload.provider_type, + base_url=payload.base_url, + api_key=payload.api_key, + api_version=payload.api_version, + enabled=payload.enabled, + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + + await reinitialize_upstreams() + return { + "id": provider.id, + "provider_type": provider.provider_type, + "base_url": provider.base_url, + "api_key": "[REDACTED]", + "api_version": provider.api_version, + "enabled": provider.enabled, + } + + +@admin_router.get( + "/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)] +) +async def get_upstream_provider(provider_id: int) -> dict[str, object]: + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + return { + "id": provider.id, + "provider_type": provider.provider_type, + "base_url": provider.base_url, + "api_key": "[REDACTED]" if provider.api_key else "", + "api_version": provider.api_version, + "enabled": provider.enabled, + } + + +@admin_router.patch( + "/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)] +) +async def update_upstream_provider( + provider_id: int, payload: UpstreamProviderUpdate +) -> dict[str, object]: + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + + if payload.provider_type is not None: + provider.provider_type = payload.provider_type + if payload.base_url is not None: + provider.base_url = payload.base_url + if payload.api_key is not None: + provider.api_key = payload.api_key + if payload.api_version is not None: + provider.api_version = payload.api_version + if payload.enabled is not None: + provider.enabled = payload.enabled + + session.add(provider) + await session.commit() + await session.refresh(provider) + + await reinitialize_upstreams() + return { + "id": provider.id, + "provider_type": provider.provider_type, + "base_url": provider.base_url, + "api_key": "[REDACTED]", + "api_version": provider.api_version, + "enabled": provider.enabled, + } + + +@admin_router.delete( + "/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)] +) +async def delete_upstream_provider(provider_id: int) -> dict[str, object]: + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + await session.delete(provider) + await session.commit() + await reinitialize_upstreams() + return {"ok": True, "deleted_id": provider_id} + + +@admin_router.get( + "/api/upstream-providers/{provider_id}/models", + dependencies=[Depends(require_admin_api)], +) +async def get_provider_models(provider_id: int) -> dict[str, object]: + from ..upstream import _instantiate_provider + + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + + db_models = await list_models( + session=session, upstream_id=provider_id, include_disabled=True + ) + + remote_models = [] + upstream_instance = _instantiate_provider(provider) + if upstream_instance: + try: + models = await upstream_instance.fetch_models() + remote_models = [m.dict() for m in models] + except Exception as e: + logger.error( + f"Failed to fetch models from {provider.provider_type}: {e}" + ) + + return { + "provider": { + "id": provider.id, + "provider_type": provider.provider_type, + "base_url": provider.base_url, + }, + "db_models": [m.dict() for m in db_models], + "remote_models": remote_models, + } + + DASHBOARD_CSS: str = """ * { margin: 0; padding: 0; box-sizing: border-box; } body { font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', sans-serif; background: #f5f7fa; color: #2c3e50; line-height: 1.6; padding: 2rem; } @@ -1591,8 +2494,8 @@ button:disabled { background: #a0aec0; cursor: not-allowed; transform: none; } @keyframes slideIn { from { transform: translateY(-20px); opacity: 0; } to { transform: translateY(0); opacity: 1; } } .close { color: #a0aec0; float: right; font-size: 28px; font-weight: bold; cursor: pointer; margin: -10px -10px 0 0; } .close:hover { color: #2d3748; } -input[type="number"], input[type="text"], select { width: 100%; padding: 10px; margin: 10px 0; border: 2px solid #e2e8f0; border-radius: 6px; font-size: 16px; transition: border 0.2s; } -input[type="number"]:focus, input[type="text"]:focus, select:focus { outline: none; border-color: #4299e1; } +input[type="number"], input[type="text"], input[type="password"], select { width: 100%; padding: 10px; margin: 10px 0; border: 2px solid #e2e8f0; border-radius: 6px; font-size: 16px; transition: border 0.2s; } +input[type="number"]:focus, input[type="text"]:focus, input[type="password"]:focus, select:focus { outline: none; border-color: #4299e1; } .warning { color: #e53e3e; font-weight: 600; margin: 10px 0; padding: 10px; background: #fff5f5; border-radius: 6px; } """ diff --git a/routstr/core/db.py b/routstr/core/db.py index 9f886791..744d5ed1 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -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( diff --git a/routstr/core/main.py b/routstr/core/main.py index 82a7a210..b0266743 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -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) diff --git a/routstr/payment/cost_caculation.py b/routstr/payment/cost_caculation.py index 50b82e53..bf4b42e9 100644 --- a/routstr/payment/cost_caculation.py +++ b/routstr/payment/cost_caculation.py @@ -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") diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index 37c7c8be..ef879b0a 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -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 diff --git a/routstr/payment/models.py b/routstr/payment/models.py index c064a8df..a5b0755e 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -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} diff --git a/routstr/proxy.py b/routstr/proxy.py index 6f8d976e..736d8bc4 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -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 diff --git a/routstr/upstream.py b/routstr/upstream.py index d9728017..d9edfa1a 100644 --- a/routstr/upstream.py +++ b/routstr/upstream.py @@ -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