diff --git a/migrations/versions/c6d7e8f9a0b1_add_slug_to_upstream_providers.py b/migrations/versions/c6d7e8f9a0b1_add_slug_to_upstream_providers.py new file mode 100644 index 00000000..c140b557 --- /dev/null +++ b/migrations/versions/c6d7e8f9a0b1_add_slug_to_upstream_providers.py @@ -0,0 +1,88 @@ +"""add slug to upstream_providers + +Revision ID: c6d7e8f9a0b1 +Revises: b5e7c9d1f3a2 +Create Date: 2026-06-29 00:00:00.000000 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +from routstr.core.provider_slugs import provider_slug_base, provider_slug_candidate + +revision = "c6d7e8f9a0b1" +down_revision = "b5e7c9d1f3a2" +branch_labels = None +depends_on = None + + +def _allocate_backfill_slug(provider_type: str, reserved_slugs: set[str]) -> str: + base = provider_slug_base(provider_type) + suffix_number = 1 + while True: + candidate = provider_slug_candidate(base, suffix_number) + if candidate not in reserved_slugs: + reserved_slugs.add(candidate) + return candidate + suffix_number += 1 + + +def _backfill_provider_slugs(conn: sa.Connection) -> None: + existing_rows = conn.execute( + sa.text( + "SELECT slug FROM upstream_providers " + "WHERE slug IS NOT NULL AND slug != ''" + ) + ) + reserved_slugs = {str(row.slug).lower() for row in existing_rows} + + rows_to_backfill = conn.execute( + sa.text( + "SELECT id, provider_type FROM upstream_providers " + "WHERE slug IS NULL OR slug = '' " + "ORDER BY id" + ) + ) + for row in rows_to_backfill: + slug = _allocate_backfill_slug(str(row.provider_type), reserved_slugs) + conn.execute( + sa.text("UPDATE upstream_providers SET slug = :slug WHERE id = :id"), + {"slug": slug, "id": row.id}, + ) + + +def upgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + columns = {c["name"] for c in inspector.get_columns("upstream_providers")} + + if "slug" not in columns: + op.add_column( + "upstream_providers", + sa.Column("slug", sa.String(), nullable=True), + ) + + _backfill_provider_slugs(conn) + + existing_indexes = {idx["name"] for idx in inspector.get_indexes("upstream_providers")} + if "ix_upstream_providers_slug" not in existing_indexes: + op.create_index( + "ix_upstream_providers_slug", + "upstream_providers", + ["slug"], + unique=True, + ) + + +def downgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + existing_indexes = {idx["name"] for idx in inspector.get_indexes("upstream_providers")} + if "ix_upstream_providers_slug" in existing_indexes: + op.drop_index("ix_upstream_providers_slug", table_name="upstream_providers") + + columns = {c["name"] for c in inspector.get_columns("upstream_providers")} + if "slug" in columns: + op.drop_column("upstream_providers", "slug") diff --git a/pyproject.toml b/pyproject.toml index 694d3b79..68f621e7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "routstr" -version = "0.4.3" +version = "0.4.4" description = "Payment proxy for your LLM endpoint using cashu and nostr." readme = "README.md" requires-python = ">=3.11" diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 3fffc8c7..b0fc080e 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1,5 +1,6 @@ import asyncio import json +import re import secrets from datetime import datetime, timezone from pathlib import Path @@ -8,6 +9,7 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Request from pydantic import BaseModel, RootModel from pydantic.v1 import ValidationError as PydanticValidationError from sqlmodel import select +from sqlmodel.ext.asyncio.session import AsyncSession from ..payment.models import _row_to_model, list_models from ..proxy import refresh_model_maps, reinitialize_upstreams @@ -29,6 +31,7 @@ from .db import ( ) from .log_manager import log_manager from .logging import get_logger +from .provider_slugs import allocate_unique_provider_slug from .settings import SettingsService, settings logger = get_logger(__name__) @@ -456,19 +459,18 @@ class ModelCreate(BaseModel): dependencies=[Depends(require_admin_api)], ) async def upsert_provider_model( - provider_id: int, payload: ModelCreate + provider_id: str, payload: ModelCreate ) -> dict[str, object]: print(payload) logger.info( f"UPSERT_PROVIDER_MODEL called: provider_id={provider_id}, model_id={payload.id}" ) 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") + provider = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) # Try to get existing model - existing_row = await session.get(ModelRow, (payload.id, provider_id)) + existing_row = await session.get(ModelRow, (payload.id, provider_pk)) if existing_row: # Update existing model @@ -524,7 +526,7 @@ async def upsert_provider_model( alias_ids=( json.dumps(payload.alias_ids) if payload.alias_ids else None ), - upstream_provider_id=provider_id, + upstream_provider_id=provider_pk, enabled=payload.enabled, forwarded_model_id=payload.forwarded_model_id or payload.id, ) @@ -543,7 +545,7 @@ async def upsert_provider_model( dependencies=[Depends(require_admin_api)], ) async def update_provider_model_legacy( - provider_id: int, model_id: str, payload: ModelCreate + provider_id: str, model_id: str, payload: ModelCreate ) -> dict[str, object]: """Legacy PATCH endpoint - redirects to upsert POST endpoint for backward compatibility.""" logger.info( @@ -556,13 +558,12 @@ async def update_provider_model_legacy( "/api/upstream-providers/{provider_id}/models/{model_id:path}", dependencies=[Depends(require_admin_api)], ) -async def get_provider_model(provider_id: int, model_id: str) -> dict[str, object]: +async def get_provider_model(provider_id: str, model_id: str) -> 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") + provider = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) - row = await session.get(ModelRow, (model_id, provider_id)) + row = await session.get(ModelRow, (model_id, provider_pk)) if not row: raise HTTPException( status_code=404, detail="Model not found for this provider" @@ -576,9 +577,11 @@ async def get_provider_model(provider_id: int, model_id: str) -> dict[str, objec "/api/upstream-providers/{provider_id}/models/{model_id:path}", dependencies=[Depends(require_admin_api)], ) -async def delete_provider_model(provider_id: int, model_id: str) -> dict[str, object]: +async def delete_provider_model(provider_id: str, model_id: str) -> dict[str, object]: async with create_session() as session: - row = await session.get(ModelRow, (model_id, provider_id)) + provider = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) + row = await session.get(ModelRow, (model_id, provider_pk)) if not row: raise HTTPException( status_code=404, detail="Model not found for this provider" @@ -593,10 +596,12 @@ async def delete_provider_model(provider_id: int, model_id: str) -> dict[str, ob "/api/upstream-providers/{provider_id}/models", dependencies=[Depends(require_admin_api)], ) -async def delete_all_provider_models(provider_id: int) -> dict[str, object]: +async def delete_all_provider_models(provider_id: str) -> dict[str, object]: async with create_session() as session: + provider = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) result = await session.exec( - select(ModelRow).where(ModelRow.upstream_provider_id == provider_id) + select(ModelRow).where(ModelRow.upstream_provider_id == provider_pk) ) # type: ignore rows = result.all() for row in rows: @@ -615,7 +620,7 @@ class BatchOverrideRequest(BaseModel): dependencies=[Depends(require_admin_api)], ) async def batch_override_provider_models( - provider_id: int, payload: BatchOverrideRequest + provider_id: str, payload: BatchOverrideRequest ) -> dict[str, object]: """Batch override models for a specific provider.""" logger.info( @@ -623,15 +628,14 @@ async def batch_override_provider_models( ) 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") + provider = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) overridden_count = 0 for model_data in payload.models: # Try to get existing model regardless of whether it's enabled or not - existing_row = await session.get(ModelRow, (model_data.id, provider_id)) + existing_row = await session.get(ModelRow, (model_data.id, provider_pk)) if existing_row: # Update existing @@ -685,7 +689,7 @@ async def batch_override_provider_models( if model_data.alias_ids else None ), - upstream_provider_id=provider_id, + upstream_provider_id=provider_pk, enabled=model_data.enabled, ) session.add(row) @@ -702,6 +706,85 @@ async def batch_override_provider_models( } +_SLUG_PATTERN = re.compile(r"^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$") + + +def _validate_slug(value: str) -> str: + candidate = value.strip().lower() + if not _SLUG_PATTERN.fullmatch(candidate): + raise HTTPException( + status_code=400, + detail=( + "slug must be 3-64 chars, lowercase letters/digits/hyphens, " + "and may not start or end with a hyphen" + ), + ) + if candidate.isdigit(): + raise HTTPException( + status_code=400, + detail="slug must not be all digits", + ) + return candidate + + +async def _ensure_unique_slug( + session: AsyncSession, slug: str, exclude_id: int | None = None +) -> None: + stmt = select(UpstreamProviderRow).where(UpstreamProviderRow.slug == slug) + result = await session.exec(stmt) + existing = result.first() + if existing and existing.id != exclude_id: + raise HTTPException( + status_code=409, + detail="Provider with this slug already exists", + ) + + +async def _get_upstream_provider_by_ref( + session: AsyncSession, provider_ref: str +) -> UpstreamProviderRow: + if provider_ref.isdigit(): + provider = await session.get(UpstreamProviderRow, int(provider_ref)) + else: + slug = _validate_slug(provider_ref) + result = await session.exec( + select(UpstreamProviderRow).where(UpstreamProviderRow.slug == slug) + ) + provider = result.first() + + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + return provider + + +def _provider_pk(provider: UpstreamProviderRow) -> int: + if provider.id is None: + raise HTTPException(status_code=500, detail="Provider has no database id") + return provider.id + + +def _serialize_provider( + provider: UpstreamProviderRow, redact_api_key: bool = True +) -> dict[str, object]: + return { + "id": provider.id, + "slug": provider.slug, + "provider_type": provider.provider_type, + "base_url": provider.base_url, + "api_key": "[REDACTED]" + if (redact_api_key and provider.api_key) + else provider.api_key + if not redact_api_key + else "", + "api_version": provider.api_version, + "enabled": provider.enabled, + "provider_fee": provider.provider_fee, + "provider_settings": json.loads(provider.provider_settings) + if provider.provider_settings + else None, + } + + class UpstreamProviderCreate(BaseModel): provider_type: str base_url: str @@ -710,6 +793,7 @@ class UpstreamProviderCreate(BaseModel): enabled: bool = True provider_fee: float = 1.01 provider_settings: dict | None = None + slug: str | None = None class UpstreamProviderUpdate(BaseModel): @@ -720,6 +804,50 @@ class UpstreamProviderUpdate(BaseModel): enabled: bool | None = None provider_fee: float | None = None provider_settings: dict | None = None + slug: str | None = None + + +class UpstreamProviderUpdateBySlug(BaseModel): + slug: str + new_slug: str | None = None + provider_type: str | None = None + base_url: str | None = None + api_key: str | None = None + api_version: str | None = None + enabled: bool | None = None + provider_fee: float | None = None + provider_settings: dict | None = None + + +async def _apply_provider_update( + session: AsyncSession, + provider: UpstreamProviderRow, + payload: UpstreamProviderUpdate, + new_slug: str | None = None, +) -> None: + if new_slug is not None: + validated = _validate_slug(new_slug) + await _ensure_unique_slug(session, validated, exclude_id=provider.id) + provider.slug = validated + + 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 + if payload.provider_fee is not None: + provider.provider_fee = payload.provider_fee + if payload.provider_settings is not None: + provider.provider_settings = json.dumps(payload.provider_settings) + + session.add(provider) + await session.commit() + await session.refresh(provider) @admin_router.get("/api/upstream-providers", dependencies=[Depends(require_admin_api)]) @@ -727,21 +855,7 @@ 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, - "provider_fee": p.provider_fee, - "provider_settings": json.loads(p.provider_settings) - if p.provider_settings - else None, - } - for p in providers - ] + return [_serialize_provider(p) for p in providers] @admin_router.post("/api/upstream-providers", dependencies=[Depends(require_admin_api)]) @@ -761,7 +875,14 @@ async def create_upstream_provider( detail="Provider with this base URL and API key already exists", ) + if payload.slug: + slug = _validate_slug(payload.slug) + await _ensure_unique_slug(session, slug) + else: + slug = await allocate_unique_provider_slug(session, payload.provider_type) + provider = UpstreamProviderRow( + slug=slug, provider_type=payload.provider_type, base_url=payload.base_url, api_key=payload.api_key, @@ -778,99 +899,81 @@ async def create_upstream_provider( await reinitialize_upstreams() await refresh_model_maps() - 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, - "provider_fee": provider.provider_fee, - "provider_settings": payload.provider_settings, - } + return _serialize_provider(provider) @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 def get_upstream_provider(provider_id: str) -> 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, - "provider_fee": provider.provider_fee, - "provider_settings": json.loads(provider.provider_settings) - if provider.provider_settings - else None, - } + provider = await _get_upstream_provider_by_ref(session, provider_id) + return _serialize_provider(provider) @admin_router.patch( "/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)] ) async def update_upstream_provider( - provider_id: int, payload: UpstreamProviderUpdate + provider_id: str, 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") + provider = await _get_upstream_provider_by_ref(session, provider_id) - 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 - if payload.provider_fee is not None: - provider.provider_fee = payload.provider_fee - if payload.provider_settings is not None: - provider.provider_settings = json.dumps(payload.provider_settings) - - session.add(provider) - await session.commit() - await session.refresh(provider) + await _apply_provider_update(session, provider, payload, new_slug=payload.slug) await reinitialize_upstreams() await refresh_model_maps() - 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, - "provider_fee": provider.provider_fee, - "provider_settings": json.loads(provider.provider_settings) - if provider.provider_settings - else None, - } + return _serialize_provider(provider) + + +@admin_router.patch( + "/api/upstream-providers", dependencies=[Depends(require_admin_api)] +) +async def update_upstream_provider_by_slug( + payload: UpstreamProviderUpdateBySlug, +) -> dict[str, object]: + lookup = _validate_slug(payload.slug) + async with create_session() as session: + result = await session.exec( + select(UpstreamProviderRow).where( + UpstreamProviderRow.slug == lookup + ) + ) + provider = result.first() + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + + update_payload = UpstreamProviderUpdate( + provider_type=payload.provider_type, + base_url=payload.base_url, + api_key=payload.api_key, + api_version=payload.api_version, + enabled=payload.enabled, + provider_fee=payload.provider_fee, + provider_settings=payload.provider_settings, + ) + await _apply_provider_update( + session, provider, update_payload, new_slug=payload.new_slug + ) + + await reinitialize_upstreams() + await refresh_model_maps() + return _serialize_provider(provider) @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 def delete_upstream_provider(provider_id: str) -> 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") + provider = await _get_upstream_provider_by_ref(session, provider_id) + deleted_id = _provider_pk(provider) await session.delete(provider) await session.commit() await reinitialize_upstreams() await refresh_model_maps() - return {"ok": True, "deleted_id": provider_id} + return {"ok": True, "deleted_id": deleted_id} @admin_router.get("/api/provider-types", dependencies=[Depends(require_admin_api)]) @@ -885,17 +988,16 @@ async def get_provider_types() -> list[dict[str, object]]: "/api/upstream-providers/{provider_id}/models", dependencies=[Depends(require_admin_api)], ) -async def get_provider_models(provider_id: int) -> dict[str, object]: +async def get_provider_models(provider_id: str) -> dict[str, object]: from ..upstream.helpers 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") + provider = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) db_models = await list_models( session=session, - upstream_id=provider_id, + upstream_id=provider_pk, include_disabled=True, apply_fees=False, ) @@ -985,13 +1087,11 @@ class TopupTokenRequest(BaseModel): dependencies=[Depends(require_admin_api)], ) async def topup_provider_with_token( - provider_id: int, payload: TopupTokenRequest + provider_id: str, payload: TopupTokenRequest ) -> dict: """Redeem a Cashu token for an upstream 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") + provider = await _get_upstream_provider_by_ref(session, provider_id) import httpx @@ -1022,15 +1122,13 @@ async def topup_provider_with_token( dependencies=[Depends(require_admin_api)], ) async def initiate_provider_topup( - provider_id: int, payload: TopupRequest + provider_id: str, payload: TopupRequest ) -> dict[str, object]: """Initiate a Lightning Network top-up for the upstream provider account.""" from ..upstream.helpers 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") + provider = await _get_upstream_provider_by_ref(session, provider_id) try: logger.info( @@ -1150,15 +1248,13 @@ async def initiate_provider_topup( "/api/upstream-providers/{provider_id}/topup/{invoice_id}/status", dependencies=[Depends(require_admin_api)], ) -async def check_topup_status(provider_id: int, invoice_id: str) -> dict[str, object]: +async def check_topup_status(provider_id: str, invoice_id: str) -> dict[str, object]: """Check the status of a Lightning Network top-up invoice.""" from ..upstream.helpers import _instantiate_provider from ..upstream.ppqai import PPQAIUpstreamProvider 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") + provider = await _get_upstream_provider_by_ref(session, provider_id) # For Routstr providers, proxy the status check if provider.provider_type == "routstr": @@ -1205,14 +1301,12 @@ async def check_topup_status(provider_id: int, invoice_id: str) -> dict[str, obj "/api/upstream-providers/{provider_id}/balance", dependencies=[Depends(require_admin_api)], ) -async def get_provider_balance(provider_id: int) -> dict[str, object]: +async def get_provider_balance(provider_id: str) -> dict[str, object]: """Get the current balance for an upstream provider account.""" from ..upstream.helpers 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") + provider = await _get_upstream_provider_by_ref(session, provider_id) # For Routstr providers, proxy the balance check if provider.provider_type == "routstr": @@ -1591,15 +1685,13 @@ async def get_lightning_invoices_api( "/api/upstream-providers/{provider_id}/routstr/refund", dependencies=[Depends(require_admin_api)], ) -async def refund_routstr_provider_balance(provider_id: int) -> dict[str, object]: +async def refund_routstr_provider_balance(provider_id: str) -> dict[str, object]: """Refund balance from an upstream Routstr provider back to the local wallet.""" from ..upstream.helpers import _instantiate_provider from ..upstream.routstr import RoutstrUpstreamProvider async with create_session() as session: - provider_row = await session.get(UpstreamProviderRow, provider_id) - if not provider_row: - raise HTTPException(status_code=404, detail="Provider not found") + provider_row = await _get_upstream_provider_by_ref(session, provider_id) if provider_row.provider_type != "routstr": raise HTTPException( diff --git a/routstr/core/db.py b/routstr/core/db.py index 282b8b93..586f467e 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -319,6 +319,12 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore ), ) id: int | None = Field(default=None, primary_key=True) + slug: str | None = Field( + default=None, + unique=True, + index=True, + description="Stable external slug used for updates via API key.", + ) provider_type: str = Field( description="Provider type: custom, openai, anthropic, azure, openrouter, etc." ) diff --git a/routstr/core/provider_slugs.py b/routstr/core/provider_slugs.py new file mode 100644 index 00000000..0efb919f --- /dev/null +++ b/routstr/core/provider_slugs.py @@ -0,0 +1,65 @@ +from __future__ import annotations + +import re +from itertools import count +from typing import Collection + +from sqlmodel import select +from sqlmodel.ext.asyncio.session import AsyncSession + +from .db import UpstreamProviderRow + +_SLUG_BASE_PATTERN = re.compile(r"[^a-z0-9]+") +_MAX_SLUG_LENGTH = 64 + + +def provider_slug_base(provider_type: str) -> str: + """Return a deterministic slug base for a provider type.""" + base = _SLUG_BASE_PATTERN.sub("-", provider_type.lower()).strip("-") + if not base: + base = "provider" + elif base.isdigit(): + base = f"provider-{base}" + elif len(base) < 3: + base = f"{base}-provider" + + if len(base) > _MAX_SLUG_LENGTH: + base = base[:_MAX_SLUG_LENGTH].rstrip("-") or "provider" + return base + + +def provider_slug_candidate(base: str, suffix_number: int) -> str: + if suffix_number == 1: + return base + + suffix = f"-{suffix_number}" + max_base_length = _MAX_SLUG_LENGTH - len(suffix) + return f"{base[:max_base_length].rstrip('-')}{suffix}" + + +async def allocate_unique_provider_slug( + session: AsyncSession, + provider_type: str, + reserved_slugs: Collection[str] = (), +) -> str: + """Allocate a stable, deterministic provider slug. + + The first provider of a type gets ``openai``; later collisions get + ``openai-2``, ``openai-3``, etc. ``reserved_slugs`` covers rows staged in + memory but not flushed yet, such as settings/env seeding. + """ + base = provider_slug_base(provider_type) + reserved = {slug.lower() for slug in reserved_slugs} + + for suffix_number in count(1): + candidate = provider_slug_candidate(base, suffix_number) + if candidate in reserved: + continue + + result = await session.exec( + select(UpstreamProviderRow).where(UpstreamProviderRow.slug == candidate) + ) + if result.first() is None: + return candidate + + raise RuntimeError("unreachable") diff --git a/routstr/core/version.py b/routstr/core/version.py index e66e645e..30ee4c6c 100644 --- a/routstr/core/version.py +++ b/routstr/core/version.py @@ -20,7 +20,7 @@ import subprocess from functools import lru_cache from pathlib import Path -BASE_VERSION = "0.4.3" +BASE_VERSION = "0.4.4" _REPO_ROOT = Path(__file__).resolve().parents[2] _GIT_TIMEOUT_SECONDS = 2.0 diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 820bd174..84f3847f 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -41,6 +41,28 @@ class CostDataError(BaseModel): code: str +def _empty_cost(cls: type[CostData] = CostData) -> CostData: + """Build an all-zero cost object — a full refund for an empty response. + + Shared by the two paths that must not bill: an upstream response with no + usage data at all, and one that reports a USD cost but carries zero tokens + in every bucket. + """ + return cls( + base_msats=0, + input_msats=0, + output_msats=0, + total_msats=0, + total_usd=0.0, + input_tokens=0, + output_tokens=0, + cache_read_input_tokens=0, + cache_creation_input_tokens=0, + cache_read_msats=0, + cache_creation_msats=0, + ) + + async def calculate_cost( response_data: dict, max_cost: int, @@ -83,19 +105,7 @@ async def calculate_cost( else None, }, ) - return MaxCostData( - base_msats=0, - input_msats=0, - output_msats=0, - total_msats=0, - total_usd=0.0, - input_tokens=0, - output_tokens=0, - cache_read_input_tokens=0, - cache_creation_input_tokens=0, - cache_read_msats=0, - cache_creation_msats=0, - ) + return _empty_cost(MaxCostData) usage_data = response_data.get("usage") or {} if not isinstance(usage_data, dict): @@ -109,6 +119,27 @@ async def calculate_cost( # Try USD cost first usd_cost = _resolve_usd_cost(usage_data, response_data) if usd_cost > 0: + truly_empty = ( + input_tokens == 0 + and output_tokens == 0 + and cache_read_tokens == 0 + and cache_creation_tokens == 0 + ) + if truly_empty: + logger.warning( + "Upstream reported a USD cost but the response carries no " + "tokens at all (input, output, cache-read and cache-creation " + "are all zero) — refunding in full rather than billing the " + "USD-derived cost for an empty response.", + extra={ + "model": response_data.get("model", "unknown"), + "usd_cost": usd_cost, + "usage_keys": sorted(usage_data.keys()) + if isinstance(usage_data, dict) + else None, + }, + ) + return _empty_cost() if input_tokens == 0 and output_tokens == 0: logger.warning( "Upstream reported a USD cost but no token counts — " diff --git a/routstr/payment/models.py b/routstr/payment/models.py index e62f7875..17e8ac64 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -85,6 +85,30 @@ class Model(BaseModel): return hash(self.id) +def litellm_cost_entry(model_id: str) -> dict | None: + """Look up ``model_id`` in litellm's bundled cost map. + + litellm ships per-model USD rates keyed by the exact OpenRouter id + (``deepseek/deepseek-chat``) or the bare model name (``gpt-4o``, + ``claude-sonnet-4-5``), so both spellings are tried. Keys are lowercase, so + a mixed-case upstream id (``deepseek-ai/DeepSeek-V4-Flash``) is retried via + a case-insensitive scan. Returns the matched cost dict, or ``None``. + """ + import litellm + + candidates = (model_id, model_id.split("/", 1)[-1]) + for key in candidates: + info = litellm.model_cost.get(key) + if isinstance(info, dict): + return info + + lowered = {c.lower() for c in candidates} + for key, info in litellm.model_cost.items(): + if isinstance(key, str) and key.lower() in lowered and isinstance(info, dict): + return info + return None + + def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing: """Fill missing cache rates from litellm's bundled cost map. @@ -92,9 +116,8 @@ def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing: for many models (most DeepSeek entries, openai/gpt-4o, ...). Without a cache rate, billing falls back to the full input rate, which overcharges cache reads (DeepSeek hits are 10x cheaper) and undercharges Anthropic - cache writes (1.25x). litellm ships per-model USD rates keyed by the exact - OpenRouter id (deepseek/deepseek-chat) or by the bare model name - (gpt-4o, claude-sonnet-4-5), so both spellings are tried. + cache writes (1.25x). The lookup (see ``litellm_cost_entry``) tries both + id spellings and a case-insensitive fallback. Rates already present (e.g. provided by OpenRouter) are authoritative and never overwritten. Unknown models are returned unchanged. @@ -104,14 +127,7 @@ def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing: if not (needs_read or needs_write): return pricing - import litellm - - info: dict | None = None - for key in (model_id, model_id.split("/", 1)[-1]): - candidate = litellm.model_cost.get(key) - if isinstance(candidate, dict): - info = candidate - break + info = litellm_cost_entry(model_id) if info is None: return pricing @@ -215,13 +231,29 @@ def _row_to_model( ) top_provider_dict = json.loads(row.top_provider) if row.top_provider else None - if apply_provider_fee and isinstance(pricing, dict): - pricing = {k: float(v) * provider_fee for k, v in pricing.items()} - if isinstance(pricing, dict) and float(pricing.get("request", 0.0)) <= 0.0: pricing["request"] = max(pricing.get("request", 0.0), 0.0) parsed_pricing = Pricing.parse_obj(pricing) + + # Fill missing cache-read/write rates from litellm's cost map BEFORE applying + # the provider fee, so they carry the same markup as every other component. + # DB-stored override pricing (e.g. generic providers) omits cache rates; + # without this, ``_row_to_model`` bills cache reads at the full input rate — + # the ``_apply_provider_fee_to_model`` path backfills, but the override path + # used for admin-configured providers did not. + # + # Key on ``forwarded_model_id`` (the actual upstream model name litellm + # prices) when set: an alias row (id="local-alias", + # forwarded_model_id="deepseek-v4-flash") would otherwise look up the alias + # and miss the cache rate. + pricing_model_id = getattr(row, "forwarded_model_id", None) or row.id + parsed_pricing = backfill_cache_pricing(pricing_model_id, parsed_pricing) + + if apply_provider_fee: + parsed_pricing = Pricing.parse_obj( + {k: float(v) * provider_fee for k, v in parsed_pricing.dict().items()} + ) model = Model( id=row.id, name=row.name, diff --git a/routstr/upstream/anthropic.py b/routstr/upstream/anthropic.py index 5e48f058..7944cbd5 100644 --- a/routstr/upstream/anthropic.py +++ b/routstr/upstream/anthropic.py @@ -24,7 +24,7 @@ class AnthropicUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def from_db_row( + def _build_from_row( cls, provider_row: "UpstreamProviderRow" ) -> "AnthropicUpstreamProvider": return cls( diff --git a/routstr/upstream/auto_topup.py b/routstr/upstream/auto_topup.py index 31397882..3582b3fc 100644 --- a/routstr/upstream/auto_topup.py +++ b/routstr/upstream/auto_topup.py @@ -97,6 +97,8 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: # Instantiate provider and check balance provider = RoutstrUpstreamProvider.from_db_row(row) + if provider is None: + return balance = await provider.get_balance() if balance is None: diff --git a/routstr/upstream/azure.py b/routstr/upstream/azure.py index 11cee17c..a693b763 100644 --- a/routstr/upstream/azure.py +++ b/routstr/upstream/azure.py @@ -38,7 +38,7 @@ class AzureUpstreamProvider(BaseUpstreamProvider): self.api_version = api_version @classmethod - def from_db_row( + def _build_from_row( cls, provider_row: "UpstreamProviderRow" ) -> "AzureUpstreamProvider | None": if not provider_row.api_version: diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 2467044c..51afb080 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -6,13 +6,12 @@ import math import traceback import uuid from collections.abc import AsyncGenerator, AsyncIterator, Iterator -from typing import Any, Mapping, cast +from typing import Any, Mapping, Self, cast import httpx from fastapi import BackgroundTasks, HTTPException, Request from fastapi.responses import Response, StreamingResponse from pydantic.v1 import BaseModel -from sqlmodel import select from ..auth import adjust_payment_for_tokens from ..core import get_logger @@ -92,6 +91,11 @@ class BaseUpstreamProvider: base_url: str api_key: str provider_fee: float = 1.05 + # Primary key of the ``upstream_providers`` row this instance was built + # from. Set by ``from_db_row`` so a live provider can re-find its own row by + # stable identity instead of its rotatable ``api_key``. ``None`` for + # instances not sourced from a row. + db_id: int | None = None _models_cache: list[Model] = [] _models_by_id: dict[str, Model] = {} @@ -106,6 +110,7 @@ class BaseUpstreamProvider: self.base_url = base_url self.api_key = api_key self.provider_fee = provider_fee + self.db_id = None self._models_cache = [] self._models_by_id = {} @@ -123,10 +128,13 @@ class BaseUpstreamProvider: return detect_litellm_prefix(self.base_url) @classmethod - def from_db_row( - cls, provider_row: "UpstreamProviderRow" - ) -> "BaseUpstreamProvider | None": - """Factory method to instantiate provider from database row. + def from_db_row(cls, provider_row: "UpstreamProviderRow") -> "Self | None": + """Instantiate a provider from a database row, carrying its identity. + + Construction itself is delegated to the ``_build_from_row`` hook (which + subclasses override to match their constructor); this wrapper stamps the + row's primary key onto the instance as ``db_id`` so the provider can + later re-find its own row by identity rather than by its ``api_key``. Args: provider_row: Database row containing provider configuration @@ -134,6 +142,19 @@ class BaseUpstreamProvider: Returns: Instantiated provider or None if instantiation fails """ + provider = cls._build_from_row(provider_row) + if provider is not None: + provider.db_id = provider_row.id + return provider + + @classmethod + def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "Self | None": + """Construct the provider instance from a row (no identity stamping). + + Overridden by subclasses whose constructors differ from the base + ``(base_url, api_key, provider_fee)`` shape. Callers should use + ``from_db_row`` instead, which also attaches ``db_id``. + """ return cls( base_url=provider_row.base_url, api_key=provider_row.api_key, @@ -4878,14 +4899,11 @@ class BaseUpstreamProvider: """Refresh the in-memory models cache from upstream API.""" try: async with create_session() as session: - stmt = select(UpstreamProviderRow).where( - UpstreamProviderRow.base_url == self.base_url, - UpstreamProviderRow.api_key == self.api_key, + provider = ( + await session.get(UpstreamProviderRow, self.db_id) + if self.db_id is not None + else None ) - result = await session.exec(stmt) - - # .first() returns the object or None if not found - provider = result.first() if not provider or not provider.id: raise HTTPException(status_code=404, detail="Provider not found") diff --git a/routstr/upstream/deepseek_v4_pricing_shim.py b/routstr/upstream/deepseek_v4_pricing_shim.py index 8720fb40..ba0c392d 100644 --- a/routstr/upstream/deepseek_v4_pricing_shim.py +++ b/routstr/upstream/deepseek_v4_pricing_shim.py @@ -3,10 +3,15 @@ litellm's bundled cost map does not yet ship ``deepseek-v4-flash`` / ``deepseek-v4-pro``. Without an entry, ``backfill_cache_pricing`` cannot find a ``cache_read_input_token_cost`` and cache reads fall back to the full input -rate — a ~60% overcharge on cache hits (DeepSeek hits are ~0.2x input). +rate — a large overcharge on cache hits (DeepSeek V4 hits are ~0.008-0.02x +input, i.e. cached tokens cost 50-120x less than regular input). This module injects the missing entries into ``litellm.model_cost`` at startup -so the existing backfill path resolves them. Rates mirror the open upstream PR +so the existing backfill path resolves them. Rates mirror the canonical +``deepseek`` provider entries now in litellm's ``model_prices`` map +(``input_cost_per_token`` is the cache-*miss* rate; +``cache_read_input_token_cost`` is the cache-*hit* rate), sourced from +https://api-docs.deepseek.com/quick_start/pricing via https://github.com/BerriAI/litellm/pull/26380 (issue https://github.com/BerriAI/litellm/issues/30430). @@ -22,21 +27,23 @@ from ..core import get_logger logger = get_logger(__name__) -# USD per token. Source: BerriAI/litellm PR #26380. +# USD per token. Mirrors the canonical ``deepseek`` provider entries in +# litellm's model_prices map (source: DeepSeek API pricing docs). Keep these in +# sync with ``litellm.model_cost["deepseek/deepseek-v4-*"]``. _DEEPSEEK_V4_RATES: dict[str, dict[str, float]] = { "deepseek-v4-flash": { "input_cost_per_token": 1.4e-07, "output_cost_per_token": 2.8e-07, - "cache_read_input_token_cost": 2.8e-08, + "cache_read_input_token_cost": 2.8e-09, "cache_creation_input_token_cost": 0.0, - "input_cost_per_token_cache_hit": 2.8e-08, + "input_cost_per_token_cache_hit": 2.8e-09, }, "deepseek-v4-pro": { - "input_cost_per_token": 1.74e-06, - "output_cost_per_token": 3.48e-06, - "cache_read_input_token_cost": 1.4e-07, + "input_cost_per_token": 4.35e-07, + "output_cost_per_token": 8.7e-07, + "cache_read_input_token_cost": 3.625e-09, "cache_creation_input_token_cost": 0.0, - "input_cost_per_token_cache_hit": 1.4e-07, + "input_cost_per_token_cache_hit": 3.625e-09, }, } diff --git a/routstr/upstream/fireworks.py b/routstr/upstream/fireworks.py index e0bb1fec..92b67838 100644 --- a/routstr/upstream/fireworks.py +++ b/routstr/upstream/fireworks.py @@ -20,7 +20,7 @@ class FireworksUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def from_db_row( + def _build_from_row( cls, provider_row: "UpstreamProviderRow" ) -> "FireworksUpstreamProvider": return cls( diff --git a/routstr/upstream/gemini.py b/routstr/upstream/gemini.py index 54bf849b..54de41a3 100644 --- a/routstr/upstream/gemini.py +++ b/routstr/upstream/gemini.py @@ -51,7 +51,7 @@ class GeminiUpstreamProvider(BaseUpstreamProvider): return self._client @classmethod - def from_db_row( + def _build_from_row( cls, provider_row: "UpstreamProviderRow" ) -> "GeminiUpstreamProvider": return cls( diff --git a/routstr/upstream/generic.py b/routstr/upstream/generic.py index 390c8372..3faf80e9 100644 --- a/routstr/upstream/generic.py +++ b/routstr/upstream/generic.py @@ -5,6 +5,12 @@ from typing import TYPE_CHECKING import httpx from .base import BaseUpstreamProvider +from .pricing_resolver import ( + FallbackPricingResolver, + ResolvedPricing, + _as_float, + estimate_context_length, +) if TYPE_CHECKING: from ..core.db import UpstreamProviderRow @@ -45,7 +51,7 @@ class GenericUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def from_db_row( + def _build_from_row( cls, provider_row: "UpstreamProviderRow" ) -> "GenericUpstreamProvider": return cls( @@ -64,6 +70,40 @@ class GenericUpstreamProvider(BaseUpstreamProvider): "platform_url": cls.platform_url, } + def _native_pricing( + self, model_id: str, model_spec: dict + ) -> ResolvedPricing | None: + """Read pricing/metadata from Venice's bespoke ``model_spec`` schema. + + Returns ``None`` when the upstream reported no *usable* native price — + absent, non-numeric, negative, or both-zero — so the caller falls + through to the shared resolution chain instead of fabricating a number + or trusting a bogus one. This mirrors the money-safety guards the + litellm and OpenRouter rungs already apply: a both-zero price would + serve the model free, a negative one would credit the caller, and a + non-numeric string would otherwise throw and drop the whole catalog. + """ + pricing_info = model_spec.get("pricing", {}) + input_usd = _as_float(pricing_info.get("input", {}).get("usd")) + output_usd = _as_float(pricing_info.get("output", {}).get("usd")) + if input_usd is None or output_usd is None: + return None + if input_usd < 0 or output_usd < 0 or (input_usd == 0 and output_usd == 0): + return None + + capabilities = model_spec.get("capabilities", {}) + input_modalities = ["text"] + if capabilities.get("supportsVision", False): + input_modalities.append("image") + + return ResolvedPricing( + prompt=input_usd / 1_000_000, + completion=output_usd / 1_000_000, + context_length=model_spec.get("availableContextTokens"), + source="native", + input_modalities=input_modalities, + ) + async def fetch_models(self) -> list[Model]: """Fetch models from upstream API using /models endpoint.""" from ..payment.models import Architecture, Model, Pricing, TopProvider @@ -78,6 +118,7 @@ class GenericUpstreamProvider(BaseUpstreamProvider): response.raise_for_status() data = response.json() + resolver = FallbackPricingResolver() models_list = [] for model_data in data.get("data", []): model_id = model_data.get("id", "") @@ -89,41 +130,44 @@ class GenericUpstreamProvider(BaseUpstreamProvider): owned_by = model_data.get("owned_by", "unknown") model_spec = model_data.get("model_spec", {}) - context_length = 4096 - if model_spec.get("availableContextTokens"): - context_length = model_spec["availableContextTokens"] - elif any( - pattern in model_id.lower() for pattern in ["32k", "32000"] - ): - context_length = 32768 - elif any( - pattern in model_id.lower() for pattern in ["16k", "16000"] - ): - context_length = 16384 - elif any(pattern in model_id.lower() for pattern in ["8k", "8000"]): - context_length = 8192 - elif "gpt-4" in model_id.lower(): - context_length = 8192 - elif "claude" in model_id.lower(): - context_length = 200000 + resolved = self._native_pricing(model_id, model_spec) + if resolved is None: + resolved = await resolver.resolve(model_id) - pricing_info = model_spec.get("pricing", {}) - input_pricing = pricing_info.get("input", {}) - output_pricing = pricing_info.get("output", {}) + if resolved is None: + # Fail closed: never invent a price. Import the model + # disabled with a warning so the operator can price it + # (the admin UI surfaces disabled remote models). + logger.warning( + f"No pricing source resolved for '{model_id}' from " + f"{self.upstream_name}; importing it disabled", + extra={"model_id": model_id, "base_url": self.base_url}, + ) + resolved = ResolvedPricing( + prompt=0.0, + completion=0.0, + context_length=None, + source="unresolved", + ) + enabled = False + else: + enabled = True - prompt_price = input_pricing.get("usd", 0.001) / 1000000 - completion_price = output_pricing.get("usd", 0.001) / 1000000 + # Prefer the source's own modality string (OpenRouter ships + # one, e.g. "text+image->text"); otherwise derive it from the + # captured input/output modalities in the same "in->out" shape + # rather than flattening vision models to "text->text". + modality = resolved.modality or ( + f"{'+'.join(resolved.input_modalities)}" + f"->{'+'.join(resolved.output_modalities)}" + ) - capabilities = model_spec.get("capabilities", {}) - input_modalities = ["text"] - output_modalities = ["text"] - - if capabilities.get("supportsVision", False): - input_modalities.append("image") - - modality = "text" - if capabilities.get("supportsVision", False): - modality = "text->text" + # A source can carry a price but no context (e.g. a litellm + # entry missing max_input_tokens); fall back to an id-based + # estimate so we never persist a zero-length window. + context_length = resolved.context_length or estimate_context_length( + model_id + ) spec_name = model_spec.get("name", model_name) description = f"{spec_name}" @@ -139,30 +183,33 @@ class GenericUpstreamProvider(BaseUpstreamProvider): context_length=context_length, architecture=Architecture( modality=modality, - input_modalities=input_modalities, - output_modalities=output_modalities, - tokenizer="unknown", - instruct_type=None, + input_modalities=resolved.input_modalities, + output_modalities=resolved.output_modalities, + tokenizer=resolved.tokenizer, + instruct_type=resolved.instruct_type, ), pricing=Pricing( - prompt=prompt_price, - completion=completion_price, + prompt=resolved.prompt, + completion=resolved.completion, request=0.0, image=0.0, web_search=0.0, internal_reasoning=0.0, - max_prompt_cost=0.001, - max_completion_cost=0.001, - max_cost=0.001, + input_cache_read=resolved.input_cache_read, + input_cache_write=resolved.input_cache_write, ), sats_pricing=None, per_request_limits=None, top_provider=TopProvider( context_length=context_length, - max_completion_tokens=context_length // 2, - is_moderated=False, + max_completion_tokens=( + resolved.max_completion_tokens + if resolved.max_completion_tokens is not None + else context_length // 2 + ), + is_moderated=bool(resolved.is_moderated), ), - enabled=True, + enabled=enabled, upstream_provider_id=None, canonical_slug=None, ) diff --git a/routstr/upstream/groq.py b/routstr/upstream/groq.py index 4020b2df..17103c35 100644 --- a/routstr/upstream/groq.py +++ b/routstr/upstream/groq.py @@ -20,7 +20,7 @@ class GroqUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def from_db_row(cls, provider_row: "UpstreamProviderRow") -> "GroqUpstreamProvider": + def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "GroqUpstreamProvider": return cls( api_key=provider_row.api_key, provider_fee=provider_row.provider_fee, diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index 9ad4d73f..c2bc2acb 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -12,6 +12,7 @@ from sqlmodel import select from ..core import get_logger from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session +from ..core.provider_slugs import allocate_unique_provider_slug from ..payment.models import Model from .base import BaseUpstreamProvider @@ -215,9 +216,6 @@ async def init_upstreams() -> list[BaseUpstreamProvider]: provider = _instantiate_provider(provider_row) if provider: - # Keep provider DB id on runtime instance so model mapping can - # bind DB overrides to the correct upstream. - setattr(provider, "db_id", provider_row.id) await provider.refresh_models_cache() logger.debug( f"Initialized {provider_row.provider_type} provider", @@ -250,6 +248,7 @@ async def _seed_providers_from_settings( providers_to_add: list[UpstreamProviderRow] = [] seeded_provider_keys: set[tuple[str, str]] = set() + reserved_slugs: set[str] = set() provider_classes_by_type = { cls.provider_type: cls @@ -279,8 +278,13 @@ async def _seed_providers_from_settings( ) ) if not result.first(): + slug = await allocate_unique_provider_slug( + session, provider_type, reserved_slugs + ) + reserved_slugs.add(slug) providers_to_add.append( UpstreamProviderRow( + slug=slug, provider_type=provider_type, base_url=base_url, api_key=api_key, @@ -299,8 +303,13 @@ async def _seed_providers_from_settings( ) ) if not result.first(): + slug = await allocate_unique_provider_slug( + session, "ollama", reserved_slugs + ) + reserved_slugs.add(slug) providers_to_add.append( UpstreamProviderRow( + slug=slug, provider_type="ollama", base_url=ollama_base_url, api_key=ollama_api_key, @@ -320,8 +329,13 @@ async def _seed_providers_from_settings( ) ) if not result.first(): + slug = await allocate_unique_provider_slug( + session, "azure", reserved_slugs + ) + reserved_slugs.add(slug) providers_to_add.append( UpstreamProviderRow( + slug=slug, provider_type="azure", base_url=base_url, api_key=api_key, @@ -342,8 +356,13 @@ async def _seed_providers_from_settings( ) ) if not result.first(): + slug = await allocate_unique_provider_slug( + session, "custom", reserved_slugs + ) + reserved_slugs.add(slug) providers_to_add.append( UpstreamProviderRow( + slug=slug, provider_type="custom", base_url=base_url, api_key=api_key, @@ -356,7 +375,7 @@ async def _seed_providers_from_settings( session.add(provider) logger.info( f"Seeding {provider.provider_type} provider", # type: ignore[str-format] - extra={"base_url": provider.base_url}, + extra={"base_url": provider.base_url, "slug": provider.slug}, ) @@ -391,9 +410,7 @@ def _instantiate_provider( return provider if provider_row.provider_type == "custom": - return BaseUpstreamProvider( - provider_row.base_url, provider_row.api_key, provider_row.provider_fee - ) + return BaseUpstreamProvider.from_db_row(provider_row) logger.error( f"Unknown provider type: {provider_row.provider_type}", diff --git a/routstr/upstream/ollama.py b/routstr/upstream/ollama.py index 74327d0b..9fed0154 100644 --- a/routstr/upstream/ollama.py +++ b/routstr/upstream/ollama.py @@ -43,7 +43,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def from_db_row( + def _build_from_row( cls, provider_row: "UpstreamProviderRow" ) -> "OllamaUpstreamProvider": return cls( diff --git a/routstr/upstream/openai.py b/routstr/upstream/openai.py index 11cc4336..f2f03cc5 100644 --- a/routstr/upstream/openai.py +++ b/routstr/upstream/openai.py @@ -20,7 +20,7 @@ class OpenAIUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def from_db_row( + def _build_from_row( cls, provider_row: "UpstreamProviderRow" ) -> "OpenAIUpstreamProvider": return cls( diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index 0c69d2cc..63995295 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -90,7 +90,7 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def from_db_row( + def _build_from_row( cls, provider_row: "UpstreamProviderRow" ) -> "OpenRouterUpstreamProvider": return cls( diff --git a/routstr/upstream/perplexity.py b/routstr/upstream/perplexity.py index 55a9116e..a088218f 100644 --- a/routstr/upstream/perplexity.py +++ b/routstr/upstream/perplexity.py @@ -23,7 +23,7 @@ class PerplexityUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def from_db_row( + def _build_from_row( cls, provider_row: "UpstreamProviderRow" ) -> "PerplexityUpstreamProvider": return cls( diff --git a/routstr/upstream/ppqai.py b/routstr/upstream/ppqai.py index 48ca9600..0824a08a 100644 --- a/routstr/upstream/ppqai.py +++ b/routstr/upstream/ppqai.py @@ -46,7 +46,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def from_db_row( + def _build_from_row( cls, provider_row: "UpstreamProviderRow" ) -> "PPQAIUpstreamProvider": return cls( @@ -229,17 +229,14 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): f"Disabling PPQ.AI provider ({self.base_url}) due to insufficient balance", extra={"error": error_message}, ) - from sqlmodel import select - from ..core.db import UpstreamProviderRow, create_session async with create_session() as session: - statement = select(UpstreamProviderRow).where( - UpstreamProviderRow.base_url == self.base_url, - UpstreamProviderRow.api_key == self.api_key, + provider = ( + await session.get(UpstreamProviderRow, self.db_id) + if self.db_id is not None + else None ) - result = await session.exec(statement) - provider = result.first() if provider: provider.enabled = False diff --git a/routstr/upstream/pricing_resolver.py b/routstr/upstream/pricing_resolver.py new file mode 100644 index 00000000..3d009fdc --- /dev/null +++ b/routstr/upstream/pricing_resolver.py @@ -0,0 +1,203 @@ +"""Shared price/metadata resolution chain for upstream model discovery. + +Most OpenAI-compatible ``/models`` responses carry no pricing. Rather than let +a provider fabricate one, this module resolves a model through decreasingly +trustworthy sources — litellm's bundled cost map (curated list prices, mirrors +provider docs), then the OpenRouter feed (resale prices, broader coverage) — +and returns ``None`` when none of them know the model, so the caller can fail +closed instead of inventing a number. + +Provider-native pricing (a gateway's own ``/models`` schema, e.g. Venice's +``model_spec``) is authoritative and handled by the provider before this chain +is consulted; only the shared fallback lives here so a later refactor can hoist +it into the base provider unchanged. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field + + +@dataclass +class ResolvedPricing: + """Per-token pricing plus whatever metadata the answering source carried. + + Prices are USD per token. ``source`` records provenance + (``native``/``litellm``/``openrouter``/``unresolved``) so later work can + surface where each price came from. + """ + + prompt: float + completion: float + context_length: int | None + source: str + modality: str | None = None + max_completion_tokens: int | None = None + input_cache_read: float = 0.0 + input_cache_write: float = 0.0 + input_modalities: list[str] = field(default_factory=lambda: ["text"]) + output_modalities: list[str] = field(default_factory=lambda: ["text"]) + tokenizer: str = "unknown" + instruct_type: str | None = None + is_moderated: bool | None = None + + +def estimate_context_length(model_id: str) -> int: + """Best-effort context window from a model id when no source reports one. + + The last rung of the fallback chain, reached only for a model whose price + resolved but whose context did not (or that imported disabled). Context is + not a billing input, so a rough id-based guess is acceptable here where a + guessed *price* never would be. + """ + lowered = model_id.lower() + if any(pattern in lowered for pattern in ["32k", "32000"]): + return 32768 + if any(pattern in lowered for pattern in ["16k", "16000"]): + return 16384 + if any(pattern in lowered for pattern in ["8k", "8000"]): + return 8192 + if "gpt-4" in lowered: + return 8192 + if "claude" in lowered: + return 200000 + return 4096 + + +def _as_float(value: object) -> float | None: + """OpenRouter reports prices as strings; coerce, ``None`` if unparseable.""" + try: + return float(value) # type: ignore[arg-type] + except (TypeError, ValueError): + return None + + +def _as_int(value: object) -> int | None: + """Coerce an already-numeric token count to ``int``, else ``None``.""" + return int(value) if isinstance(value, (int, float)) else None + + +def _from_litellm(model_id: str) -> ResolvedPricing | None: + # Lazy import so the resolver stays import-light and shares the exact + # lookup semantics used by cache-rate backfill. + from ..payment.models import litellm_cost_entry + + info = litellm_cost_entry(model_id) + if info is None: + return None + + prompt = info.get("input_cost_per_token") + completion = info.get("output_cost_per_token") + if not isinstance(prompt, (int, float)) or not isinstance(completion, (int, float)): + return None + # A both-zero entry is litellm listing a model without a real price (free + # moderation/rerank tiers do this) — treating 0/0 as resolved would serve + # the model for free. Reject it (and any negative) so the caller falls + # through, mirroring async_fetch_openrouter_models' _has_valid_pricing. + if prompt < 0 or completion < 0 or (prompt == 0 and completion == 0): + return None + + input_modalities = ["text"] + if info.get("supports_vision"): + input_modalities.append("image") + + return ResolvedPricing( + prompt=float(prompt), + completion=float(completion), + # max_input_tokens is the context window; max_tokens is litellm's + # completion cap (it tracks max_output_tokens for ~94% of models), so + # it is never a context source. A missing window falls to the id-based + # estimate downstream rather than borrowing the output cap. + context_length=_as_int(info.get("max_input_tokens")), + source="litellm", + max_completion_tokens=_as_int(info.get("max_output_tokens")), + input_cache_read=float(info.get("cache_read_input_token_cost") or 0.0), + input_cache_write=float(info.get("cache_creation_input_token_cost") or 0.0), + input_modalities=input_modalities, + ) + + +def _match_openrouter(model_id: str, feed: list[dict]) -> dict | None: + """Find ``model_id`` in the OpenRouter feed, exact id before bare tail. + + Bare-tail matching (``deepseek-chat`` ↔ ``deepseek/deepseek-chat``) is a + looser, lower-trust match — OpenRouter fans a model out across resellers — + so an exact id match always wins first. When several entries share the bare + tail, the one with the highest *combined* (prompt + completion) per-token + cost wins: the choice must be deterministic (not feed-order-dependent) and + money-safe whichever way traffic leans, since undercharging is the hazard. + Ranking on prompt alone could pick an entry that is cheap on input but dear + on output. The live feed has no such collisions today; this only governs + the latent case. + """ + bare = model_id.split("/", 1)[-1] + exact = next((m for m in feed if m.get("id") == model_id), None) + if exact is not None: + return exact + matches = [m for m in feed if m.get("id", "").split("/", 1)[-1] == bare] + if not matches: + return None + + def _combined_cost(m: dict) -> float: + pricing = m.get("pricing", {}) + return (_as_float(pricing.get("prompt")) or 0.0) + ( + _as_float(pricing.get("completion")) or 0.0 + ) + + return max(matches, key=_combined_cost) + + +def _from_openrouter(model_id: str, feed: list[dict]) -> ResolvedPricing | None: + entry = _match_openrouter(model_id, feed) + if entry is None: + return None + + pricing = entry.get("pricing", {}) + prompt = _as_float(pricing.get("prompt")) + completion = _as_float(pricing.get("completion")) + if prompt is None or completion is None: + return None + + architecture = entry.get("architecture", {}) + top_provider = entry.get("top_provider", {}) + + return ResolvedPricing( + prompt=prompt, + completion=completion, + context_length=_as_int(entry.get("context_length")), + source="openrouter", + modality=architecture.get("modality"), + max_completion_tokens=_as_int(top_provider.get("max_completion_tokens")), + input_cache_read=_as_float(pricing.get("input_cache_read")) or 0.0, + input_cache_write=_as_float(pricing.get("input_cache_write")) or 0.0, + input_modalities=architecture.get("input_modalities") or ["text"], + output_modalities=architecture.get("output_modalities") or ["text"], + tokenizer=architecture.get("tokenizer") or "unknown", + instruct_type=architecture.get("instruct_type"), + is_moderated=top_provider.get("is_moderated"), + ) + + +class FallbackPricingResolver: + """Resolves models via litellm → OpenRouter for one discovery pass. + + The OpenRouter catalog is fetched at most once and only when a model + actually misses litellm, so a provider full of litellm-known models never + touches the network. Instantiate one per ``fetch_models`` call. + """ + + def __init__(self) -> None: + self._openrouter_feed: list[dict] | None = None + + async def resolve(self, model_id: str) -> ResolvedPricing | None: + """Resolve ``model_id``; ``None`` if no source knows it.""" + resolved = _from_litellm(model_id) + if resolved is not None: + return resolved + + if self._openrouter_feed is None: + # Lazy import so tests can patch the feed at its source. + from ..payment.models import async_fetch_openrouter_models + + self._openrouter_feed = await async_fetch_openrouter_models() + return _from_openrouter(model_id, self._openrouter_feed) diff --git a/routstr/upstream/routstr.py b/routstr/upstream/routstr.py index c9962301..0371946a 100644 --- a/routstr/upstream/routstr.py +++ b/routstr/upstream/routstr.py @@ -55,7 +55,7 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider): return path.lstrip("/") @classmethod - def from_db_row( + def _build_from_row( cls, provider_row: "UpstreamProviderRow" ) -> "RoutstrUpstreamProvider": import json diff --git a/routstr/upstream/xai.py b/routstr/upstream/xai.py index b46676cb..58caaba0 100644 --- a/routstr/upstream/xai.py +++ b/routstr/upstream/xai.py @@ -21,7 +21,7 @@ class XAIUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def from_db_row(cls, provider_row: "UpstreamProviderRow") -> "XAIUpstreamProvider": + def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "XAIUpstreamProvider": return cls( api_key=provider_row.api_key, provider_fee=provider_row.provider_fee, diff --git a/tests/integration/test_admin_model_cache_pricing.py b/tests/integration/test_admin_model_cache_pricing.py new file mode 100644 index 00000000..4329f497 --- /dev/null +++ b/tests/integration/test_admin_model_cache_pricing.py @@ -0,0 +1,229 @@ +import json +from datetime import datetime, timedelta, timezone +from unittest.mock import patch + +import pytest +from httpx import AsyncClient +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.admin import admin_sessions +from routstr.core.db import ModelRow, UpstreamProviderRow +from routstr.payment.cost_calculation import CostData, calculate_cost +from routstr.proxy import get_model_instance, reinitialize_upstreams + + +def _admin_headers() -> dict[str, str]: + token = "test-admin-cache-pricing-token" + admin_sessions[token] = int( + (datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp() + ) + return {"Authorization": f"Bearer {token}"} + + +def _model_payload( + provider_id: int, + *, + cache_read: float, + cache_write: float, +) -> dict[str, object]: + return { + "id": "custom-cache-model", + "name": "Custom Cache Model", + "description": "custom model with explicit cache pricing", + "created": 0, + "context_length": 128000, + "architecture": { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, + }, + "pricing": { + "prompt": 1.4e-7, + "completion": 2.8e-7, + "input_cache_read": cache_read, + "input_cache_write": cache_write, + "request": 0.0, + "image": 0.0, + "web_search": 0.0, + "internal_reasoning": 0.0, + }, + "per_request_limits": None, + "top_provider": None, + "upstream_provider_id": provider_id, + "canonical_slug": None, + "alias_ids": [], + "enabled": True, + "forwarded_model_id": "custom-cache-model", + } + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_admin_provider_model_api_persists_cache_pricing_on_create_and_update( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + provider = UpstreamProviderRow( + provider_type="generic", + base_url="https://custom-upstream.example/v1", + api_key="test-key", + provider_fee=1.0, + ) + integration_session.add(provider) + await integration_session.commit() + await integration_session.refresh(provider) + assert provider.id is not None + await reinitialize_upstreams() + + headers = _admin_headers() + create_payload = _model_payload( + provider.id, + cache_read=2.8e-9, + cache_write=3.5e-9, + ) + + with patch("routstr.payment.models.sats_usd_price", return_value=1e-6): + create_response = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/models", + headers=headers, + json=create_payload, + ) + + assert create_response.status_code == 200 + create_body = create_response.json() + assert create_body["pricing"]["input_cache_read"] == pytest.approx(2.8e-9) + assert create_body["pricing"]["input_cache_write"] == pytest.approx(3.5e-9) + + row = await integration_session.get(ModelRow, ("custom-cache-model", provider.id)) + assert row is not None + stored_pricing = json.loads(row.pricing) + assert stored_pricing["input_cache_read"] == pytest.approx(2.8e-9) + assert stored_pricing["input_cache_write"] == pytest.approx(3.5e-9) + + update_payload = _model_payload( + provider.id, + cache_read=1.25e-9, + cache_write=4.5e-9, + ) + + with patch("routstr.payment.models.sats_usd_price", return_value=1e-6): + update_response = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/models", + headers=headers, + json=update_payload, + ) + + assert update_response.status_code == 200 + update_body = update_response.json() + assert update_body["pricing"]["input_cache_read"] == pytest.approx(1.25e-9) + assert update_body["pricing"]["input_cache_write"] == pytest.approx(4.5e-9) + + await integration_session.refresh(row) + updated_pricing = json.loads(row.pricing) + assert updated_pricing["input_cache_read"] == pytest.approx(1.25e-9) + assert updated_pricing["input_cache_write"] == pytest.approx(4.5e-9) + + model = get_model_instance("custom-cache-model") + assert model is not None + assert model.sats_pricing is not None + assert model.sats_pricing.input_cache_read == pytest.approx(0.00125) + assert model.sats_pricing.input_cache_write == pytest.approx(0.0045) + + with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=1e-6): + cost = await calculate_cost( + { + "model": "custom-cache-model", + "usage": { + "prompt_tokens": 1000, + "completion_tokens": 100, + "prompt_tokens_details": {"cached_tokens": 800}, + }, + }, + max_cost=1_000_000, + ) + + assert isinstance(cost, CostData) + assert cost.input_tokens == 200 + assert cost.cache_read_input_tokens == 800 + assert cost.cache_read_msats == 1000 + assert cost.output_msats == 28000 + assert cost.input_msats == 29000 + assert cost.total_msats == 57000 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_upstream_response_cost_uses_model_cache_pricing( + integration_session: AsyncSession, + patched_db_engine: None, +) -> None: + """A model's configured cache price must discount upstream cached-token usage.""" + provider = UpstreamProviderRow( + provider_type="generic", + base_url="https://cache-priced-upstream.example/v1", + api_key="test-key", + provider_fee=1.0, + ) + integration_session.add(provider) + await integration_session.commit() + await integration_session.refresh(provider) + assert provider.id is not None + + row = ModelRow( + id="cache-priced-model", + name="Cache Priced Model", + description="model seeded with explicit cache pricing", + created=0, + context_length=128000, + architecture=json.dumps( + { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, + } + ), + pricing=json.dumps( + { + "prompt": 1.4e-7, + "completion": 2.8e-7, + "input_cache_read": 1.25e-9, + "input_cache_write": 4.5e-9, + } + ), + upstream_provider_id=provider.id, + enabled=True, + forwarded_model_id="cache-priced-model", + ) + integration_session.add(row) + await integration_session.commit() + + with patch("routstr.payment.models.sats_usd_price", return_value=1e-6): + await reinitialize_upstreams() + + with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=1e-6): + cost = await calculate_cost( + { + "model": "cache-priced-model", + "usage": { + "prompt_tokens": 1000, + "completion_tokens": 100, + "prompt_tokens_details": {"cached_tokens": 800}, + }, + }, + max_cost=1_000_000, + ) + + assert isinstance(cost, CostData) + # prompt 1.4e-7 USD/token -> 0.14 sats/token -> 140 msats/token + # cache 1.25e-9 USD/token -> 0.00125 sats/token -> 1.25 msats/token + # completion 2.8e-7 USD/token -> 0.28 sats/token -> 280 msats/token + assert cost.input_tokens == 200 + assert cost.cache_read_input_tokens == 800 + assert cost.cache_read_msats == 1000 + assert cost.input_msats == 29000 + assert cost.output_msats == 28000 + assert cost.total_msats == 57000 diff --git a/tests/integration/test_provider_self_lookup.py b/tests/integration/test_provider_self_lookup.py new file mode 100644 index 00000000..5745ab65 --- /dev/null +++ b/tests/integration/test_provider_self_lookup.py @@ -0,0 +1,128 @@ +"""A live upstream provider resolves its OWN database row by stable identity +(its primary key), not by its mutable/secret ``api_key``. + +Today ``from_db_row`` drops ``provider_row.id`` and the two self-referential +paths — PPQ.AI's insufficient-balance self-disable and the base +``refresh_models_cache`` — re-find their own row with +``WHERE base_url == self.base_url AND api_key == self.api_key``. That uses a +rotatable secret as a self-handle: the moment the row's key changes underneath a +live object (a rotation racing an in-flight request), the object can no longer +find itself. These tests pin the invariant that a provider looks itself up by +identity, so the lookup survives a key change (and, later, key encryption). +""" + +from unittest.mock import AsyncMock, patch + +import pytest + +from routstr.core.db import UpstreamProviderRow +from routstr.upstream.ppqai import PPQAIUpstreamProvider + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_provider_object_carries_its_persistent_identity( + integration_session: object, + patched_db_engine: None, +) -> None: + """``from_db_row`` gives the in-memory object its row's identity (``db_id``).""" + row = UpstreamProviderRow( + provider_type="ppqai", + base_url="https://api.ppq.ai", + api_key="sk-original", + enabled=True, + provider_fee=1.0, + ) + integration_session.add(row) # type: ignore[attr-defined] + await integration_session.commit() # type: ignore[attr-defined] + await integration_session.refresh(row) # type: ignore[attr-defined] + + provider = PPQAIUpstreamProvider.from_db_row(row) + assert provider is not None + + assert provider.db_id == row.id + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_self_disable_targets_own_row_after_key_rotation( + integration_session: object, + patched_db_engine: None, +) -> None: + """PPQ.AI self-disable must disable *its* row even after the key rotated. + + RED (current): the object holds the pre-rotation key, so the + ``(base_url, api_key)`` lookup misses the row → the provider is never + disabled. GREEN: lookup by ``id`` finds it and disables it. + """ + row = UpstreamProviderRow( + provider_type="ppqai", + base_url="https://api.ppq.ai", + api_key="sk-original", + enabled=True, + provider_fee=1.0, + ) + integration_session.add(row) # type: ignore[attr-defined] + await integration_session.commit() # type: ignore[attr-defined] + await integration_session.refresh(row) # type: ignore[attr-defined] + + provider = PPQAIUpstreamProvider.from_db_row(row) # captures sk-original + assert provider is not None + + # Key is rotated in the DB while `provider` is still live. + row.api_key = "sk-rotated" + integration_session.add(row) # type: ignore[attr-defined] + await integration_session.commit() # type: ignore[attr-defined] + + with patch("routstr.proxy.reinitialize_upstreams", new=AsyncMock()): + await provider.on_upstream_error_redirect(402, "Insufficient balance") + + await integration_session.refresh(row) # type: ignore[attr-defined] + assert row.enabled is False + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_refresh_models_cache_finds_own_row_after_key_rotation( + integration_session: object, + patched_db_engine: None, +) -> None: + """``refresh_models_cache`` must resolve its own row after a key rotation. + + ``refresh_models_cache`` swallows every exception (it only logs), so the + observable proof it found its row is that it reaches ``list_models`` — which + is called with the row's ``id`` only *after* the row is resolved. RED + (current): the stale-key ``(base_url, api_key)`` lookup returns nothing, the + method raises ``404`` internally and returns before ``list_models`` is ever + called. GREEN: lookup by ``id`` finds the row and ``list_models`` runs for + that ``id``. + """ + row = UpstreamProviderRow( + provider_type="ppqai", + base_url="https://api.ppq.ai", + api_key="sk-original", + enabled=True, + provider_fee=1.0, + ) + integration_session.add(row) # type: ignore[attr-defined] + await integration_session.commit() # type: ignore[attr-defined] + await integration_session.refresh(row) # type: ignore[attr-defined] + row_id = row.id + + provider = PPQAIUpstreamProvider.from_db_row(row) # captures sk-original + assert provider is not None + + row.api_key = "sk-rotated" + integration_session.add(row) # type: ignore[attr-defined] + await integration_session.commit() # type: ignore[attr-defined] + + list_models_mock = AsyncMock(return_value=[]) + with ( + patch.object(provider, "fetch_models", new=AsyncMock(return_value=[])), + patch("routstr.upstream.base.list_models", new=list_models_mock), + ): + await provider.refresh_models_cache() + + list_models_mock.assert_awaited_once() + assert list_models_mock.await_args is not None + assert list_models_mock.await_args.kwargs["upstream_id"] == row_id diff --git a/tests/unit/test_cache_pricing.py b/tests/unit/test_cache_pricing.py index c9af6e2d..cfb7e9a6 100644 --- a/tests/unit/test_cache_pricing.py +++ b/tests/unit/test_cache_pricing.py @@ -81,6 +81,21 @@ def test_backfill_strips_vendor_prefix_for_litellm_lookup() -> None: assert result.input_cache_read == expected +def test_backfill_case_insensitive_lookup() -> None: + """A generic upstream may report a mixed-case id + (deepseek-ai/DeepSeek-V4-Flash); litellm keys are lowercase. The + case-insensitive fallback still resolves the cache rate.""" + pricing = Pricing(prompt=1.4e-07, completion=2.8e-07) + + result = backfill_cache_pricing("deepseek-ai/DeepSeek-V4-Flash", pricing) + + expected = litellm.model_cost["deepseek-v4-flash"][ + "cache_read_input_token_cost" + ] + assert result.input_cache_read == expected + assert result.input_cache_read < pricing.prompt # sanity: it's a discount + + def test_backfill_fills_cache_write_rate() -> None: """Anthropic cache writes cost more than input (1.25x); billing them at the input rate undercharges. litellm carries the write rate.""" @@ -136,6 +151,93 @@ def test_provider_fee_applies_to_backfilled_cache_rates() -> None: assert adjusted.pricing.prompt == pytest.approx(2.8e-07 * 2.0) +def test_row_to_model_backfills_cache_rate() -> None: + """The DB-override path (admin-configured providers, e.g. a generic + upstream) stores pricing without cache rates. ``_row_to_model`` must + backfill them from litellm just like ``_apply_provider_fee_to_model``, + otherwise cache reads bill at the full input rate.""" + import json + + from routstr.core.db import ModelRow + from routstr.payment.models import _row_to_model + + row = ModelRow( + id="deepseek-v4-flash", + name="deepseek-v4-flash", + created=0, + description="", + context_length=1000000, + architecture=json.dumps( + { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, + } + ), + # Stored pricing omits input_cache_read (generic provider never sets it). + pricing=json.dumps({"prompt": 1.4e-07, "completion": 2.8e-07}), + enabled=True, + upstream_provider_id=1, + ) + + with patch( + "routstr.payment.models.sats_usd_price", return_value=5.0e-5 + ): + model = _row_to_model(row, apply_provider_fee=True, provider_fee=1.0) + + litellm_read = litellm.model_cost["deepseek-v4-flash"][ + "cache_read_input_token_cost" + ] + assert model.pricing.input_cache_read == pytest.approx(litellm_read) + assert model.pricing.input_cache_read < model.pricing.prompt # a discount + assert model.sats_pricing is not None + assert model.sats_pricing.input_cache_read > 0 + + +def test_row_to_model_backfills_via_forwarded_model_id() -> None: + """An alias row (id != forwarded_model_id) must backfill cache rates from + the *forwarded* model name — the real upstream model litellm prices — + not the alias id, which litellm doesn't know.""" + import json + + from routstr.core.db import ModelRow + from routstr.payment.models import _row_to_model + + row = ModelRow( + id="local-alias", # litellm has no such key + name="local-alias", + created=0, + description="", + context_length=1000000, + architecture=json.dumps( + { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, + } + ), + pricing=json.dumps({"prompt": 1.4e-07, "completion": 2.8e-07}), + enabled=True, + upstream_provider_id=1, + forwarded_model_id="deepseek-v4-flash", + ) + + with patch( + "routstr.payment.models.sats_usd_price", return_value=5.0e-5 + ): + model = _row_to_model(row, apply_provider_fee=True, provider_fee=1.0) + + litellm_read = litellm.model_cost["deepseek-v4-flash"][ + "cache_read_input_token_cost" + ] + assert model.pricing.input_cache_read == pytest.approx(litellm_read) + assert model.pricing.input_cache_read < model.pricing.prompt + + # ============================================================================ # calculate_cost — cached tokens billed at cache rates # ============================================================================ diff --git a/tests/unit/test_cost_calculation_caching.py b/tests/unit/test_cost_calculation_caching.py index 41d85385..93c68877 100644 --- a/tests/unit/test_cost_calculation_caching.py +++ b/tests/unit/test_cost_calculation_caching.py @@ -404,6 +404,67 @@ async def test_deepseek_malformed_hit_tokens_coerce_to_zero() -> None: assert result.cache_read_input_tokens == 0 +# ============================================================================ +# Truly-empty response with a non-zero USD cost → full refund +# +# When an upstream reports a USD cost but the response carries NO tokens at all +# (input, output, cache-read and cache-creation all zero), billing the +# USD-derived cost charges the user for nothing. Refund in full. The gate is +# tightened relative to PR #489: a cache-read/-creation-only turn legitimately +# reports zero prompt/completion tokens with a real cost and must still bill. +# ============================================================================ +@pytest.mark.asyncio +async def test_truly_empty_usd_cost_response_is_refunded( + mock_fixed_pricing: None, +) -> None: + """0 input + 0 output + 0 cache tokens with a non-zero USD cost → refund.""" + response = { + "model": "gpt-4", + "usage": { + "prompt_tokens": 0, + "completion_tokens": 0, + "total_cost": 0.01, # non-zero USD cost despite no tokens + }, + } + result = await calculate_cost(response, max_cost=100000) + + assert isinstance(result, CostData) + assert result.total_msats == 0 # full refund + assert result.input_msats == 0 + assert result.output_msats == 0 + assert result.total_usd == 0.0 + assert result.input_tokens == 0 + assert result.output_tokens == 0 + assert result.cache_read_input_tokens == 0 + assert result.cache_creation_input_tokens == 0 + + +@pytest.mark.asyncio +async def test_cache_read_only_usd_cost_response_is_billed( + mock_fixed_pricing: None, +) -> None: + """Cache-read-only turn (0 prompt/completion, non-zero cost) still bills.""" + response = { + "model": "claude-3-5-sonnet", + "usage": { + "input_tokens": 0, + "output_tokens": 0, + "cache_read_input_tokens": 1000, # real cached usage + "cache_creation_input_tokens": 0, + "total_cost": 0.01, # non-zero USD cost + }, + } + result = await calculate_cost(response, max_cost=100000) + + assert isinstance(result, CostData) + # NOT refunded — the USD cost is billed in full. Pinning the exact value + # guards against any future regression that would over-refund a cache-only + # turn (the bug in PR #489, which refunded whenever prompt+completion == 0). + assert result.total_msats == 200000 + assert result.total_usd == 0.01 + assert result.cache_read_input_tokens == 1000 + + # ============================================================================ # Test 13: Missing Usage Block # ============================================================================ diff --git a/tests/unit/test_provider_slug_migration.py b/tests/unit/test_provider_slug_migration.py new file mode 100644 index 00000000..c44faefb --- /dev/null +++ b/tests/unit/test_provider_slug_migration.py @@ -0,0 +1,55 @@ +from __future__ import annotations + +import importlib.util +from pathlib import Path + +import sqlalchemy as sa + +_MIGRATION_PATH = ( + Path(__file__).resolve().parents[2] + / "migrations" + / "versions" + / "c6d7e8f9a0b1_add_slug_to_upstream_providers.py" +) +_spec = importlib.util.spec_from_file_location("provider_slug_migration", _MIGRATION_PATH) +assert _spec is not None and _spec.loader is not None +migration = importlib.util.module_from_spec(_spec) +_spec.loader.exec_module(migration) + + +def test_slug_migration_backfill_uses_api_safe_deterministic_slugs() -> None: + engine = sa.create_engine("sqlite:///:memory:") + with engine.begin() as conn: + conn.execute( + sa.text( + "CREATE TABLE upstream_providers (" + "id INTEGER PRIMARY KEY, " + "provider_type VARCHAR NOT NULL, " + "slug VARCHAR NULL" + ")" + ) + ) + conn.execute( + sa.text( + "INSERT INTO upstream_providers (id, provider_type, slug) VALUES " + "(1, 'OpenAI Compatible', NULL), " + "(2, 'OpenAI Compatible', ''), " + "(3, '123', NULL), " + "(4, 'x', NULL), " + "(5, 'anthropic', 'anthropic')" + ) + ) + + migration._backfill_provider_slugs(conn) + + rows = conn.execute( + sa.text("SELECT id, slug FROM upstream_providers ORDER BY id") + ).all() + + assert rows == [ + (1, "openai-compatible"), + (2, "openai-compatible-2"), + (3, "provider-123"), + (4, "x-provider"), + (5, "anthropic"), + ] diff --git a/tests/unit/test_provider_slugs.py b/tests/unit/test_provider_slugs.py new file mode 100644 index 00000000..0d446fc1 --- /dev/null +++ b/tests/unit/test_provider_slugs.py @@ -0,0 +1,144 @@ +from __future__ import annotations + +import pytest +from sqlalchemy.ext.asyncio import create_async_engine +from sqlmodel import SQLModel, select +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.admin import _get_upstream_provider_by_ref +from routstr.core.db import UpstreamProviderRow +from routstr.core.provider_slugs import ( + allocate_unique_provider_slug, + provider_slug_base, +) +from routstr.upstream.helpers import _seed_providers_from_settings + + +@pytest.mark.asyncio +async def test_allocate_unique_provider_slug_is_deterministic_with_suffixes() -> None: + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + async with AsyncSession(engine) as session: + session.add( + UpstreamProviderRow( + slug="openai", + provider_type="openai", + base_url="https://api.openai.com/v1", + api_key="key-1", + ) + ) + await session.commit() + + assert await allocate_unique_provider_slug(session, "openai") == "openai-2" + assert ( + await allocate_unique_provider_slug(session, "openai", {"openai-2"}) + == "openai-3" + ) + + await engine.dispose() + + +def test_provider_slug_base_sanitizes_provider_type() -> None: + assert provider_slug_base("OpenAI Compatible") == "openai-compatible" + assert provider_slug_base("!!!") == "provider" + assert provider_slug_base("AI") == "ai-provider" + assert provider_slug_base("123") == "provider-123" + + +@pytest.mark.asyncio +async def test_provider_ref_lookup_accepts_existing_numeric_ids_and_slugs() -> None: + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + async with AsyncSession(engine) as session: + provider = UpstreamProviderRow( + slug="openai", + provider_type="openai", + base_url="https://api.openai.com/v1", + api_key="key-1", + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + + by_id = await _get_upstream_provider_by_ref(session, str(provider.id)) + by_slug = await _get_upstream_provider_by_ref(session, "openai") + + assert by_id.id == provider.id + assert by_slug.id == provider.id + + await engine.dispose() + + +@pytest.mark.asyncio +async def test_seed_providers_from_settings_sets_deterministic_slug( + monkeypatch: pytest.MonkeyPatch, +) -> None: + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + monkeypatch.setenv("OPENAI_API_KEY", "seeded-openai-key") + + class SettingsStub: + chat_completions_api_version: str | None = None + upstream_base_url: str | None = None + upstream_api_key: str = "" + + async with AsyncSession(engine) as session: + await _seed_providers_from_settings(session, SettingsStub()) # type: ignore[arg-type] + await session.commit() + + result = await session.exec(select(UpstreamProviderRow)) + providers: list[UpstreamProviderRow] = list(result.all()) + + assert [(p.provider_type, p.slug) for p in providers] == [("openai", "openai")] + + await engine.dispose() + + +@pytest.mark.asyncio +async def test_seed_providers_from_settings_keeps_slug_stable_on_reseed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + monkeypatch.setenv("OPENAI_API_KEY", "seeded-openai-key") + + class SettingsStub: + chat_completions_api_version: str | None = None + upstream_base_url: str | None = None + upstream_api_key: str = "" + + async with AsyncSession(engine) as session: + session.add( + UpstreamProviderRow( + slug="openai", + provider_type="openai", + base_url="https://example.invalid/v1", + api_key="other-key", + ) + ) + await session.commit() + + await _seed_providers_from_settings(session, SettingsStub()) # type: ignore[arg-type] + await session.commit() + await _seed_providers_from_settings(session, SettingsStub()) # type: ignore[arg-type] + await session.commit() + + result = await session.exec( + select(UpstreamProviderRow).order_by(UpstreamProviderRow.slug) + ) + providers: list[UpstreamProviderRow] = list(result.all()) + + assert [(p.provider_type, p.slug) for p in providers] == [ + ("openai", "openai"), + ("openai", "openai-2"), + ] + + await engine.dispose() diff --git a/tests/unit/test_upstream_generic.py b/tests/unit/test_upstream_generic.py new file mode 100644 index 00000000..39daff5c --- /dev/null +++ b/tests/unit/test_upstream_generic.py @@ -0,0 +1,524 @@ +"""Unit tests for ``GenericUpstreamProvider.fetch_models`` price/metadata resolution. + +A generic upstream is any OpenAI-compatible API. Most (DeepSeek, OpenAI, +Groq, ...) answer ``/models`` with bare ``{id, object, owned_by}`` entries that +carry *no* pricing. The provider must not fabricate a price for those: it +resolves through native ``model_spec`` (Venice's bespoke schema) → litellm's +bundled cost map → the OpenRouter feed, and only when every source misses does +it import the model **disabled** with a warning rather than invent a number. + +These tests drive that behaviour through the public ``fetch_models`` API. The +``/models`` HTTP call is faked at ``httpx.AsyncClient``; the OpenRouter feed is +patched at its source (``routstr.payment.models.async_fetch_openrouter_models``) +so the resolver's lazy import picks up the stub. litellm's real bundled cost map +is used unmocked — the DeepSeek rates it ships are the assertion's ground truth. +""" + +from __future__ import annotations + +import logging +from typing import Any +from unittest.mock import AsyncMock, patch + +import pytest + +from routstr.upstream.generic import GenericUpstreamProvider + + +class _FakeResponse: + def __init__(self, payload: dict[str, Any]) -> None: + self._payload = payload + + def raise_for_status(self) -> None: + return None + + def json(self) -> dict[str, Any]: + return self._payload + + +class _FakeAsyncClient: + """Stand-in for ``httpx.AsyncClient`` returning a canned ``/models`` body.""" + + def __init__(self, payload: dict[str, Any]) -> None: + self._payload = payload + + async def __aenter__(self) -> "_FakeAsyncClient": + return self + + async def __aexit__(self, *exc: object) -> bool: + return False + + async def get(self, url: str, headers: dict[str, str] | None = None) -> _FakeResponse: + return _FakeResponse(self._payload) + + +def _patch_models_endpoint(payload: dict[str, Any]) -> Any: + """Patch the provider's ``httpx.AsyncClient`` to serve ``payload``.""" + return patch( + "routstr.upstream.generic.httpx.AsyncClient", + lambda *args, **kwargs: _FakeAsyncClient(payload), + ) + + +def _model_by_id(models: list[Any], model_id: str) -> Any: + return next(m for m in models if m.id == model_id) + + +# --------------------------------------------------------------------------- +# native model_spec (Venice) — must keep resolving, and capture its metadata +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_native_model_spec_resolves_and_captures_metadata() -> None: + """A Venice-shaped ``model_spec`` is authoritative: prices/context come + straight from it and vision capability becomes an image input modality.""" + payload = { + "data": [ + { + "id": "venice-llama", + "owned_by": "venice", + "model_spec": { + "name": "Venice Llama", + "availableContextTokens": 65536, + "pricing": { + "input": {"usd": 0.5}, + "output": {"usd": 1.5}, + }, + "capabilities": {"supportsVision": True}, + }, + } + ] + } + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=[]) + with patch( + "routstr.payment.models.async_fetch_openrouter_models", or_feed + ): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + model = _model_by_id(models, "venice-llama") + assert model.enabled is True + assert model.pricing.prompt == pytest.approx(0.5 / 1_000_000) + assert model.pricing.completion == pytest.approx(1.5 / 1_000_000) + assert model.context_length == 65536 + assert "image" in model.architecture.input_modalities + # Vision capability must be reflected in the combined modality string, not + # flattened to "text->text". + assert model.architecture.modality == "text+image->text" + # A native price never needs the OpenRouter feed. + or_feed.assert_not_awaited() + + +# --------------------------------------------------------------------------- +# native model_spec validation — a bogus native price is not authoritative +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_native_both_zero_price_falls_through_to_litellm() -> None: + """A native ``model_spec`` that prices both tokens at 0 is not a real price + (the same free-tier trap the litellm/OpenRouter rungs already reject). It + must not be treated as authoritative and served free; the resolver falls + through, so a litellm-known model lands on litellm's real rate instead.""" + payload = { + "data": [ + { + "id": "deepseek-chat", + "owned_by": "deepseek", + "model_spec": { + "pricing": {"input": {"usd": 0}, "output": {"usd": 0}}, + }, + } + ] + } + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=[]) + with patch("routstr.payment.models.async_fetch_openrouter_models", or_feed): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + model = _model_by_id(models, "deepseek-chat") + assert model.enabled is True + assert model.pricing.prompt == pytest.approx(2.8e-07) + assert model.pricing.completion == pytest.approx(4.2e-07) + + +@pytest.mark.asyncio +async def test_native_negative_price_falls_through_to_litellm() -> None: + """A negative native price would credit the caller's balance on every + request (a fund drain, not a discount). Reject it like any other unusable + price and fall through to the chain.""" + payload = { + "data": [ + { + "id": "deepseek-chat", + "owned_by": "deepseek", + "model_spec": { + "pricing": {"input": {"usd": -0.5}, "output": {"usd": -1.5}}, + }, + } + ] + } + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=[]) + with patch("routstr.payment.models.async_fetch_openrouter_models", or_feed): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + model = _model_by_id(models, "deepseek-chat") + assert model.enabled is True + assert model.pricing.prompt == pytest.approx(2.8e-07) + assert model.pricing.completion == pytest.approx(4.2e-07) + + +@pytest.mark.asyncio +async def test_native_non_numeric_price_does_not_break_catalog() -> None: + """A non-numeric native price (``"free"``) must not raise while parsing — + an unguarded ``"free" / 1_000_000`` throws and the outer catch drops the + *entire* provider catalog. It has to fail closed for that one model while + every other model in the same response still resolves.""" + payload = { + "data": [ + { + "id": "broken-price", + "owned_by": "mystery", + "model_spec": { + "pricing": {"input": {"usd": "free"}, "output": {"usd": "free"}}, + }, + }, + {"id": "deepseek-chat", "object": "model", "owned_by": "deepseek"}, + ] + } + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=[]) + with patch("routstr.payment.models.async_fetch_openrouter_models", or_feed): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + # One malformed entry must not empty the catalog. + assert {m.id for m in models} == {"broken-price", "deepseek-chat"} + broken = _model_by_id(models, "broken-price") + assert broken.enabled is False + healthy = _model_by_id(models, "deepseek-chat") + assert healthy.enabled is True + assert healthy.pricing.prompt == pytest.approx(2.8e-07) + + +# --------------------------------------------------------------------------- +# litellm rescue — the money-critical case (DeepSeek bare /models) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_bare_deepseek_resolves_via_litellm() -> None: + """DeepSeek's ``/models`` carries no pricing. The old code fabricated + ``$0.001`` + ctx 4096; the resolver must instead pull DeepSeek's real + rates from litellm's bundled cost map (``$0.28``/``$0.42`` per 1M, ctx + 131072) and keep the model enabled.""" + payload = { + "data": [ + {"id": "deepseek-chat", "object": "model", "owned_by": "deepseek"}, + ] + } + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=[]) + with patch( + "routstr.payment.models.async_fetch_openrouter_models", or_feed + ): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + model = _model_by_id(models, "deepseek-chat") + assert model.enabled is True + assert model.pricing.prompt == pytest.approx(2.8e-07) + assert model.pricing.completion == pytest.approx(4.2e-07) + assert model.context_length == 131072 + # Richer metadata than the two base prices is captured too. + assert model.pricing.input_cache_read == pytest.approx(2.8e-08) + assert model.top_provider is not None + assert model.top_provider.max_completion_tokens == 8192 + # litellm answered, so the OpenRouter feed is never consulted. + or_feed.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_litellm_zero_price_entry_fails_closed( + caplog: pytest.LogCaptureFixture, +) -> None: + """A litellm entry that lists a model but prices it at 0/0 (free-tier + moderation/rerank models do this) is not a real price — treating it as one + would silently serve the model for free. The resolver must reject a both-zero + litellm hit and fall through, so the model imports disabled, not at $0.""" + payload = { + "data": [ + {"id": "omni-moderation-latest", "object": "model", "owned_by": "openai"}, + ] + } + + gen_logger = logging.getLogger("routstr.upstream.generic") + gen_logger.addHandler(caplog.handler) + try: + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=[]) + with patch( + "routstr.payment.models.async_fetch_openrouter_models", or_feed + ): + models = await GenericUpstreamProvider( + base_url="http://x" + ).fetch_models() + finally: + gen_logger.removeHandler(caplog.handler) + + model = _model_by_id(models, "omni-moderation-latest") + assert model.enabled is False + assert model.pricing.prompt == 0.0 + assert model.pricing.completion == 0.0 + assert any( + "omni-moderation-latest" in rec.getMessage() + for rec in caplog.records + if rec.levelno >= logging.WARNING + ) + + +@pytest.mark.asyncio +async def test_litellm_output_cap_not_used_as_context() -> None: + """litellm's ``max_tokens`` is the completion cap, not the context window + (it tracks ``max_output_tokens`` for ~94% of models). When a model reports + no ``max_input_tokens``, the resolver must not smuggle the output cap in as + the context window; it falls back to the id-based estimate instead, while + ``max_tokens`` still feeds the completion limit.""" + payload = { + "data": [ + { + "id": "gemini/gemini-gemma-2-9b-it", + "object": "model", + "owned_by": "google", + }, + ] + } + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=[]) + with patch( + "routstr.payment.models.async_fetch_openrouter_models", or_feed + ): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + model = _model_by_id(models, "gemini/gemini-gemma-2-9b-it") + assert model.enabled is True + # litellm gives this model max_input_tokens=None, max_tokens=8192 (an output + # cap). Context must come from the estimate (4096), never the 8192 cap. + assert model.context_length == 4096 + # The 8192 output cap still lands where it belongs: the completion limit. + assert model.top_provider is not None + assert model.top_provider.max_completion_tokens == 8192 + + +# --------------------------------------------------------------------------- +# OpenRouter fallback — litellm misses, OR carries a full payload +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_unknown_to_litellm_resolves_via_openrouter() -> None: + """A model litellm has never heard of still resolves if the OpenRouter + feed lists it, pulling price + context from that entry.""" + payload = { + "data": [ + {"id": "exotic/model-9000", "object": "model", "owned_by": "exotic"}, + ] + } + or_entry = { + "id": "exotic/model-9000", + "name": "Exotic 9000", + "context_length": 65536, + "architecture": { + "modality": "text+image->text", + "input_modalities": ["text", "image"], + "output_modalities": ["text"], + "tokenizer": "Other", + "instruct_type": None, + }, + "pricing": {"prompt": "0.000005", "completion": "0.000010"}, + "top_provider": { + "context_length": 65536, + "max_completion_tokens": 4096, + "is_moderated": False, + }, + } + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=[or_entry]) + with patch( + "routstr.payment.models.async_fetch_openrouter_models", or_feed + ): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + model = _model_by_id(models, "exotic/model-9000") + assert model.enabled is True + assert model.pricing.prompt == pytest.approx(5e-06) + assert model.pricing.completion == pytest.approx(1e-05) + assert model.context_length == 65536 + # The feed's own modality string is carried through verbatim, not recomputed. + assert model.architecture.modality == "text+image->text" + or_feed.assert_awaited() + + +@pytest.mark.asyncio +async def test_openrouter_bare_tail_collision_picks_highest_price() -> None: + """When a bare model id matches several OpenRouter entries by tail + (``model`` ↔ ``a/model``, ``b/model``), the match must be deterministic and + money-safe: pick the highest-priced candidate regardless of feed order, so + ordering can never leave the node charging below true cost. (The live feed + has zero such collisions today; this guards the latent case.)""" + payload = { + "data": [ + {"id": "zzz-phantom-model", "object": "model", "owned_by": "mystery"}, + ] + } + # Same bare tail, different resellers; the pricier one is listed *second* + # so a first-wins match would pick the cheaper (undercharging) entry. + or_feed = AsyncMock( + return_value=[ + { + "id": "cheapco/zzz-phantom-model", + "context_length": 8192, + "pricing": {"prompt": "0.000001", "completion": "0.000002"}, + }, + { + "id": "premiumco/zzz-phantom-model", + "context_length": 8192, + "pricing": {"prompt": "0.000009", "completion": "0.000010"}, + }, + ] + ) + + with _patch_models_endpoint(payload): + with patch("routstr.payment.models.async_fetch_openrouter_models", or_feed): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + model = _model_by_id(models, "zzz-phantom-model") + assert model.pricing.prompt == pytest.approx(9e-06) + assert model.pricing.completion == pytest.approx(1e-05) + + +@pytest.mark.asyncio +async def test_openrouter_bare_tail_tie_breaks_on_combined_cost() -> None: + """The bare-tail tie-break must weigh *both* rates, not prompt alone. + Given two colliding entries where one is cheaper on prompt but far dearer + on completion, ranking by prompt would pick the entry that undercharges + output-heavy traffic. Pick the highest *combined* per-token cost so the + money-safe choice holds whichever way the traffic leans.""" + payload = { + "data": [ + {"id": "yyy-phantom-model", "object": "model", "owned_by": "mystery"}, + ] + } + # dear-overall is listed first with the *lower* prompt, so a prompt-only max + # would wrongly pick the second (cheaper-overall) entry. + or_feed = AsyncMock( + return_value=[ + { + "id": "dearco/yyy-phantom-model", + "context_length": 8192, + "pricing": {"prompt": "0.000001", "completion": "0.000100"}, + }, + { + "id": "cheapco/yyy-phantom-model", + "context_length": 8192, + "pricing": {"prompt": "0.000009", "completion": "0.000002"}, + }, + ] + ) + + with _patch_models_endpoint(payload): + with patch("routstr.payment.models.async_fetch_openrouter_models", or_feed): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + model = _model_by_id(models, "yyy-phantom-model") + assert model.pricing.prompt == pytest.approx(1e-06) + assert model.pricing.completion == pytest.approx(1e-04) + + +@pytest.mark.asyncio +async def test_openrouter_feed_fetched_once_per_discovery() -> None: + """Two models both missing litellm must share a single OpenRouter fetch — + the feed is not re-downloaded per model.""" + payload = { + "data": [ + {"id": "exotic/model-a", "object": "model", "owned_by": "exotic"}, + {"id": "exotic/model-b", "object": "model", "owned_by": "exotic"}, + ] + } + or_feed = AsyncMock( + return_value=[ + { + "id": "exotic/model-a", + "context_length": 8192, + "pricing": {"prompt": "0.000001", "completion": "0.000002"}, + }, + { + "id": "exotic/model-b", + "context_length": 8192, + "pricing": {"prompt": "0.000003", "completion": "0.000004"}, + }, + ] + ) + + with _patch_models_endpoint(payload): + with patch("routstr.payment.models.async_fetch_openrouter_models", or_feed): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + assert {m.id for m in models} == {"exotic/model-a", "exotic/model-b"} + assert or_feed.await_count == 1 + + +# --------------------------------------------------------------------------- +# fail closed — no source resolves → disabled + warned, never fabricated +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_unresolvable_model_fails_closed( + caplog: pytest.LogCaptureFixture, +) -> None: + """When native, litellm and OpenRouter all miss, the model is imported + disabled with a warning naming it — and no price is invented (the old + ``$0.001`` placeholder is gone).""" + payload = { + "data": [ + { + "id": "nobody-has-priced-this-xyz", + "object": "model", + "owned_by": "mystery", + }, + ] + } + + # routstr loggers set propagate=False, so caplog's root handler misses + # them; attach its handler to the provider logger directly. + gen_logger = logging.getLogger("routstr.upstream.generic") + gen_logger.addHandler(caplog.handler) + try: + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=[]) + with patch( + "routstr.payment.models.async_fetch_openrouter_models", or_feed + ): + models = await GenericUpstreamProvider( + base_url="http://x" + ).fetch_models() + finally: + gen_logger.removeHandler(caplog.handler) + + model = _model_by_id(models, "nobody-has-priced-this-xyz") + assert model.enabled is False + assert model.pricing.prompt == 0.0 + assert model.pricing.completion == 0.0 + assert any( + "nobody-has-priced-this-xyz" in rec.getMessage() + for rec in caplog.records + if rec.levelno >= logging.WARNING + ) diff --git a/ui/components/add-provider-model-dialog.tsx b/ui/components/add-provider-model-dialog.tsx index c9f6984d..ddb0a210 100644 --- a/ui/components/add-provider-model-dialog.tsx +++ b/ui/components/add-provider-model-dialog.tsx @@ -68,6 +68,8 @@ const FormSchema = z.object({ upstream_provider_id: z.string().default(''), input_cost: z.coerce.number().min(0).default(0), output_cost: z.coerce.number().min(0).default(0), + cache_read_cost: z.coerce.number().min(0).default(0), + cache_write_cost: z.coerce.number().min(0).default(0), request_cost: z.coerce.number().min(0).default(0), image_cost: z.coerce.number().min(0).default(0), web_search_cost: z.coerce.number().min(0).default(0), @@ -125,6 +127,8 @@ export function AddProviderModelDialog({ upstream_provider_id: '', input_cost: 0, output_cost: 0, + cache_read_cost: 0, + cache_write_cost: 0, request_cost: 0, image_cost: 0, web_search_cost: 0, @@ -190,6 +194,8 @@ export function AddProviderModelDialog({ : initialData.upstream_provider_id?.toString() || '', input_cost: pricing?.prompt ?? 0, output_cost: pricing?.completion ?? 0, + cache_read_cost: pricing?.input_cache_read ?? 0, + cache_write_cost: pricing?.input_cache_write ?? 0, request_cost: pricing?.request ?? 0, image_cost: pricing?.image ?? 0, web_search_cost: pricing?.web_search ?? 0, @@ -231,6 +237,8 @@ export function AddProviderModelDialog({ upstream_provider_id: '', input_cost: 0, output_cost: 0, + cache_read_cost: 0, + cache_write_cost: 0, request_cost: 0, image_cost: 0, web_search_cost: 0, @@ -294,6 +302,8 @@ export function AddProviderModelDialog({ ); form.setValue('input_cost', pricing?.prompt ?? 0); form.setValue('output_cost', pricing?.completion ?? 0); + form.setValue('cache_read_cost', pricing?.input_cache_read ?? 0); + form.setValue('cache_write_cost', pricing?.input_cache_write ?? 0); form.setValue('request_cost', pricing?.request ?? 0); form.setValue('image_cost', pricing?.image ?? 0); form.setValue('web_search_cost', pricing?.web_search ?? 0); @@ -365,6 +375,8 @@ export function AddProviderModelDialog({ pricing: { prompt: data.input_cost, completion: data.output_cost, + input_cache_read: data.cache_read_cost, + input_cache_write: data.cache_write_cost, request: data.request_cost, image: data.image_cost, web_search: data.web_search_cost, @@ -810,6 +822,38 @@ export function AddProviderModelDialog({ )} /> + ( + + Cache Read Cost + + + + + Discounted cached-input read price per 1M tokens. + + + + )} + /> + ( + + Cache Write Cost + + + + + Cached-input creation price per 1M tokens. + + + + )} + /> )} +
+ + + setFormData((prev) => ({ + ...prev, + slug: e.target.value || undefined, + })) + } + placeholder='e.g. openai-prod' + /> +

+ Stable external key used to update this provider via the admin API. +

+
+
{ const val = result[field]; if (val !== undefined && val !== null) { @@ -161,6 +166,8 @@ export class AdminService { convertField('prompt'); convertField('completion'); + convertField('input_cache_read'); + convertField('input_cache_write'); // Other fields (request, image, etc.) are already flat fees (per item) // so we do NOT scale them. @@ -174,7 +181,7 @@ export class AdminService { if (!pricing) return pricing; const result = { ...pricing }; - // Only prompt and completion are per-1M in UI and need scaling down to per-token + // Token-priced fields are per-1M in the UI and need scaling down to per-token. const convertField = (field: string) => { const val = result[field]; if (val !== undefined && val !== null) { @@ -187,6 +194,8 @@ export class AdminService { convertField('prompt'); convertField('completion'); + convertField('input_cache_read'); + convertField('input_cache_write'); // Other fields stay as flat fees