Compare commits

..
Author SHA1 Message Date
9qeklajc 87efed0021 Merge branch 'v0.3.0' into claude-code
# Conflicts:
#	routstr/upstream/base.py
2026-01-12 20:53:27 +01:00
Shroominic 3dfbd3815c fix typing 2026-01-12 17:47:14 +08:00
Shroominic 6199f6467b mock-upstream-with-testnut-mint 2026-01-12 17:47:14 +08:00
Shroominic 9960e5596e bump v0.2.2 2026-01-12 17:45:40 +08:00
shroominicandShroominic a516a10737 Merge pull request #292 from Routstr/fix-not-enough-inputs-to-melt
Fix not enough inputs to melt
2026-01-12 17:45:40 +08:00
shroominicandShroominic be9b1da832 Merge pull request #295 from Routstr/update-ui-deps
update vuln deps
2026-01-12 17:45:40 +08:00
Shroominic e4dd0aceae fix not enough inputs to melt bug 2026-01-12 17:45:40 +08:00
9qeklajcandShroominic 30170c2ec6 Merge pull request #291 from Routstr/fix-reserved-balance
Fix reserved balance
2026-01-12 17:45:40 +08:00
9qeklajcandShroominic 825bd38d8e fix test 2026-01-12 17:45:40 +08:00
9qeklajcandShroominic d3b3152520 Merge pull request #289 from Routstr/282-better-filltering
#282 more filter options
2026-01-12 17:45:40 +08:00
Shroominic 653b51452a bump delay to make sure its not taking more time to reset 2026-01-12 17:45:40 +08:00
9qeklajcandShroominic aa00664cd6 fmt 2026-01-12 17:45:40 +08:00
9qeklajcandShroominic ea677cf66b Merge pull request #284 from Routstr/274-do-not-charge-when-empty-content
#274 do not charge user for empty response by upstream
2026-01-12 17:45:40 +08:00
9qeklajcandShroominic 11868f9180 #282 more filter options 2026-01-12 17:45:40 +08:00
Shroominic daae4cd2cb lint pls 2026-01-12 17:45:40 +08:00
9qeklajcandShroominic 2cb8d2d744 update build script 2026-01-12 17:45:40 +08:00
9qeklajcandShroominic 22f7d198a6 #274 do not change user for empty response by upstream 2026-01-12 17:45:40 +08:00
Shroominic 73af5e23c7 fix typing 2026-01-12 17:45:40 +08:00
9qeklajcandShroominic eb5af32fac update vuln deps 2026-01-12 17:45:40 +08:00
Shroominic 0f9df3ca77 rm not needed generic 2026-01-12 17:45:40 +08:00
Shroominic 3b1b3da847 fix linting 2026-01-12 17:45:40 +08:00
Shroominic a34583d2ca fix other potential reserved balance problems 2026-01-12 17:45:40 +08:00
Shroominic 4eda9eaf1b finalize_without_usage when client disconnects 2026-01-12 17:45:40 +08:00
9qeklajc b9879cbea7 claude code integration 2026-01-10 23:32:45 +01:00
33 changed files with 1371 additions and 1155 deletions
-45
View File
@@ -1,45 +0,0 @@
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.")
-42
View File
@@ -1,42 +0,0 @@
"""
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")
+1 -1
View File
@@ -1,6 +1,6 @@
[project] [project]
name = "routstr" name = "routstr"
version = "0.3.0" version = "0.2.2"
description = "Payment proxy for your LLM endpoint using cashu and nostr." description = "Payment proxy for your LLM endpoint using cashu and nostr."
readme = "README.md" readme = "README.md"
requires-python = ">=3.11" requires-python = ">=3.11"
+96 -52
View File
@@ -84,26 +84,93 @@ def get_provider_penalty(provider: "BaseUpstreamProvider") -> float:
return penalty 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( def create_model_mappings(
upstreams: list["BaseUpstreamProvider"], upstreams: list["BaseUpstreamProvider"],
overrides_by_id: dict[str, tuple], overrides_by_id: dict[str, tuple],
disabled_model_ids: set[str], disabled_model_ids: set[str],
) -> tuple[ ) -> tuple[dict[str, "Model"], dict[str, "BaseUpstreamProvider"], dict[str, "Model"]]:
dict[str, "Model"], dict[str, list["BaseUpstreamProvider"]], dict[str, "Model"]
]:
"""Create optimal model mappings based on cost and provider preferences. """Create optimal model mappings based on cost and provider preferences.
This is the main entry point for the algorithm. It processes all upstream providers This is the main entry point for the algorithm. It processes all upstream providers
and creates three mappings based on cost optimization: and creates three mappings based on cost optimization:
1. model_instances: alias -> Model (all model aliases mapped to their Model objects) 1. model_instances: alias -> Model (all model aliases mapped to their Model objects)
2. provider_map: alias -> List[UpstreamProvider] (sorted list of providers for each alias) 2. provider_map: alias -> UpstreamProvider (which provider to use for each alias)
3. unique_models: base_id -> Model (unique models without provider prefixes) 3. unique_models: base_id -> Model (unique models without provider prefixes)
The algorithm: The algorithm:
- Processes non-OpenRouter providers first (they're typically cheaper) - Processes non-OpenRouter providers first (they're typically cheaper)
- Then processes OpenRouter models (they can still win if cheaper) - Then processes OpenRouter models (they can still win if cheaper)
- For each model alias, collects all candidates and sorts them by priority and cost. - For each model alias, uses should_prefer_model() to select the best provider
Args: Args:
upstreams: List of all upstream provider instances upstreams: List of all upstream provider instances
@@ -116,7 +183,8 @@ def create_model_mappings(
from .payment.models import _row_to_model from .payment.models import _row_to_model
from .upstream.helpers import resolve_model_alias from .upstream.helpers import resolve_model_alias
candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {} model_instances: dict[str, "Model"] = {}
provider_map: dict[str, "BaseUpstreamProvider"] = {}
unique_models: dict[str, "Model"] = {} unique_models: dict[str, "Model"] = {}
# Separate OpenRouter from other providers # Separate OpenRouter from other providers
@@ -134,14 +202,24 @@ def create_model_mappings(
"""Get base model ID by removing provider prefix.""" """Get base model ID by removing provider prefix."""
return model_id.split("/", 1)[1] if "/" in model_id else model_id return model_id.split("/", 1)[1] if "/" in model_id else model_id
def _add_candidate( def _maybe_set_alias(
alias: str, model: "Model", provider: "BaseUpstreamProvider" alias: str, model: "Model", provider: "BaseUpstreamProvider"
) -> None: ) -> None:
"""Add candidate model/provider for an alias.""" """Set alias to model/provider if not set or if new model is preferred."""
alias_lower = alias.lower() alias_lower = alias.lower()
if alias_lower not in candidates: existing_model = model_instances.get(alias_lower)
candidates[alias_lower] = [] if not existing_model:
candidates[alias_lower].append((model, provider)) # 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
def process_provider_models( def process_provider_models(
upstream: "BaseUpstreamProvider", is_openrouter: bool = False upstream: "BaseUpstreamProvider", is_openrouter: bool = False
@@ -188,55 +266,21 @@ def create_model_mappings(
# Try to set each alias # Try to set each alias
for alias in aliases: for alias in aliases:
_add_candidate(alias, model_to_use, upstream) _maybe_set_alias(alias, model_to_use, upstream)
# Process non-OpenRouter providers first # Process non-OpenRouter providers first (they're typically cheaper)
for upstream in other_upstreams: for upstream in other_upstreams:
process_provider_models(upstream, is_openrouter=False) process_provider_models(upstream, is_openrouter=False)
# Process OpenRouter last # Process OpenRouter last - models only win if they're cheaper or better matched
if openrouter: if openrouter:
process_provider_models(openrouter, is_openrouter=True) process_provider_models(openrouter, is_openrouter=True)
# Sort candidates and build final maps # Log provider distribution
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] = {} provider_counts: dict[str, int] = {}
for providers in provider_map.values(): for provider in provider_map.values():
if providers: provider_name = getattr(provider, "upstream_name", "unknown")
provider = providers[0] provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1
provider_name = getattr(provider, "upstream_name", "unknown")
provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1
logger.debug( logger.debug(
f"Updated model mappings with ({len(unique_models)} unique models and {len(model_instances)} aliases)", f"Updated model mappings with ({len(unique_models)} unique models and {len(model_instances)} aliases)",
+40 -166
View File
@@ -286,55 +286,30 @@ 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( async def pay_for_request(
key: ApiKey, cost_per_request: int, session: AsyncSession key: ApiKey, cost_per_request: int, session: AsyncSession
) -> int: ) -> int:
"""Process payment for a request.""" """Process payment for a request."""
billing_key = await get_billing_key(key, session)
logger.info( logger.info(
"Processing payment for request", "Processing payment for request",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...", "current_balance": key.balance,
"current_balance": billing_key.balance,
"required_cost": cost_per_request, "required_cost": cost_per_request,
"sufficient_balance": billing_key.balance >= cost_per_request, "sufficient_balance": key.balance >= cost_per_request,
}, },
) )
if billing_key.total_balance < cost_per_request: if key.total_balance < cost_per_request:
logger.warning( logger.warning(
"Insufficient balance for request", "Insufficient balance for request",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...", "balance": key.balance,
"balance": billing_key.balance, "reserved_balance": key.reserved_balance,
"reserved_balance": billing_key.reserved_balance,
"required": cost_per_request, "required": cost_per_request,
"shortfall": cost_per_request - billing_key.total_balance, "shortfall": cost_per_request - key.total_balance,
}, },
) )
@@ -342,7 +317,7 @@ async def pay_for_request(
status_code=402, status_code=402,
detail={ detail={
"error": { "error": {
"message": f"Insufficient balance: {cost_per_request} mSats required. {billing_key.total_balance} available. (reserved: {billing_key.reserved_balance})", "message": f"Insufficient balance: {cost_per_request} mSats required. {key.total_balance} available. (reserved: {key.reserved_balance})",
"type": "insufficient_quota", "type": "insufficient_quota",
"code": "insufficient_balance", "code": "insufficient_balance",
} }
@@ -353,16 +328,15 @@ async def pay_for_request(
"Charging base cost for request", "Charging base cost for request",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"cost": cost_per_request, "cost": cost_per_request,
"balance_before": billing_key.balance, "balance_before": key.balance,
}, },
) )
# Charge the base cost for the request atomically to avoid race conditions # Charge the base cost for the request atomically to avoid race conditions
stmt = ( stmt = (
update(ApiKey) update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key) .where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.balance) - col(ApiKey.reserved_balance) >= cost_per_request) .where(col(ApiKey.balance) - col(ApiKey.reserved_balance) >= cost_per_request)
.values( .values(
reserved_balance=col(ApiKey.reserved_balance) + cost_per_request, reserved_balance=col(ApiKey.reserved_balance) + cost_per_request,
@@ -370,16 +344,6 @@ async def pay_for_request(
) )
) )
result = await session.exec(stmt) # type: ignore[call-overload] 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() await session.commit()
if result.rowcount == 0: if result.rowcount == 0:
@@ -387,9 +351,8 @@ async def pay_for_request(
"Concurrent request depleted balance", "Concurrent request depleted balance",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"required_cost": cost_per_request, "required_cost": cost_per_request,
"current_balance": billing_key.balance, "current_balance": key.balance,
}, },
) )
@@ -398,26 +361,23 @@ async def pay_for_request(
status_code=402, status_code=402,
detail={ detail={
"error": { "error": {
"message": f"Insufficient balance: {cost_per_request} mSats required. {billing_key.balance} available.", "message": f"Insufficient balance: {cost_per_request} mSats required. {key.balance} available.",
"type": "insufficient_quota", "type": "insufficient_quota",
"code": "insufficient_balance", "code": "insufficient_balance",
} }
}, },
) )
await session.refresh(billing_key) await session.refresh(key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
logger.info( logger.info(
"Payment processed successfully", "Payment processed successfully",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"charged_amount": cost_per_request, "charged_amount": cost_per_request,
"new_balance": billing_key.balance, "new_balance": key.balance,
"total_spent": billing_key.total_spent, "total_spent": key.total_spent,
"total_requests": billing_key.total_requests, "total_requests": key.total_requests,
}, },
) )
@@ -427,11 +387,9 @@ async def pay_for_request(
async def revert_pay_for_request( async def revert_pay_for_request(
key: ApiKey, session: AsyncSession, cost_per_request: int key: ApiKey, session: AsyncSession, cost_per_request: int
) -> None: ) -> None:
billing_key = await get_billing_key(key, session)
stmt = ( stmt = (
update(ApiKey) update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key) .where(col(ApiKey.hashed_key) == key.hashed_key)
.values( .values(
reserved_balance=col(ApiKey.reserved_balance) - cost_per_request, reserved_balance=col(ApiKey.reserved_balance) - cost_per_request,
total_requests=col(ApiKey.total_requests) - 1, total_requests=col(ApiKey.total_requests) - 1,
@@ -439,40 +397,27 @@ async def revert_pay_for_request(
) )
result = await session.exec(stmt) # type: ignore[call-overload] 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() await session.commit()
if result.rowcount == 0: if result.rowcount == 0:
logger.error( logger.error(
"Failed to revert payment - insufficient reserved balance", "Failed to revert payment - insufficient reserved balance",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"cost_to_revert": cost_per_request, "cost_to_revert": cost_per_request,
"current_reserved_balance": billing_key.reserved_balance, "current_reserved_balance": key.reserved_balance,
}, },
) )
raise HTTPException( raise HTTPException(
status_code=402, status_code=402,
detail={ detail={
"error": { "error": {
"message": f"failed to revert request payment: {cost_per_request} mSats required. {billing_key.balance} available.", "message": f"failed to revert request payment: {cost_per_request} mSats required. {key.balance} available.",
"type": "payment_error", "type": "payment_error",
"code": "payment_error", "code": "payment_error",
} }
}, },
) )
await session.refresh(billing_key) await session.refresh(key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
async def adjust_payment_for_tokens( async def adjust_payment_for_tokens(
@@ -483,17 +428,15 @@ async def adjust_payment_for_tokens(
This is called after the initial payment and the upstream request is complete. This is called after the initial payment and the upstream request is complete.
Returns cost data to be included in the response. Returns cost data to be included in the response.
""" """
billing_key = await get_billing_key(key, session)
model = response_data.get("model", "unknown") model = response_data.get("model", "unknown")
logger.debug( logger.debug(
"Starting payment adjustment for tokens", "Starting payment adjustment for tokens",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model, "model": model,
"deducted_max_cost": deducted_max_cost, "deducted_max_cost": deducted_max_cost,
"current_balance": billing_key.balance, "current_balance": key.balance,
"has_usage": "usage" in response_data, "has_usage": "usage" in response_data,
}, },
) )
@@ -503,10 +446,8 @@ async def adjust_payment_for_tokens(
try: try:
release_stmt = ( release_stmt = (
update(ApiKey) update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key) .where(col(ApiKey.hashed_key) == key.hashed_key)
.values( .values(reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost)
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost
)
) )
await session.exec(release_stmt) # type: ignore[call-overload] await session.exec(release_stmt) # type: ignore[call-overload]
await session.commit() await session.commit()
@@ -514,18 +455,13 @@ async def adjust_payment_for_tokens(
"Released reservation without charging (fallback)", "Released reservation without charging (fallback)",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"deducted_max_cost": deducted_max_cost, "deducted_max_cost": deducted_max_cost,
}, },
) )
except Exception as e: except Exception as e:
logger.error( logger.error(
"Failed to release reservation in fallback", "Failed to release reservation in fallback",
extra={ extra={"error": str(e), "key_hash": key.hashed_key[:8] + "..."},
"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): match await calculate_cost(response_data, deducted_max_cost, session):
@@ -534,7 +470,6 @@ async def adjust_payment_for_tokens(
"Using max cost data (no token adjustment)", "Using max cost data (no token adjustment)",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model, "model": model,
"max_cost": cost.total_msats, "max_cost": cost.total_msats,
}, },
@@ -542,7 +477,7 @@ async def adjust_payment_for_tokens(
# Finalize by releasing reservation and charging max cost # Finalize by releasing reservation and charging max cost
finalize_stmt = ( finalize_stmt = (
update(ApiKey) update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key) .where(col(ApiKey.hashed_key) == key.hashed_key)
.values( .values(
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost,
balance=col(ApiKey.balance) - cost.total_msats, balance=col(ApiKey.balance) - cost.total_msats,
@@ -550,41 +485,27 @@ async def adjust_payment_for_tokens(
) )
) )
result = await session.exec(finalize_stmt) # type: ignore[call-overload] 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() await session.commit()
if result.rowcount == 0: if result.rowcount == 0:
logger.error( logger.error(
"Failed to finalize max-cost payment - retrying reservation release", "Failed to finalize max-cost payment - retrying reservation release",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"deducted_max_cost": deducted_max_cost, "deducted_max_cost": deducted_max_cost,
"current_reserved_balance": billing_key.reserved_balance, "current_reserved_balance": key.reserved_balance,
"total_cost": cost.total_msats, "total_cost": cost.total_msats,
"model": model, "model": model,
}, },
) )
await release_reservation_only() await release_reservation_only()
else: else:
await session.refresh(billing_key) await session.refresh(key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
logger.info( logger.info(
"Max cost payment finalized", "Max cost payment finalized",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"charged_amount": cost.total_msats, "charged_amount": cost.total_msats,
"new_balance": billing_key.balance, "new_balance": key.balance,
"model": model, "model": model,
}, },
) )
@@ -600,7 +521,6 @@ async def adjust_payment_for_tokens(
"Calculated token-based cost", "Calculated token-based cost",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model, "model": model,
"token_cost": cost.total_msats, "token_cost": cost.total_msats,
"deducted_max_cost": deducted_max_cost, "deducted_max_cost": deducted_max_cost,
@@ -613,15 +533,11 @@ async def adjust_payment_for_tokens(
if cost_difference == 0: if cost_difference == 0:
logger.debug( logger.debug(
"Finalizing with exact reserved cost", "Finalizing with exact reserved cost",
extra={ extra={"key_hash": key.hashed_key[:8] + "...", "model": model},
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model,
},
) )
finalize_stmt = ( finalize_stmt = (
update(ApiKey) update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key) .where(col(ApiKey.hashed_key) == key.hashed_key)
.values( .values(
reserved_balance=col(ApiKey.reserved_balance) reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost, - deducted_max_cost,
@@ -630,20 +546,8 @@ async def adjust_payment_for_tokens(
) )
) )
await session.exec(finalize_stmt) # type: ignore[call-overload] 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.commit()
await session.refresh(billing_key) await session.refresh(key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
return cost.dict() return cost.dict()
# this should never happen why do we handle this??? # this should never happen why do we handle this???
@@ -653,17 +557,16 @@ async def adjust_payment_for_tokens(
"Additional charge required for token usage", "Additional charge required for token usage",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"additional_charge": cost_difference, "additional_charge": cost_difference,
"current_balance": billing_key.balance, "current_balance": key.balance,
"sufficient_balance": billing_key.balance >= cost_difference, "sufficient_balance": key.balance >= cost_difference,
"model": model, "model": model,
}, },
) )
finalize_stmt = ( finalize_stmt = (
update(ApiKey) update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key) .where(col(ApiKey.hashed_key) == key.hashed_key)
.values( .values(
reserved_balance=col(ApiKey.reserved_balance) reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost, - deducted_max_cost,
@@ -672,31 +575,18 @@ async def adjust_payment_for_tokens(
) )
) )
result = await session.exec(finalize_stmt) # type: ignore[call-overload] 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() await session.commit()
if result.rowcount: if result.rowcount:
cost.total_msats = total_cost_msats cost.total_msats = total_cost_msats
await session.refresh(billing_key) await session.refresh(key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
logger.info( logger.info(
"Finalized payment with additional charge", "Finalized payment with additional charge",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"charged_amount": total_cost_msats, "charged_amount": total_cost_msats,
"new_balance": billing_key.balance, "new_balance": key.balance,
"model": model, "model": model,
}, },
) )
@@ -705,7 +595,6 @@ async def adjust_payment_for_tokens(
"Failed to finalize additional charge - releasing reservation", "Failed to finalize additional charge - releasing reservation",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"attempted_charge": total_cost_msats, "attempted_charge": total_cost_msats,
"model": model, "model": model,
}, },
@@ -718,16 +607,15 @@ async def adjust_payment_for_tokens(
"Refunding excess payment", "Refunding excess payment",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"refund_amount": refund, "refund_amount": refund,
"current_balance": billing_key.balance, "current_balance": key.balance,
"model": model, "model": model,
}, },
) )
refund_stmt = ( refund_stmt = (
update(ApiKey) update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key) .where(col(ApiKey.hashed_key) == key.hashed_key)
.values( .values(
reserved_balance=col(ApiKey.reserved_balance) reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost, - deducted_max_cost,
@@ -736,16 +624,6 @@ async def adjust_payment_for_tokens(
) )
) )
result = await session.exec(refund_stmt) # type: ignore[call-overload] 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() await session.commit()
if result.rowcount == 0: if result.rowcount == 0:
@@ -753,9 +631,8 @@ async def adjust_payment_for_tokens(
"Failed to finalize payment - releasing reservation", "Failed to finalize payment - releasing reservation",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"deducted_max_cost": deducted_max_cost, "deducted_max_cost": deducted_max_cost,
"current_reserved_balance": billing_key.reserved_balance, "current_reserved_balance": key.reserved_balance,
"total_cost": total_cost_msats, "total_cost": total_cost_msats,
"model": model, "model": model,
}, },
@@ -763,17 +640,14 @@ async def adjust_payment_for_tokens(
await release_reservation_only() await release_reservation_only()
else: else:
cost.total_msats = total_cost_msats cost.total_msats = total_cost_msats
await session.refresh(billing_key) await session.refresh(key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
logger.info( logger.info(
"Refund processed successfully", "Refund processed successfully",
extra={ extra={
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"refunded_amount": refund, "refunded_amount": refund,
"new_balance": billing_key.balance, "new_balance": key.balance,
"final_cost": cost.total_msats, "final_cost": cost.total_msats,
"model": model, "model": model,
}, },
+13 -83
View File
@@ -32,28 +32,14 @@ async def get_key_from_header(
) )
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": 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 # TODO: remove this endpoint when frontend is updated
@router.get("/", include_in_schema=False) @router.get("/", include_in_schema=False)
async def account_info( async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict:
key: ApiKey = Depends(get_key_from_header), return {
session: AsyncSession = Depends(get_session), "api_key": "sk-" + key.hashed_key,
) -> dict: "balance": key.balance,
return await get_balance_info(key, session) "reserved": key.reserved_balance,
}
# TODO: Implement POST /v1/wallet/create endpoint # TODO: Implement POST /v1/wallet/create endpoint
@@ -80,11 +66,12 @@ async def create_balance(
@router.get("/info") @router.get("/info")
async def wallet_info( async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict:
key: ApiKey = Depends(get_key_from_header), return {
session: AsyncSession = Depends(get_session), "api_key": "sk-" + key.hashed_key,
) -> dict: "balance": key.balance,
return await get_balance_info(key, session) "reserved": key.reserved_balance,
}
class TopupRequest(BaseModel): class TopupRequest(BaseModel):
@@ -98,10 +85,6 @@ async def topup_wallet_endpoint(
key: ApiKey = Depends(get_key_from_header), key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session), session: AsyncSession = Depends(get_session),
) -> dict[str, int]: ) -> dict[str, int]:
from .auth import get_billing_key
billing_key = await get_billing_key(key, session)
if topup_request is not None: if topup_request is not None:
cashu_token = topup_request.cashu_token cashu_token = topup_request.cashu_token
if cashu_token is None: if cashu_token is None:
@@ -111,7 +94,7 @@ async def topup_wallet_endpoint(
if len(cashu_token) < 10 or "cashu" not in cashu_token: if len(cashu_token) < 10 or "cashu" not in cashu_token:
raise HTTPException(status_code=400, detail="Invalid token format") raise HTTPException(status_code=400, detail="Invalid token format")
try: try:
amount_msats = await credit_balance(cashu_token, billing_key, session) amount_msats = await credit_balance(cashu_token, key, session)
except ValueError as e: except ValueError as e:
error_msg = str(e) error_msg = str(e)
if "already spent" in error_msg.lower(): if "already spent" in error_msg.lower():
@@ -172,12 +155,6 @@ async def refund_wallet_endpoint(
key: ApiKey = await validate_bearer_key(bearer_value, session) key: ApiKey = await validate_bearer_key(bearer_value, session)
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 remaining_balance_msats: int = key.total_balance
if key.refund_currency == "sat": if key.refund_currency == "sat":
@@ -251,53 +228,6 @@ async def donate(token: str, ref: str | None = None) -> str:
return "Invalid token." 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( @router.api_route(
"/{path:path}", "/{path:path}",
methods=["GET", "POST", "PUT", "DELETE"], methods=["GET", "POST", "PUT", "DELETE"],
+1 -2
View File
@@ -124,7 +124,7 @@ async def partial_apikeys(request: Request) -> str:
rows = "".join( rows = "".join(
[ [
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>" 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>"
for key in api_keys for key in api_keys
] ]
) )
@@ -158,7 +158,6 @@ async def get_temporary_balances_api(request: Request) -> list[dict[str, object]
"total_requests": key.total_requests, "total_requests": key.total_requests,
"refund_address": key.refund_address, "refund_address": key.refund_address,
"key_expiry_time": key.key_expiry_time, "key_expiry_time": key.key_expiry_time,
"parent_key_hash": key.parent_key_hash,
} }
for key in api_keys for key in api_keys
] ]
-3
View File
@@ -47,9 +47,6 @@ class ApiKey(SQLModel, table=True): # type: ignore
default=None, default=None,
description="Currency of the cashu-token", description="Currency of the cashu-token",
) )
parent_key_hash: str | None = Field(
default=None, foreign_key="api_keys.hashed_key", index=True
)
@property @property
def total_balance(self) -> int: def total_balance(self) -> int:
-9
View File
@@ -6,15 +6,6 @@ from .logging import get_logger
logger = get_logger(__name__) 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: async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse:
"""Handle HTTP exceptions and include request ID in response.""" """Handle HTTP exceptions and include request ID in response."""
request_id = getattr(request.state, "request_id", "unknown") request_id = getattr(request.state, "request_id", "unknown")
+6 -8
View File
@@ -13,7 +13,10 @@ from starlette.exceptions import HTTPException
from ..balance import balance_router, deprecated_wallet_router from ..balance import balance_router, deprecated_wallet_router
from ..discovery import providers_cache_refresher, providers_router from ..discovery import providers_cache_refresher, providers_router
from ..nip91 import announce_provider 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 ..payment.price import update_prices_periodically
from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically
from ..wallet import periodic_payout from ..wallet import periodic_payout
@@ -30,9 +33,9 @@ setup_logging()
logger = get_logger(__name__) logger = get_logger(__name__)
if os.getenv("VERSION_SUFFIX") is not None: if os.getenv("VERSION_SUFFIX") is not None:
__version__ = f"0.3.0-{os.getenv('VERSION_SUFFIX')}" __version__ = f"0.2.2-{os.getenv('VERSION_SUFFIX')}"
else: else:
__version__ = "0.3.0" __version__ = "0.2.2"
@asynccontextmanager @asynccontextmanager
@@ -63,11 +66,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
await reset_all_reserved_balances(session) 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 # Apply app metadata from settings
try: try:
app.title = s.name app.title = s.name
+9 -3
View File
@@ -6,7 +6,7 @@ import os
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Any from typing import Any
from pydantic.v1 import BaseModel, BaseSettings, Field from pydantic.v1 import BaseModel, BaseSettings, Field, validator
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
@@ -37,6 +37,13 @@ class Settings(BaseSettings):
# Cashu # Cashu
cashu_mints: list[str] = Field(default_factory=list, env="CASHU_MINTS") 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") receive_ln_address: str = Field(default="", env="RECEIVE_LN_ADDRESS")
primary_mint: str = Field(default="", env="PRIMARY_MINT_URL") primary_mint: str = Field(default="", env="PRIMARY_MINT_URL")
primary_mint_unit: str = Field(default="sat", env="PRIMARY_MINT_UNIT") primary_mint_unit: str = Field(default="sat", env="PRIMARY_MINT_UNIT")
@@ -52,7 +59,6 @@ class Settings(BaseSettings):
exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE") exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE")
upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE") upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE")
tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE") 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 # Minimum per-request charge in millisatoshis when model pricing is free/zero
min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT") min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT")
reset_reserved_balance_on_startup: bool = Field( reset_reserved_balance_on_startup: bool = Field(
@@ -63,7 +69,7 @@ class Settings(BaseSettings):
cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS") 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") tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL")
providers_refresh_interval_seconds: int = Field( providers_refresh_interval_seconds: int = Field(
default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS" default=300, env="PROVIDERS_REFRESH_INTERVAL_SECONDS"
) )
pricing_refresh_interval_seconds: int = Field( pricing_refresh_interval_seconds: int = Field(
default=120, env="PRICING_REFRESH_INTERVAL_SECONDS" default=120, env="PRICING_REFRESH_INTERVAL_SECONDS"
+1 -4
View File
@@ -6,7 +6,7 @@ from typing import Any
import httpx import httpx
import websockets import websockets
from fastapi import APIRouter, HTTPException from fastapi import APIRouter
from .core.logging import get_logger from .core.logging import get_logger
from .core.settings import settings from .core.settings import settings
@@ -389,9 +389,6 @@ async def get_providers(
Return cached providers. If include_json, return provider+health; otherwise provider only. Return cached providers. If include_json, return provider+health; otherwise provider only.
Optional filter by pubkey. 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() cache = await get_cache()
if not cache: if not cache:
await refresh_providers_cache(pubkey=pubkey) await refresh_providers_cache(pubkey=pubkey)
-6
View File
@@ -15,7 +15,6 @@ class CostData(BaseModel):
input_msats: int input_msats: int
output_msats: int output_msats: int
total_msats: int total_msats: int
total_usd: float = 0.0
class MaxCostData(CostData): class MaxCostData(CostData):
@@ -62,7 +61,6 @@ async def calculate_cost( # todo: can be sync
input_msats=0, input_msats=0,
output_msats=0, output_msats=0,
total_msats=0, total_msats=0,
total_usd=0.0,
) )
usage_data = response_data["usage"] usage_data = response_data["usage"]
@@ -103,7 +101,6 @@ async def calculate_cost( # todo: can be sync
input_msats=-1, # Cost field doesn't break down by token type input_msats=-1, # Cost field doesn't break down by token type
output_msats=-1, output_msats=-1,
total_msats=cost_in_msats, total_msats=cost_in_msats,
total_usd=usd_cost,
) )
except Exception as e: except Exception as e:
logger.warning( logger.warning(
@@ -213,7 +210,6 @@ async def calculate_cost( # todo: can be sync
output_msats = round(output_tokens / 1000 * MSATS_PER_1K_OUTPUT_TOKENS, 3) output_msats = round(output_tokens / 1000 * MSATS_PER_1K_OUTPUT_TOKENS, 3)
token_based_cost = math.ceil(input_msats + output_msats) token_based_cost = math.ceil(input_msats + output_msats)
total_usd = (token_based_cost / 1000.0) * sats_usd_price()
logger.info( logger.info(
"Calculated token-based cost", "Calculated token-based cost",
@@ -223,7 +219,6 @@ async def calculate_cost( # todo: can be sync
"input_cost_msats": input_msats, "input_cost_msats": input_msats,
"output_cost_msats": output_msats, "output_cost_msats": output_msats,
"total_cost_msats": token_based_cost, "total_cost_msats": token_based_cost,
"total_usd": total_usd,
"model": response_data.get("model", "unknown"), "model": response_data.get("model", "unknown"),
}, },
) )
@@ -233,5 +228,4 @@ async def calculate_cost( # todo: can be sync
input_msats=int(input_msats), input_msats=int(input_msats),
output_msats=int(output_msats), output_msats=int(output_msats),
total_msats=token_based_cost, total_msats=token_based_cost,
total_usd=total_usd,
) )
-7
View File
@@ -253,11 +253,6 @@ async def list_models(
else 1.01, else 1.01,
) )
for r in rows 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
)
] ]
@@ -270,8 +265,6 @@ async def get_model_by_id(
if not row or not row.enabled: if not row or not row.enabled:
return None return None
provider = await session.get(UpstreamProviderRow, provider_id) 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 provider_fee = provider.provider_fee if provider else 1.01
return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee) return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee)
+6 -20
View File
@@ -79,29 +79,15 @@ async def _fetch_btc_usd_price() -> float:
"""Fetch the lowest BTC/USD price from multiple exchanges.""" """Fetch the lowest BTC/USD price from multiple exchanges."""
async with httpx.AsyncClient(timeout=30.0) as client: async with httpx.AsyncClient(timeout=30.0) as client:
try: try:
tasks = [ prices = await asyncio.gather(
asyncio.create_task(_kraken_btc_usd(client)), _kraken_btc_usd(client),
asyncio.create_task(_coinbase_btc_usd(client)), _coinbase_btc_usd(client),
asyncio.create_task(_binance_btc_usdt(client)), _binance_btc_usdt(client),
] )
valid_prices: list[float] = [] valid_prices = [price for price in prices if price is not None]
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: if not valid_prices:
logger.error("No valid BTC prices obtained from any exchange") logger.error("No valid BTC prices obtained from any exchange")
raise ValueError("Unable to fetch BTC price from any exchange") raise ValueError("Unable to fetch BTC price from any exchange")
return min(valid_prices) return min(valid_prices)
except Exception as e: except Exception as e:
logger.error( logger.error(
+76 -178
View File
@@ -16,7 +16,7 @@ from .core.db import (
create_session, create_session,
get_session, get_session,
) )
from .core.exceptions import UpstreamError from .core.settings import settings
from .payment.helpers import ( from .payment.helpers import (
calculate_discounted_max_cost, calculate_discounted_max_cost,
check_token_balance, check_token_balance,
@@ -26,15 +26,14 @@ from .payment.helpers import (
from .payment.models import Model from .payment.models import Model
from .upstream import BaseUpstreamProvider from .upstream import BaseUpstreamProvider
from .upstream.helpers import init_upstreams from .upstream.helpers import init_upstreams
from .wallet import deserialize_token_from_string
logger = get_logger(__name__) logger = get_logger(__name__)
proxy_router = APIRouter() proxy_router = APIRouter()
_upstreams: list[BaseUpstreamProvider] = [] _upstreams: list[BaseUpstreamProvider] = []
_model_instances: dict[str, Model] = {} # All aliases -> Model _model_instances: dict[str, Model] = {} # All aliases -> Model
_provider_map: dict[ _provider_map: dict[str, BaseUpstreamProvider] = {} # All aliases -> Provider
str, list[BaseUpstreamProvider]
] = {} # All aliases -> List[Provider]
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates) _unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
@@ -71,8 +70,8 @@ def get_model_instance(model_id: str) -> Model | None:
return _model_instances.get(model_id.lower()) return _model_instances.get(model_id.lower())
def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None: def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None:
"""Get UpstreamProvider list for model ID from global cache.""" """Get UpstreamProvider for model ID from global cache."""
return _provider_map.get(model_id.lower()) return _provider_map.get(model_id.lower())
@@ -99,8 +98,6 @@ async def refresh_model_maps() -> None:
disabled_model_ids: set[str] = set() disabled_model_ids: set[str] = set()
for provider in provider_rows: for provider in provider_rows:
if not provider.enabled:
continue
for model in provider.models: for model in provider.models:
if model.enabled: if model.enabled:
overrides_by_id[model.id] = (model, provider.provider_fee) overrides_by_id[model.id] = (model, provider.provider_fee)
@@ -151,14 +148,32 @@ async def proxy(
else: else:
model_id = request_body_dict.get("model", "unknown") 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) model_obj = get_model_instance(model_id)
if not model_obj: if not model_obj:
return create_error_response( return create_error_response(
"invalid_model", f"Model '{model_id}' not found", 400, request=request "invalid_model", f"Model '{model_id}' not found", 400, request=request
) )
upstreams = get_provider_for_model(model_id) upstream = get_provider_for_model(model_id)
if not upstreams: if not upstream:
return create_error_response( return create_error_response(
"invalid_model", "invalid_model",
f"No provider found for model '{model_id}'", f"No provider found for model '{model_id}'",
@@ -166,10 +181,6 @@ async def proxy(
request=request, 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( _max_cost_for_model = await get_max_cost_for_model(
model=model_id, session=session, model_obj=model_obj model=model_id, session=session, model_obj=model_obj
) )
@@ -179,31 +190,14 @@ async def proxy(
check_token_balance(headers, request_body_dict, max_cost_for_model) check_token_balance(headers, request_body_dict, max_cost_for_model)
if x_cashu := headers.get("x-cashu", None): if x_cashu := headers.get("x-cashu", None):
last_error = None if is_responses_api:
for i, upstream in enumerate(upstreams): return await upstream.handle_x_cashu_responses(
try: request, x_cashu, path, max_cost_for_model, model_obj
if is_responses_api: )
return await upstream.handle_x_cashu_responses( else:
request, x_cashu, path, max_cost_for_model, model_obj return await upstream.handle_x_cashu(
) 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): elif auth := headers.get("authorization", None):
key = await get_bearer_token_key(headers, path, session, auth) key = await get_bearer_token_key(headers, path, session, auth)
@@ -218,152 +212,56 @@ async def proxy(
) )
logger.debug("Processing unauthenticated GET request", extra={"path": path}) logger.debug("Processing unauthenticated GET request", extra={"path": path})
headers = upstream.prepare_headers(dict(request.headers))
last_error_response = None return await upstream.forward_get_request(request, path, headers)
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: if request_body_dict:
await pay_for_request(key, max_cost_for_model, session) await pay_for_request(key, max_cost_for_model, session)
for i, upstream in enumerate(upstreams): headers = upstream.prepare_headers(dict(request.headers))
headers = upstream.prepare_headers(dict(request.headers))
try: if is_responses_api:
if is_responses_api: response = await upstream.forward_responses_request(
response = await upstream.forward_responses_request( request,
request, path,
path, headers,
headers, request_body,
request_body, key,
key, max_cost_for_model,
max_cost_for_model, session,
session, model_obj,
model_obj, )
) else:
else: response = await upstream.forward_request(
response = await upstream.forward_request( request,
request, path,
path, headers,
headers, request_body,
request_body, key,
key, max_cost_for_model,
max_cost_for_model, session,
session, model_obj,
model_obj, )
)
if response.status_code != 200: if response.status_code != 200:
# Check if we should retry (502 Upstream Error or 429 Rate Limit) await revert_pay_for_request(key, session, max_cost_for_model)
should_retry = response.status_code in [502, 429, 400, 401, 403, 404] logger.warning(
if should_retry and i < len(upstreams) - 1: "Upstream request failed, revert payment",
error_message = "" extra={
try: "status_code": response.status_code,
if hasattr(response, "body"): "path": path,
body_bytes = response.body "key_hash": key.hashed_key[:8] + "...",
data = json.loads(body_bytes) "key_balance": key.balance,
if "error" in data: "max_cost_for_model": max_cost_for_model,
error_data = data["error"] "upstream_headers": response.headers
if isinstance(error_data, dict): if hasattr(response, "headers")
error_message = error_data.get("message", "") else None,
elif isinstance(error_data, str): },
error_message = error_data )
except Exception: # Return the mapped error response generated earlier rather than masking with 502
pass return response
await upstream.on_upstream_error_redirect( return response
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( async def get_bearer_token_key(
+494 -236
View File
@@ -15,7 +15,6 @@ from pydantic import BaseModel
from ..auth import adjust_payment_for_tokens from ..auth import adjust_payment_for_tokens
from ..core import get_logger from ..core import get_logger
from ..core.db import ApiKey, AsyncSession, create_session from ..core.db import ApiKey, AsyncSession, create_session
from ..core.exceptions import UpstreamError
if TYPE_CHECKING: if TYPE_CHECKING:
from ..core.db import UpstreamProviderRow from ..core.db import UpstreamProviderRow
@@ -341,20 +340,6 @@ class BaseUpstreamProvider:
message = preview[:500] message = preview[:500]
return message, upstream_code 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( async def map_upstream_error_response(
self, request: Request, path: str, upstream_response: httpx.Response self, request: Request, path: str, upstream_response: httpx.Response
) -> Response: ) -> Response:
@@ -451,111 +436,164 @@ class BaseUpstreamProvider:
async def stream_with_cost( async def stream_with_cost(
max_cost_for_model: int, max_cost_for_model: int,
) -> AsyncGenerator[bytes, None]: ) -> AsyncGenerator[bytes, None]:
stored_chunks: list[bytes] = []
usage_finalized: bool = False usage_finalized: bool = False
last_model_seen: str | None = None last_model_seen: str | None = None
usage_chunk_data: dict | None = None
done_seen: bool = False
async def finalize_db_only() -> None: async def finalize_without_usage() -> bytes | None:
nonlocal usage_finalized nonlocal usage_finalized
if usage_finalized: if usage_finalized:
return return None
async with create_session() as new_session: async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key) fresh_key = await new_session.get(key.__class__, key.hashed_key)
if not fresh_key: if not fresh_key:
return logger.warning(
try: "Key not found when finalizing streaming payment",
await adjust_payment_for_tokens( extra={"key_hash": key.hashed_key[:8] + "..."},
fresh_key,
{"model": last_model_seen or "unknown", "usage": None},
new_session,
max_cost_for_model,
) )
usage_finalized = True usage_finalized = True
except Exception: return None
pass try:
fallback: dict = {
"model": last_model_seen or "unknown",
"usage": None,
}
cost_data = await adjust_payment_for_tokens(
fresh_key, fallback, new_session, max_cost_for_model
)
usage_finalized = True
logger.info(
"Finalized streaming payment without explicit usage",
extra={
"key_hash": key.hashed_key[:8] + "...",
"cost_data": cost_data,
"balance_after_adjustment": fresh_key.balance,
},
)
return f"data: {json.dumps({'cost': cost_data})}\n\n".encode()
except Exception as cost_error:
logger.error(
"Error finalizing payment without usage",
extra={
"error": str(cost_error),
"error_type": type(cost_error).__name__,
"key_hash": key.hashed_key[:8] + "...",
},
)
usage_finalized = True
return None
try: try:
async for chunk in response.aiter_bytes(): async for chunk in response.aiter_bytes():
# Split chunk into SSE events stored_chunks.append(chunk)
parts = re.split(b"data: ", chunk) try:
for i, part in enumerate(parts): for part in re.split(b"data: ", chunk):
if not part: if not part or part.strip() in (b"[DONE]", b""):
continue continue
stripped_part = part.strip()
if not stripped_part:
continue
if stripped_part == b"[DONE]":
done_seen = True
continue
try:
obj = json.loads(part)
if isinstance(obj, dict):
if obj.get("model"):
last_model_seen = str(obj.get("model"))
if isinstance(obj.get("usage"), dict):
# Hold this chunk back to merge cost later
usage_chunk_data = obj
continue
except json.JSONDecodeError:
pass
prefix = (
b"data: " if (i > 0 or chunk.startswith(b"data: ")) else b""
)
yield prefix + part
# Stream finished, process usage if found
if usage_chunk_data:
async with create_session() as session:
fresh_key = await session.get(key.__class__, key.hashed_key)
if fresh_key:
try: try:
cost_data = await adjust_payment_for_tokens( obj = json.loads(part)
fresh_key, if isinstance(obj, dict) and obj.get("model"):
usage_chunk_data, last_model_seen = str(obj.get("model"))
session, except json.JSONDecodeError:
max_cost_for_model, pass
) except Exception:
# Merge cost into usage pass
usage_chunk_data["usage"]["cost"] = cost_data.get(
"total_usd", 0.0 yield chunk
)
# Keep detailed cost in metadata logger.debug(
usage_chunk_data["metadata"] = usage_chunk_data.get( "Streaming completed, analyzing usage data",
"metadata", {} extra={
) "key_hash": key.hashed_key[:8] + "...",
usage_chunk_data["metadata"]["routstr"] = { "chunks_count": len(stored_chunks),
"cost": cost_data },
} )
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
usage_finalized = True for i in range(len(stored_chunks) - 1, -1, -1):
except Exception: chunk = stored_chunks[i]
# Fallback: yield original usage chunk if adjustment fails if not chunk:
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() continue
try:
events = re.split(b"data: ", chunk)
for event_data in events:
if not event_data or event_data.strip() in (b"[DONE]", b""):
continue
try:
data = json.loads(event_data)
if isinstance(data, dict) and data.get("model"):
last_model_seen = str(data.get("model"))
if isinstance(data, dict) and isinstance(
data.get("usage"), dict
):
async with create_session() as new_session:
fresh_key = await new_session.get(
key.__class__, key.hashed_key
)
if fresh_key:
try:
cost_data = (
await adjust_payment_for_tokens(
fresh_key,
data,
new_session,
max_cost_for_model,
)
)
usage_finalized = True
logger.info(
"Payment adjustment completed for streaming",
extra={
"key_hash": key.hashed_key[:8]
+ "...",
"cost_data": cost_data,
"model": last_model_seen,
"balance_after_adjustment": fresh_key.balance,
},
)
yield f"data: {json.dumps({'cost': cost_data})}\n\n".encode()
except Exception as cost_error:
logger.error(
"Error adjusting payment for streaming tokens",
extra={
"error": str(cost_error),
"error_type": type(
cost_error
).__name__,
"key_hash": key.hashed_key[:8]
+ "...",
},
)
break
except json.JSONDecodeError:
continue
except Exception as e:
logger.error(
"Error processing streaming response chunk",
extra={
"error": str(e),
"error_type": type(e).__name__,
"key_hash": key.hashed_key[:8] + "...",
},
)
if not usage_finalized: if not usage_finalized:
await finalize_db_only() maybe_cost_event = await finalize_without_usage()
if maybe_cost_event is not None:
if done_seen: yield maybe_cost_event
yield b"data: [DONE]\n\n"
except Exception as stream_error: except Exception as stream_error:
logger.warning( logger.warning(
"Streaming interrupted; finalizing in background", "Streaming interrupted; finalizing without usage",
extra={ extra={
"error": str(stream_error), "error": str(stream_error),
"error_type": type(stream_error).__name__,
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
}, },
) )
raise raise
finally: finally:
if not usage_finalized: if not usage_finalized:
await finalize_db_only() await finalize_without_usage()
# Remove inaccurate encoding headers from upstream response # Remove inaccurate encoding headers from upstream response
response_headers = dict(response.headers) response_headers = dict(response.headers)
@@ -595,7 +633,6 @@ class BaseUpstreamProvider:
}, },
) )
content: bytes | None = None
try: try:
content = await response.aread() content = await response.aread()
response_json = json.loads(content) response_json = json.loads(content)
@@ -612,14 +649,6 @@ class BaseUpstreamProvider:
cost_data = await adjust_payment_for_tokens( cost_data = await adjust_payment_for_tokens(
key, response_json, session, deducted_max_cost key, response_json, session, deducted_max_cost
) )
# Merge cost into usage for OpenCode
if "usage" in response_json:
response_json["usage"]["cost"] = cost_data.get("total_usd", 0.0)
# Keep detailed cost
response_json["metadata"] = response_json.get("metadata", {})
response_json["metadata"]["routstr"] = {"cost": cost_data}
response_json["cost"] = cost_data response_json["cost"] = cost_data
logger.info( logger.info(
@@ -705,135 +734,180 @@ class BaseUpstreamProvider:
async def stream_with_responses_cost( async def stream_with_responses_cost(
max_cost_for_model: int, max_cost_for_model: int,
) -> AsyncGenerator[bytes, None]: ) -> AsyncGenerator[bytes, None]:
stored_chunks: list[bytes] = []
usage_finalized: bool = False usage_finalized: bool = False
last_model_seen: str | None = None last_model_seen: str | None = None
reasoning_tokens: int = 0 reasoning_tokens: int = 0
usage_chunk_data: dict | None = None
done_seen: bool = False
async def finalize_db_only() -> None: async def finalize_without_usage() -> bytes | None:
nonlocal usage_finalized nonlocal usage_finalized
if usage_finalized: if usage_finalized:
return return None
async with create_session() as new_session: async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key) fresh_key = await new_session.get(key.__class__, key.hashed_key)
if not fresh_key: if not fresh_key:
return logger.warning(
try: "Key not found when finalizing Responses API streaming payment",
await adjust_payment_for_tokens( extra={"key_hash": key.hashed_key[:8] + "..."},
fresh_key,
{"model": last_model_seen or "unknown", "usage": None},
new_session,
max_cost_for_model,
) )
usage_finalized = True usage_finalized = True
except Exception: return None
pass try:
fallback: dict = {
"model": last_model_seen or "unknown",
"usage": None,
}
cost_data = await adjust_payment_for_tokens(
fresh_key, fallback, new_session, max_cost_for_model
)
usage_finalized = True
logger.info(
"Finalized Responses API streaming payment without explicit usage",
extra={
"key_hash": key.hashed_key[:8] + "...",
"cost_data": cost_data,
"balance_after_adjustment": fresh_key.balance,
},
)
return f"data: {json.dumps({'cost': cost_data})}\\n\\n".encode()
except Exception as cost_error:
logger.error(
"Error finalizing Responses API payment without usage",
extra={
"error": str(cost_error),
"error_type": type(cost_error).__name__,
"key_hash": key.hashed_key[:8] + "...",
},
)
usage_finalized = True
return None
try: try:
async for chunk in response.aiter_bytes(): async for chunk in response.aiter_bytes():
# Split chunk into SSE events stored_chunks.append(chunk)
parts = re.split(b"data: ", chunk) try:
for i, part in enumerate(parts): for part in re.split(b"data: ", chunk):
if not part: if not part or part.strip() in (b"[DONE]", b""):
continue continue
stripped_part = part.strip()
if not stripped_part:
continue
if stripped_part == b"[DONE]":
done_seen = True
continue
try:
obj = json.loads(part)
if isinstance(obj, dict):
if obj.get("model"):
last_model_seen = str(obj.get("model"))
# Track reasoning tokens for Responses API
if usage := obj.get("usage", {}):
if (
isinstance(usage, dict)
and "reasoning_tokens" in usage
):
reasoning_tokens += usage.get(
"reasoning_tokens", 0
)
# Responses API usage is in response.completed/incomplete events
chunk_type = obj.get("type", "")
if chunk_type in (
"response.completed",
"response.incomplete",
):
usage_chunk_data = obj
continue
except json.JSONDecodeError:
pass
prefix = (
b"data: " if (i > 0 or chunk.startswith(b"data: ")) else b""
)
yield prefix + part
# Stream finished, process usage if found
if usage_chunk_data:
async with create_session() as session:
fresh_key = await session.get(key.__class__, key.hashed_key)
if fresh_key:
try: try:
cost_data = await adjust_payment_for_tokens( obj = json.loads(part)
fresh_key, if isinstance(obj, dict):
usage_chunk_data, if obj.get("model"):
session, last_model_seen = str(obj.get("model"))
max_cost_for_model,
)
# Merge cost into usage chunk
if (
"response" in usage_chunk_data
and "usage" in usage_chunk_data["response"]
):
usage_chunk_data["response"]["usage"]["cost"] = (
cost_data.get("total_usd", 0.0)
)
elif "usage" in usage_chunk_data:
usage_chunk_data["usage"]["cost"] = cost_data.get(
"total_usd", 0.0
)
# Keep detailed cost in metadata # Track reasoning tokens for Responses API
usage_chunk_data["metadata"] = usage_chunk_data.get( if usage := obj.get("usage", {}):
"metadata", {} if (
) isinstance(usage, dict)
usage_chunk_data["metadata"]["routstr"] = { and "reasoning_tokens" in usage
"cost": cost_data ):
} reasoning_tokens += usage.get(
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() "reasoning_tokens", 0
usage_finalized = True )
except Exception: except json.JSONDecodeError:
# Fallback: yield original usage chunk if adjustment fails pass
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() except Exception:
pass
yield chunk
logger.debug(
"Responses API streaming completed, analyzing usage data",
extra={
"key_hash": key.hashed_key[:8] + "...",
"chunks_count": len(stored_chunks),
"reasoning_tokens": reasoning_tokens,
},
)
# Process final usage data
for i in range(len(stored_chunks) - 1, -1, -1):
chunk = stored_chunks[i]
if not chunk:
continue
try:
events = re.split(b"data: ", chunk)
for event_data in events:
if not event_data or event_data.strip() in (b"[DONE]", b""):
continue
try:
data = json.loads(event_data)
if isinstance(data, dict) and data.get("model"):
last_model_seen = str(data.get("model"))
if isinstance(data, dict) and isinstance(
data.get("usage"), dict
):
# Include reasoning tokens in usage calculation
async with create_session() as new_session:
fresh_key = await new_session.get(
key.__class__, key.hashed_key
)
if fresh_key:
try:
cost_data = (
await adjust_payment_for_tokens(
fresh_key,
data,
new_session,
max_cost_for_model,
)
)
usage_finalized = True
logger.info(
"Payment adjustment completed for Responses API streaming",
extra={
"key_hash": key.hashed_key[:8]
+ "...",
"cost_data": cost_data,
"model": last_model_seen,
"reasoning_tokens": reasoning_tokens,
"balance_after_adjustment": fresh_key.balance,
},
)
yield f"data: {json.dumps({'cost': cost_data})}\\n\\n".encode()
except Exception as cost_error:
logger.error(
"Error adjusting payment for Responses API streaming tokens",
extra={
"error": str(cost_error),
"error_type": type(
cost_error
).__name__,
"key_hash": key.hashed_key[:8]
+ "...",
},
)
break
except json.JSONDecodeError:
continue
except Exception as e:
logger.error(
"Error processing Responses API streaming response chunk",
extra={
"error": str(e),
"error_type": type(e).__name__,
"key_hash": key.hashed_key[:8] + "...",
},
)
if not usage_finalized: if not usage_finalized:
await finalize_db_only() maybe_cost_event = await finalize_without_usage()
if maybe_cost_event is not None:
if done_seen: yield maybe_cost_event
yield b"data: [DONE]\n\n"
except Exception as stream_error: except Exception as stream_error:
logger.warning( logger.warning(
"Responses API streaming interrupted; finalizing in background", "Responses API streaming interrupted; finalizing without usage",
extra={ extra={
"error": str(stream_error), "error": str(stream_error),
"error_type": type(stream_error).__name__,
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
}, },
) )
raise raise
finally: finally:
if not usage_finalized: if not usage_finalized:
await finalize_db_only() await finalize_without_usage()
# Remove inaccurate encoding headers from upstream response # Remove inaccurate encoding headers from upstream response
response_headers = dict(response.headers) response_headers = dict(response.headers)
@@ -873,7 +947,6 @@ class BaseUpstreamProvider:
}, },
) )
content: bytes | None = None
try: try:
content = await response.aread() content = await response.aread()
response_json = json.loads(content) response_json = json.loads(content)
@@ -893,14 +966,6 @@ class BaseUpstreamProvider:
cost_data = await adjust_payment_for_tokens( cost_data = await adjust_payment_for_tokens(
key, response_json, session, deducted_max_cost key, response_json, session, deducted_max_cost
) )
# Merge cost into usage for OpenCode
if "usage" in response_json:
response_json["usage"]["cost"] = cost_data.get("total_usd", 0.0)
# Keep detailed cost
response_json["metadata"] = response_json.get("metadata", {})
response_json["metadata"]["routstr"] = {"cost": cost_data}
response_json["cost"] = cost_data response_json["cost"] = cost_data
logger.info( logger.info(
@@ -999,6 +1064,170 @@ class BaseUpstreamProvider:
}, },
) )
async def handle_streaming_messages_completion(
self, response: httpx.Response, key: ApiKey, max_cost_for_model: int
) -> StreamingResponse:
async def stream_with_cost(
max_cost_for_model: int,
) -> AsyncGenerator[bytes, None]:
stored_chunks: list[bytes] = []
usage_finalized: bool = False
last_model_seen: str | None = None
input_tokens: int = 0
output_tokens: int = 0
async def finalize_without_usage() -> bytes | None:
nonlocal usage_finalized
if usage_finalized:
return None
async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if not fresh_key:
usage_finalized = True
return None
try:
fallback: dict = {
"model": last_model_seen or "unknown",
"usage": None,
}
cost_data = await adjust_payment_for_tokens(
fresh_key, fallback, new_session, max_cost_for_model
)
usage_finalized = True
return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
except Exception:
usage_finalized = True
return None
try:
async for chunk in response.aiter_bytes():
stored_chunks.append(chunk)
try:
decoded_chunk = chunk.decode("utf-8", errors="ignore")
for line in decoded_chunk.split("\n"):
if line.startswith("data: "):
try:
data = json.loads(line[6:])
if isinstance(data, dict):
msg = data.get("message", {})
if msg and msg.get("model"):
last_model_seen = str(msg.get("model"))
if usage := msg.get("usage"):
input_tokens += usage.get("input_tokens", 0)
output_tokens += usage.get(
"output_tokens", 0
)
if usage := data.get("usage"):
input_tokens += usage.get("input_tokens", 0)
output_tokens += usage.get(
"output_tokens", 0
)
except json.JSONDecodeError:
pass
except Exception:
pass
yield chunk
usage_data = {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
}
if input_tokens > 0 or output_tokens > 0:
async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if fresh_key:
try:
combined_data = {
"model": last_model_seen or "unknown",
"usage": usage_data,
}
cost_data = await adjust_payment_for_tokens(
fresh_key,
combined_data,
new_session,
max_cost_for_model,
)
usage_finalized = True
yield f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
except Exception:
pass
if not usage_finalized:
maybe_cost_event = await finalize_without_usage()
if maybe_cost_event is not None:
yield maybe_cost_event
except Exception:
if not usage_finalized:
await finalize_without_usage()
raise
finally:
if not usage_finalized:
await finalize_without_usage()
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return StreamingResponse(
stream_with_cost(max_cost_for_model),
status_code=response.status_code,
headers=response_headers,
)
async def handle_non_streaming_messages_completion(
self,
response: httpx.Response,
key: ApiKey,
session: AsyncSession,
deducted_max_cost: int,
path: str,
) -> Response:
try:
content = await response.aread()
response_json = json.loads(content)
if path.endswith("count_tokens") and "usage" not in response_json:
input_tokens = response_json.get("input_tokens", 0)
response_json["usage"] = {"input_tokens": input_tokens}
cost_data = await adjust_payment_for_tokens(
key, response_json, session, deducted_max_cost
)
response_json["cost"] = cost_data
allowed_headers = {
"content-type",
"cache-control",
"date",
"vary",
"access-control-allow-origin",
"access-control-allow-methods",
"access-control-allow-headers",
"access-control-allow-credentials",
"access-control-expose-headers",
"access-control-max-age",
}
response_headers = {
k: v
for k, v in response.headers.items()
if k.lower() in allowed_headers
}
return Response(
content=json.dumps(response_json).encode(),
status_code=response.status_code,
headers=response_headers,
media_type="application/json",
)
except Exception:
raise
async def forward_request( async def forward_request(
self, self,
request: Request, request: Request,
@@ -1083,14 +1312,6 @@ class BaseUpstreamProvider:
) )
if response.status_code != 200: 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: try:
mapped_error = await self.map_upstream_error_response( mapped_error = await self.map_upstream_error_response(
request, path, response request, path, response
@@ -1100,7 +1321,54 @@ class BaseUpstreamProvider:
await client.aclose() await client.aclose()
return mapped_error return mapped_error
if path.endswith("chat/completions") or path.endswith("embeddings"): if (
path.endswith("chat/completions")
or path.endswith("embeddings")
or path.endswith("messages")
or path.endswith("messages/count_tokens")
):
if path.endswith("messages"):
client_wants_streaming = False
if request_body:
try:
request_data = json.loads(request_body)
client_wants_streaming = request_data.get("stream", False)
except json.JSONDecodeError:
pass
content_type = response.headers.get("content-type", "")
upstream_is_streaming = "text/event-stream" in content_type
is_streaming = client_wants_streaming and upstream_is_streaming
if is_streaming and response.status_code == 200:
result = await self.handle_streaming_messages_completion(
response, key, max_cost_for_model
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
background_tasks.add_task(client.aclose)
result.background = background_tasks
return result
if response.status_code == 200:
try:
return await self.handle_non_streaming_messages_completion(
response, key, session, max_cost_for_model, path
)
finally:
await response.aclose()
await client.aclose()
if path.endswith("messages/count_tokens"):
if response.status_code == 200:
try:
return await self.handle_non_streaming_messages_completion(
response, key, session, max_cost_for_model, path
)
finally:
await response.aclose()
await client.aclose()
if path.endswith("chat/completions"): if path.endswith("chat/completions"):
client_wants_streaming = False client_wants_streaming = False
if request_body: if request_body:
@@ -1181,9 +1449,6 @@ class BaseUpstreamProvider:
background=background_tasks, background=background_tasks,
) )
except UpstreamError:
raise
except httpx.RequestError as exc: except httpx.RequestError as exc:
await client.aclose() await client.aclose()
error_type = type(exc).__name__ error_type = type(exc).__name__
@@ -1211,7 +1476,9 @@ class BaseUpstreamProvider:
else: else:
error_message = f"Error connecting to upstream service: {error_type}" error_message = f"Error connecting to upstream service: {error_type}"
raise UpstreamError(error_message, status_code=502) return create_error_response(
"upstream_error", error_message, 502, request=request
)
except Exception as exc: except Exception as exc:
await client.aclose() await client.aclose()
@@ -1324,14 +1591,6 @@ class BaseUpstreamProvider:
) )
if response.status_code != 200: 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: try:
mapped_error = await self.map_upstream_error_response( mapped_error = await self.map_upstream_error_response(
request, path, response request, path, response
@@ -1399,9 +1658,6 @@ class BaseUpstreamProvider:
background=background_tasks, background=background_tasks,
) )
except UpstreamError:
raise
except httpx.RequestError as exc: except httpx.RequestError as exc:
await client.aclose() await client.aclose()
error_type = type(exc).__name__ error_type = type(exc).__name__
@@ -1429,7 +1685,9 @@ class BaseUpstreamProvider:
else: else:
error_message = f"Error connecting to upstream service: {error_type}" error_message = f"Error connecting to upstream service: {error_type}"
raise UpstreamError(error_message, status_code=502) return create_error_response(
"upstream_error", error_message, 502, request=request
)
except Exception as exc: except Exception as exc:
await client.aclose() await client.aclose()
+265
View File
@@ -0,0 +1,265 @@
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
+8 -2
View File
@@ -102,8 +102,6 @@ async def get_all_models_with_overrides(
) )
for row in override_rows for row in override_rows
if row.upstream_provider_id is not None 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] = {} all_models: dict[str, Model] = {}
@@ -220,6 +218,14 @@ async def init_upstreams() -> list[BaseUpstreamProvider]:
results = await asyncio.gather(*tasks) results = await asyncio.gather(*tasks)
upstreams = [p for p in results if p is not None] 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 return upstreams
-31
View File
@@ -196,37 +196,6 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
) )
return [] 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]: async def create_account(self) -> dict[str, object]:
"""Create a new PPQ.AI account. """Create a new PPQ.AI account.
+2
View File
@@ -313,6 +313,8 @@ async def periodic_payout() -> None:
try: try:
async with db.create_session() as session: async with db.create_session() as session:
for mint_url in settings.cashu_mints: for mint_url in settings.cashu_mints:
if mint_url == "https://testnut.cashu.space":
continue
for unit in ["sat", "msat"]: for unit in ["sat", "msat"]:
wallet = await get_wallet(mint_url, unit) wallet = await get_wallet(mint_url, unit)
proofs = get_proofs_per_mint_and_unit( proofs = get_proofs_per_mint_and_unit(
-137
View File
@@ -1,137 +0,0 @@
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)
+24
View File
@@ -10,6 +10,7 @@ from httpx import AsyncClient
from .utils import ( from .utils import (
CashuTokenGenerator, CashuTokenGenerator,
PerformanceValidator,
ResponseValidator, ResponseValidator,
) )
@@ -158,7 +159,30 @@ async def test_error_handling(
assert response.status_code == 401 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 @pytest.mark.integration
+28 -33
View File
@@ -9,7 +9,6 @@ import gc
import statistics import statistics
import time import time
from typing import Any, Dict, List from typing import Any, Dict, List
from unittest.mock import patch
import psutil import psutil
import pytest import pytest
@@ -106,46 +105,42 @@ class TestPerformanceBaseline:
("GET", "/v1/wallet/info", authenticated_client, None), ("GET", "/v1/wallet/info", authenticated_client, None),
] ]
# Enable provider discovery for this test # Warm up
with patch( for _ in range(10):
"routstr.core.settings.settings.providers_refresh_interval_seconds", 300 await integration_client.get("/")
):
# Warm up
for _ in range(10):
await integration_client.get("/")
# Test each endpoint # Test each endpoint
for method, path, client, data in endpoints: for method, path, client, data in endpoints:
response_times = [] response_times = []
for i in range(100): for i in range(100):
start = time.time() start = time.time()
if method == "GET": if method == "GET":
response = await client.get(path) response = await client.get(path)
else: else:
response = await client.post(path, json=data) response = await client.post(path, json=data)
duration = time.time() - start duration = time.time() - start
response_times.append(duration * 1000) # Convert to ms 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: if i % 10 == 0:
metrics.record_system_metrics() metrics.record_system_metrics()
# Verify 95th percentile < 500ms # Verify 95th percentile < 500ms
p95 = sorted(response_times)[int(len(response_times) * 0.95)] p95 = sorted(response_times)[int(len(response_times) * 0.95)]
assert p95 < 500, ( assert p95 < 500, (
f"{method} {path} p95 response time {p95}ms exceeds 500ms limit" f"{method} {path} p95 response time {p95}ms exceeds 500ms limit"
) )
print(f"\n{method} {path}:") print(f"\n{method} {path}:")
print(f" Mean: {statistics.mean(response_times):.2f}ms") print(f" Mean: {statistics.mean(response_times):.2f}ms")
print(f" P95: {p95:.2f}ms") print(f" P95: {p95:.2f}ms")
print( print(
f" P99: {sorted(response_times)[int(len(response_times) * 0.99)]:.2f}ms" f" P99: {sorted(response_times)[int(len(response_times) * 0.99)]:.2f}ms"
) )
@pytest.mark.integration @pytest.mark.integration
+42 -11
View File
@@ -3,7 +3,7 @@ Integration tests for provider management functionality.
Tests GET /v1/providers/ endpoint for listing and managing providers. Tests GET /v1/providers/ endpoint for listing and managing providers.
""" """
from typing import Any, Generator from typing import Any
from unittest.mock import patch from unittest.mock import patch
import pytest import pytest
@@ -11,7 +11,7 @@ from httpx import AsyncClient
from routstr.discovery import _PROVIDERS_CACHE from routstr.discovery import _PROVIDERS_CACHE
from .utils import ResponseValidator from .utils import PerformanceValidator, ResponseValidator
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
@@ -19,15 +19,6 @@ def _clear_providers_cache() -> None:
_PROVIDERS_CACHE.clear() _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.integration
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_providers_endpoint_default_response( async def test_providers_endpoint_default_response(
@@ -527,6 +518,46 @@ async def test_providers_endpoint_response_format(
assert isinstance(data_json["providers"], list) 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.integration
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_providers_endpoint_concurrent_requests( async def test_providers_endpoint_concurrent_requests(
@@ -18,6 +18,7 @@ from routstr.core.db import ApiKey
from .utils import ( from .utils import (
ConcurrencyTester, ConcurrencyTester,
PerformanceValidator,
) )
@@ -550,7 +551,39 @@ async def test_proxy_get_concurrent_requests(
assert response.status_code == 200 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 @pytest.mark.integration
@@ -15,6 +15,7 @@ from httpx import ASGITransport, AsyncClient
from .utils import ( from .utils import (
ConcurrencyTester, ConcurrencyTester,
PerformanceValidator,
) )
@@ -289,7 +290,55 @@ async def test_proxy_post_unauthorized_access(integration_client: AsyncClient) -
assert response.status_code in [400, 401] 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 @pytest.mark.integration
+27 -1
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. Tests GET /v1/wallet/ and GET /v1/wallet/info endpoints with various scenarios.
""" """
import time
from datetime import datetime, timedelta from datetime import datetime, timedelta
from typing import Any from typing import Any
@@ -367,4 +367,30 @@ async def test_wallet_info_with_special_characters_in_headers(
# Note: Current implementation doesn't return refund_address in response # 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
+38
View File
@@ -537,4 +537,42 @@ async def test_refund_with_expired_key(
assert response.json()["recipient"] == "expired@ln.address" 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
+98
View File
@@ -10,6 +10,7 @@ os.environ["UPSTREAM_API_KEY"] = "test"
from routstr.algorithm import ( # noqa: E402 from routstr.algorithm import ( # noqa: E402
calculate_model_cost_score, calculate_model_cost_score,
get_provider_penalty, get_provider_penalty,
should_prefer_model,
) )
from routstr.payment.models import Architecture, Model, Pricing # noqa: E402 from routstr.payment.models import Architecture, Model, Pricing # noqa: E402
@@ -99,3 +100,100 @@ def test_get_provider_penalty_openrouter() -> None:
provider = create_test_provider("openrouter", "https://openrouter.ai/api/v1") provider = create_test_provider("openrouter", "https://openrouter.ai/api/v1")
penalty = get_provider_penalty(provider) penalty = get_provider_penalty(provider)
assert penalty == 1.001 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")
+13 -73
View File
@@ -60,11 +60,7 @@ export function TemporaryBalances({
let totalRequests = 0; let totalRequests = 0;
balances.forEach((balance) => { balances.forEach((balance) => {
// Only count parents for total balance to avoid double counting totalBalance += balance.balance || 0;
// since child keys use parent balance
if (!balance.parent_key_hash) {
totalBalance += balance.balance || 0;
}
totalSpent += balance.total_spent || 0; totalSpent += balance.total_spent || 0;
totalRequests += balance.total_requests || 0; totalRequests += balance.total_requests || 0;
}); });
@@ -76,34 +72,6 @@ export function TemporaryBalances({
? calculateTotals(data) ? calculateTotals(data)
: { totalBalance: 0, totalSpent: 0, totalRequests: 0 }; : { 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 ( return (
<> <>
<Card className='h-full w-full shadow-sm'> <Card className='h-full w-full shadow-sm'>
@@ -214,37 +182,22 @@ export function TemporaryBalances({
<div className='text-right'>Expiry Time</div> <div className='text-right'>Expiry Time</div>
</div> </div>
{hierarchicalData.length > 0 ? ( {filteredData.length > 0 ? (
hierarchicalData.map((balance, index) => ( filteredData.map((balance, index) => (
<div <div
key={index} key={index}
className={cn( className={cn(
'hover:bg-muted/50 border-t p-3 text-sm transition-colors', 'hover:bg-muted/50 border-t p-3 text-sm transition-colors',
balance.balance === 0 && balance.balance === 0 && 'opacity-60'
!balance.isChild &&
'opacity-60',
balance.isChild &&
'ml-4 border-l-2 border-l-blue-200 bg-blue-50/30'
)} )}
> >
{/* Desktop Layout */} {/* Desktop Layout */}
<div className='hidden grid-cols-6 gap-2 md:grid'> <div className='hidden grid-cols-6 gap-2 md:grid'>
<div className='flex max-w-48 items-center gap-2 truncate font-mono text-xs break-all'> <div className='max-w-32 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} {balance.hashed_key}
</div> </div>
<div className='text-right font-mono'> <div className='text-right font-mono'>
{balance.isChild ? ( {formatBalance(balance.balance)}
<span className='text-muted-foreground italic'>
(Parent)
</span>
) : (
formatBalance(balance.balance)
)}
</div> </div>
<div className='text-right font-mono'> <div className='text-right font-mono'>
{formatBalance(balance.total_spent)} {formatBalance(balance.total_spent)}
@@ -273,20 +226,13 @@ export function TemporaryBalances({
{/* Mobile Layout */} {/* Mobile Layout */}
<div className='space-y-3 md:hidden'> <div className='space-y-3 md:hidden'>
<div className='flex items-center justify-between'> <div className='space-y-1'>
<div className='space-y-1'> <span className='text-muted-foreground text-xs font-medium'>
<span className='text-muted-foreground text-xs font-medium'> Key
{balance.isChild ? 'Child Key' : 'Key'} </span>
</span> <div className='font-mono text-xs break-all'>
<div className='font-mono text-xs break-all'> {balance.hashed_key}
{balance.hashed_key}
</div>
</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>
<div className='grid grid-cols-2 gap-3'> <div className='grid grid-cols-2 gap-3'>
@@ -295,13 +241,7 @@ export function TemporaryBalances({
Balance Balance
</div> </div>
<div className='truncate font-mono text-sm'> <div className='truncate font-mono text-sm'>
{balance.isChild ? ( {formatBalance(balance.balance)}
<span className='text-muted-foreground text-xs italic'>
(Uses Parent)
</span>
) : (
formatBalance(balance.balance)
)}
</div> </div>
</div> </div>
<div className='space-y-1'> <div className='space-y-1'>
-1
View File
@@ -926,7 +926,6 @@ export const TemporaryBalanceSchema = z.object({
total_requests: z.number(), total_requests: z.number(),
refund_address: z.string().nullable(), refund_address: z.string().nullable(),
key_expiry_time: z.number().nullable(), key_expiry_time: z.number().nullable(),
parent_key_hash: z.string().nullable().optional(),
}); });
export type TemporaryBalance = z.infer<typeof TemporaryBalanceSchema>; export type TemporaryBalance = z.infer<typeof TemporaryBalanceSchema>;
Generated
+1 -1
View File
@@ -1878,7 +1878,7 @@ wheels = [
[[package]] [[package]]
name = "routstr" name = "routstr"
version = "0.3.0" version = "0.2.2"
source = { editable = "." } source = { editable = "." }
dependencies = [ dependencies = [
{ name = "aiosqlite" }, { name = "aiosqlite" },