Compare commits

...

48 Commits

Author SHA1 Message Date
Shroominic
8d3c064b29 ruff fix 2026-01-23 09:47:48 +08:00
shroominic
4ac96ade5f Merge branch 'v0.3.0' into provider-fallback 2026-01-23 09:45:14 +08:00
shroominic
180a469399 Merge pull request #304 from Routstr/introduce-child-key
Introduce child key
2026-01-23 09:35:53 +08:00
shroominic
3465a44d0e routstr v0.2.2 - fixfix
v0.2.2 - release summary
--------------------------
#292 - Fix not enough inputs to melt
#295 - Update UI dependencies
#291 - Fix reserved balance
#289 - Better filtering options (#282)
#284 - Do not charge for empty content (#274)
#298 - Fix Provider Balance display in dashboard
#299 - fix refunds not accounting reserved balance
#300 - reset reserved balance on startup (optional)
#303 - ignore disabled provider
#301 - optimize price fetching
2026-01-15 11:21:22 +08:00
shroominic
a4259af38f Merge pull request #301 from Routstr/optimize-price-fetching
optimize price fetching
2026-01-15 11:13:14 +08:00
shroominic
0d07dd0cdb Merge pull request #303 from Routstr/ignore-disabled-provider
ignore disabled provider
2026-01-15 11:12:16 +08:00
Shroominic
24015ebec1 Revert "Merge remote-tracking branch 'origin/mock-upstream-with-testnut-mint' into v0.2.2"
This reverts commit 5a4ba60072, reversing
changes made to eed5bc5b04.
2026-01-13 06:31:24 +08:00
9qeklajc
db021866d8 fmt u 2026-01-10 21:42:21 +01:00
9qeklajc
d4339287be chore: add type annotations to example and test files 2026-01-10 21:35:16 +01:00
9qeklajc
88bcc0edcb fmt 2026-01-10 21:34:48 +01:00
9qeklajc
917a4d32b1 ui: visualize parent-child relationship in balances page 2026-01-10 20:10:59 +01:00
9qeklajc
daf17f51ab fix: move child-key route before catch-all and fix indentation 2026-01-10 20:10:31 +01:00
9qeklajc
fc042c768c support child keys mapped to parent balance 2026-01-10 19:53:27 +01:00
9qeklajc
3bc38937e8 ignore disabled provider 2026-01-10 18:39:24 +01:00
Shroominic
9229b87b70 optimize price fetching 2026-01-10 17:37:05 +08:00
shroominic
367265b9fe Merge pull request #300 from Routstr/reset-reserved-balance-on-startup
Reset reserved balance on startup
2026-01-10 08:50:32 +08:00
shroominic
0b3ccb5fb0 Merge pull request #299 from Routstr/urgent-reserved-balance-fix
Reserved balance fix
2026-01-10 08:50:20 +08:00
shroominic
dbd43f52fb Merge pull request #298 from Routstr/fix-provider-balance
Fix provider balance
2026-01-10 08:49:23 +08:00
Shroominic
39657ed64f reset reserved balance on startup 2026-01-09 17:36:04 +08:00
Shroominic
493b4f0f1f urgent reserved balance fix 2026-01-09 17:32:35 +08:00
Shroominic
cb36189db3 fmt 2026-01-09 17:04:09 +08:00
Shroominic
d2487f42b0 fix tests 2026-01-09 17:02:35 +08:00
Shroominic
5255fce7b2 todo comment 2026-01-09 16:55:24 +08:00
Shroominic
ca7e8bec71 fmt 2026-01-09 16:53:25 +08:00
Shroominic
c4cc09d61e fix provider balance not displaying when 0 2026-01-09 16:50:10 +08:00
Shroominic
ee668ee93b handle specific errors eg deactivate ppq on insufficient balance 2026-01-09 16:49:03 +08:00
Shroominic
cf8b990fc7 fix specify error codes retry 2026-01-09 16:10:14 +08:00
Shroominic
78dd74845b fix merge 2026-01-07 17:44:28 +01:00
Shroominic
e7f4c98475 ranked provider fallback on upstream errors 2026-01-07 17:39:05 +01:00
9qeklajc
fad792068e Merge pull request #287 from Routstr/provider-discovery-default
Disable provider discovery by default
2026-01-06 19:39:09 +01:00
Shroominic
21d363f6aa bump v0.2.2 2026-01-06 19:26:10 +01:00
shroominic
5e21f6ccbc Merge pull request #292 from Routstr/fix-not-enough-inputs-to-melt
Fix not enough inputs to melt
2026-01-06 19:24:33 +01:00
shroominic
1e2d130022 Merge pull request #285 from Routstr/admin-info
Add admin password warning
2026-01-06 19:14:53 +01:00
Shroominic
1e21dce735 change to warning log 2026-01-06 18:23:01 +01:00
shroominic
b7603dcf69 Merge pull request #295 from Routstr/update-ui-deps
update vuln deps
2026-01-06 17:38:13 +01:00
9qeklajc
7dccfa745f fix test 2026-01-06 12:02:27 +01:00
9qeklajc
54d5118980 fmt 2026-01-06 11:50:12 +01:00
9qeklajc
7723ab4a95 update build script 2026-01-06 11:48:53 +01:00
9qeklajc
86c022d8db update vuln deps 2026-01-06 11:46:12 +01:00
9qeklajc
21ae22abec Merge pull request #293 from Routstr/rm-performance-tests
test: Remove flaky performance requirement tests
2026-01-05 22:00:46 +01:00
9qeklajc
3a939d0dd1 Merge pull request #291 from Routstr/fix-reserved-balance
Fix reserved balance
2026-01-05 21:59:46 +01:00
Shroominic
50eabafa57 remove performance tests due to unpredictable behaviour 2026-01-05 12:12:04 +01:00
Shroominic
d192a6a6b4 fix not enough inputs to melt bug 2026-01-05 00:34:58 +01:00
Shroominic
7d829af681 fix types 2026-01-05 00:06:53 +01:00
Shroominic
9e9bc5bff8 fix pytests 2026-01-03 22:49:46 +01:00
Shroominic
6d780ef96d feat: disable provider discovery by default 2026-01-03 22:12:55 +01:00
Shroominic
b70b94b9b4 feat: add admin password warning log 2026-01-03 22:12:23 +01:00
shroominic
f1fa7d094f routstr/v0.2.1
v0.2.1
2025-12-27 22:08:51 +01:00
41 changed files with 1562 additions and 10710 deletions

View File

@@ -59,25 +59,30 @@ jobs:
- name: Checkout code
uses: actions/checkout@v4
- name: Setup pnpm
uses: pnpm/action-setup@v4
with:
version: 10
- name: Setup Node.js
uses: actions/setup-node@v4
with:
node-version: "18"
cache: "npm"
cache-dependency-path: ui/package-lock.json
cache: "pnpm"
cache-dependency-path: ui/pnpm-lock.yaml
- name: Install UI dependencies
working-directory: ./ui
run: npm ci
run: pnpm install --frozen-lockfile
- name: Run UI format check
working-directory: ./ui
run: npm run format-check
run: pnpm run format-check
- name: Run UI linting
working-directory: ./ui
run: npm run lint
run: pnpm run lint
- name: Run UI build
working-directory: ./ui
run: npm run build
run: pnpm run build

View File

@@ -0,0 +1,45 @@
import json
import sys
import httpx
def create_child_keys(base_url: str, api_key: str, count: int = 3) -> list[str]:
headers = {"Authorization": f"Bearer {api_key}"}
print(f"Requesting {count} child keys from {base_url}...")
child_keys = []
for i in range(count):
try:
response = httpx.post(f"{base_url}/v1/balance/child-key", headers=headers)
if response.status_code == 200:
data = response.json()
child_keys.append(data["api_key"])
print(
f" [{i + 1}] Created: {data['api_key']} (Cost: {data['cost_msats']} msats)"
)
else:
print(f" [{i + 1}] Failed: {response.status_code} - {response.text}")
except Exception as e:
print(f" [{i + 1}] Error: {str(e)}")
return child_keys
if __name__ == "__main__":
if len(sys.argv) < 2:
print("Usage: python create_child_keys.py <api_key_or_cashu_token> [base_url]")
sys.exit(1)
auth_key = sys.argv[1]
base_url = sys.argv[2] if len(sys.argv) > 2 else "http://localhost:8000"
keys = create_child_keys(base_url, auth_key)
if keys:
print("\nSuccessfully created child keys:")
print(json.dumps(keys, indent=2))
else:
print("\nNo child keys were created.")

View File

@@ -0,0 +1,42 @@
"""
Revision ID: a86e5348850b
Revises: b9667ffc5701
Create Date: 2026-01-10 18:57:48.475781
"""
import sqlalchemy as sa
import sqlmodel
from alembic import op
# revision identifiers, used by Alembic.
revision = "a86e5348850b"
down_revision = "b9667ffc5701"
branch_labels = None
depends_on = None
def upgrade() -> None:
# Use batch_alter_table for SQLite compatibility
with op.batch_alter_table("api_keys", schema=None) as batch_op:
batch_op.add_column(
sa.Column(
"parent_key_hash", sqlmodel.sql.sqltypes.AutoString(), nullable=True
)
)
batch_op.create_index(
batch_op.f("ix_api_keys_parent_key_hash"), ["parent_key_hash"], unique=False
)
batch_op.create_foreign_key(
"fk_api_keys_parent_key_hash",
"api_keys",
["parent_key_hash"],
["hashed_key"],
)
def downgrade() -> None:
with op.batch_alter_table("api_keys", schema=None) as batch_op:
batch_op.drop_constraint("fk_api_keys_parent_key_hash", type_="foreignkey")
batch_op.drop_index(batch_op.f("ix_api_keys_parent_key_hash"))
batch_op.drop_column("parent_key_hash")

View File

@@ -1,6 +1,6 @@
[project]
name = "routstr"
version = "0.2.1"
version = "0.2.2"
description = "Payment proxy for your LLM endpoint using cashu and nostr."
readme = "README.md"
requires-python = ">=3.11"

View File

@@ -84,93 +84,26 @@ def get_provider_penalty(provider: "BaseUpstreamProvider") -> float:
return penalty
def should_prefer_model(
candidate_model: "Model",
candidate_provider: "BaseUpstreamProvider",
current_model: "Model",
current_provider: "BaseUpstreamProvider",
alias: str,
) -> bool:
"""Determine if candidate model should replace current model for an alias.
This is the core decision function for model prioritization. It considers:
1. Alias matching quality (exact match vs. canonical slug match)
2. Model cost (lower is better)
3. Provider penalties (e.g., slight preference against OpenRouter)
Args:
candidate_model: The new model being considered
candidate_provider: Provider offering the candidate model
current_model: The currently selected model for this alias
current_provider: Provider offering the current model
alias: The model alias being mapped
Returns:
True if candidate should replace current, False otherwise
"""
def get_base_model_id(model_id: str) -> str:
"""Get base model ID by removing provider prefix."""
return model_id.split("/", 1)[1] if "/" in model_id else model_id
def alias_priority(model: "Model") -> int:
"""Rank how strong the mapping of alias->model is.
Highest priority when alias exactly equals the model ID without provider prefix.
Next when alias equals canonical slug without prefix. Otherwise lowest.
"""
model_base = get_base_model_id(model.id)
if model_base == alias:
return 3
if model.canonical_slug:
canonical_base = get_base_model_id(model.canonical_slug)
if canonical_base == alias:
return 2
return 1
candidate_alias_priority = alias_priority(candidate_model)
current_alias_priority = alias_priority(current_model)
# If candidate has better alias match, prefer it regardless of cost
if candidate_alias_priority > current_alias_priority:
return True
# If current has better alias match, keep it regardless of cost
if current_alias_priority > candidate_alias_priority:
return False
# Same alias priority - compare costs
candidate_cost = calculate_model_cost_score(candidate_model)
current_cost = calculate_model_cost_score(current_model)
# Apply provider penalties
candidate_adjusted = candidate_cost * get_provider_penalty(candidate_provider)
current_adjusted = current_cost * get_provider_penalty(current_provider)
# Prefer lower adjusted cost
should_replace = candidate_adjusted < current_adjusted
return should_replace
def create_model_mappings(
upstreams: list["BaseUpstreamProvider"],
overrides_by_id: dict[str, tuple],
disabled_model_ids: set[str],
) -> tuple[dict[str, "Model"], dict[str, "BaseUpstreamProvider"], dict[str, "Model"]]:
) -> tuple[
dict[str, "Model"], dict[str, list["BaseUpstreamProvider"]], dict[str, "Model"]
]:
"""Create optimal model mappings based on cost and provider preferences.
This is the main entry point for the algorithm. It processes all upstream providers
and creates three mappings based on cost optimization:
1. model_instances: alias -> Model (all model aliases mapped to their Model objects)
2. provider_map: alias -> UpstreamProvider (which provider to use for each alias)
2. provider_map: alias -> List[UpstreamProvider] (sorted list of providers for each alias)
3. unique_models: base_id -> Model (unique models without provider prefixes)
The algorithm:
- Processes non-OpenRouter providers first (they're typically cheaper)
- Then processes OpenRouter models (they can still win if cheaper)
- For each model alias, uses should_prefer_model() to select the best provider
- For each model alias, collects all candidates and sorts them by priority and cost.
Args:
upstreams: List of all upstream provider instances
@@ -183,8 +116,7 @@ def create_model_mappings(
from .payment.models import _row_to_model
from .upstream.helpers import resolve_model_alias
model_instances: dict[str, "Model"] = {}
provider_map: dict[str, "BaseUpstreamProvider"] = {}
candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {}
unique_models: dict[str, "Model"] = {}
# Separate OpenRouter from other providers
@@ -202,24 +134,14 @@ 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 _maybe_set_alias(
def _add_candidate(
alias: str, model: "Model", provider: "BaseUpstreamProvider"
) -> None:
"""Set alias to model/provider if not set or if new model is preferred."""
"""Add candidate model/provider for an alias."""
alias_lower = alias.lower()
existing_model = model_instances.get(alias_lower)
if not existing_model:
# No existing mapping, set it
model_instances[alias_lower] = model
provider_map[alias_lower] = provider
else:
# Check if candidate should replace existing
existing_provider = provider_map[alias_lower]
if should_prefer_model(
model, provider, existing_model, existing_provider, alias
):
model_instances[alias_lower] = model
provider_map[alias_lower] = provider
if alias_lower not in candidates:
candidates[alias_lower] = []
candidates[alias_lower].append((model, provider))
def process_provider_models(
upstream: "BaseUpstreamProvider", is_openrouter: bool = False
@@ -266,21 +188,55 @@ def create_model_mappings(
# Try to set each alias
for alias in aliases:
_maybe_set_alias(alias, model_to_use, upstream)
_add_candidate(alias, model_to_use, upstream)
# Process non-OpenRouter providers first (they're typically cheaper)
# Process non-OpenRouter providers first
for upstream in other_upstreams:
process_provider_models(upstream, is_openrouter=False)
# Process OpenRouter last - models only win if they're cheaper or better matched
# Process OpenRouter last
if openrouter:
process_provider_models(openrouter, is_openrouter=True)
# Log provider distribution
# Sort candidates and build final maps
model_instances: dict[str, "Model"] = {}
provider_map: dict[str, list["BaseUpstreamProvider"]] = {}
def alias_priority(model: "Model", alias: str) -> int:
"""Rank how strong the mapping of alias->model is."""
model_base = get_base_model_id(model.id)
if model_base == alias:
return 3
if model.canonical_slug:
canonical_base = get_base_model_id(model.canonical_slug)
if canonical_base == alias:
return 2
return 1
for alias, items in candidates.items():
# Sort key: (priority DESC, cost ASC)
# Using negative cost for DESC sort overall to keep high priority first
def sort_key(item: tuple["Model", "BaseUpstreamProvider"]) -> tuple[int, float]:
model, provider = item
priority = alias_priority(model, alias)
cost = calculate_model_cost_score(model)
penalty = get_provider_penalty(provider)
adjusted_cost = cost * penalty
return (priority, -adjusted_cost)
items.sort(key=sort_key, reverse=True)
best_model, best_provider = items[0]
model_instances[alias] = best_model
provider_map[alias] = [p for _, p in items]
# Log provider distribution (using top provider for stats)
provider_counts: dict[str, int] = {}
for provider in provider_map.values():
provider_name = getattr(provider, "upstream_name", "unknown")
provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1
for providers in provider_map.values():
if providers:
provider = providers[0]
provider_name = getattr(provider, "upstream_name", "unknown")
provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1
logger.debug(
f"Updated model mappings with ({len(unique_models)} unique models and {len(model_instances)} aliases)",

View File

@@ -286,30 +286,55 @@ async def validate_bearer_key(
)
async def get_billing_key(key: ApiKey, session: AsyncSession) -> ApiKey:
"""Returns the key that should be charged for the request."""
if key.parent_key_hash:
parent = await session.get(ApiKey, key.parent_key_hash)
if parent:
# We want to keep the total_requests and total_spent on the child key
# but use the balance and reserved_balance of the parent.
# However, pay_for_request updates reserved_balance and total_requests.
# To stay simple, we charge the parent's balance and update parent's total_requests.
return parent
else:
logger.error(
"Parent key not found for child key",
extra={
"child_key_hash": key.hashed_key[:8] + "...",
"parent_key_hash": key.parent_key_hash[:8] + "...",
},
)
return key
async def pay_for_request(
key: ApiKey, cost_per_request: int, session: AsyncSession
) -> int:
"""Process payment for a request."""
billing_key = await get_billing_key(key, session)
logger.info(
"Processing payment for request",
extra={
"key_hash": key.hashed_key[:8] + "...",
"current_balance": key.balance,
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"current_balance": billing_key.balance,
"required_cost": cost_per_request,
"sufficient_balance": key.balance >= cost_per_request,
"sufficient_balance": billing_key.balance >= cost_per_request,
},
)
if key.total_balance < cost_per_request:
if billing_key.total_balance < cost_per_request:
logger.warning(
"Insufficient balance for request",
extra={
"key_hash": key.hashed_key[:8] + "...",
"balance": key.balance,
"reserved_balance": key.reserved_balance,
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"required": cost_per_request,
"shortfall": cost_per_request - key.total_balance,
"shortfall": cost_per_request - billing_key.total_balance,
},
)
@@ -317,7 +342,7 @@ async def pay_for_request(
status_code=402,
detail={
"error": {
"message": f"Insufficient balance: {cost_per_request} mSats required. {key.total_balance} available. (reserved: {key.reserved_balance})",
"message": f"Insufficient balance: {cost_per_request} mSats required. {billing_key.total_balance} available. (reserved: {billing_key.reserved_balance})",
"type": "insufficient_quota",
"code": "insufficient_balance",
}
@@ -328,15 +353,16 @@ async def pay_for_request(
"Charging base cost for request",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"cost": cost_per_request,
"balance_before": key.balance,
"balance_before": billing_key.balance,
},
)
# Charge the base cost for the request atomically to avoid race conditions
stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.where(col(ApiKey.balance) - col(ApiKey.reserved_balance) >= cost_per_request)
.values(
reserved_balance=col(ApiKey.reserved_balance) + cost_per_request,
@@ -344,6 +370,16 @@ async def pay_for_request(
)
)
result = await session.exec(stmt) # type: ignore[call-overload]
# Also increment total_requests on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(total_requests=col(ApiKey.total_requests) + 1)
)
await session.exec(child_stmt) # type: ignore[call-overload]
await session.commit()
if result.rowcount == 0:
@@ -351,8 +387,9 @@ async def pay_for_request(
"Concurrent request depleted balance",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"required_cost": cost_per_request,
"current_balance": key.balance,
"current_balance": billing_key.balance,
},
)
@@ -361,23 +398,26 @@ async def pay_for_request(
status_code=402,
detail={
"error": {
"message": f"Insufficient balance: {cost_per_request} mSats required. {key.balance} available.",
"message": f"Insufficient balance: {cost_per_request} mSats required. {billing_key.balance} available.",
"type": "insufficient_quota",
"code": "insufficient_balance",
}
},
)
await session.refresh(key)
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
logger.info(
"Payment processed successfully",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"charged_amount": cost_per_request,
"new_balance": key.balance,
"total_spent": key.total_spent,
"total_requests": key.total_requests,
"new_balance": billing_key.balance,
"total_spent": billing_key.total_spent,
"total_requests": billing_key.total_requests,
},
)
@@ -387,9 +427,11 @@ async def pay_for_request(
async def revert_pay_for_request(
key: ApiKey, session: AsyncSession, cost_per_request: int
) -> None:
billing_key = await get_billing_key(key, session)
stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=col(ApiKey.reserved_balance) - cost_per_request,
total_requests=col(ApiKey.total_requests) - 1,
@@ -397,27 +439,40 @@ async def revert_pay_for_request(
)
result = await session.exec(stmt) # type: ignore[call-overload]
# Also decrement total_requests on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(total_requests=col(ApiKey.total_requests) - 1)
)
await session.exec(child_stmt) # type: ignore[call-overload]
await session.commit()
if result.rowcount == 0:
logger.error(
"Failed to revert payment - insufficient reserved balance",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"cost_to_revert": cost_per_request,
"current_reserved_balance": key.reserved_balance,
"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. {key.balance} available.",
"message": f"failed to revert request payment: {cost_per_request} mSats required. {billing_key.balance} available.",
"type": "payment_error",
"code": "payment_error",
}
},
)
await session.refresh(key)
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
async def adjust_payment_for_tokens(
@@ -428,15 +483,17 @@ async def adjust_payment_for_tokens(
This is called after the initial payment and the upstream request is complete.
Returns cost data to be included in the response.
"""
billing_key = await get_billing_key(key, session)
model = response_data.get("model", "unknown")
logger.debug(
"Starting payment adjustment for tokens",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model,
"deducted_max_cost": deducted_max_cost,
"current_balance": key.balance,
"current_balance": billing_key.balance,
"has_usage": "usage" in response_data,
},
)
@@ -446,8 +503,10 @@ async def adjust_payment_for_tokens(
try:
release_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost
)
)
await session.exec(release_stmt) # type: ignore[call-overload]
await session.commit()
@@ -455,13 +514,18 @@ async def adjust_payment_for_tokens(
"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",
extra={"error": str(e), "key_hash": key.hashed_key[:8] + "..."},
extra={
"error": str(e),
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
},
)
match await calculate_cost(response_data, deducted_max_cost, session):
@@ -470,6 +534,7 @@ async def adjust_payment_for_tokens(
"Using max cost data (no token adjustment)",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model,
"max_cost": cost.total_msats,
},
@@ -477,7 +542,7 @@ async def adjust_payment_for_tokens(
# Finalize by releasing reservation and charging max cost
finalize_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost,
balance=col(ApiKey.balance) - cost.total_msats,
@@ -485,27 +550,41 @@ async def adjust_payment_for_tokens(
)
)
result = await session.exec(finalize_stmt) # type: ignore[call-overload]
# Also update total_spent on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(total_spent=col(ApiKey.total_spent) + cost.total_msats)
)
await session.exec(child_stmt) # type: ignore[call-overload]
await session.commit()
if result.rowcount == 0:
logger.error(
"Failed to finalize max-cost payment - retrying reservation release",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"deducted_max_cost": deducted_max_cost,
"current_reserved_balance": key.reserved_balance,
"current_reserved_balance": billing_key.reserved_balance,
"total_cost": cost.total_msats,
"model": model,
},
)
await release_reservation_only()
else:
await session.refresh(key)
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
logger.info(
"Max cost payment finalized",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"charged_amount": cost.total_msats,
"new_balance": key.balance,
"new_balance": billing_key.balance,
"model": model,
},
)
@@ -521,6 +600,7 @@ async def adjust_payment_for_tokens(
"Calculated token-based cost",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model,
"token_cost": cost.total_msats,
"deducted_max_cost": deducted_max_cost,
@@ -533,11 +613,15 @@ async def adjust_payment_for_tokens(
if cost_difference == 0:
logger.debug(
"Finalizing with exact reserved cost",
extra={"key_hash": key.hashed_key[:8] + "...", "model": model},
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model,
},
)
finalize_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
@@ -546,8 +630,20 @@ async def adjust_payment_for_tokens(
)
)
await session.exec(finalize_stmt) # type: ignore[call-overload]
# Also update total_spent on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(total_spent=col(ApiKey.total_spent) + total_cost_msats)
)
await session.exec(child_stmt) # type: ignore[call-overload]
await session.commit()
await session.refresh(key)
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
return cost.dict()
# this should never happen why do we handle this???
@@ -557,16 +653,17 @@ async def adjust_payment_for_tokens(
"Additional charge required for token usage",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"additional_charge": cost_difference,
"current_balance": key.balance,
"sufficient_balance": key.balance >= cost_difference,
"current_balance": billing_key.balance,
"sufficient_balance": billing_key.balance >= cost_difference,
"model": model,
},
)
finalize_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
@@ -575,18 +672,31 @@ async def adjust_payment_for_tokens(
)
)
result = await session.exec(finalize_stmt) # type: ignore[call-overload]
# Also update total_spent on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(total_spent=col(ApiKey.total_spent) + total_cost_msats)
)
await session.exec(child_stmt) # type: ignore[call-overload]
await session.commit()
if result.rowcount:
cost.total_msats = total_cost_msats
await session.refresh(key)
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
logger.info(
"Finalized payment with additional charge",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"charged_amount": total_cost_msats,
"new_balance": key.balance,
"new_balance": billing_key.balance,
"model": model,
},
)
@@ -595,6 +705,7 @@ async def adjust_payment_for_tokens(
"Failed to finalize additional charge - releasing reservation",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"attempted_charge": total_cost_msats,
"model": model,
},
@@ -607,15 +718,16 @@ async def adjust_payment_for_tokens(
"Refunding excess payment",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"refund_amount": refund,
"current_balance": key.balance,
"current_balance": billing_key.balance,
"model": model,
},
)
refund_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
@@ -624,6 +736,16 @@ async def adjust_payment_for_tokens(
)
)
result = await session.exec(refund_stmt) # type: ignore[call-overload]
# Also update total_spent on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(total_spent=col(ApiKey.total_spent) + total_cost_msats)
)
await session.exec(child_stmt) # type: ignore[call-overload]
await session.commit()
if result.rowcount == 0:
@@ -631,8 +753,9 @@ async def adjust_payment_for_tokens(
"Failed to finalize payment - releasing reservation",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"deducted_max_cost": deducted_max_cost,
"current_reserved_balance": key.reserved_balance,
"current_reserved_balance": billing_key.reserved_balance,
"total_cost": total_cost_msats,
"model": model,
},
@@ -640,14 +763,17 @@ async def adjust_payment_for_tokens(
await release_reservation_only()
else:
cost.total_msats = total_cost_msats
await session.refresh(key)
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
logger.info(
"Refund processed successfully",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"refunded_amount": refund,
"new_balance": key.balance,
"new_balance": billing_key.balance,
"final_cost": cost.total_msats,
"model": model,
},

View File

@@ -32,16 +32,30 @@ async def get_key_from_header(
)
# TODO: remove this endpoint when frontend is updated
@router.get("/", include_in_schema=False)
async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict:
async def get_balance_info(key: ApiKey, session: AsyncSession) -> dict:
from .auth import get_billing_key
billing_key = await get_billing_key(key, session)
return {
"api_key": "sk-" + key.hashed_key,
"balance": key.balance,
"reserved": key.reserved_balance,
"balance": billing_key.balance,
"reserved": billing_key.reserved_balance,
"is_child": key.parent_key_hash is not None,
"parent_key": "sk-" + key.parent_key_hash if key.parent_key_hash else None,
"total_requests": key.total_requests,
"total_spent": key.total_spent,
}
# TODO: remove this endpoint when frontend is updated
@router.get("/", include_in_schema=False)
async def account_info(
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
) -> dict:
return await get_balance_info(key, session)
# TODO: Implement POST /v1/wallet/create endpoint
# This endpoint should accept:
# - cashu_token (required): The eCash token to deposit
@@ -66,12 +80,11 @@ async def create_balance(
@router.get("/info")
async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict:
return {
"api_key": "sk-" + key.hashed_key,
"balance": key.balance,
"reserved": key.reserved_balance,
}
async def wallet_info(
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
) -> dict:
return await get_balance_info(key, session)
class TopupRequest(BaseModel):
@@ -85,6 +98,10 @@ async def topup_wallet_endpoint(
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
) -> dict[str, int]:
from .auth import get_billing_key
billing_key = await get_billing_key(key, session)
if topup_request is not None:
cashu_token = topup_request.cashu_token
if cashu_token is None:
@@ -94,7 +111,7 @@ async def topup_wallet_endpoint(
if len(cashu_token) < 10 or "cashu" not in cashu_token:
raise HTTPException(status_code=400, detail="Invalid token format")
try:
amount_msats = await credit_balance(cashu_token, key, session)
amount_msats = await credit_balance(cashu_token, billing_key, session)
except ValueError as e:
error_msg = str(e)
if "already spent" in error_msg.lower():
@@ -154,7 +171,14 @@ async def refund_wallet_endpoint(
return cached
key: ApiKey = await validate_bearer_key(bearer_value, session)
remaining_balance_msats: int = key.balance
if key.parent_key_hash:
raise HTTPException(
status_code=400,
detail="Cannot refund child key. Please refund the parent key instead.",
)
remaining_balance_msats: int = key.total_balance
if key.refund_currency == "sat":
remaining_balance = remaining_balance_msats // 1000
@@ -227,6 +251,53 @@ async def donate(token: str, ref: str | None = None) -> str:
return "Invalid token."
@router.post("/child-key")
async def create_child_key(
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
) -> dict:
"""Creates a child API key that uses the parent's balance."""
# Check if this is already a child key
if key.parent_key_hash:
raise HTTPException(
status_code=400,
detail="Cannot create a child key for another child key.",
)
cost = settings.child_key_cost
if key.total_balance < cost:
raise HTTPException(
status_code=402,
detail=f"Insufficient balance to create child key. {cost} mSats required.",
)
# Deduct cost from parent
key.balance -= cost
key.total_spent += cost
session.add(key)
# Generate new key
import secrets
new_key_raw = secrets.token_hex(32)
new_key_hash = new_key_raw # We use the raw key as the hash for sk- keys
child_key = ApiKey(
hashed_key=new_key_hash,
balance=0,
parent_key_hash=key.hashed_key,
)
session.add(child_key)
await session.commit()
return {
"api_key": "sk-" + new_key_hash,
"cost_msats": cost,
"parent_balance": key.balance,
}
@router.api_route(
"/{path:path}",
methods=["GET", "POST", "PUT", "DELETE"],

View File

@@ -124,7 +124,7 @@ async def partial_apikeys(request: Request) -> str:
rows = "".join(
[
f"<tr><td>{key.hashed_key}</td><td>{key.balance}</td><td>{key.total_spent}</td><td>{key.total_requests}</td><td>{key.refund_address}</td><td>{fmt_time(key.key_expiry_time)}</td></tr>"
f"<tr><td>{key.hashed_key}{' <br><small>(Child of ' + key.parent_key_hash[:8] + '...)</small>' if key.parent_key_hash else ''}</td><td>{key.balance}</td><td>{key.total_spent}</td><td>{key.total_requests}</td><td>{key.refund_address}</td><td>{fmt_time(key.key_expiry_time)}</td></tr>"
for key in api_keys
]
)
@@ -158,6 +158,7 @@ async def get_temporary_balances_api(request: Request) -> list[dict[str, object]
"total_requests": key.total_requests,
"refund_address": key.refund_address,
"key_expiry_time": key.key_expiry_time,
"parent_key_hash": key.parent_key_hash,
}
for key in api_keys
]

View File

@@ -6,7 +6,7 @@ from typing import AsyncGenerator
from alembic import command
from alembic.config import Config
from sqlalchemy.ext.asyncio.engine import create_async_engine
from sqlmodel import Field, Relationship, SQLModel, func, select
from sqlmodel import Field, Relationship, SQLModel, func, select, update
from sqlmodel.ext.asyncio.session import AsyncSession
from .logging import get_logger
@@ -47,12 +47,23 @@ class ApiKey(SQLModel, table=True): # type: ignore
default=None,
description="Currency of the cashu-token",
)
parent_key_hash: str | None = Field(
default=None, foreign_key="api_keys.hashed_key", index=True
)
@property
def total_balance(self) -> int:
return self.balance - self.reserved_balance
async def reset_all_reserved_balances(session: AsyncSession) -> None:
logger.info("Resetting all reserved balances to 0")
stmt = update(ApiKey).values(reserved_balance=0)
await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
logger.info("Reserved balances reset successfully")
class ModelRow(SQLModel, table=True): # type: ignore
__tablename__ = "models"
id: str = Field(primary_key=True)

View File

@@ -6,6 +6,15 @@ from .logging import get_logger
logger = get_logger(__name__)
class UpstreamError(Exception):
"""Exception raised when an upstream provider fails."""
def __init__(self, message: str, status_code: int = 502):
self.message = message
self.status_code = status_code
super().__init__(message)
async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse:
"""Handle HTTP exceptions and include request ID in response."""
request_id = getattr(request.state, "request_id", "unknown")

View File

@@ -13,10 +13,7 @@ from starlette.exceptions import HTTPException
from ..balance import balance_router, deprecated_wallet_router
from ..discovery import providers_cache_refresher, providers_router
from ..nip91 import announce_provider
from ..payment.models import (
models_router,
update_sats_pricing,
)
from ..payment.models import models_router, update_sats_pricing
from ..payment.price import update_prices_periodically
from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically
from ..wallet import periodic_payout
@@ -33,9 +30,9 @@ setup_logging()
logger = get_logger(__name__)
if os.getenv("VERSION_SUFFIX") is not None:
__version__ = f"0.2.1-{os.getenv('VERSION_SUFFIX')}"
__version__ = f"0.2.2-{os.getenv('VERSION_SUFFIX')}"
else:
__version__ = "0.2.1"
__version__ = "0.2.2"
@asynccontextmanager
@@ -61,6 +58,15 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
# Initialize application settings (env -> computed -> DB precedence)
async with create_session() as session:
s = await SettingsService.initialize(session)
if s.reset_reserved_balance_on_startup:
from .db import reset_all_reserved_balances
await reset_all_reserved_balances(session)
if not s.admin_password:
logger.warning(
f"Admin password is not set. Visit {s.http_url or 'http://localhost:8000'}/admin to set the password."
)
# Apply app metadata from settings
try:

View File

@@ -6,7 +6,7 @@ import os
from datetime import datetime, timezone
from typing import Any
from pydantic.v1 import BaseModel, BaseSettings, Field, validator
from pydantic.v1 import BaseModel, BaseSettings, Field
from sqlmodel.ext.asyncio.session import AsyncSession
@@ -37,13 +37,6 @@ class Settings(BaseSettings):
# Cashu
cashu_mints: list[str] = Field(default_factory=list, env="CASHU_MINTS")
@validator("cashu_mints", pre=True, each_item=True)
def normalize_mint_url(cls, v: str) -> str:
if isinstance(v, str):
return v.rstrip("/")
return v
receive_ln_address: str = Field(default="", env="RECEIVE_LN_ADDRESS")
primary_mint: str = Field(default="", env="PRIMARY_MINT_URL")
primary_mint_unit: str = Field(default="sat", env="PRIMARY_MINT_UNIT")
@@ -59,14 +52,18 @@ class Settings(BaseSettings):
exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE")
upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE")
tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE")
child_key_cost: int = Field(default=1000, env="CHILD_KEY_COST")
# Minimum per-request charge in millisatoshis when model pricing is free/zero
min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT")
reset_reserved_balance_on_startup: bool = Field(
default=True, env="RESET_RESERVED_BALANCE_ON_STARTUP"
) # deactivate in horizontal scaling setups
# Network
cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS")
tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL")
providers_refresh_interval_seconds: int = Field(
default=300, env="PROVIDERS_REFRESH_INTERVAL_SECONDS"
default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS"
)
pricing_refresh_interval_seconds: int = Field(
default=120, env="PRICING_REFRESH_INTERVAL_SECONDS"

View File

@@ -6,7 +6,7 @@ from typing import Any
import httpx
import websockets
from fastapi import APIRouter
from fastapi import APIRouter, HTTPException
from .core.logging import get_logger
from .core.settings import settings
@@ -389,6 +389,9 @@ async def get_providers(
Return cached providers. If include_json, return provider+health; otherwise provider only.
Optional filter by pubkey.
"""
if settings.providers_refresh_interval_seconds == 0:
raise HTTPException(status_code=404, detail="Provider discovery is disabled")
cache = await get_cache()
if not cache:
await refresh_providers_cache(pubkey=pubkey)

View File

@@ -283,7 +283,7 @@ async def raw_send_to_lnurl(
f"({min_sendable_sat} - {max_sendable_sat} {unit})"
)
estimated_fees_sat = int(max(math.ceil((amount_msat / 1000) * 0.01), 2))
estimated_fees_sat = int(max(math.ceil((amount_msat / 1000) * 0.01), 2)) + 1
estimated_fees_msat = estimated_fees_sat * 1000
final_amount = amount_msat - estimated_fees_msat

View File

@@ -253,6 +253,11 @@ async def list_models(
else 1.01,
)
for r in rows
if include_disabled
or (
r.upstream_provider_id in providers_by_id
and providers_by_id[r.upstream_provider_id].enabled
)
]
@@ -265,6 +270,8 @@ async def get_model_by_id(
if not row or not row.enabled:
return None
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider or not provider.enabled:
return None
provider_fee = provider.provider_fee if provider else 1.01
return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee)

View File

@@ -79,15 +79,29 @@ async def _fetch_btc_usd_price() -> float:
"""Fetch the lowest BTC/USD price from multiple exchanges."""
async with httpx.AsyncClient(timeout=30.0) as client:
try:
prices = await asyncio.gather(
_kraken_btc_usd(client),
_coinbase_btc_usd(client),
_binance_btc_usdt(client),
)
valid_prices = [price for price in prices if price is not None]
tasks = [
asyncio.create_task(_kraken_btc_usd(client)),
asyncio.create_task(_coinbase_btc_usd(client)),
asyncio.create_task(_binance_btc_usdt(client)),
]
valid_prices: list[float] = []
for future in asyncio.as_completed(tasks):
price = await future
if price is not None:
valid_prices.append(price)
if len(valid_prices) >= 2:
break
for task in tasks:
if not task.done():
task.cancel()
if not valid_prices:
logger.error("No valid BTC prices obtained from any exchange")
raise ValueError("Unable to fetch BTC price from any exchange")
return min(valid_prices)
except Exception as e:
logger.error(

View File

@@ -16,7 +16,7 @@ from .core.db import (
create_session,
get_session,
)
from .core.settings import settings
from .core.exceptions import UpstreamError
from .payment.helpers import (
calculate_discounted_max_cost,
check_token_balance,
@@ -26,14 +26,15 @@ from .payment.helpers import (
from .payment.models import Model
from .upstream import BaseUpstreamProvider
from .upstream.helpers import init_upstreams
from .wallet import deserialize_token_from_string
logger = get_logger(__name__)
proxy_router = APIRouter()
_upstreams: list[BaseUpstreamProvider] = []
_model_instances: dict[str, Model] = {} # All aliases -> Model
_provider_map: dict[str, BaseUpstreamProvider] = {} # All aliases -> Provider
_provider_map: dict[
str, list[BaseUpstreamProvider]
] = {} # All aliases -> List[Provider]
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
@@ -70,8 +71,8 @@ def get_model_instance(model_id: str) -> Model | None:
return _model_instances.get(model_id.lower())
def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None:
"""Get UpstreamProvider for model ID from global cache."""
def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None:
"""Get UpstreamProvider list for model ID from global cache."""
return _provider_map.get(model_id.lower())
@@ -98,6 +99,8 @@ async def refresh_model_maps() -> None:
disabled_model_ids: set[str] = set()
for provider in provider_rows:
if not provider.enabled:
continue
for model in provider.models:
if model.enabled:
overrides_by_id[model.id] = (model, provider.provider_fee)
@@ -148,32 +151,14 @@ async def proxy(
else:
model_id = request_body_dict.get("model", "unknown")
if "https://testnut.cashu.space" in settings.cashu_mints:
try:
token_str = None
if x_cashu_header := headers.get("x-cashu"):
token_str = x_cashu_header
elif auth_header := headers.get("authorization"):
parts = auth_header.split(" ")
if len(parts) > 1 and not parts[1].startswith("sk-"):
token_str = parts[1]
if token_str:
token_obj = deserialize_token_from_string(token_str)
if token_obj.mint == "https://testnut.cashu.space":
model_id = "mock/gpt-420-mock"
request_body_dict["model"] = model_id
except Exception:
pass
model_obj = get_model_instance(model_id)
if not model_obj:
return create_error_response(
"invalid_model", f"Model '{model_id}' not found", 400, request=request
)
upstream = get_provider_for_model(model_id)
if not upstream:
upstreams = get_provider_for_model(model_id)
if not upstreams:
return create_error_response(
"invalid_model",
f"No provider found for model '{model_id}'",
@@ -181,6 +166,10 @@ async def proxy(
request=request,
)
# todo figure out cost calculation since fallback provider is usually not the same price
# Use first provider for initial checks/cost calculation
# primary_upstream = upstreams[0]
_max_cost_for_model = await get_max_cost_for_model(
model=model_id, session=session, model_obj=model_obj
)
@@ -190,14 +179,31 @@ async def proxy(
check_token_balance(headers, request_body_dict, max_cost_for_model)
if x_cashu := headers.get("x-cashu", None):
if is_responses_api:
return await upstream.handle_x_cashu_responses(
request, x_cashu, path, max_cost_for_model, model_obj
)
else:
return await upstream.handle_x_cashu(
request, x_cashu, path, max_cost_for_model, model_obj
)
last_error = None
for i, upstream in enumerate(upstreams):
try:
if is_responses_api:
return await upstream.handle_x_cashu_responses(
request, x_cashu, path, max_cost_for_model, model_obj
)
else:
return await upstream.handle_x_cashu(
request, x_cashu, path, max_cost_for_model, model_obj
)
except UpstreamError as e:
logger.warning(
f"Upstream {upstream.provider_type} failed (x-cashu): {e}"
)
if i == len(upstreams) - 1:
last_error = e
continue
return create_error_response(
"upstream_error",
str(last_error) if last_error else "All upstreams failed",
502,
request=request,
)
elif auth := headers.get("authorization", None):
key = await get_bearer_token_key(headers, path, session, auth)
@@ -212,56 +218,152 @@ async def proxy(
)
logger.debug("Processing unauthenticated GET request", extra={"path": path})
headers = upstream.prepare_headers(dict(request.headers))
return await upstream.forward_get_request(request, path, headers)
last_error_response = None
for i, upstream in enumerate(upstreams):
try:
headers = upstream.prepare_headers(dict(request.headers))
response = await upstream.forward_get_request(request, path, headers)
if response.status_code in [502, 429] and i < len(upstreams) - 1:
error_message = ""
try:
if hasattr(response, "body"):
body_bytes = response.body
data = json.loads(body_bytes)
if "error" in data:
error_data = data["error"]
if isinstance(error_data, dict):
error_message = error_data.get("message", "")
elif isinstance(error_data, str):
error_message = error_data
except Exception:
pass
await upstream.on_upstream_error_redirect(
response.status_code, error_message
)
logger.warning(
f"Upstream {upstream.provider_type} returned {response.status_code} (GET), trying next provider",
extra={
"status_code": response.status_code,
"upstream": upstream.provider_type,
},
)
continue
return response
except UpstreamError as e:
logger.warning(f"Upstream {upstream.provider_type} failed (GET): {e}")
if i == len(upstreams) - 1:
last_error_response = create_error_response(
"upstream_error", str(e), 502, request=request
)
continue
return last_error_response or create_error_response(
"upstream_error", "All upstreams failed", 502, request=request
)
if request_body_dict:
await pay_for_request(key, max_cost_for_model, session)
headers = upstream.prepare_headers(dict(request.headers))
for i, upstream in enumerate(upstreams):
headers = upstream.prepare_headers(dict(request.headers))
if is_responses_api:
response = await upstream.forward_responses_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
else:
response = await upstream.forward_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
try:
if is_responses_api:
response = await upstream.forward_responses_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
else:
response = await upstream.forward_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
if response.status_code != 200:
await revert_pay_for_request(key, session, max_cost_for_model)
logger.warning(
"Upstream request failed, revert payment",
extra={
"status_code": response.status_code,
"path": path,
"key_hash": key.hashed_key[:8] + "...",
"key_balance": key.balance,
"max_cost_for_model": max_cost_for_model,
"upstream_headers": response.headers
if hasattr(response, "headers")
else None,
},
)
# Return the mapped error response generated earlier rather than masking with 502
return response
if response.status_code != 200:
# Check if we should retry (502 Upstream Error or 429 Rate Limit)
should_retry = response.status_code in [502, 429, 400, 401, 403, 404]
if should_retry and i < len(upstreams) - 1:
error_message = ""
try:
if hasattr(response, "body"):
body_bytes = response.body
data = json.loads(body_bytes)
if "error" in data:
error_data = data["error"]
if isinstance(error_data, dict):
error_message = error_data.get("message", "")
elif isinstance(error_data, str):
error_message = error_data
except Exception:
pass
return response
await upstream.on_upstream_error_redirect(
response.status_code, error_message
)
logger.warning(
f"Upstream {upstream.provider_type} returned {response.status_code}, trying next provider",
extra={
"status_code": response.status_code,
"upstream": upstream.provider_type,
},
)
continue
# 4xx error (user error), or other non-retryable error, or last provider failed
await revert_pay_for_request(key, session, max_cost_for_model)
logger.warning(
"Upstream request failed, revert payment",
extra={
"status_code": response.status_code,
"path": path,
"key_hash": key.hashed_key[:8] + "...",
"key_balance": key.balance,
"max_cost_for_model": max_cost_for_model,
"upstream_headers": response.headers
if hasattr(response, "headers")
else None,
},
)
return response
return response
except UpstreamError as e:
logger.warning(
f"Upstream {upstream.provider_type} failed: {e}",
extra={"retry": i < len(upstreams) - 1},
)
# If this was the last provider
if i == len(upstreams) - 1:
await revert_pay_for_request(key, session, max_cost_for_model)
return create_error_response(
"upstream_error", str(e), 502, request=request
)
# Otherwise loop continues to next provider
continue
# Should not be reached given logic above
return create_error_response(
"upstream_error", "All upstreams failed", 502, request=request
)
async def get_bearer_token_key(

View File

@@ -15,6 +15,7 @@ from pydantic import BaseModel
from ..auth import adjust_payment_for_tokens
from ..core import get_logger
from ..core.db import ApiKey, AsyncSession, create_session
from ..core.exceptions import UpstreamError
if TYPE_CHECKING:
from ..core.db import UpstreamProviderRow
@@ -340,6 +341,20 @@ class BaseUpstreamProvider:
message = preview[:500]
return message, upstream_code
async def on_upstream_error_redirect(
self, status_code: int, error_message: str
) -> None:
"""Hook called when the proxy redirects to another provider due to an error.
Subclasses can implement this to perform actions like disabling the provider
if it's out of balance.
Args:
status_code: The HTTP status code returned by the upstream
error_message: The error message extracted from the upstream response
"""
pass
async def map_upstream_error_response(
self, request: Request, path: str, upstream_response: httpx.Response
) -> Response:
@@ -1148,6 +1163,14 @@ class BaseUpstreamProvider:
)
if response.status_code != 200:
if response.status_code >= 500:
await response.aclose()
await client.aclose()
raise UpstreamError(
f"Upstream returned status {response.status_code}",
status_code=response.status_code,
)
try:
mapped_error = await self.map_upstream_error_response(
request, path, response
@@ -1238,6 +1261,9 @@ class BaseUpstreamProvider:
background=background_tasks,
)
except UpstreamError:
raise
except httpx.RequestError as exc:
await client.aclose()
error_type = type(exc).__name__
@@ -1265,9 +1291,7 @@ class BaseUpstreamProvider:
else:
error_message = f"Error connecting to upstream service: {error_type}"
return create_error_response(
"upstream_error", error_message, 502, request=request
)
raise UpstreamError(error_message, status_code=502)
except Exception as exc:
await client.aclose()
@@ -1380,6 +1404,14 @@ class BaseUpstreamProvider:
)
if response.status_code != 200:
if response.status_code >= 500:
await response.aclose()
await client.aclose()
raise UpstreamError(
f"Upstream returned status {response.status_code}",
status_code=response.status_code,
)
try:
mapped_error = await self.map_upstream_error_response(
request, path, response
@@ -1447,6 +1479,9 @@ class BaseUpstreamProvider:
background=background_tasks,
)
except UpstreamError:
raise
except httpx.RequestError as exc:
await client.aclose()
error_type = type(exc).__name__
@@ -1474,9 +1509,7 @@ class BaseUpstreamProvider:
else:
error_message = f"Error connecting to upstream service: {error_type}"
return create_error_response(
"upstream_error", error_message, 502, request=request
)
raise UpstreamError(error_message, status_code=502)
except Exception as exc:
await client.aclose()

View File

@@ -1,265 +0,0 @@
import asyncio
import json
import random
from typing import AsyncIterator
from fastapi import Request
from fastapi.responses import Response, StreamingResponse
from ..core.db import ApiKey, AsyncSession
from ..payment.models import Architecture, Model, Pricing
from .base import BaseUpstreamProvider
class MockUpstreamProvider(BaseUpstreamProvider):
"""Fack Mock Upstream provider specifically for Testing."""
provider_type = "mock"
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:
if path.endswith("chat/completions"):
is_streaming = False
if request_body:
request_data = json.loads(request_body)
is_streaming = request_data.get("stream", False)
if is_streaming:
async def fake_streaming_response(
chunk_size: int | None = None,
) -> AsyncIterator[bytes]:
suffix = random.randint(1000, 9999)
req_id = f"gen-mock-stream-{suffix}"
created = 1766138895
model = "mock/gpt-420-mock"
def make_chunk(
delta: dict,
finish_reason: str | None = None,
usage: dict | None = None,
) -> bytes:
chunk = {
"id": req_id,
"provider": "MockProvider",
"model": model,
"object": "chat.completion.chunk",
"created": created,
"choices": [
{
"index": 0,
"delta": delta,
"finish_reason": finish_reason,
"native_finish_reason": "completed"
if finish_reason
else None,
"logprobs": None,
}
],
}
if usage:
chunk["usage"] = usage
return f"data: {json.dumps(chunk)}\n\n".encode()
# 1. Initial chunk
yield make_chunk({"role": "assistant", "content": ""})
await asyncio.sleep(0.02)
# 2. Reasoning chunks
reasoning_tokens = ["Mock", " reason", "ing", "..."]
for token in reasoning_tokens:
delta = {
"role": "assistant",
"content": "",
"reasoning": token,
"reasoning_details": [
{
"type": "reasoning.summary",
"summary": token,
"format": "openai-responses-v1",
"index": 0,
}
],
}
yield make_chunk(delta)
await asyncio.sleep(0.03)
# 3. Content chunks
content_tokens = ["This", " is", " a", " mock", " stream", "."]
for token in content_tokens:
yield make_chunk({"role": "assistant", "content": token})
await asyncio.sleep(0.03)
# 4. Finish chunk
yield make_chunk(
{"role": "assistant", "content": ""}, finish_reason="stop"
)
# 5. Usage chunk
usage_data = {
"prompt_tokens": 10,
"completion_tokens": 20,
"total_tokens": 30,
"cost": 0.001,
"is_byok": False,
"prompt_tokens_details": {
"cached_tokens": 0,
"audio_tokens": 0,
"video_tokens": 0,
},
"cost_details": {
"upstream_inference_cost": None,
"upstream_inference_prompt_cost": 0,
"upstream_inference_completions_cost": 0.001,
},
"completion_tokens_details": {
"reasoning_tokens": 10,
"image_tokens": 0,
},
}
usage_chunk = {
"id": req_id,
"provider": "MockProvider",
"model": model,
"object": "chat.completion.chunk",
"created": created,
"choices": [
{
"index": 0,
"delta": {"role": "assistant", "content": ""},
"finish_reason": None,
"native_finish_reason": None,
"logprobs": None,
}
],
"usage": usage_data,
}
yield f"data: {json.dumps(usage_chunk)}\n\n".encode()
# 6. DONE
yield b"data: [DONE]\n\n"
# 7. Cost
cost_chunk = {
"cost": {
"base_msats": 0,
"input_msats": 2,
"output_msats": 10,
"total_msats": 12,
}
}
yield f"data: {json.dumps(cost_chunk)}\n\n".encode()
return StreamingResponse(
fake_streaming_response(),
200,
)
else:
suffix = random.randint(1000, 9999)
content_dict = {
"id": f"gen-mock-{suffix}",
"provider": "MockProvider",
"model": "mock/gpt-5-mini",
"object": "chat.completion",
"created": 1766138655,
"choices": [
{
"logprobs": None,
"finish_reason": "length",
"native_finish_reason": "max_output_tokens",
"index": 0,
"message": {
"role": "assistant",
"content": f"Mock Content {suffix}",
"refusal": None,
"reasoning": f"Mock Reasoning {suffix}",
"reasoning_details": [
{
"format": "openai-responses-v1",
"index": 0,
"type": "reasoning.summary",
"summary": f"Mock Summary {suffix}",
},
{
"id": f"rs_mock_{suffix}",
"format": "openai-responses-v1",
"index": 0,
"type": "reasoning.encrypted",
"data": "mock_encrypted_data",
},
],
},
}
],
"usage": {
"prompt_tokens": 10,
"completion_tokens": 10,
"total_tokens": 20,
"cost": 0,
"is_byok": False,
"prompt_tokens_details": {
"cached_tokens": 0,
"audio_tokens": 0,
"video_tokens": 0,
},
"cost_details": {
"upstream_inference_cost": None,
"upstream_inference_prompt_cost": 0,
"upstream_inference_completions_cost": 0,
},
"completion_tokens_details": {
"reasoning_tokens": 5,
"image_tokens": 0,
},
},
"cost": {
"base_msats": 0,
"input_msats": 0,
"output_msats": 0,
"total_msats": 0,
},
}
return Response(json.dumps(content_dict).encode(), 200)
elif path.endswith("embeddings"):
raise NotImplementedError
elif path.endswith("responses"):
raise NotImplementedError
else:
raise NotImplementedError
async def fetch_models(self) -> list[Model]:
return [
Model(
id="mock/gpt-420-mock",
name="mock/gpt-420-mock",
created=0,
description="mock model for testing",
context_length=8192,
architecture=Architecture(
modality="text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="",
instruct_type=None,
),
pricing=Pricing(prompt=0.01, completion=0.01),
),
]
def transform_model_name(self, model_id: str) -> str:
return "fake-model"
async def get_balance(self) -> float | None:
return 420.69

View File

@@ -102,6 +102,8 @@ async def get_all_models_with_overrides(
)
for row in override_rows
if row.upstream_provider_id is not None
and row.upstream_provider_id in providers_by_id
and providers_by_id[row.upstream_provider_id].enabled
}
all_models: dict[str, Model] = {}
@@ -218,14 +220,6 @@ async def init_upstreams() -> list[BaseUpstreamProvider]:
results = await asyncio.gather(*tasks)
upstreams = [p for p in results if p is not None]
if "https://testnut.cashu.space" in settings.cashu_mints:
from .fake import MockUpstreamProvider
mock_provider = MockUpstreamProvider("mock", "mock")
await mock_provider.refresh_models_cache()
upstreams.append(mock_provider)
logger.info("Initialized MockUpstreamProvider for testnut mint")
return upstreams

View File

@@ -196,6 +196,37 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
)
return []
async def on_upstream_error_redirect(
self, status_code: int, error_message: str
) -> None:
if "insufficient balance" in error_message.lower():
logger.warning(
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
and UpstreamProviderRow.api_key == self.api_key
)
result = await session.exec(statement)
provider = result.first()
if provider:
provider.enabled = False
session.add(provider)
await session.commit()
# Trigger re-initialization of providers
# Import here to avoid circular dependency
from ..proxy import reinitialize_upstreams
await reinitialize_upstreams()
async def create_account(self) -> dict[str, object]:
"""Create a new PPQ.AI account.

View File

@@ -81,7 +81,7 @@ async def swap_to_primary_mint(
amount_msat = token_amount
else:
raise ValueError("Invalid unit")
estimated_fee_sat = math.ceil(max(amount_msat // 1000 * 0.01, 2))
estimated_fee_sat = math.ceil(max(amount_msat // 1000 * 0.01, 2)) + 1
amount_msat_after_fee = amount_msat - estimated_fee_sat * 1000
primary_wallet = await get_wallet(settings.primary_mint, settings.primary_mint_unit)
@@ -313,8 +313,6 @@ async def periodic_payout() -> None:
try:
async with db.create_session() as session:
for mint_url in settings.cashu_mints:
if mint_url == "https://testnut.cashu.space":
continue
for unit in ["sat", "msat"]:
wallet = await get_wallet(mint_url, unit)
proofs = get_proofs_per_mint_and_unit(

View File

@@ -0,0 +1,137 @@
import secrets
from typing import Any
import pytest
from fastapi import HTTPException
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.auth import adjust_payment_for_tokens, pay_for_request
from routstr.balance import create_child_key
from routstr.core.db import ApiKey
from routstr.core.settings import settings
@pytest.mark.asyncio
async def test_child_key_flow(integration_session: AsyncSession) -> None:
# 1. Create a parent key with balance
parent_raw = "parent_test_key_" + secrets.token_hex(4)
parent_key = ApiKey(
hashed_key=parent_raw,
balance=10000, # 10 sats
)
integration_session.add(parent_key)
await integration_session.commit()
await integration_session.refresh(parent_key)
# Mock settings
settings.child_key_cost = 1000 # 1 sat
# 2. Call create_child_key
result = await create_child_key(parent_key, integration_session)
assert "api_key" in result
assert result["cost_msats"] == 1000
assert result["parent_balance"] == 9000
child_key_raw = result["api_key"][3:] # remove sk-
# 3. Verify child key exists in DB
child_key_db = await integration_session.get(ApiKey, child_key_raw)
assert child_key_db is not None
assert child_key_db.parent_key_hash == parent_key.hashed_key
assert child_key_db.balance == 0
# 4. Test payment with child key
cost = 500
await pay_for_request(child_key_db, cost, integration_session)
# Refresh keys
await integration_session.refresh(parent_key)
await integration_session.refresh(child_key_db)
# Parent should be charged
assert parent_key.reserved_balance == 500
assert parent_key.total_requests == 1
# Child should have total_requests incremented
assert child_key_db.total_requests == 1
# 5. Test adjustment
response_data = {"model": "test-model", "usage": {"total_tokens": 10}}
# Mock calculate_cost
import routstr.auth
from routstr.payment.cost_calculation import CostData
async def mock_calculate_cost(*args: Any, **kwargs: Any) -> CostData:
return CostData(
base_msats=0, input_msats=200, output_msats=200, total_msats=400
)
# Patch calculate_cost
original_calculate_cost = routstr.auth.calculate_cost
routstr.auth.calculate_cost = mock_calculate_cost
try:
adjustment = await adjust_payment_for_tokens(
child_key_db, response_data, integration_session, 500
)
assert adjustment["total_msats"] == 400
# Refresh keys
await integration_session.refresh(parent_key)
await integration_session.refresh(child_key_db)
# Parent should have updated balance and total_spent
assert parent_key.reserved_balance == 0
assert parent_key.balance == 9000 - 400
assert (
parent_key.total_spent == 1400
) # 1000 for child key creation + 400 for request
# Child should also have total_spent updated
assert child_key_db.total_spent == 400
finally:
routstr.auth.calculate_cost = original_calculate_cost
@pytest.mark.asyncio
async def test_child_key_insufficient_balance(
integration_session: AsyncSession,
) -> None:
parent_key = ApiKey(
hashed_key="poor_parent_" + secrets.token_hex(4),
balance=500,
)
integration_session.add(parent_key)
await integration_session.commit()
await integration_session.refresh(parent_key)
settings.child_key_cost = 1000
with pytest.raises(HTTPException) as exc:
await create_child_key(parent_key, integration_session)
assert exc.value.status_code == 402
@pytest.mark.asyncio
async def test_child_key_cannot_create_child(integration_session: AsyncSession) -> None:
parent_key = ApiKey(
hashed_key="parent_" + secrets.token_hex(4),
balance=10000,
)
child_key = ApiKey(
hashed_key="child_" + secrets.token_hex(4),
balance=0,
parent_key_hash=parent_key.hashed_key,
)
integration_session.add(parent_key)
integration_session.add(child_key)
await integration_session.commit()
await integration_session.refresh(child_key)
with pytest.raises(HTTPException) as exc:
await create_child_key(child_key, integration_session)
assert exc.value.status_code == 400
assert "Cannot create a child key for another child key" in str(exc.value.detail)

View File

@@ -10,7 +10,6 @@ from httpx import AsyncClient
from .utils import (
CashuTokenGenerator,
PerformanceValidator,
ResponseValidator,
)
@@ -159,30 +158,7 @@ async def test_error_handling(
assert response.status_code == 401
@pytest.mark.integration
@pytest.mark.asyncio
async def test_performance_requirements(integration_client: AsyncClient) -> None:
"""Test that endpoints meet performance requirements"""
validator = PerformanceValidator()
# Test info endpoint performance
for i in range(50):
start = validator.start_timing("info_endpoint")
response = await integration_client.get("/")
validator.end_timing("info_endpoint", start)
assert response.status_code == 200
# Validate 95th percentile is under 500ms
result = validator.validate_response_time(
"info_endpoint", max_duration=0.5, percentile=0.95
)
assert result["valid"], (
f"Performance requirement failed: "
f"95th percentile was {result['percentile_time']:.3f}s "
f"(required < {result['max_allowed']}s)"
)
@pytest.mark.integration

View File

@@ -9,6 +9,7 @@ import gc
import statistics
import time
from typing import Any, Dict, List
from unittest.mock import patch
import psutil
import pytest
@@ -105,42 +106,46 @@ class TestPerformanceBaseline:
("GET", "/v1/wallet/info", authenticated_client, None),
]
# Warm up
for _ in range(10):
await integration_client.get("/")
# Enable provider discovery for this test
with patch(
"routstr.core.settings.settings.providers_refresh_interval_seconds", 300
):
# Warm up
for _ in range(10):
await integration_client.get("/")
# Test each endpoint
for method, path, client, data in endpoints:
response_times = []
# Test each endpoint
for method, path, client, data in endpoints:
response_times = []
for i in range(100):
start = time.time()
for i in range(100):
start = time.time()
if method == "GET":
response = await client.get(path)
else:
response = await client.post(path, json=data)
if method == "GET":
response = await client.get(path)
else:
response = await client.post(path, json=data)
duration = time.time() - start
response_times.append(duration * 1000) # Convert to ms
duration = time.time() - start
response_times.append(duration * 1000) # Convert to ms
assert response.status_code in [200, 201]
assert response.status_code in [200, 201]
if i % 10 == 0:
metrics.record_system_metrics()
if i % 10 == 0:
metrics.record_system_metrics()
# Verify 95th percentile < 500ms
p95 = sorted(response_times)[int(len(response_times) * 0.95)]
assert p95 < 500, (
f"{method} {path} p95 response time {p95}ms exceeds 500ms limit"
)
# Verify 95th percentile < 500ms
p95 = sorted(response_times)[int(len(response_times) * 0.95)]
assert p95 < 500, (
f"{method} {path} p95 response time {p95}ms exceeds 500ms limit"
)
print(f"\n{method} {path}:")
print(f" Mean: {statistics.mean(response_times):.2f}ms")
print(f" P95: {p95:.2f}ms")
print(
f" P99: {sorted(response_times)[int(len(response_times) * 0.99)]:.2f}ms"
)
print(f"\n{method} {path}:")
print(f" Mean: {statistics.mean(response_times):.2f}ms")
print(f" P95: {p95:.2f}ms")
print(
f" P99: {sorted(response_times)[int(len(response_times) * 0.99)]:.2f}ms"
)
@pytest.mark.integration

View File

@@ -3,7 +3,7 @@ Integration tests for provider management functionality.
Tests GET /v1/providers/ endpoint for listing and managing providers.
"""
from typing import Any
from typing import Any, Generator
from unittest.mock import patch
import pytest
@@ -11,7 +11,7 @@ from httpx import AsyncClient
from routstr.discovery import _PROVIDERS_CACHE
from .utils import PerformanceValidator, ResponseValidator
from .utils import ResponseValidator
@pytest.fixture(autouse=True)
@@ -19,6 +19,15 @@ def _clear_providers_cache() -> None:
_PROVIDERS_CACHE.clear()
@pytest.fixture(autouse=True)
def _enable_provider_discovery() -> Generator[None, Any, Any]:
"""Enable provider discovery for all tests in this module"""
with patch(
"routstr.core.settings.settings.providers_refresh_interval_seconds", 300
):
yield
@pytest.mark.integration
@pytest.mark.asyncio
async def test_providers_endpoint_default_response(
@@ -518,46 +527,6 @@ async def test_providers_endpoint_response_format(
assert isinstance(data_json["providers"], list)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_providers_endpoint_performance(integration_client: AsyncClient) -> None:
"""Test providers endpoint meets performance requirements"""
# Mock quick responses to avoid network delays
mock_events: list[dict[str, Any]] = [
{
"id": f"event{i}",
"content": f"Provider: http://provider{i}.onion",
"created_at": 1234567890 + i,
}
for i in range(5)
]
validator = PerformanceValidator()
with patch(
"routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events
):
with patch("routstr.discovery.fetch_provider_health") as mock_fetch:
mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}}
# Test multiple requests
for i in range(10):
start = validator.start_timing("providers_endpoint")
response = await integration_client.get("/v1/providers/")
validator.end_timing("providers_endpoint", start)
assert response.status_code == 200
# Validate performance (should be fast with mocked dependencies)
perf_result = validator.validate_response_time(
"providers_endpoint",
max_duration=2.0, # Allow more time since it involves multiple operations
percentile=0.95,
)
assert perf_result["valid"], f"Performance requirement failed: {perf_result}"
@pytest.mark.integration
@pytest.mark.asyncio
async def test_providers_endpoint_concurrent_requests(

View File

@@ -18,7 +18,6 @@ from routstr.core.db import ApiKey
from .utils import (
ConcurrencyTester,
PerformanceValidator,
)
@@ -551,39 +550,7 @@ async def test_proxy_get_concurrent_requests(
assert response.status_code == 200
@pytest.mark.integration
@pytest.mark.asyncio
async def test_proxy_get_performance_requirements(
integration_client: AsyncClient, authenticated_client: AsyncClient
) -> None:
"""Test that GET proxy requests meet performance requirements"""
validator = PerformanceValidator()
with patch("httpx.AsyncClient.request") as mock_request:
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.json = MagicMock(return_value={"performance": "test"})
mock_response.text = '{"performance": "test"}'
mock_response.iter_bytes = AsyncMock(return_value=[b'{"performance": "test"}'])
mock_request.return_value = mock_response
# Test multiple requests for performance measurement
for i in range(20):
start = validator.start_timing("proxy_get")
response = await authenticated_client.get(f"/v1/perf-test-{i}")
validator.end_timing("proxy_get", start)
assert response.status_code == 200
# Validate performance requirements
perf_result = validator.validate_response_time(
"proxy_get",
max_duration=1.0, # Should complete within 1 second
percentile=0.95,
)
assert perf_result["valid"], f"Performance requirement failed: {perf_result}"
@pytest.mark.integration

View File

@@ -15,7 +15,6 @@ from httpx import ASGITransport, AsyncClient
from .utils import (
ConcurrencyTester,
PerformanceValidator,
)
@@ -290,55 +289,7 @@ async def test_proxy_post_unauthorized_access(integration_client: AsyncClient) -
assert response.status_code in [400, 401]
@pytest.mark.integration
@pytest.mark.asyncio
async def test_proxy_post_performance(
integration_client: AsyncClient, authenticated_client: AsyncClient
) -> None:
"""Test POST endpoint performance requirements"""
test_payload = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Performance test"}],
}
validator = PerformanceValidator()
with patch("httpx.AsyncClient.send") as mock_send:
# Mock fast responses
async def mock_iter_bytes(*args: Any, **kwargs: Any) -> Any:
yield b'{"choices": [{"message": {"content": "Fast"}}], "usage": {"total_tokens": 5}}'
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
response_data = {
"choices": [{"message": {"content": "Fast"}}],
"usage": {"total_tokens": 5},
}
mock_response.text = json.dumps(response_data)
mock_response.json = AsyncMock(return_value=response_data)
mock_response.iter_bytes = mock_iter_bytes
mock_response.aiter_bytes = mock_iter_bytes
mock_send.return_value = mock_response
# Run multiple requests for performance measurement
for i in range(20):
start = validator.start_timing("proxy_post")
response = await authenticated_client.post(
"/v1/chat/completions", json=test_payload
)
validator.end_timing("proxy_post", start)
assert response.status_code == 200
# Validate performance
perf_result = validator.validate_response_time(
"proxy_post",
max_duration=1.5, # Allow slightly more time for POST
percentile=0.95,
)
assert perf_result["valid"], f"Performance requirement failed: {perf_result}"
@pytest.mark.integration

View File

@@ -3,7 +3,7 @@ Integration tests for wallet information retrieval endpoints.
Tests GET /v1/wallet/ and GET /v1/wallet/info endpoints with various scenarios.
"""
import time
from datetime import datetime, timedelta
from typing import Any
@@ -367,30 +367,4 @@ async def test_wallet_info_with_special_characters_in_headers(
# Note: Current implementation doesn't return refund_address in response
@pytest.mark.integration
@pytest.mark.asyncio
@pytest.mark.slow
async def test_wallet_endpoints_performance(authenticated_client: AsyncClient) -> None:
"""Test wallet endpoints meet performance requirements"""
# Warm up
await authenticated_client.get("/v1/wallet/")
# Measure response times
response_times = []
for _ in range(50):
start_time = time.time()
response = await authenticated_client.get("/v1/wallet/")
end_time = time.time()
assert response.status_code == 200
response_times.append(end_time - start_time)
# Calculate statistics
avg_time = sum(response_times) / len(response_times)
max_time = max(response_times)
# Performance assertions
assert avg_time < 0.1 # Average should be under 100ms
assert max_time < 0.5 # No request should take more than 500ms

View File

@@ -537,42 +537,4 @@ async def test_refund_with_expired_key(
assert response.json()["recipient"] == "expired@ln.address"
@pytest.mark.integration
@pytest.mark.asyncio
@pytest.mark.slow
async def test_refund_performance(
integration_client: AsyncClient, testmint_wallet: Any
) -> None:
"""Test refund endpoint performance"""
import time
# Create multiple API keys
api_keys = []
for i in range(10):
token = await testmint_wallet.mint_tokens(100 + i)
# Use cashu token as Bearer auth to create API key
integration_client.headers["Authorization"] = f"Bearer {token}"
response = await integration_client.get("/v1/wallet/info")
assert response.status_code == 200
api_keys.append(response.json()["api_key"])
# Measure refund times
refund_times = []
for api_key in api_keys:
integration_client.headers["Authorization"] = f"Bearer {api_key}"
start_time = time.time()
response = await integration_client.post("/v1/wallet/refund")
end_time = time.time()
assert response.status_code == 200
refund_times.append(end_time - start_time)
# Performance assertions
avg_time = sum(refund_times) / len(refund_times)
max_time = max(refund_times)
assert avg_time < 0.5 # Average under 500ms
assert max_time < 1.0 # No refund takes more than 1 second

View File

@@ -10,7 +10,6 @@ os.environ["UPSTREAM_API_KEY"] = "test"
from routstr.algorithm import ( # noqa: E402
calculate_model_cost_score,
get_provider_penalty,
should_prefer_model,
)
from routstr.payment.models import Architecture, Model, Pricing # noqa: E402
@@ -100,100 +99,3 @@ 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_should_prefer_model_cheaper_wins() -> None:
"""Test that cheaper model is preferred."""
cheap_model = create_test_model("cheap", prompt_price=0.001, completion_price=0.002)
expensive_model = create_test_model(
"expensive", prompt_price=0.03, completion_price=0.06
)
provider1 = create_test_provider("provider1")
provider2 = create_test_provider("provider2")
# Cheaper model should win
assert should_prefer_model(
cheap_model, provider1, expensive_model, provider2, "test-alias"
)
# More expensive model should not win
assert not should_prefer_model(
expensive_model, provider2, cheap_model, provider1, "test-alias"
)
def test_should_prefer_model_exact_match_wins() -> None:
"""Test that exact alias match beats cheaper price."""
# Make model IDs match the alias differently
exact_match = create_test_model(
"test-model", prompt_price=0.03, completion_price=0.06
)
no_match = create_test_model(
"other-model", prompt_price=0.001, completion_price=0.002
)
provider1 = create_test_provider("provider1")
provider2 = create_test_provider("provider2")
# Exact match should win even though it's more expensive
assert should_prefer_model(
exact_match, provider1, no_match, provider2, "test-model"
)
def test_should_prefer_model_openrouter_slight_penalty() -> None:
"""Test that OpenRouter has slight penalty compared to other providers."""
model1 = create_test_model("model1", prompt_price=0.001, completion_price=0.002)
model2 = create_test_model("model2", prompt_price=0.001, completion_price=0.002)
regular_provider = create_test_provider("regular", "http://provider.com")
openrouter_provider = create_test_provider(
"openrouter", "https://openrouter.ai/api/v1"
)
# Regular provider should be preferred over OpenRouter at same cost
assert should_prefer_model(
model1, regular_provider, model2, openrouter_provider, "test-alias"
)
# OpenRouter should not replace regular provider at same cost
assert not should_prefer_model(
model2, openrouter_provider, model1, regular_provider, "test-alias"
)
def test_should_prefer_model_openrouter_can_win_if_cheaper() -> None:
"""Test that OpenRouter can still win if significantly cheaper."""
cheap_model = create_test_model(
"cheap", prompt_price=0.0001, completion_price=0.0002
)
expensive_model = create_test_model(
"expensive", prompt_price=0.03, completion_price=0.06
)
regular_provider = create_test_provider("regular", "http://provider.com")
openrouter_provider = create_test_provider(
"openrouter", "https://openrouter.ai/api/v1"
)
# OpenRouter should win if it's much cheaper (even with penalty)
assert should_prefer_model(
cheap_model,
openrouter_provider,
expensive_model,
regular_provider,
"test-alias",
)
def test_should_prefer_model_same_cost_first_wins() -> None:
"""Test that when costs are identical, current model is kept."""
model1 = create_test_model("model1", prompt_price=0.001, completion_price=0.002)
model2 = create_test_model("model2", prompt_price=0.001, completion_price=0.002)
provider1 = create_test_provider("provider1")
provider2 = create_test_provider("provider2")
# When costs are equal, should not replace
assert not should_prefer_model(model2, provider2, model1, provider1, "test-alias")

View File

@@ -4,17 +4,22 @@ FROM base AS deps
RUN apk add --no-cache libc6-compat
WORKDIR /app
COPY package.json yarn.lock* package-lock.json* pnpm-lock.yaml* .npmrc* ./
RUN npm i
RUN corepack enable pnpm && corepack prepare pnpm@latest --activate
COPY package.json pnpm-lock.yaml* ./
RUN pnpm install --frozen-lockfile
FROM base AS builder
WORKDIR /app
RUN corepack enable pnpm && corepack prepare pnpm@latest --activate
COPY --from=deps /app/node_modules ./node_modules
COPY . .
ENV NEXT_TELEMETRY_DISABLED=1
RUN npm run build
RUN pnpm run build
FROM base AS runner
WORKDIR /app

View File

@@ -6,12 +6,14 @@ RUN apk add --no-cache libc6-compat
WORKDIR /app
# Copy package files
COPY package.json package-lock.json* pnpm-lock.yaml* ./
RUN npm ci
COPY package.json pnpm-lock.yaml* ./
RUN corepack enable pnpm && corepack prepare pnpm@latest --activate
RUN pnpm install --frozen-lockfile
# Build the UI
FROM base AS builder
WORKDIR /app
RUN corepack enable pnpm && corepack prepare pnpm@latest --activate
COPY --from=deps /app/node_modules ./node_modules
COPY . .
@@ -27,7 +29,7 @@ ENV NODE_ENV=production
ENV NEXT_TELEMETRY_DISABLED=1
# Build the application
RUN npm run build && \
RUN pnpm run build && \
echo "UI build completed at $(date)"
# Use the builder stage as the final stage

View File

@@ -208,7 +208,12 @@ function ProviderBalance({
return <Skeleton className='h-9 w-24' />;
}
if (error || !balanceData?.ok || !balanceData.balance_data) {
if (
error ||
!balanceData?.ok ||
balanceData.balance_data === undefined ||
balanceData.balance_data === null
) {
return null;
}

View File

@@ -60,7 +60,11 @@ export function TemporaryBalances({
let totalRequests = 0;
balances.forEach((balance) => {
totalBalance += balance.balance || 0;
// Only count parents for total balance to avoid double counting
// since child keys use parent balance
if (!balance.parent_key_hash) {
totalBalance += balance.balance || 0;
}
totalSpent += balance.total_spent || 0;
totalRequests += balance.total_requests || 0;
});
@@ -72,6 +76,34 @@ export function TemporaryBalances({
? calculateTotals(data)
: { totalBalance: 0, totalSpent: 0, totalRequests: 0 };
// Group parents and children
const hierarchicalData = (() => {
if (!data) return [];
const parents = filteredData.filter((item) => !item.parent_key_hash);
const result: (TemporaryBalance & { isChild?: boolean })[] = [];
parents.forEach((parent) => {
result.push(parent);
const children = data.filter(
(item) => item.parent_key_hash === parent.hashed_key
);
children.forEach((child) => {
result.push({ ...child, isChild: true });
});
});
// Add children whose parents didn't match the search or aren't in the list
const orphans = filteredData.filter(
(item) =>
item.parent_key_hash &&
!result.some((r) => r.hashed_key === item.hashed_key)
);
result.push(...orphans.map((o) => ({ ...o, isChild: true })));
return result;
})();
return (
<>
<Card className='h-full w-full shadow-sm'>
@@ -182,22 +214,37 @@ export function TemporaryBalances({
<div className='text-right'>Expiry Time</div>
</div>
{filteredData.length > 0 ? (
filteredData.map((balance, index) => (
{hierarchicalData.length > 0 ? (
hierarchicalData.map((balance, index) => (
<div
key={index}
className={cn(
'hover:bg-muted/50 border-t p-3 text-sm transition-colors',
balance.balance === 0 && 'opacity-60'
balance.balance === 0 &&
!balance.isChild &&
'opacity-60',
balance.isChild &&
'ml-4 border-l-2 border-l-blue-200 bg-blue-50/30'
)}
>
{/* Desktop Layout */}
<div className='hidden grid-cols-6 gap-2 md:grid'>
<div className='max-w-32 truncate font-mono text-xs break-all'>
<div className='flex max-w-48 items-center gap-2 truncate font-mono text-xs break-all'>
{balance.isChild && (
<span className='rounded bg-blue-100 px-1 py-0.5 text-[10px] font-bold text-blue-700 uppercase'>
Child
</span>
)}
{balance.hashed_key}
</div>
<div className='text-right font-mono'>
{formatBalance(balance.balance)}
{balance.isChild ? (
<span className='text-muted-foreground italic'>
(Parent)
</span>
) : (
formatBalance(balance.balance)
)}
</div>
<div className='text-right font-mono'>
{formatBalance(balance.total_spent)}
@@ -226,13 +273,20 @@ export function TemporaryBalances({
{/* Mobile Layout */}
<div className='space-y-3 md:hidden'>
<div className='space-y-1'>
<span className='text-muted-foreground text-xs font-medium'>
Key
</span>
<div className='font-mono text-xs break-all'>
{balance.hashed_key}
<div className='flex items-center justify-between'>
<div className='space-y-1'>
<span className='text-muted-foreground text-xs font-medium'>
{balance.isChild ? 'Child Key' : 'Key'}
</span>
<div className='font-mono text-xs break-all'>
{balance.hashed_key}
</div>
</div>
{balance.isChild && (
<span className='rounded bg-blue-100 px-1.5 py-0.5 text-[10px] font-bold text-blue-700 uppercase'>
Child
</span>
)}
</div>
<div className='grid grid-cols-2 gap-3'>
@@ -241,7 +295,13 @@ export function TemporaryBalances({
Balance
</div>
<div className='truncate font-mono text-sm'>
{formatBalance(balance.balance)}
{balance.isChild ? (
<span className='text-muted-foreground text-xs italic'>
(Uses Parent)
</span>
) : (
formatBalance(balance.balance)
)}
</div>
</div>
<div className='space-y-1'>

View File

@@ -926,6 +926,7 @@ export const TemporaryBalanceSchema = z.object({
total_requests: z.number(),
refund_address: z.string().nullable(),
key_expiry_time: z.number().nullable(),
parent_key_hash: z.string().nullable().optional(),
});
export type TemporaryBalance = z.infer<typeof TemporaryBalanceSchema>;

9257
ui/package-lock.json generated

File diff suppressed because it is too large Load Diff

View File

@@ -15,6 +15,7 @@
"@dnd-kit/core": "^6.3.1",
"@dnd-kit/modifiers": "^9.0.0",
"@dnd-kit/sortable": "^10.0.0",
"@dnd-kit/utilities": "^3.2.2",
"@hookform/resolvers": "^5.0.1",
"@radix-ui/react-alert-dialog": "^1.1.10",
"@radix-ui/react-avatar": "^1.1.6",
@@ -40,17 +41,17 @@
"@radix-ui/react-toggle": "^1.1.6",
"@radix-ui/react-toggle-group": "^1.1.6",
"@radix-ui/react-tooltip": "^1.2.3",
"@tanstack/react-query": "^5.74.4",
"@tanstack/react-query": "^5.90.16",
"@tanstack/react-table": "^8.21.3",
"axios": "^1.13.2",
"class-variance-authority": "^0.7.1",
"clsx": "^2.1.1",
"cmdk": "^1.1.1",
"date-fns": "^3.6.0",
"date-fns": "^4.1.0",
"embla-carousel-react": "^8.6.0",
"input-otp": "^1.4.2",
"lucide-react": "^0.501.0",
"next": "15.3.1",
"lucide-react": "^0.562.0",
"next": "15.5.9",
"next-themes": "^0.4.6",
"qrcode": "^1.5.4",
"qrcode.react": "^4.2.0",
@@ -67,21 +68,21 @@
"zustand": "^5.0.3"
},
"devDependencies": {
"@eslint/eslintrc": "^3",
"@tailwindcss/postcss": "^4",
"@tanstack/react-query-devtools": "^5.74.4",
"@eslint/eslintrc": "^3.3.3",
"@tailwindcss/postcss": "^4.1.18",
"@tanstack/react-query-devtools": "^5.91.2",
"@types/node": "^20",
"@types/qrcode": "^1.5.6",
"@types/react": "^19",
"@types/react-dom": "^19",
"eslint": "^9.25.0",
"eslint-config-next": "15.3.1",
"eslint-config-next": "15.5.9",
"eslint-config-prettier": "^10.1.2",
"eslint-plugin-prettier": "^5.2.6",
"eslint-plugin-react": "^7.37.5",
"prettier": "^3.5.3",
"prettier-plugin-tailwindcss": "^0.6.11",
"tailwindcss": "^4",
"tailwindcss": "^4.1.18",
"typescript": "^5"
}
}

1088
ui/pnpm-lock.yaml generated

File diff suppressed because it is too large Load Diff

View File

@@ -22,6 +22,12 @@
"@/*": ["./*"]
}
},
"include": ["next-env.d.ts", "**/*.ts", "**/*.tsx", ".next/types/**/*.ts"],
"include": [
"next-env.d.ts",
"**/*.ts",
"**/*.tsx",
".next/types/**/*.ts",
".next/dev/types/**/*.ts"
],
"exclude": ["node_modules"]
}

2
uv.lock generated
View File

@@ -1878,7 +1878,7 @@ wheels = [
[[package]]
name = "routstr"
version = "0.2.1"
version = "0.2.2"
source = { editable = "." }
dependencies = [
{ name = "aiosqlite" },