Compare commits

...
Author SHA1 Message Date
9qeklajc b535e88282 make claude fix cashu migration issue 2026-03-09 22:02:02 +01:00
9qeklajcandGitHub 5788d63892 Merge pull request #394 from Routstr/fix-refund-token
fix refund token multiple times
2026-03-06 23:06:14 +01:00
9qeklajc d0c7cc6bd9 fix refund token multiple times 2026-03-06 23:02:47 +01:00
9qeklajcandGitHub 40096867d9 Merge pull request #390 from Routstr/fix-negative-reserve
fix negative reserve balance
2026-03-06 18:02:13 +01:00
9qeklajc e002c0b66f clean up 2026-03-06 17:59:49 +01:00
9qeklajcandGitHub 608549d051 Merge pull request #357 from Routstr/fix/azure-kimi-routing-v040-on-v0.4.0
Fix Azure Kimi routing and DB override model mapping
2026-03-06 16:58:57 +01:00
9qeklajcandGitHub 72f389c8d2 Merge pull request #354 from Routstr/fix/azure-remote-model-filter-by-id-v040
fix(admin): filter provider remote models by id
2026-03-06 16:58:02 +01:00
9qeklajcandGitHub 4bce04ada4 Merge pull request #391 from Routstr/fix/cashu-402-not-wrapped
Fix Cashu insufficient-balance responses being wrapped as 401
2026-03-03 23:56:20 +01:00
9qeklajcandGitHub f1e2448620 Merge pull request #388 from Routstr/fix/skip-auth-preflight-balance
Skip preflight balance checks for Authorization tokens
2026-03-03 23:53:41 +01:00
redshift 28340a152c fix cashu auth to preserve insufficient balance status 2026-03-03 22:20:42 +00:00
redshift e28f6118e7 Merge v0.4.0 into fix/skip-auth-preflight-balance 2026-03-03 15:03:31 +00:00
9qeklajc 9fb6f54d12 fix negative reserve balance 2026-03-03 15:29:07 +01:00
9qeklajcandGitHub 3cb8d7b5dd Merge pull request #377 from Routstr/split-all
split all spaces
2026-03-03 00:47:54 +01:00
9qeklajcandGitHub ede076b881 Merge pull request #374 from Routstr/add-cost-details
Add cost details
2026-03-03 00:47:40 +01:00
redshift 449a0951f9 skip preflight balance checks for Authorization tokens 2026-03-02 14:39:55 +00:00
redshiftand9qeklajc e794b09614 Add sats cost and remaining balance to response 2026-02-22 15:44:37 +01:00
Evan Yang ce9834d7ec fix: harden azure routing and model override mapping 2026-02-10 19:15:25 +08:00
9qeklajc 6ebe73f2f7 Fix Azure Kimi routing and DB override model mapping 2026-02-10 07:32:54 +00:00
Evan Yang 98aecb08f9 fix(admin): filter provider remote models by id 2026-02-10 01:04:14 +08:00
16 changed files with 862 additions and 156 deletions
+98
View File
@@ -118,6 +118,13 @@ def create_model_mappings(
candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {}
unique_models: dict[str, "Model"] = {}
seen_model_provider: set[tuple[str, str]] = set()
providers_by_db_id: dict[int, "BaseUpstreamProvider"] = {}
for upstream in upstreams:
db_id = getattr(upstream, "db_id", None)
if isinstance(db_id, int):
providers_by_db_id[db_id] = upstream
# Separate OpenRouter from other providers
openrouter: "BaseUpstreamProvider" | None = None
@@ -134,6 +141,16 @@ def create_model_mappings(
"""Get base model ID by removing provider prefix."""
return model_id.split("/", 1)[1] if "/" in model_id else model_id
def get_provider_identity(upstream: "BaseUpstreamProvider") -> str:
"""Get a stable provider identity used for deduplication."""
db_id = getattr(upstream, "db_id", None)
if isinstance(db_id, int):
return f"db:{db_id}"
provider_type = str(getattr(upstream, "provider_type", "") or "").lower()
base_url = str(getattr(upstream, "base_url", "") or "").lower()
return f"{provider_type}|{base_url}"
def _add_candidate(
alias: str, model: "Model", provider: "BaseUpstreamProvider"
) -> None:
@@ -148,6 +165,7 @@ def create_model_mappings(
) -> None:
"""Process all models from a given provider."""
upstream_prefix = getattr(upstream, "upstream_name", None)
provider_key = get_provider_identity(upstream)
for model in upstream.get_cached_models():
if not model.enabled or model.id in disabled_model_ids:
@@ -189,6 +207,7 @@ def create_model_mappings(
# Try to set each alias
for alias in aliases:
_add_candidate(alias, model_to_use, upstream)
seen_model_provider.add((model_to_use.id.lower(), provider_key))
# Process non-OpenRouter providers first
for upstream in other_upstreams:
@@ -198,6 +217,85 @@ def create_model_mappings(
if openrouter:
process_provider_models(openrouter, is_openrouter=True)
# Include enabled DB overrides even when provider discovery misses models.
# This is important for deployment-based providers like Azure.
for model_id, override_data in overrides_by_id.items():
if model_id in disabled_model_ids:
continue
override_row, provider_fee = override_data
upstream_provider_id = getattr(override_row, "upstream_provider_id", None)
if not isinstance(upstream_provider_id, int):
continue
upstream_for_override = providers_by_db_id.get(upstream_provider_id)
if upstream_for_override is None:
continue
provider_key = get_provider_identity(upstream_for_override)
dedupe_key = (model_id.lower(), provider_key)
if dedupe_key in seen_model_provider:
continue
try:
model_to_use = _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee
)
except Exception as exc:
logger.warning(
"Skipping invalid model override while building model mappings",
extra={
"model_id": model_id,
"upstream_provider_id": upstream_provider_id,
"error": str(exc),
"error_type": type(exc).__name__,
},
)
continue
if not model_to_use.enabled:
continue
base_id = get_base_model_id(model_to_use.id)
is_openrouter = (
getattr(upstream_for_override, "base_url", "")
== "https://openrouter.ai/api/v1"
)
if not is_openrouter or base_id not in unique_models:
unique_model = model_to_use.copy(
update={
"id": base_id,
"upstream_provider_id": upstream_for_override.provider_type,
}
)
unique_models[base_id] = unique_model
try:
aliases = resolve_model_alias(
model_to_use.id,
model_to_use.canonical_slug,
alias_ids=model_to_use.alias_ids,
)
except Exception as exc:
logger.warning(
"Skipping model aliases for invalid override model",
extra={
"model_id": model_id,
"upstream_provider_id": upstream_provider_id,
"error": str(exc),
"error_type": type(exc).__name__,
},
)
continue
upstream_prefix = getattr(upstream_for_override, "upstream_name", None)
if upstream_prefix and "/" not in model_to_use.id:
prefixed_id = f"{upstream_prefix}/{model_to_use.id}"
if prefixed_id not in aliases:
aliases.append(prefixed_id)
for alias in aliases:
_add_candidate(alias, model_to_use, upstream_for_override)
seen_model_provider.add(dedupe_key)
# Sort candidates and build final maps
model_instances: dict[str, "Model"] = {}
provider_map: dict[str, list["BaseUpstreamProvider"]] = {}
+34 -34
View File
@@ -348,6 +348,8 @@ async def validate_bearer_key(
)
return new_key
except HTTPException:
raise
except Exception as e:
logger.error(
"Cashu token redemption failed",
@@ -588,12 +590,15 @@ async def pay_for_request(
async def revert_pay_for_request(
key: ApiKey, session: AsyncSession, cost_per_request: int
) -> None:
) -> bool:
"""Revert a previously reserved payment. Returns True if revert succeeded,
False if the reservation was already released (prevents negative reserved_balance)."""
billing_key = await get_billing_key(key, session)
stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.where(col(ApiKey.reserved_balance) >= cost_per_request)
.values(
reserved_balance=col(ApiKey.reserved_balance) - cost_per_request,
total_requests=col(ApiKey.total_requests) - 1,
@@ -607,6 +612,7 @@ async def revert_pay_for_request(
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.reserved_balance) >= cost_per_request)
.values(
total_requests=col(ApiKey.total_requests) - 1,
reserved_balance=col(ApiKey.reserved_balance) - cost_per_request,
@@ -616,8 +622,8 @@ async def revert_pay_for_request(
await session.commit()
if result.rowcount == 0:
logger.error(
"Failed to revert payment - insufficient reserved balance",
logger.warning(
"Revert skipped - reservation already released (no-op to prevent negative reserved_balance)",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
@@ -625,19 +631,11 @@ async def revert_pay_for_request(
"current_reserved_balance": billing_key.reserved_balance,
},
)
raise HTTPException(
status_code=402,
detail={
"error": {
"message": f"failed to revert request payment: {cost_per_request} mSats required. {billing_key.balance} available.",
"type": "payment_error",
"code": "payment_error",
}
},
)
return False
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
return True
async def adjust_payment_for_tokens(
@@ -669,17 +667,19 @@ async def adjust_payment_for_tokens(
release_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.where(col(ApiKey.reserved_balance) >= deducted_max_cost)
.values(
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost
)
)
await session.exec(release_stmt) # type: ignore[call-overload]
result = await session.exec(release_stmt) # type: ignore[call-overload]
# Also release on child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_release_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.reserved_balance) >= deducted_max_cost)
.values(
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost
@@ -688,14 +688,24 @@ async def adjust_payment_for_tokens(
await session.exec(child_release_stmt) # type: ignore[call-overload]
await session.commit()
logger.warning(
"Released reservation without charging (fallback)",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"deducted_max_cost": deducted_max_cost,
},
)
if result.rowcount == 0: # type: ignore[union-attr]
logger.warning(
"Release reservation skipped - already released (no-op to prevent negative reserved_balance)",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"deducted_max_cost": deducted_max_cost,
},
)
else:
logger.warning(
"Released reservation without charging (fallback)",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"deducted_max_cost": deducted_max_cost,
},
)
except Exception as e:
logger.error(
"Failed to release reservation in fallback",
@@ -997,18 +1007,8 @@ async def adjust_payment_for_tokens(
}
},
)
# Fallback: should not reach here, but release reservation just in case
logger.error(
"Unexpected fallback in adjust_payment_for_tokens - releasing reservation",
extra={"key_hash": key.hashed_key[:8] + "...", "model": model},
)
await release_reservation_only()
return {
"base_msats": deducted_max_cost,
"input_msats": 0,
"output_msats": 0,
"total_msats": deducted_max_cost,
}
# All calculate_cost variants are handled above.
raise AssertionError("Unreachable: unhandled calculate_cost result")
async def periodic_key_reset() -> None:
+3 -2
View File
@@ -205,8 +205,9 @@ async def refund_wallet_endpoint(
key: ApiKey = await validate_bearer_key(bearer_value, session)
if cached := await _refund_cache_get(bearer_value):
return cached
if key.total_balance <= 0:
if cached := await _refund_cache_get(bearer_value):
return cached
if key.parent_key_hash:
raise HTTPException(
+1 -3
View File
@@ -730,9 +730,7 @@ async def get_provider_models(provider_id: int) -> dict[str, object]:
)
db_model_ids = {model.id for model in db_models}
filtered_remote_models = [
m for m in upstream_models if m.name not in db_model_ids
]
filtered_remote_models = [m for m in upstream_models if m.id not in db_model_ids]
return {
"provider": {
+125 -10
View File
@@ -183,11 +183,112 @@ async def create_session() -> AsyncGenerator[AsyncSession, None]:
yield session
def _get_table_columns(cursor: sqlite3.Cursor, table: str) -> set[str]:
"""Return the set of column names for a table, or empty set if table doesn't exist."""
cursor.execute(
"SELECT name FROM sqlite_master WHERE type='table' AND name=?", (table,)
)
if not cursor.fetchone():
return set()
cursor.execute(f"PRAGMA table_info({table})")
return {row[1] for row in cursor.fetchall()}
def _detect_schema_version(cursor: sqlite3.Cursor) -> int:
"""
Detect the actual schema version of a cashu wallet database by inspecting
which columns/tables exist, matching them to migration milestones.
Cashu wallet migrations (m000-m015) add columns non-idempotently.
We detect the highest migration that has already been fully applied.
"""
proofs_cols = _get_table_columns(cursor, "proofs")
keysets_cols = _get_table_columns(cursor, "keysets")
if not proofs_cols:
return 0 # no proofs table means nothing has run
# Walk backwards from the highest migration to find the actual version.
# Each check tests for schema artifacts introduced by that migration.
# m015: mints table
mints_cols = _get_table_columns(cursor, "mints")
if mints_cols:
return 15
# m014: bolt11_mint_quotes.key column
mint_quotes_cols = _get_table_columns(cursor, "bolt11_mint_quotes")
if "key" in mint_quotes_cols:
return 14
# m013: bolt11_mint_quotes table
if mint_quotes_cols:
return 13
# m012: keysets.input_fee_ppk
if "input_fee_ppk" in keysets_cols:
return 12
# m011: keysets.unit
if "unit" in keysets_cols:
return 11
# m010: proofs.dleq, proofs.mint_id
if "mint_id" in proofs_cols:
return 10
if "dleq" in proofs_cols:
return 10
# m009: keysets.counter, proofs.derivation_path
if "derivation_path" in proofs_cols:
return 9
# m008: keysets.public_keys
if "public_keys" in keysets_cols:
return 8
# m007: nostr table
nostr_cols = _get_table_columns(cursor, "nostr")
if nostr_cols:
return 7
# m006: invoices table
invoices_cols = _get_table_columns(cursor, "invoices")
if invoices_cols:
return 6
# m005: proofs.id column, keysets table
if "id" in proofs_cols and keysets_cols:
return 5
# m004: p2sh table
p2sh_cols = _get_table_columns(cursor, "p2sh")
if p2sh_cols:
return 4
# m003: proofs.send_id
if "send_id" in proofs_cols:
return 3
# m002: proofs.reserved
if "reserved" in proofs_cols:
return 2
# m001: proofs table exists
return 1
def fix_cashu_migrations() -> None:
"""
Fixes Cashu wallet migrations that are not idempotent.
This specifically addresses the 'duplicate column name: public_keys' error
in the keysets table of Cashu's internal SQLite databases.
Cashu's migration runner uses ALTER TABLE ADD COLUMN without checking
if the column already exists. If the dbversions table is out of sync
with the actual schema, re-running migrations crashes.
This function detects the real schema version by inspecting existing
columns/tables and updates dbversions accordingly so cashu skips
already-applied migrations.
"""
project_root = pathlib.Path(__file__).resolve().parents[2]
wallet_dir = project_root / ".wallet"
@@ -202,21 +303,35 @@ def fix_cashu_migrations() -> None:
conn = sqlite3.connect(db_file)
cursor = conn.cursor()
# Check if keysets table exists
detected_version = _detect_schema_version(cursor)
if detected_version == 0:
conn.close()
continue
# Ensure dbversions table and row exist
cursor.execute(
"SELECT name FROM sqlite_master WHERE type='table' AND name='keysets'"
"SELECT name FROM sqlite_master WHERE type='table' AND name='dbversions'"
)
if not cursor.fetchone():
conn.close()
continue
# Check if public_keys column exists
cursor.execute("PRAGMA table_info(keysets)")
columns = [info[1] for info in cursor.fetchall()]
cursor.execute("SELECT version FROM dbversions WHERE db = 'wallet'")
row = cursor.fetchone()
if row is None:
conn.close()
continue
if "public_keys" not in columns:
logger.info(f"Adding missing public_keys column to {db_file.name}")
cursor.execute("ALTER TABLE keysets ADD COLUMN public_keys TEXT")
recorded_version = row[0]
if recorded_version < detected_version:
logger.info(
f"Fixing {db_file.name}: dbversions says v{recorded_version} "
f"but schema is v{detected_version}"
)
cursor.execute(
"UPDATE dbversions SET version = ? WHERE db = 'wallet'",
(detected_version,),
)
conn.commit()
conn.close()
+4 -7
View File
@@ -29,16 +29,13 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N
},
)
elif auth := headers.get("authorization", None):
parts = auth.split()
cashu_token = parts[1] if len(parts) > 1 else ""
logger.debug(
"Using Authorization header token",
"Skipping preflight token balance check for Authorization header",
extra={
"token_preview": cashu_token[:20] + "..."
if len(cashu_token) > 20
else cashu_token
"auth_preview": auth[:20] + "..." if len(auth) > 20 else auth,
},
)
return
else:
logger.error("No authentication token provided")
raise HTTPException(status_code=401, detail="Unauthorized")
@@ -76,7 +73,7 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N
if max_cost_for_model > amount_msat:
raise HTTPException(
status_code=413,
status_code=402,
detail={
"reason": "Insufficient balance",
"amount_required_msat": max_cost_for_model,
+5 -1
View File
@@ -300,9 +300,13 @@ async def proxy(
session,
model_obj,
)
except UpstreamError:
# Let the outer UpstreamError handler manage retry/revert
raise
except Exception as e:
# Unexpected error (not an upstream failure) — revert and propagate
logger.error(
"Upstream request failed, ensuring payment is reverted",
"Unexpected error in upstream request, reverting payment",
extra={
"error": str(e),
"error_type": type(e).__name__,
+45 -11
View File
@@ -4,6 +4,7 @@ from .base import BaseUpstreamProvider
if TYPE_CHECKING:
from ..core.db import UpstreamProviderRow
from ..payment.models import Model
class AzureUpstreamProvider(BaseUpstreamProvider):
@@ -58,19 +59,52 @@ class AzureUpstreamProvider(BaseUpstreamProvider):
"platform_url": cls.platform_url,
}
def prepare_headers(self, request_headers: dict) -> dict:
"""Prepare headers for Azure OpenAI, adding api-key."""
headers = super().prepare_headers(request_headers)
if self.api_key:
headers["api-key"] = self.api_key
headers.pop("Authorization", None)
headers.pop("authorization", None)
return headers
def prepare_params(
self, path: str, query_params: Mapping[str, str] | None
) -> Mapping[str, str]:
"""Prepare query parameters for Azure OpenAI, adding API version.
Args:
path: Request path
query_params: Original query parameters from the client
Returns:
Query parameters dict with Azure API version added for chat completions
"""
"""Prepare query parameters for Azure OpenAI, adding API version."""
params = dict(query_params or {})
if path.endswith("chat/completions"):
params["api-version"] = self.api_version
version = (self.api_version or "").replace("\ufeff", "").strip()
if not version or version.lower() == "v1":
version = "2024-02-15-preview"
params["api-version"] = version
return params
def normalize_request_path(
self, path: str, model_obj: "Model | None" = None
) -> str:
"""Build Azure deployment-specific request path."""
clean_path = super().normalize_request_path(path, model_obj).lstrip("/")
if model_obj is None:
return clean_path
deployment_id = getattr(
model_obj, "canonical_slug", None
) or self.transform_model_name(model_obj.id)
deployment_id = deployment_id.split("/")[-1]
return f"openai/deployments/{deployment_id}/{clean_path}"
def get_request_base_url(
self, path: str, model_obj: "Model | None" = None
) -> str:
"""Use endpoint root, stripping accidental /openai/v1 suffix if present."""
base_url = self.base_url.rstrip("/")
marker = "/openai/v1"
if marker in base_url:
base_url = base_url.split(marker, 1)[0].rstrip("/")
return base_url
def transform_model_name(self, model_id: str) -> str:
"""Extract deployment name from model ID."""
if "/" in model_id:
return model_id.split("/")[-1]
return model_id
+99 -34
View File
@@ -13,7 +13,7 @@ from fastapi.responses import Response, StreamingResponse
from pydantic import BaseModel
from sqlmodel import select
from ..auth import adjust_payment_for_tokens, revert_pay_for_request
from ..auth import adjust_payment_for_tokens
from ..core import get_logger
from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow, create_session
from ..core.exceptions import UpstreamError
@@ -197,6 +197,23 @@ class BaseUpstreamProvider:
"""
return model_id
def normalize_request_path(self, path: str, model_obj: Model | None = None) -> str:
"""Normalize request path before forwarding to upstream."""
if path.startswith("v1/"):
return path.replace("v1/", "", 1)
return path
def get_request_base_url(
self, path: str, model_obj: Model | None = None
) -> str:
"""Get upstream base URL used when building forwarding URL."""
return self.base_url.rstrip("/")
def build_request_url(self, path: str, model_obj: Model | None = None) -> str:
"""Build full upstream URL from normalized path."""
clean_path = path.lstrip("/")
return f"{self.get_request_base_url(path, model_obj)}/{clean_path}"
def prepare_responses_request_body(
self, body: bytes | None, model_obj: Model
) -> bytes | None:
@@ -523,10 +540,17 @@ class BaseUpstreamProvider:
session,
max_cost_for_model,
)
remaining_balance_msats = fresh_key.balance
# Merge cost into usage
usage_chunk_data["usage"]["cost"] = cost_data.get(
"total_usd", 0.0
)
usage_chunk_data["usage"]["cost_sats"] = (
cost_data.get("total_msats", 0) // 1000
)
usage_chunk_data["usage"]["remaining_balance_msats"] = (
remaining_balance_msats
)
# Keep detailed cost in metadata
usage_chunk_data["metadata"] = usage_chunk_data.get(
"metadata", {}
@@ -534,6 +558,12 @@ class BaseUpstreamProvider:
usage_chunk_data["metadata"]["routstr"] = {
"cost": cost_data
}
usage_chunk_data["metadata"]["routstr"]["cost"][
"sats_cost"
] = cost_data.get("total_msats", 0) // 1000
usage_chunk_data["metadata"]["routstr"]["cost"][
"remaining_balance_msats"
] = remaining_balance_msats
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
usage_finalized = True
except Exception as e:
@@ -624,14 +654,31 @@ class BaseUpstreamProvider:
key, response_json, session, deducted_max_cost
)
await session.refresh(key)
remaining_balance_msats = key.balance
# Merge cost into usage for OpenCode
if "usage" in response_json:
response_json["usage"]["cost"] = cost_data.get("total_usd", 0.0)
response_json["usage"]["cost_sats"] = (
cost_data.get("total_msats", 0) // 1000
)
response_json["usage"]["remaining_balance_msats"] = (
remaining_balance_msats
)
# Keep detailed cost
response_json["metadata"] = response_json.get("metadata", {})
response_json["metadata"]["routstr"] = {"cost": cost_data}
response_json["metadata"]["routstr"]["cost"]["sats_cost"] = (
cost_data.get("total_msats", 0) // 1000
)
response_json["metadata"]["routstr"]["cost"]["remaining_balance_msats"] = (
remaining_balance_msats
)
response_json["cost"] = cost_data
response_json["cost"]["sats_cost"] = cost_data.get("total_msats", 0) // 1000
response_json["cost"]["remaining_balance_msats"] = remaining_balance_msats
logger.info(
"Payment adjustment completed for non-streaming",
@@ -801,6 +848,7 @@ class BaseUpstreamProvider:
session,
max_cost_for_model,
)
remaining_balance_msats = fresh_key.balance
# Merge cost into usage chunk
if (
"response" in usage_chunk_data
@@ -809,10 +857,22 @@ class BaseUpstreamProvider:
usage_chunk_data["response"]["usage"]["cost"] = (
cost_data.get("total_usd", 0.0)
)
usage_chunk_data["response"]["usage"][
"cost_sats"
] = cost_data.get("total_msats", 0) // 1000
usage_chunk_data["response"]["usage"][
"remaining_balance_msats"
] = remaining_balance_msats
elif "usage" in usage_chunk_data:
usage_chunk_data["usage"]["cost"] = cost_data.get(
"total_usd", 0.0
)
usage_chunk_data["usage"]["cost_sats"] = (
cost_data.get("total_msats", 0) // 1000
)
usage_chunk_data["usage"][
"remaining_balance_msats"
] = remaining_balance_msats
# Keep detailed cost in metadata
usage_chunk_data["metadata"] = usage_chunk_data.get(
@@ -821,6 +881,12 @@ class BaseUpstreamProvider:
usage_chunk_data["metadata"]["routstr"] = {
"cost": cost_data
}
usage_chunk_data["metadata"]["routstr"]["cost"][
"sats_cost"
] = cost_data.get("total_msats", 0) // 1000
usage_chunk_data["metadata"]["routstr"]["cost"][
"remaining_balance_msats"
] = remaining_balance_msats
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
usage_finalized = True
except Exception:
@@ -905,14 +971,31 @@ class BaseUpstreamProvider:
key, response_json, session, deducted_max_cost
)
await session.refresh(key)
remaining_balance_msats = key.balance
# Merge cost into usage for OpenCode
if "usage" in response_json:
response_json["usage"]["cost"] = cost_data.get("total_usd", 0.0)
response_json["usage"]["cost_sats"] = (
cost_data.get("total_msats", 0) // 1000
)
response_json["usage"]["remaining_balance_msats"] = (
remaining_balance_msats
)
# Keep detailed cost
response_json["metadata"] = response_json.get("metadata", {})
response_json["metadata"]["routstr"] = {"cost": cost_data}
response_json["metadata"]["routstr"]["cost"]["sats_cost"] = (
cost_data.get("total_msats", 0) // 1000
)
response_json["metadata"]["routstr"]["cost"]["remaining_balance_msats"] = (
remaining_balance_msats
)
response_json["cost"] = cost_data
response_json["cost"]["sats_cost"] = cost_data.get("total_msats", 0) // 1000
response_json["cost"]["remaining_balance_msats"] = remaining_balance_msats
logger.info(
"Payment adjustment completed for non-streaming Responses API",
@@ -1035,10 +1118,8 @@ class BaseUpstreamProvider:
Returns:
Response or StreamingResponse from upstream with cost tracking
"""
if path.startswith("v1/"):
path = path.replace("v1/", "")
url = f"{self.base_url}/{path}"
path = self.normalize_request_path(path, model_obj)
url = self.build_request_url(path, model_obj)
transformed_body = self.prepare_request_body(request_body, model_obj)
@@ -1213,8 +1294,7 @@ class BaseUpstreamProvider:
},
)
await revert_pay_for_request(key, session, max_cost_for_model)
# Don't revert here — proxy.py owns payment revert to avoid double-revert
if isinstance(exc, httpx.ConnectError):
error_message = "Unable to connect to upstream service"
elif isinstance(exc, httpx.TimeoutException):
@@ -1244,13 +1324,9 @@ class BaseUpstreamProvider:
},
)
await revert_pay_for_request(key, session, max_cost_for_model)
return create_error_response(
"internal_error",
"An unexpected server error occurred",
500,
request=request,
# Don't revert here — proxy.py owns payment revert to avoid double-revert
raise UpstreamError(
"An unexpected server error occurred", status_code=500
)
async def forward_responses_request(
@@ -1279,11 +1355,8 @@ class BaseUpstreamProvider:
Returns:
Response or StreamingResponse from upstream with cost tracking
"""
# Remove v1/ prefix if present for Responses API
if path.startswith("v1/"):
path = path.replace("v1/", "")
url = f"{self.base_url}/{path}"
path = self.normalize_request_path(path, model_obj)
url = self.build_request_url(path, model_obj)
transformed_body = self.prepare_responses_request_body(request_body, model_obj)
@@ -1435,8 +1508,7 @@ class BaseUpstreamProvider:
},
)
await revert_pay_for_request(key, session, max_cost_for_model)
# Don't revert here — proxy.py owns payment revert to avoid double-revert
if isinstance(exc, httpx.ConnectError):
error_message = "Unable to connect to upstream service"
elif isinstance(exc, httpx.TimeoutException):
@@ -1466,13 +1538,9 @@ class BaseUpstreamProvider:
},
)
await revert_pay_for_request(key, session, max_cost_for_model)
return create_error_response(
"internal_error",
"An unexpected server error occurred",
500,
request=request,
# Don't revert here — proxy.py owns payment revert to avoid double-revert
raise UpstreamError(
"An unexpected server error occurred", status_code=500
)
async def forward_get_request(
@@ -1491,10 +1559,8 @@ class BaseUpstreamProvider:
Returns:
StreamingResponse from upstream
"""
if path.startswith("v1/"):
path = path.replace("v1/", "")
url = f"{self.base_url}/{path}"
path = self.normalize_request_path(path)
url = self.build_request_url(path)
logger.info(
"Forwarding GET request to upstream",
@@ -3027,7 +3093,7 @@ class BaseUpstreamProvider:
UpstreamProviderRow.api_key == self.api_key
)
result = await session.exec(stmt)
# .first() returns the object or None if not found
provider = result.first()
if not provider or not provider.id:
@@ -3049,7 +3115,6 @@ class BaseUpstreamProvider:
models.append(found_db_model)
models_with_fees = [self._apply_provider_fee_to_model(m) for m in models]
print([mode.id for mode in models_with_fees])
try:
sats_to_usd = sats_usd_price()
+8
View File
@@ -259,7 +259,15 @@ class GeminiUpstreamProvider(BaseUpstreamProvider):
cost_data = await adjust_payment_for_tokens(
key, openai_format_response, session, max_cost_for_model
)
await session.refresh(key)
remaining_balance_msats = key.balance
openai_format_response["cost"] = cost_data
openai_format_response["cost"]["sats_cost"] = (
cost_data.get("total_msats", 0) // 1000
)
openai_format_response["cost"]["remaining_balance_msats"] = (
remaining_balance_msats
)
logger.info(
"Gemini non-streaming payment completed",
+3
View File
@@ -205,6 +205,9 @@ 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",
+6 -35
View File
@@ -3,13 +3,11 @@ from __future__ import annotations
from typing import TYPE_CHECKING
import httpx
from fastapi import Request
from fastapi.responses import Response, StreamingResponse
from .base import BaseUpstreamProvider
if TYPE_CHECKING:
from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow
from ..core.db import UpstreamProviderRow
from ..payment.models import Model
from ..core.logging import get_logger
@@ -67,38 +65,11 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
"""Strip 'ollama/' prefix for Ollama API compatibility."""
return model_id.removeprefix("ollama/")
async def forward_request(
self,
request: Request,
path: str,
headers: dict,
request_body: bytes | None,
key: ApiKey,
max_cost_for_model: int,
session: AsyncSession,
model_obj: Model,
) -> Response | StreamingResponse:
"""Override to use OpenAI-compatible endpoint for proxy requests."""
if path.startswith("v1/"):
path = path.replace("v1/", "")
original_base_url = self.base_url
self.base_url = f"{self.base_url}/v1"
try:
result = await super().forward_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
return result
finally:
self.base_url = original_base_url
def get_request_base_url(
self, path: str, model_obj: Model | None = None
) -> str:
"""Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint."""
return f"{self.base_url.rstrip('/')}/v1"
async def fetch_models(self) -> list[Model]:
"""Fetch models from Ollama API using /api/tags endpoint."""
+48 -6
View File
@@ -2,6 +2,7 @@ import asyncio
import math
from typing import TypedDict
import httpx
from cashu.core.base import Proof, Token
from cashu.wallet.helpers import deserialize_token_from_string
from cashu.wallet.wallet import Wallet
@@ -13,6 +14,41 @@ from .payment.lnurl import raw_send_to_lnurl
logger = get_logger(__name__)
# Cache of mint_url -> list of supported unit strings
_mint_units_cache: dict[str, list[str]] = {}
async def get_mint_units(mint_url: str) -> list[str]:
"""Query a mint's /v1/keysets endpoint to discover its supported units."""
if mint_url in _mint_units_cache:
return _mint_units_cache[mint_url]
try:
url = f"{mint_url.rstrip('/')}/v1/keysets"
async with httpx.AsyncClient(timeout=10) as client:
resp = await client.get(url)
resp.raise_for_status()
data = resp.json()
units = list(
{
ks["unit"]
for ks in data.get("keysets", [])
if ks.get("active", False)
}
)
if not units:
logger.warning(f"No active keysets found for {mint_url}, defaulting to sat")
units = ["sat"]
_mint_units_cache[mint_url] = units
logger.info(f"Mint {mint_url} supports units: {units}")
return units
except Exception as e:
logger.warning(f"Failed to query keysets for {mint_url}: {e}, defaulting to sat")
_mint_units_cache[mint_url] = ["sat"]
return ["sat"]
async def get_balance(unit: str) -> int:
wallet = await get_wallet(settings.primary_mint, unit)
@@ -224,8 +260,6 @@ async def fetch_all_balances(
- Total user balance in sats
- Owner balance in sats (wallet - user)
"""
if units is None:
units = ["sat", "msat"]
async def fetch_balance(
session: db.AsyncSession, mint_url: str, unit: str
@@ -261,12 +295,19 @@ async def fetch_all_balances(
}
return error_result
# Create tasks for all mint/unit combinations
# Discover supported units for each mint, then create tasks
mint_units_map: dict[str, list[str]] = {}
for mint_url in settings.cashu_mints:
if units is not None:
mint_units_map[mint_url] = units
else:
mint_units_map[mint_url] = await get_mint_units(mint_url)
async with db.create_session() as session:
tasks = [
fetch_balance(session, mint_url, unit)
for mint_url in settings.cashu_mints
for unit in units
for mint_url, mint_units in mint_units_map.items()
for unit in mint_units
]
# Run all tasks concurrently
@@ -313,7 +354,8 @@ async def periodic_payout() -> None:
try:
async with db.create_session() as session:
for mint_url in settings.cashu_mints:
for unit in ["sat", "msat"]:
mint_supported_units = await get_mint_units(mint_url)
for unit in mint_supported_units:
wallet = await get_wallet(mint_url, unit)
proofs = get_proofs_per_mint_and_unit(
wallet, mint_url, unit, not_reserved=True
@@ -133,13 +133,16 @@ async def test_reserved_balance_with_successful_requests(
@pytest.mark.asyncio
async def test_insufficient_reserved_balance_for_revert(
async def test_revert_with_zero_reserved_balance_is_noop(
integration_session: AsyncSession,
) -> None:
"""Test revert_pay_for_request behavior with insufficient reserved balance."""
"""Test that revert_pay_for_request is a no-op when reserved_balance is 0.
Previously this would drive reserved_balance negative. With the floor guard,
it should return False and leave reserved_balance at 0.
"""
from routstr.auth import revert_pay_for_request
# Create key with zero reserved balance
unique_key = f"test_revert_key_{uuid.uuid4().hex[:8]}"
test_key = ApiKey(
hashed_key=unique_key,
@@ -149,17 +152,211 @@ async def test_insufficient_reserved_balance_for_revert(
integration_session.add(test_key)
await integration_session.commit()
# Try to revert more than available
# Note: Current implementation allows reserved_balance to go negative
await revert_pay_for_request(test_key, integration_session, 100)
# Try to revert more than available — should be a no-op
result = await revert_pay_for_request(test_key, integration_session, 100)
# Refresh to get updated values
await integration_session.refresh(test_key)
# Current implementation allows negative reserved balance
assert test_key.reserved_balance == -100, (
f"Expected reserved_balance to be -100, got: {test_key.reserved_balance}"
assert result is False, "Revert should return False when reservation already released"
assert test_key.reserved_balance == 0, (
f"Reserved balance should remain 0, got: {test_key.reserved_balance}"
)
assert test_key.total_requests == -1, (
f"Expected total_requests to be -1, got: {test_key.total_requests}"
assert test_key.total_requests == 0, (
f"Total requests should remain 0, got: {test_key.total_requests}"
)
@pytest.mark.asyncio
async def test_revert_with_sufficient_reserved_balance_succeeds(
integration_session: AsyncSession,
) -> None:
"""Test that revert_pay_for_request works correctly when there is enough reserved balance."""
from routstr.auth import revert_pay_for_request
unique_key = f"test_revert_ok_{uuid.uuid4().hex[:8]}"
test_key = ApiKey(
hashed_key=unique_key,
balance=5000,
reserved_balance=500,
total_requests=3,
)
integration_session.add(test_key)
await integration_session.commit()
result = await revert_pay_for_request(test_key, integration_session, 500)
await integration_session.refresh(test_key)
assert result is True, "Revert should return True on success"
assert test_key.reserved_balance == 0, (
f"Reserved balance should be 0, got: {test_key.reserved_balance}"
)
assert test_key.total_requests == 2, (
f"Total requests should be 2, got: {test_key.total_requests}"
)
assert test_key.balance == 5000, "Balance should not change on revert"
@pytest.mark.asyncio
async def test_revert_partial_reserved_balance_is_noop(
integration_session: AsyncSession,
) -> None:
"""Test that reverting more than the current reserved_balance is a no-op."""
from routstr.auth import revert_pay_for_request
unique_key = f"test_revert_partial_{uuid.uuid4().hex[:8]}"
test_key = ApiKey(
hashed_key=unique_key,
balance=5000,
reserved_balance=50,
total_requests=1,
)
integration_session.add(test_key)
await integration_session.commit()
# Try to revert 500 when only 50 is reserved — should be no-op
result = await revert_pay_for_request(test_key, integration_session, 500)
await integration_session.refresh(test_key)
assert result is False, "Revert should fail when cost > reserved_balance"
assert test_key.reserved_balance == 50, (
f"Reserved balance should stay at 50, got: {test_key.reserved_balance}"
)
assert test_key.total_requests == 1, (
f"Total requests should stay at 1, got: {test_key.total_requests}"
)
@pytest.mark.asyncio
async def test_double_revert_prevented(
integration_session: AsyncSession,
) -> None:
"""Test that calling revert twice doesn't drive reserved_balance negative.
This simulates the double-revert scenario where both upstream/base.py
and proxy.py attempt to revert the same reservation.
"""
from routstr.auth import revert_pay_for_request
unique_key = f"test_double_revert_{uuid.uuid4().hex[:8]}"
test_key = ApiKey(
hashed_key=unique_key,
balance=10000,
reserved_balance=500,
total_requests=5,
)
integration_session.add(test_key)
await integration_session.commit()
# First revert — should succeed
result1 = await revert_pay_for_request(test_key, integration_session, 500)
await integration_session.refresh(test_key)
assert result1 is True
assert test_key.reserved_balance == 0
assert test_key.total_requests == 4
# Second revert of the same amount — should be no-op
result2 = await revert_pay_for_request(test_key, integration_session, 500)
await integration_session.refresh(test_key)
assert result2 is False, "Second revert should be a no-op"
assert test_key.reserved_balance == 0, (
f"Reserved balance should stay 0, got: {test_key.reserved_balance}"
)
assert test_key.total_requests == 4, (
f"Total requests should stay 4, got: {test_key.total_requests}"
)
@pytest.mark.asyncio
async def test_sequential_reverts_never_go_negative(
integration_session: AsyncSession,
) -> None:
"""Test that multiple reverts don't cause negative reserved_balance.
Simulates the double-revert scenario where multiple code paths
attempt to revert the same reservation.
"""
from routstr.auth import revert_pay_for_request
unique_key = f"test_multi_revert_{uuid.uuid4().hex[:8]}"
test_key = ApiKey(
hashed_key=unique_key,
balance=10000,
reserved_balance=500,
total_requests=5,
)
integration_session.add(test_key)
await integration_session.commit()
# Run 5 sequential reverts for the same 500 reservation
results = []
for _ in range(5):
r = await revert_pay_for_request(test_key, integration_session, 500)
results.append(r)
await integration_session.refresh(test_key)
# Exactly one should succeed, rest should be no-ops
success_count = sum(1 for r in results if r is True)
assert success_count == 1, (
f"Exactly one revert should succeed, got {success_count} successes"
)
assert test_key.reserved_balance == 0, (
f"Reserved balance should be 0, got: {test_key.reserved_balance}"
)
assert test_key.reserved_balance >= 0, (
f"Reserved balance went negative: {test_key.reserved_balance}"
)
@pytest.mark.asyncio
async def test_child_key_revert_floor_guard(
integration_session: AsyncSession,
) -> None:
"""Test that child key reserved_balance also has floor guard on revert."""
from routstr.auth import revert_pay_for_request
parent_key_hash = f"test_parent_{uuid.uuid4().hex[:8]}"
child_key_hash = f"test_child_{uuid.uuid4().hex[:8]}"
parent_key = ApiKey(
hashed_key=parent_key_hash,
balance=10000,
reserved_balance=500,
total_requests=3,
)
child_key = ApiKey(
hashed_key=child_key_hash,
balance=0,
reserved_balance=500,
total_requests=3,
parent_key_hash=parent_key_hash,
)
integration_session.add(parent_key)
integration_session.add(child_key)
await integration_session.commit()
# First revert succeeds
result1 = await revert_pay_for_request(child_key, integration_session, 500)
await integration_session.refresh(parent_key)
await integration_session.refresh(child_key)
assert result1 is True
assert parent_key.reserved_balance == 0
assert child_key.reserved_balance == 0
# Second revert is a no-op for both parent and child
result2 = await revert_pay_for_request(child_key, integration_session, 500)
await integration_session.refresh(parent_key)
await integration_session.refresh(child_key)
assert result2 is False
assert parent_key.reserved_balance == 0, (
f"Parent reserved_balance should stay 0, got: {parent_key.reserved_balance}"
)
assert child_key.reserved_balance == 0, (
f"Child reserved_balance should stay 0, got: {child_key.reserved_balance}"
)
+92 -1
View File
@@ -1,14 +1,18 @@
"""Tests for the model prioritization algorithm."""
import os
from types import SimpleNamespace
from unittest.mock import Mock
import pytest
# Set required env vars before importing
os.environ["UPSTREAM_BASE_URL"] = "http://test"
os.environ["UPSTREAM_API_KEY"] = "test"
from routstr.algorithm import ( # noqa: E402
calculate_model_cost_score,
create_model_mappings,
get_provider_penalty,
)
from routstr.payment.models import Architecture, Model, Pricing # noqa: E402
@@ -45,11 +49,21 @@ def create_test_model(
)
def create_test_provider(name: str, base_url: str = "http://test.com") -> Mock:
def create_test_provider(
name: str,
base_url: str = "http://test.com",
*,
db_id: int | None = None,
models: list[Model] | None = None,
upstream_name: str | None = None,
) -> Mock:
"""Helper to create a test provider mock."""
provider = Mock()
provider.provider_type = name
provider.base_url = base_url
provider.db_id = db_id
provider.upstream_name = upstream_name or name
provider.get_cached_models.return_value = models or []
return provider
@@ -99,3 +113,80 @@ def test_get_provider_penalty_openrouter() -> None:
provider = create_test_provider("openrouter", "https://openrouter.ai/api/v1")
penalty = get_provider_penalty(provider)
assert penalty == 1.001
def test_create_model_mappings_includes_db_override_for_missing_cached_model(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Model overrides should still map when provider discovery misses the model."""
provider = create_test_provider(
"azure",
"https://example.openai.azure.com/openai/v1",
db_id=7,
models=[],
)
override_model = create_test_model("azure/gpt-4o")
override_model.canonical_slug = "azure-deployment"
def fake_row_to_model(*args, **kwargs) -> Model: # type: ignore[no-untyped-def]
return override_model
monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model)
override_row = SimpleNamespace(id="azure/gpt-4o", upstream_provider_id=7, enabled=True)
model_instances, provider_map, unique_models = create_model_mappings(
upstreams=[provider],
overrides_by_id={"azure/gpt-4o": (override_row, 1.01)},
disabled_model_ids=set(),
)
assert "azure/gpt-4o" in model_instances
assert provider_map["azure/gpt-4o"] == [provider]
assert "gpt-4o" in unique_models
def test_create_model_mappings_dedupes_with_provider_identity_not_provider_type(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Different provider instances of same type should both survive dedupe."""
provider_a_model = create_test_model(
"azure/gpt-4o", prompt_price=0.01, completion_price=0.01
)
provider_a = create_test_provider(
"azure",
"https://a.openai.azure.com/openai/v1",
db_id=1,
models=[provider_a_model],
upstream_name="azure-a",
)
provider_b = create_test_provider(
"azure",
"https://b.openai.azure.com/openai/v1",
db_id=2,
models=[],
upstream_name="azure-b",
)
override_model = create_test_model(
"azure/gpt-4o", prompt_price=0.001, completion_price=0.001
)
override_model.canonical_slug = "azure-b-deployment"
def fake_row_to_model(*args, **kwargs) -> Model: # type: ignore[no-untyped-def]
return override_model
monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model)
override_row = SimpleNamespace(id="azure/gpt-4o", upstream_provider_id=2, enabled=True)
_, provider_map, _ = create_model_mappings(
upstreams=[provider_a, provider_b],
overrides_by_id={"azure/gpt-4o": (override_row, 1.01)},
disabled_model_ids=set(),
)
providers_for_alias = provider_map["azure/gpt-4o"]
assert provider_a in providers_for_alias
assert provider_b in providers_for_alias
assert len(providers_for_alias) == 2
+82
View File
@@ -0,0 +1,82 @@
"""Tests for Azure upstream provider request normalization."""
from routstr.payment.models import Architecture, Model, Pricing
from routstr.upstream.azure import AzureUpstreamProvider
def create_test_model(model_id: str, canonical_slug: str | None = None) -> Model:
return Model(
id=model_id,
name=model_id,
created=0,
description="test",
context_length=8192,
architecture=Architecture(
modality="text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="test",
instruct_type=None,
),
pricing=Pricing(
prompt=0.001,
completion=0.001,
request=0.0,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
),
canonical_slug=canonical_slug,
)
def test_prepare_headers_uses_azure_api_key_header() -> None:
provider = AzureUpstreamProvider(
base_url="https://example.openai.azure.com",
api_key="azure-key",
api_version="2024-02-15-preview",
)
headers = provider.prepare_headers({"Authorization": "Bearer user-token"})
assert headers["api-key"] == "azure-key"
assert "Authorization" not in headers
assert "authorization" not in headers
def test_prepare_params_normalizes_azure_api_version() -> None:
provider = AzureUpstreamProvider(
base_url="https://example.openai.azure.com",
api_key="azure-key",
api_version="\ufeff v1 ",
)
params = provider.prepare_params("chat/completions", {})
assert params["api-version"] == "2024-02-15-preview"
def test_normalize_request_path_includes_deployment_id() -> None:
provider = AzureUpstreamProvider(
base_url="https://example.openai.azure.com/openai/v1",
api_key="azure-key",
api_version="2024-02-15-preview",
)
model = create_test_model("azure/gpt-4o", canonical_slug="deploy-gpt4o")
path = provider.normalize_request_path("v1/chat/completions", model)
assert path == "openai/deployments/deploy-gpt4o/chat/completions"
def test_get_request_base_url_strips_openai_v1_suffix() -> None:
provider = AzureUpstreamProvider(
base_url="https://example.openai.azure.com/openai/v1",
api_key="azure-key",
api_version="2024-02-15-preview",
)
model = create_test_model("azure/gpt-4o")
base_url = provider.get_request_base_url("chat/completions", model)
assert base_url == "https://example.openai.azure.com"