From e972b627583bbeb06b3dbc29c603415c69a54715 Mon Sep 17 00:00:00 2001 From: Shroominic <34897716+shroominic@users.noreply.github.com> Date: Sat, 20 Jun 2026 10:53:58 +0800 Subject: [PATCH 01/20] fix(payment): refund truly-empty responses with a non-zero USD cost When an upstream reports a non-zero USD cost but the response carries no tokens at all (input, output, cache-read and cache-creation all zero), the USD path billed the full USD-derived amount for an empty response. Return an all-zero cost (full refund) for that case only. The gate is tightened relative to the superseded PR #489: a cache-read- or cache-creation-only turn legitimately reports zero prompt/completion tokens with a real cost and must still be billed, so the refund only fires when every token bucket is zero. Adds unit tests covering both the refund and the still-billed cache-only path. Co-Authored-By: Claude Opus 4.8 (1M context) --- .gitignore | 5 ++ routstr/payment/cost_calculation.py | 33 +++++++++++ tests/unit/test_cost_calculation_caching.py | 61 +++++++++++++++++++++ 3 files changed, 99 insertions(+) diff --git a/.gitignore b/.gitignore index f9db7ffb..c845a15c 100644 --- a/.gitignore +++ b/.gitignore @@ -38,3 +38,8 @@ proof_backups *.todo ui_out + +# local cashu wallet state (never commit) +.wallet/ +*.sqlite3-shm +*.sqlite3-wal diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 27b04235..4540f113 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -109,6 +109,39 @@ 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 CostData( + 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, + ) if input_tokens == 0 and output_tokens == 0: logger.warning( "Upstream reported a USD cost but no token counts — " 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 # ============================================================================ From a33ea5da0719f142a64d43e20791323c3995f44a Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 27 Jun 2026 15:48:13 +0200 Subject: [PATCH 02/20] update-pricing --- routstr/payment/models.py | 37 ++++++++++-- routstr/upstream/deepseek_v4_pricing_shim.py | 25 +++++--- tests/unit/test_cache_pricing.py | 60 ++++++++++++++++++++ 3 files changed, 108 insertions(+), 14 deletions(-) diff --git a/routstr/payment/models.py b/routstr/payment/models.py index e62f7875..db477b93 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -94,7 +94,10 @@ def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing: 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. + (gpt-4o, claude-sonnet-4-5), so both spellings are tried. litellm keys are + lowercase, but a generic upstream may report a mixed-case id + (``deepseek-ai/DeepSeek-V4-Flash``); an exact match is attempted first, then + a case-insensitive fallback so such ids still resolve. Rates already present (e.g. provided by OpenRouter) are authoritative and never overwritten. Unknown models are returned unchanged. @@ -106,12 +109,26 @@ def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing: import litellm + candidates = (model_id, model_id.split("/", 1)[-1]) info: dict | None = None - for key in (model_id, model_id.split("/", 1)[-1]): + for key in candidates: candidate = litellm.model_cost.get(key) if isinstance(candidate, dict): info = candidate break + if info is None: + # Case-insensitive fallback: a mixed-case upstream id (e.g. + # ``deepseek-ai/DeepSeek-V4-Flash``) won't match litellm's lowercase + # keys exactly. Build a lowercased index once and retry. + lowered = {c.lower() for c in candidates} + for key, candidate in litellm.model_cost.items(): + if ( + isinstance(key, str) + and key.lower() in lowered + and isinstance(candidate, dict) + ): + info = candidate + break if info is None: return pricing @@ -215,13 +232,23 @@ 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. + parsed_pricing = backfill_cache_pricing(row.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/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/tests/unit/test_cache_pricing.py b/tests/unit/test_cache_pricing.py index c9af6e2d..14bb884c 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,51 @@ 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 + + # ============================================================================ # calculate_cost — cached tokens billed at cache rates # ============================================================================ From 0c7373675f7a4f0ff2a27f1c352ab99e6f931388 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 27 Jun 2026 22:53:35 +0200 Subject: [PATCH 03/20] solidify logic --- routstr/payment/models.py | 8 +++++- tests/unit/test_cache_pricing.py | 42 ++++++++++++++++++++++++++++++++ 2 files changed, 49 insertions(+), 1 deletion(-) diff --git a/routstr/payment/models.py b/routstr/payment/models.py index db477b93..e6ea643d 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -243,7 +243,13 @@ def _row_to_model( # 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. - parsed_pricing = backfill_cache_pricing(row.id, parsed_pricing) + # + # 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( diff --git a/tests/unit/test_cache_pricing.py b/tests/unit/test_cache_pricing.py index 14bb884c..cfb7e9a6 100644 --- a/tests/unit/test_cache_pricing.py +++ b/tests/unit/test_cache_pricing.py @@ -196,6 +196,48 @@ def test_row_to_model_backfills_cache_rate() -> 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 # ============================================================================ From 9fcf870e3f08d9bbcf55f27aeeb068b3c10b76b4 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 16 May 2026 23:00:01 +0200 Subject: [PATCH 04/20] add provider slug --- ...e8f9a0b1_add_slug_to_upstream_providers.py | 55 +++ routstr/core/admin.py | 363 ++++++++++++------ routstr/core/db.py | 6 + ui/components/provider-form-fields.tsx | 20 + ui/lib/api/services/admin.ts | 3 + 5 files changed, 320 insertions(+), 127 deletions(-) create mode 100644 migrations/versions/c6d7e8f9a0b1_add_slug_to_upstream_providers.py 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..198537d8 --- /dev/null +++ b/migrations/versions/c6d7e8f9a0b1_add_slug_to_upstream_providers.py @@ -0,0 +1,55 @@ +"""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 + +revision = "c6d7e8f9a0b1" +down_revision = "b5e7c9d1f3a2" +branch_labels = None +depends_on = None + + +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), + ) + + op.execute( + "UPDATE upstream_providers " + "SET slug = LOWER(provider_type) || '-' || CAST(id AS TEXT) " + "WHERE slug IS NULL OR slug = ''" + ) + + 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/routstr/core/admin.py b/routstr/core/admin.py index 3fffc8c7..6e42b511 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 @@ -456,19 +458,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 +525,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 +544,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 +557,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 +576,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 +595,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 +619,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 +627,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 +688,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 +705,90 @@ 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 + + +def _generate_slug(provider_type: str) -> str: + base = re.sub(r"[^a-z0-9]+", "-", provider_type.lower()).strip("-") or "provider" + return f"{base}-{secrets.token_hex(3)}" + + +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 +797,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 +808,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 +859,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 +879,27 @@ 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: + for _ in range(8): + slug = _generate_slug(payload.provider_type) + existing = await session.exec( + select(UpstreamProviderRow).where( + UpstreamProviderRow.slug == slug + ) + ) + if existing.first() is None: + break + else: + raise HTTPException( + status_code=500, + detail="Could not generate a unique slug", + ) + provider = UpstreamProviderRow( + slug=slug, provider_type=payload.provider_type, base_url=payload.base_url, api_key=payload.api_key, @@ -778,99 +916,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 +1005,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 +1104,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 +1139,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 +1265,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 +1318,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 +1702,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/ui/components/provider-form-fields.tsx b/ui/components/provider-form-fields.tsx index 46a7d2ef..16e6e523 100644 --- a/ui/components/provider-form-fields.tsx +++ b/ui/components/provider-form-fields.tsx @@ -118,6 +118,26 @@ export function ProviderFormFields({ /> )} +
+ + + 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. +

+
+
Date: Wed, 1 Jul 2026 12:53:31 +0200 Subject: [PATCH 05/20] refactor(upstream): resolve a provider's own row by primary key MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A live upstream provider re-finds its own database row in two places — PPQ.AI's insufficient-balance self-disable and the base refresh_models_cache — with WHERE base_url == self.base_url AND api_key == self.api_key. That uses a rotatable secret as a self-handle: if the row's key rotates under a live object, it can no longer find itself. Carry the row's primary key on the instance as db_id, stamped centrally by from_db_row via a _build_from_row construction hook that subclasses override, and look the row up with session.get(UpstreamProviderRow, db_id). This also closes a latent gap where providers built outside the init path (auto-topup) never received db_id. Co-Authored-By: Claude Opus 4.8 --- routstr/upstream/anthropic.py | 2 +- routstr/upstream/auto_topup.py | 2 + routstr/upstream/azure.py | 2 +- routstr/upstream/base.py | 44 ++++-- routstr/upstream/fireworks.py | 2 +- routstr/upstream/gemini.py | 2 +- routstr/upstream/generic.py | 2 +- routstr/upstream/groq.py | 2 +- routstr/upstream/helpers.py | 7 +- routstr/upstream/ollama.py | 2 +- routstr/upstream/openai.py | 2 +- routstr/upstream/openrouter.py | 2 +- routstr/upstream/perplexity.py | 2 +- routstr/upstream/ppqai.py | 13 +- routstr/upstream/routstr.py | 2 +- routstr/upstream/xai.py | 2 +- .../integration/test_provider_self_lookup.py | 128 ++++++++++++++++++ 17 files changed, 179 insertions(+), 39 deletions(-) create mode 100644 tests/integration/test_provider_self_lookup.py 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/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..83b9f565 100644 --- a/routstr/upstream/generic.py +++ b/routstr/upstream/generic.py @@ -45,7 +45,7 @@ class GenericUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def from_db_row( + def _build_from_row( cls, provider_row: "UpstreamProviderRow" ) -> "GenericUpstreamProvider": return cls( 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..f40a3a2a 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -215,9 +215,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", @@ -391,9 +388,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/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_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 From 770565601698d2c5df434f0f641dcea8673b0fb0 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 1 Jul 2026 16:40:31 +0200 Subject: [PATCH 06/20] resolve review comment --- routstr/core/admin.py | 21 +----- routstr/core/provider_slugs.py | 65 +++++++++++++++++ routstr/upstream/helpers.py | 24 +++++- tests/unit/test_provider_slugs.py | 117 ++++++++++++++++++++++++++++++ 4 files changed, 207 insertions(+), 20 deletions(-) create mode 100644 routstr/core/provider_slugs.py create mode 100644 tests/unit/test_provider_slugs.py diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 6e42b511..b0fc080e 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -31,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__) @@ -726,11 +727,6 @@ def _validate_slug(value: str) -> str: return candidate -def _generate_slug(provider_type: str) -> str: - base = re.sub(r"[^a-z0-9]+", "-", provider_type.lower()).strip("-") or "provider" - return f"{base}-{secrets.token_hex(3)}" - - async def _ensure_unique_slug( session: AsyncSession, slug: str, exclude_id: int | None = None ) -> None: @@ -883,20 +879,7 @@ async def create_upstream_provider( slug = _validate_slug(payload.slug) await _ensure_unique_slug(session, slug) else: - for _ in range(8): - slug = _generate_slug(payload.provider_type) - existing = await session.exec( - select(UpstreamProviderRow).where( - UpstreamProviderRow.slug == slug - ) - ) - if existing.first() is None: - break - else: - raise HTTPException( - status_code=500, - detail="Could not generate a unique slug", - ) + slug = await allocate_unique_provider_slug(session, payload.provider_type) provider = UpstreamProviderRow( slug=slug, diff --git a/routstr/core/provider_slugs.py b/routstr/core/provider_slugs.py new file mode 100644 index 00000000..cd0d0410 --- /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 _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 = _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/upstream/helpers.py b/routstr/upstream/helpers.py index 9ad4d73f..a79b8dc2 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 @@ -250,6 +251,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 +281,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 +306,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 +332,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 +359,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 +378,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}, ) diff --git a/tests/unit/test_provider_slugs.py b/tests/unit/test_provider_slugs.py new file mode 100644 index 00000000..e33fbeff --- /dev/null +++ b/tests/unit/test_provider_slugs.py @@ -0,0 +1,117 @@ +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.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_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() From 17bc94959781a8243d2ba2f2ffeece84e17bad3b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 1 Jul 2026 16:53:52 +0200 Subject: [PATCH 07/20] add test --- tests/unit/test_provider_slugs.py | 27 +++++++++++++++++++++++++++ 1 file changed, 27 insertions(+) diff --git a/tests/unit/test_provider_slugs.py b/tests/unit/test_provider_slugs.py index e33fbeff..0d446fc1 100644 --- a/tests/unit/test_provider_slugs.py +++ b/tests/unit/test_provider_slugs.py @@ -5,6 +5,7 @@ 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, @@ -46,6 +47,32 @@ def test_provider_slug_base_sanitizes_provider_type() -> None: 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, From f588147b41c9b9f8be309d362a9ae382f0b0f868 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 1 Jul 2026 17:01:32 +0200 Subject: [PATCH 08/20] use slug base --- ...e8f9a0b1_add_slug_to_upstream_providers.py | 43 +++++++++++++++-- routstr/core/provider_slugs.py | 4 +- tests/unit/test_provider_slug_migration.py | 47 +++++++++++++++++++ 3 files changed, 87 insertions(+), 7 deletions(-) create mode 100644 tests/unit/test_provider_slug_migration.py diff --git a/migrations/versions/c6d7e8f9a0b1_add_slug_to_upstream_providers.py b/migrations/versions/c6d7e8f9a0b1_add_slug_to_upstream_providers.py index 198537d8..c140b557 100644 --- a/migrations/versions/c6d7e8f9a0b1_add_slug_to_upstream_providers.py +++ b/migrations/versions/c6d7e8f9a0b1_add_slug_to_upstream_providers.py @@ -10,12 +10,49 @@ 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) @@ -27,11 +64,7 @@ def upgrade() -> None: sa.Column("slug", sa.String(), nullable=True), ) - op.execute( - "UPDATE upstream_providers " - "SET slug = LOWER(provider_type) || '-' || CAST(id AS TEXT) " - "WHERE slug IS NULL OR slug = ''" - ) + _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: diff --git a/routstr/core/provider_slugs.py b/routstr/core/provider_slugs.py index cd0d0410..0efb919f 100644 --- a/routstr/core/provider_slugs.py +++ b/routstr/core/provider_slugs.py @@ -28,7 +28,7 @@ def provider_slug_base(provider_type: str) -> str: return base -def _slug_candidate(base: str, suffix_number: int) -> str: +def provider_slug_candidate(base: str, suffix_number: int) -> str: if suffix_number == 1: return base @@ -52,7 +52,7 @@ async def allocate_unique_provider_slug( reserved = {slug.lower() for slug in reserved_slugs} for suffix_number in count(1): - candidate = _slug_candidate(base, suffix_number) + candidate = provider_slug_candidate(base, suffix_number) if candidate in reserved: continue diff --git a/tests/unit/test_provider_slug_migration.py b/tests/unit/test_provider_slug_migration.py new file mode 100644 index 00000000..47adf75f --- /dev/null +++ b/tests/unit/test_provider_slug_migration.py @@ -0,0 +1,47 @@ +from __future__ import annotations + +import importlib + +import sqlalchemy as sa + +migration = importlib.import_module( + "migrations.versions.c6d7e8f9a0b1_add_slug_to_upstream_providers" +) + + +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"), + ] From 5bedbd129f9ade740f586b45b1823f5a4ef93009 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 1 Jul 2026 17:08:28 +0200 Subject: [PATCH 09/20] use file path --- tests/unit/test_provider_slug_migration.py | 14 +++++++++++--- 1 file changed, 11 insertions(+), 3 deletions(-) diff --git a/tests/unit/test_provider_slug_migration.py b/tests/unit/test_provider_slug_migration.py index 47adf75f..c44faefb 100644 --- a/tests/unit/test_provider_slug_migration.py +++ b/tests/unit/test_provider_slug_migration.py @@ -1,12 +1,20 @@ from __future__ import annotations -import importlib +import importlib.util +from pathlib import Path import sqlalchemy as sa -migration = importlib.import_module( - "migrations.versions.c6d7e8f9a0b1_add_slug_to_upstream_providers" +_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: From ffde661d9367c6145ac4403a1e9a2ea4c8d77efe Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Fri, 3 Jul 2026 14:42:17 +0200 Subject: [PATCH 10/20] refactor(payment): extract litellm_cost_entry lookup helper Pull the litellm model_cost lookup (both id spellings + case-insensitive fallback) out of backfill_cache_pricing into a reusable litellm_cost_entry helper. No behaviour change; the upstream price resolver reuses the same lookup semantics. Co-Authored-By: Claude Opus 4.8 --- routstr/payment/models.py | 55 +++++++++++++++++++-------------------- 1 file changed, 27 insertions(+), 28 deletions(-) diff --git a/routstr/payment/models.py b/routstr/payment/models.py index e6ea643d..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,12 +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. litellm keys are - lowercase, but a generic upstream may report a mixed-case id - (``deepseek-ai/DeepSeek-V4-Flash``); an exact match is attempted first, then - a case-insensitive fallback so such ids still resolve. + 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. @@ -107,28 +127,7 @@ def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing: if not (needs_read or needs_write): return pricing - import litellm - - candidates = (model_id, model_id.split("/", 1)[-1]) - info: dict | None = None - for key in candidates: - candidate = litellm.model_cost.get(key) - if isinstance(candidate, dict): - info = candidate - break - if info is None: - # Case-insensitive fallback: a mixed-case upstream id (e.g. - # ``deepseek-ai/DeepSeek-V4-Flash``) won't match litellm's lowercase - # keys exactly. Build a lowercased index once and retry. - lowered = {c.lower() for c in candidates} - for key, candidate in litellm.model_cost.items(): - if ( - isinstance(key, str) - and key.lower() in lowered - and isinstance(candidate, dict) - ): - info = candidate - break + info = litellm_cost_entry(model_id) if info is None: return pricing From 2a27fb239e5cfbc1eade95ed7ca8360f6924208d Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Fri, 3 Jul 2026 14:42:33 +0200 Subject: [PATCH 11/20] fix(upstream): resolve generic provider pricing instead of fabricating it MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A generic (OpenAI-compatible) upstream whose /models response omits pricing was silently defaulting to $0.001/M tokens and a 4096 context window. For a provider like DeepSeek that reports no price, this undercharged real usage by ~280x — a direct money leak — while presenting a plausible-looking price. Resolve each model through trust-ordered sources instead: the provider's native schema (Venice's model_spec) first, then litellm's bundled cost map (curated list prices), then the OpenRouter feed. Capture the richer metadata those sources carry (cache rates, modalities, max output tokens, context) rather than only price and context. When no source knows the model, import it disabled with a warning rather than invent a number, so an operator can price it before it serves traffic. Context has no trustworthy source of last resort, but it is not a billing input, so a model priced without a reported context window falls back to an id-based estimate. The whole source-incomplete fallback (price chain + context estimate) lives in one pricing_resolver module so it can later be hoisted into the base provider unchanged. Co-Authored-By: Claude Opus 4.8 --- routstr/upstream/generic.py | 126 +++++++----- routstr/upstream/pricing_resolver.py | 177 +++++++++++++++++ tests/unit/test_upstream_generic.py | 275 +++++++++++++++++++++++++++ 3 files changed, 534 insertions(+), 44 deletions(-) create mode 100644 routstr/upstream/pricing_resolver.py create mode 100644 tests/unit/test_upstream_generic.py diff --git a/routstr/upstream/generic.py b/routstr/upstream/generic.py index 83b9f565..5f388825 100644 --- a/routstr/upstream/generic.py +++ b/routstr/upstream/generic.py @@ -5,6 +5,11 @@ from typing import TYPE_CHECKING import httpx from .base import BaseUpstreamProvider +from .pricing_resolver import ( + FallbackPricingResolver, + ResolvedPricing, + estimate_context_length, +) if TYPE_CHECKING: from ..core.db import UpstreamProviderRow @@ -64,6 +69,35 @@ 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 native price (the common + case for bare OpenAI-compatible ``/models`` responses), so the caller + falls through to the shared resolution chain instead of fabricating a + number. + """ + pricing_info = model_spec.get("pricing", {}) + input_usd = pricing_info.get("input", {}).get("usd") + output_usd = pricing_info.get("output", {}).get("usd") + if input_usd is None or output_usd is None: + 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 +112,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 +124,41 @@ 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 + modality = ( + "text->text" + if "image" in resolved.input_modalities + else "text" + ) - 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 +174,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/pricing_resolver.py b/routstr/upstream/pricing_resolver.py new file mode 100644 index 00000000..4b06febe --- /dev/null +++ b/routstr/upstream/pricing_resolver.py @@ -0,0 +1,177 @@ +"""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 + 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 + + input_modalities = ["text"] + if info.get("supports_vision"): + input_modalities.append("image") + + return ResolvedPricing( + prompt=float(prompt), + completion=float(completion), + context_length=_as_int(info.get("max_input_tokens") or info.get("max_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. + """ + 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 + return next( + (m for m in feed if m.get("id", "").split("/", 1)[-1] == bare), None + ) + + +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", + 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/tests/unit/test_upstream_generic.py b/tests/unit/test_upstream_generic.py new file mode 100644 index 00000000..3c312460 --- /dev/null +++ b/tests/unit/test_upstream_generic.py @@ -0,0 +1,275 @@ +"""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 + # A native price never needs the OpenRouter feed. + or_feed.assert_not_awaited() + + +# --------------------------------------------------------------------------- +# 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() + + +# --------------------------------------------------------------------------- +# 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->text", + "input_modalities": ["text"], + "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 + or_feed.assert_awaited() + + +@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 + ) From ee5ced10c7788616230911c9db706257481ea377 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 4 Jul 2026 20:13:01 +0200 Subject: [PATCH 12/20] refactor(payment): extract shared _empty_cost helper; drop .gitignore change Address PR #561 review: - Extract the all-zero refund cost object into _empty_cost(), reused by both the no-usage-data path and the truly-empty USD-cost path. - Revert the .gitignore additions; those belong in a separate PR. --- .gitignore | 5 --- routstr/payment/cost_calculation.py | 50 ++++++++++++++--------------- 2 files changed, 24 insertions(+), 31 deletions(-) diff --git a/.gitignore b/.gitignore index c845a15c..f9db7ffb 100644 --- a/.gitignore +++ b/.gitignore @@ -38,8 +38,3 @@ proof_backups *.todo ui_out - -# local cashu wallet state (never commit) -.wallet/ -*.sqlite3-shm -*.sqlite3-wal diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 4540f113..86dbbcae 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): @@ -129,19 +139,7 @@ async def calculate_cost( else None, }, ) - return CostData( - 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() if input_tokens == 0 and output_tokens == 0: logger.warning( "Upstream reported a USD cost but no token counts — " From 12c4c1f030235e431a612da9788b98c3ed6b0e74 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 4 Jul 2026 21:49:05 +0200 Subject: [PATCH 13/20] make cache prices configurable --- .../test_admin_model_cache_pricing.py | 229 ++++++++++++++++++ ui/components/add-provider-model-dialog.tsx | 44 ++++ ui/lib/api/services/admin.ts | 10 +- 3 files changed, 281 insertions(+), 2 deletions(-) create mode 100644 tests/integration/test_admin_model_cache_pricing.py 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/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. + + + + )} + /> { const val = result[field]; if (val !== undefined && val !== null) { @@ -164,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. @@ -177,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) { @@ -190,6 +194,8 @@ export class AdminService { convertField('prompt'); convertField('completion'); + convertField('input_cache_read'); + convertField('input_cache_write'); // Other fields stay as flat fees From 9f04525c826f2bc37b65efe5dc7f405a8cfe0a84 Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Sun, 5 Jul 2026 15:39:03 +0200 Subject: [PATCH 14/20] fix(upstream): reject unusable litellm prices instead of resolving them A litellm cost entry that lists a model but prices both tokens at 0 (free moderation/rerank tiers) was resolved as a real price, importing the model enabled and served for free. Reject a both-zero (or negative) litellm hit so the caller falls through to the next source or fails closed, mirroring the _has_valid_pricing check the OpenRouter feed already applies. Co-Authored-By: Claude Opus 4.8 --- routstr/upstream/pricing_resolver.py | 6 +++++ tests/unit/test_upstream_generic.py | 39 ++++++++++++++++++++++++++++ 2 files changed, 45 insertions(+) diff --git a/routstr/upstream/pricing_resolver.py b/routstr/upstream/pricing_resolver.py index 4b06febe..80b27fc8 100644 --- a/routstr/upstream/pricing_resolver.py +++ b/routstr/upstream/pricing_resolver.py @@ -89,6 +89,12 @@ def _from_litellm(model_id: str) -> ResolvedPricing | None: 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"): diff --git a/tests/unit/test_upstream_generic.py b/tests/unit/test_upstream_generic.py index 3c312460..edc448a6 100644 --- a/tests/unit/test_upstream_generic.py +++ b/tests/unit/test_upstream_generic.py @@ -145,6 +145,45 @@ async def test_bare_deepseek_resolves_via_litellm() -> None: 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 + ) + + # --------------------------------------------------------------------------- # OpenRouter fallback — litellm misses, OR carries a full payload # --------------------------------------------------------------------------- From ccee76b31eb4e7c61c36595fdf5e9581c4547c1a Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Sun, 5 Jul 2026 15:39:41 +0200 Subject: [PATCH 15/20] fix(upstream): source litellm context from max_input_tokens only litellm's max_tokens is the completion cap (it equals max_output_tokens for ~94% of models), not the context window, so falling back to it overstated the context as the output limit. Take context from max_input_tokens alone; a model that reports none falls through to the id-based estimate downstream, which is honest about being a guess rather than mislabelling the output cap. Co-Authored-By: Claude Opus 4.8 --- routstr/upstream/pricing_resolver.py | 6 ++++- tests/unit/test_upstream_generic.py | 34 ++++++++++++++++++++++++++++ 2 files changed, 39 insertions(+), 1 deletion(-) diff --git a/routstr/upstream/pricing_resolver.py b/routstr/upstream/pricing_resolver.py index 80b27fc8..f338c12e 100644 --- a/routstr/upstream/pricing_resolver.py +++ b/routstr/upstream/pricing_resolver.py @@ -103,7 +103,11 @@ def _from_litellm(model_id: str) -> ResolvedPricing | None: return ResolvedPricing( prompt=float(prompt), completion=float(completion), - context_length=_as_int(info.get("max_input_tokens") or info.get("max_tokens")), + # 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), diff --git a/tests/unit/test_upstream_generic.py b/tests/unit/test_upstream_generic.py index edc448a6..28bec0ff 100644 --- a/tests/unit/test_upstream_generic.py +++ b/tests/unit/test_upstream_generic.py @@ -184,6 +184,40 @@ async def test_litellm_zero_price_entry_fails_closed( ) +@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 # --------------------------------------------------------------------------- From ae9748db026ffaa877d2e68c07dfe45a2b866864 Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Sun, 5 Jul 2026 15:40:43 +0200 Subject: [PATCH 16/20] fix(upstream): break OpenRouter bare-tail ties by highest price When a bare model id matched several OpenRouter entries by tail, the first feed entry won, making the resolved price depend on feed order and risking an undercharge if a cheaper reseller happened to sort first. Pick the highest-priced candidate instead: deterministic and money-safe, since undercharging is the hazard. An exact id match still wins ahead of any tail match. Measured against the live feed there are zero bare-tail collisions today, so this only governs the latent case. Co-Authored-By: Claude Opus 4.8 --- routstr/upstream/pricing_resolver.py | 14 +++++++--- tests/unit/test_upstream_generic.py | 38 ++++++++++++++++++++++++++++ 2 files changed, 49 insertions(+), 3 deletions(-) diff --git a/routstr/upstream/pricing_resolver.py b/routstr/upstream/pricing_resolver.py index f338c12e..f956334c 100644 --- a/routstr/upstream/pricing_resolver.py +++ b/routstr/upstream/pricing_resolver.py @@ -121,14 +121,22 @@ def _match_openrouter(model_id: str, feed: list[dict]) -> dict | None: 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. + so an exact id match always wins first. When several entries share the bare + tail, the highest-priced one wins: the choice must be deterministic (not + feed-order-dependent) and money-safe, since undercharging is the hazard. + 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 - return next( - (m for m in feed if m.get("id", "").split("/", 1)[-1] == bare), None + matches = [m for m in feed if m.get("id", "").split("/", 1)[-1] == bare] + if not matches: + return None + return max( + matches, + key=lambda m: _as_float(m.get("pricing", {}).get("prompt")) or 0.0, ) diff --git a/tests/unit/test_upstream_generic.py b/tests/unit/test_upstream_generic.py index 28bec0ff..65c9ada2 100644 --- a/tests/unit/test_upstream_generic.py +++ b/tests/unit/test_upstream_generic.py @@ -266,6 +266,44 @@ async def test_unknown_to_litellm_resolves_via_openrouter() -> None: 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_feed_fetched_once_per_discovery() -> None: """Two models both missing litellm must share a single OpenRouter fetch — From ec0fd1143b2d880c0eb182b688b687d4f08bf6d7 Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Sun, 5 Jul 2026 15:40:57 +0200 Subject: [PATCH 17/20] fix(upstream): preserve source modality instead of flattening it The generic provider computed modality as "text->text" for any vision model and "text" otherwise, discarding the source's own modality and mislabelling image-capable models. Carry OpenRouter's modality string through verbatim, and when a source doesn't supply one, derive it from the captured input/output modalities in the same "in->out" shape (e.g. "text+image->text"). Co-Authored-By: Claude Opus 4.8 --- routstr/upstream/generic.py | 11 +++++++---- routstr/upstream/pricing_resolver.py | 2 ++ tests/unit/test_upstream_generic.py | 9 +++++++-- 3 files changed, 16 insertions(+), 6 deletions(-) diff --git a/routstr/upstream/generic.py b/routstr/upstream/generic.py index 5f388825..0a040b9c 100644 --- a/routstr/upstream/generic.py +++ b/routstr/upstream/generic.py @@ -147,10 +147,13 @@ class GenericUpstreamProvider(BaseUpstreamProvider): else: enabled = True - modality = ( - "text->text" - if "image" in resolved.input_modalities - else "text" + # 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)}" ) # A source can carry a price but no context (e.g. a litellm diff --git a/routstr/upstream/pricing_resolver.py b/routstr/upstream/pricing_resolver.py index f956334c..f5c1cb32 100644 --- a/routstr/upstream/pricing_resolver.py +++ b/routstr/upstream/pricing_resolver.py @@ -31,6 +31,7 @@ class ResolvedPricing: 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 @@ -159,6 +160,7 @@ def _from_openrouter(model_id: str, feed: list[dict]) -> ResolvedPricing | None: 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, diff --git a/tests/unit/test_upstream_generic.py b/tests/unit/test_upstream_generic.py index 65c9ada2..1bef8867 100644 --- a/tests/unit/test_upstream_generic.py +++ b/tests/unit/test_upstream_generic.py @@ -104,6 +104,9 @@ async def test_native_model_spec_resolves_and_captures_metadata() -> None: 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() @@ -237,8 +240,8 @@ async def test_unknown_to_litellm_resolves_via_openrouter() -> None: "name": "Exotic 9000", "context_length": 65536, "architecture": { - "modality": "text->text", - "input_modalities": ["text"], + "modality": "text+image->text", + "input_modalities": ["text", "image"], "output_modalities": ["text"], "tokenizer": "Other", "instruct_type": None, @@ -263,6 +266,8 @@ async def test_unknown_to_litellm_resolves_via_openrouter() -> None: 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() From e3ac06342f2c52bd8041cc34481accdd9ae143cb Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 6 Jul 2026 23:00:41 +0200 Subject: [PATCH 18/20] bump version --- pyproject.toml | 2 +- routstr/core/version.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) 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/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 From 64d9460711aec1dc7424f4c2ded44728bf6c3bc6 Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Tue, 7 Jul 2026 12:24:01 +0200 Subject: [PATCH 19/20] fix(upstream): validate native model_spec prices before trusting them MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The generic provider treated any Venice-style model_spec.pricing as authoritative after only a None check, so a native both-zero price served the model free, a negative one credited the caller on every request, and a non-numeric string threw while parsing — the outer catch then dropped the provider's entire catalog. Coerce both native prices through the resolver's _as_float and reject absent / non-numeric / negative / both-zero values, falling through to the shared litellm→OpenRouter→fail-closed chain instead. This extends the same money-safety guard the litellm and OpenRouter rungs already apply to the native source, and keeps one malformed entry from emptying the catalog. Co-Authored-By: Claude Opus 4.8 --- routstr/upstream/generic.py | 18 ++++-- tests/unit/test_upstream_generic.py | 95 +++++++++++++++++++++++++++++ 2 files changed, 107 insertions(+), 6 deletions(-) diff --git a/routstr/upstream/generic.py b/routstr/upstream/generic.py index 0a040b9c..3faf80e9 100644 --- a/routstr/upstream/generic.py +++ b/routstr/upstream/generic.py @@ -8,6 +8,7 @@ from .base import BaseUpstreamProvider from .pricing_resolver import ( FallbackPricingResolver, ResolvedPricing, + _as_float, estimate_context_length, ) @@ -74,16 +75,21 @@ class GenericUpstreamProvider(BaseUpstreamProvider): ) -> ResolvedPricing | None: """Read pricing/metadata from Venice's bespoke ``model_spec`` schema. - Returns ``None`` when the upstream reported no native price (the common - case for bare OpenAI-compatible ``/models`` responses), so the caller - falls through to the shared resolution chain instead of fabricating a - number. + 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 = pricing_info.get("input", {}).get("usd") - output_usd = pricing_info.get("output", {}).get("usd") + 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"] diff --git a/tests/unit/test_upstream_generic.py b/tests/unit/test_upstream_generic.py index 1bef8867..4ff6f589 100644 --- a/tests/unit/test_upstream_generic.py +++ b/tests/unit/test_upstream_generic.py @@ -111,6 +111,101 @@ async def test_native_model_spec_resolves_and_captures_metadata() -> None: 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) # --------------------------------------------------------------------------- From 2baea4149e9a9f38f938c5cc9b1dc60df69d081c Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Tue, 7 Jul 2026 12:25:09 +0200 Subject: [PATCH 20/20] fix(upstream): break OpenRouter bare-tail ties on combined token cost MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The bare-tail tie-break ranked candidates by prompt price alone, so two entries sharing a tail where one is cheaper on prompt but far dearer on completion could resolve to the entry that undercharges output-heavy traffic — contradicting the "highest-priced wins for money safety" promise. Rank by the combined prompt + completion per-token cost instead, so the choice stays deterministic and money-safe whichever way traffic leans. As before there are zero bare-tail collisions in the live feed, so this changes no resolved price today; it only governs the latent case. Co-Authored-By: Claude Opus 4.8 --- routstr/upstream/pricing_resolver.py | 22 ++++++++++------ tests/unit/test_upstream_generic.py | 38 ++++++++++++++++++++++++++++ 2 files changed, 52 insertions(+), 8 deletions(-) diff --git a/routstr/upstream/pricing_resolver.py b/routstr/upstream/pricing_resolver.py index f5c1cb32..3d009fdc 100644 --- a/routstr/upstream/pricing_resolver.py +++ b/routstr/upstream/pricing_resolver.py @@ -123,10 +123,12 @@ def _match_openrouter(model_id: str, feed: list[dict]) -> dict | None: 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 highest-priced one wins: the choice must be deterministic (not - feed-order-dependent) and money-safe, since undercharging is the hazard. - The live feed has no such collisions today; this only governs the latent - case. + 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) @@ -135,10 +137,14 @@ def _match_openrouter(model_id: str, feed: list[dict]) -> dict | None: matches = [m for m in feed if m.get("id", "").split("/", 1)[-1] == bare] if not matches: return None - return max( - matches, - key=lambda m: _as_float(m.get("pricing", {}).get("prompt")) or 0.0, - ) + + 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: diff --git a/tests/unit/test_upstream_generic.py b/tests/unit/test_upstream_generic.py index 4ff6f589..39daff5c 100644 --- a/tests/unit/test_upstream_generic.py +++ b/tests/unit/test_upstream_generic.py @@ -404,6 +404,44 @@ async def test_openrouter_bare_tail_collision_picks_highest_price() -> None: 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 —