diff --git a/compose.yml b/compose.yml index 2e7559a9..28f9c5d8 100644 --- a/compose.yml +++ b/compose.yml @@ -6,7 +6,7 @@ services: context: ./ui dockerfile: Dockerfile.build args: - NEXT_PUBLIC_API_URL: ${NEXT_PUBLIC_API_URL:-http://127.0.0.1:8000} + # NEXT_PUBLIC_API_URL: ${NEXT_PUBLIC_API_URL:-http://127.0.0.1:8000} NEXT_PUBLIC_ADMIN_API_KEY: ${NEXT_PUBLIC_ADMIN_API_KEY:-} volumes: - ./ui_out:/output diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index 2e9c6aba..952c9373 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -434,6 +434,42 @@ Authorization: Bearer sk-... } ``` +### Create Child Key + +Creates one or more child API keys that share the parent's balance. Each child key creation costs a fixed amount (configurable). + +```http +POST /v1/balance/child-key +Authorization: Bearer sk-... +``` + +**Request Body:** + +```json +{ + "count": 1 +} +``` + +**Parameters:** + +| Parameter | Type | Required | Default | Description | +|-----------|------|----------|---------|-------------| +| `count` | integer | Yes | - | Number of child keys to create (1-50) | + +**Response:** + +```json +{ + "api_keys": ["sk-abc...", "sk-def..."], + "count": 2, + "cost_msats": 2000, + "cost_sats": 2, + "parent_balance": 98000, + "parent_balance_sats": 98 +} +``` + ## Provider Discovery ## Admin Settings diff --git a/examples/create_child_keys.py b/examples/create_child_keys.py new file mode 100644 index 00000000..24f4556d --- /dev/null +++ b/examples/create_child_keys.py @@ -0,0 +1,45 @@ +import json +import sys + +import httpx + + +def create_child_keys(base_url: str, api_key: str, count: int = 3) -> list[str]: + headers = {"Authorization": f"Bearer {api_key}"} + + print(f"Requesting {count} child keys from {base_url}...") + + child_keys = [] + + for i in range(count): + try: + response = httpx.post(f"{base_url}/v1/balance/child-key", headers=headers) + if response.status_code == 200: + data = response.json() + child_keys.append(data["api_key"]) + print( + f" [{i + 1}] Created: {data['api_key']} (Cost: {data['cost_msats']} msats)" + ) + else: + print(f" [{i + 1}] Failed: {response.status_code} - {response.text}") + except Exception as e: + print(f" [{i + 1}] Error: {str(e)}") + + return child_keys + + +if __name__ == "__main__": + if len(sys.argv) < 2: + print("Usage: python create_child_keys.py [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.") diff --git a/migrations/versions/a86e5348850b_.py b/migrations/versions/a86e5348850b_.py new file mode 100644 index 00000000..12c35e41 --- /dev/null +++ b/migrations/versions/a86e5348850b_.py @@ -0,0 +1,42 @@ +""" + +Revision ID: a86e5348850b +Revises: b9667ffc5701 +Create Date: 2026-01-10 18:57:48.475781 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +# revision identifiers, used by Alembic. +revision = "a86e5348850b" +down_revision = "b9667ffc5701" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # Use batch_alter_table for SQLite compatibility + with op.batch_alter_table("api_keys", schema=None) as batch_op: + batch_op.add_column( + sa.Column( + "parent_key_hash", sqlmodel.sql.sqltypes.AutoString(), nullable=True + ) + ) + batch_op.create_index( + batch_op.f("ix_api_keys_parent_key_hash"), ["parent_key_hash"], unique=False + ) + batch_op.create_foreign_key( + "fk_api_keys_parent_key_hash", + "api_keys", + ["parent_key_hash"], + ["hashed_key"], + ) + + +def downgrade() -> None: + with op.batch_alter_table("api_keys", schema=None) as batch_op: + batch_op.drop_constraint("fk_api_keys_parent_key_hash", type_="foreignkey") + batch_op.drop_index(batch_op.f("ix_api_keys_parent_key_hash")) + batch_op.drop_column("parent_key_hash") diff --git a/migrations/versions/c2d3e4f5a6b7_upstream_provider_base_url_api_key_unique.py b/migrations/versions/c2d3e4f5a6b7_upstream_provider_base_url_api_key_unique.py new file mode 100644 index 00000000..5132be66 --- /dev/null +++ b/migrations/versions/c2d3e4f5a6b7_upstream_provider_base_url_api_key_unique.py @@ -0,0 +1,118 @@ +"""make upstream provider base_url + api_key unique + +Revision ID: c2d3e4f5a6b7 +Revises: a86e5348850b +Create Date: 2026-01-25 00:00:00.000000 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +revision = "c2d3e4f5a6b7" +down_revision = "a86e5348850b" +branch_labels = None +depends_on = None + + +def _recreate_table_sqlite(add_base_url_unique: bool) -> None: + conn = op.get_bind() + existing_tables = { + row[0] + for row in conn.exec_driver_sql( + "SELECT name FROM sqlite_master WHERE type='table'" + ).fetchall() + } + if "upstream_providers_old" in existing_tables: + if "upstream_providers" in existing_tables: + op.drop_table("upstream_providers_old") + else: + op.execute( + "ALTER TABLE upstream_providers_old RENAME TO upstream_providers" + ) + existing_tables.add("upstream_providers") + if "upstream_providers" not in existing_tables: + return + + constraints = [ + sa.UniqueConstraint( + "base_url", + "api_key", + name="uq_upstream_providers_base_url_api_key", + ) + ] + if add_base_url_unique: + constraints.append( + sa.UniqueConstraint("base_url", name="uq_upstream_providers_base_url") + ) + + op.execute("ALTER TABLE upstream_providers RENAME TO upstream_providers_old") + op.create_table( + "upstream_providers", + sa.Column( + "id", sa.Integer(), primary_key=True, nullable=False, autoincrement=True + ), + sa.Column("provider_type", sa.String(), nullable=False), + sa.Column("base_url", sa.String(), nullable=False), + sa.Column("api_key", sa.String(), nullable=False), + sa.Column("api_version", sa.String(), nullable=True), + sa.Column("enabled", sa.Boolean(), nullable=False), + sa.Column("provider_fee", sa.Float(), nullable=False, server_default="1.01"), + *constraints, + ) + op.execute( + "INSERT INTO upstream_providers (id, provider_type, base_url, api_key, api_version, enabled, provider_fee) " + "SELECT id, provider_type, base_url, api_key, api_version, enabled, provider_fee " + "FROM upstream_providers_old" + ) + op.drop_table("upstream_providers_old") + + +def upgrade() -> None: + conn = op.get_bind() + if conn.dialect.name == "sqlite": + _recreate_table_sqlite(add_base_url_unique=False) + return + + inspector = sa.inspect(conn) + for constraint in inspector.get_unique_constraints("upstream_providers"): + name = constraint.get("name") + if constraint.get("column_names") == ["base_url"] and name: + op.drop_constraint( + name, + "upstream_providers", + type_="unique", + ) + index_names = {idx["name"] for idx in inspector.get_indexes("upstream_providers")} + if "ix_upstream_providers_base_url" in index_names: + op.drop_index("ix_upstream_providers_base_url", table_name="upstream_providers") + op.create_unique_constraint( + "uq_upstream_providers_base_url_api_key", + "upstream_providers", + ["base_url", "api_key"], + ) + + +def downgrade() -> None: + conn = op.get_bind() + if conn.dialect.name == "sqlite": + _recreate_table_sqlite(add_base_url_unique=True) + return + + op.drop_constraint( + "uq_upstream_providers_base_url_api_key", + "upstream_providers", + type_="unique", + ) + op.create_unique_constraint( + "uq_upstream_providers_base_url", + "upstream_providers", + ["base_url"], + ) + op.create_index( + "ix_upstream_providers_base_url", + "upstream_providers", + ["base_url"], + unique=True, + ) diff --git a/pyproject.toml b/pyproject.toml index 04dfb147..71f69194 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "routstr" -version = "0.2.2" +version = "0.3.0" description = "Payment proxy for your LLM endpoint using cashu and nostr." readme = "README.md" requires-python = ">=3.11" diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 537f07f5..7dffcae1 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -84,93 +84,26 @@ def get_provider_penalty(provider: "BaseUpstreamProvider") -> float: return penalty -def should_prefer_model( - candidate_model: "Model", - candidate_provider: "BaseUpstreamProvider", - current_model: "Model", - current_provider: "BaseUpstreamProvider", - alias: str, -) -> bool: - """Determine if candidate model should replace current model for an alias. - - This is the core decision function for model prioritization. It considers: - 1. Alias matching quality (exact match vs. canonical slug match) - 2. Model cost (lower is better) - 3. Provider penalties (e.g., slight preference against OpenRouter) - - Args: - candidate_model: The new model being considered - candidate_provider: Provider offering the candidate model - current_model: The currently selected model for this alias - current_provider: Provider offering the current model - alias: The model alias being mapped - - Returns: - True if candidate should replace current, False otherwise - """ - - def get_base_model_id(model_id: str) -> str: - """Get base model ID by removing provider prefix.""" - return model_id.split("/", 1)[1] if "/" in model_id else model_id - - def alias_priority(model: "Model") -> int: - """Rank how strong the mapping of alias->model is. - - Highest priority when alias exactly equals the model ID without provider prefix. - Next when alias equals canonical slug without prefix. Otherwise lowest. - """ - model_base = get_base_model_id(model.id) - if model_base == alias: - return 3 - if model.canonical_slug: - canonical_base = get_base_model_id(model.canonical_slug) - if canonical_base == alias: - return 2 - return 1 - - candidate_alias_priority = alias_priority(candidate_model) - current_alias_priority = alias_priority(current_model) - - # If candidate has better alias match, prefer it regardless of cost - if candidate_alias_priority > current_alias_priority: - return True - - # If current has better alias match, keep it regardless of cost - if current_alias_priority > candidate_alias_priority: - return False - - # Same alias priority - compare costs - candidate_cost = calculate_model_cost_score(candidate_model) - current_cost = calculate_model_cost_score(current_model) - - # Apply provider penalties - candidate_adjusted = candidate_cost * get_provider_penalty(candidate_provider) - current_adjusted = current_cost * get_provider_penalty(current_provider) - - # Prefer lower adjusted cost - should_replace = candidate_adjusted < current_adjusted - - return should_replace - - def create_model_mappings( upstreams: list["BaseUpstreamProvider"], overrides_by_id: dict[str, tuple], disabled_model_ids: set[str], -) -> tuple[dict[str, "Model"], dict[str, "BaseUpstreamProvider"], dict[str, "Model"]]: +) -> tuple[ + dict[str, "Model"], dict[str, list["BaseUpstreamProvider"]], dict[str, "Model"] +]: """Create optimal model mappings based on cost and provider preferences. This is the main entry point for the algorithm. It processes all upstream providers and creates three mappings based on cost optimization: 1. model_instances: alias -> Model (all model aliases mapped to their Model objects) - 2. provider_map: alias -> UpstreamProvider (which provider to use for each alias) + 2. provider_map: alias -> List[UpstreamProvider] (sorted list of providers for each alias) 3. unique_models: base_id -> Model (unique models without provider prefixes) The algorithm: - Processes non-OpenRouter providers first (they're typically cheaper) - Then processes OpenRouter models (they can still win if cheaper) - - For each model alias, uses should_prefer_model() to select the best provider + - For each model alias, collects all candidates and sorts them by priority and cost. Args: upstreams: List of all upstream provider instances @@ -183,8 +116,7 @@ def create_model_mappings( from .payment.models import _row_to_model from .upstream.helpers import resolve_model_alias - model_instances: dict[str, "Model"] = {} - provider_map: dict[str, "BaseUpstreamProvider"] = {} + candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {} unique_models: dict[str, "Model"] = {} # Separate OpenRouter from other providers @@ -202,24 +134,14 @@ def create_model_mappings( """Get base model ID by removing provider prefix.""" return model_id.split("/", 1)[1] if "/" in model_id else model_id - def _maybe_set_alias( + def _add_candidate( alias: str, model: "Model", provider: "BaseUpstreamProvider" ) -> None: - """Set alias to model/provider if not set or if new model is preferred.""" + """Add candidate model/provider for an alias.""" alias_lower = alias.lower() - existing_model = model_instances.get(alias_lower) - if not existing_model: - # No existing mapping, set it - model_instances[alias_lower] = model - provider_map[alias_lower] = provider - else: - # Check if candidate should replace existing - existing_provider = provider_map[alias_lower] - if should_prefer_model( - model, provider, existing_model, existing_provider, alias - ): - model_instances[alias_lower] = model - provider_map[alias_lower] = provider + if alias_lower not in candidates: + candidates[alias_lower] = [] + candidates[alias_lower].append((model, provider)) def process_provider_models( upstream: "BaseUpstreamProvider", is_openrouter: bool = False @@ -266,21 +188,55 @@ def create_model_mappings( # Try to set each alias for alias in aliases: - _maybe_set_alias(alias, model_to_use, upstream) + _add_candidate(alias, model_to_use, upstream) - # Process non-OpenRouter providers first (they're typically cheaper) + # Process non-OpenRouter providers first for upstream in other_upstreams: process_provider_models(upstream, is_openrouter=False) - # Process OpenRouter last - models only win if they're cheaper or better matched + # Process OpenRouter last if openrouter: process_provider_models(openrouter, is_openrouter=True) - # Log provider distribution + # Sort candidates and build final maps + model_instances: dict[str, "Model"] = {} + provider_map: dict[str, list["BaseUpstreamProvider"]] = {} + + def alias_priority(model: "Model", alias: str) -> int: + """Rank how strong the mapping of alias->model is.""" + model_base = get_base_model_id(model.id) + if model_base == alias: + return 3 + if model.canonical_slug: + canonical_base = get_base_model_id(model.canonical_slug) + if canonical_base == alias: + return 2 + return 1 + + for alias, items in candidates.items(): + # Sort key: (priority DESC, cost ASC) + # Using negative cost for DESC sort overall to keep high priority first + def sort_key(item: tuple["Model", "BaseUpstreamProvider"]) -> tuple[int, float]: + model, provider = item + priority = alias_priority(model, alias) + cost = calculate_model_cost_score(model) + penalty = get_provider_penalty(provider) + adjusted_cost = cost * penalty + return (priority, -adjusted_cost) + + items.sort(key=sort_key, reverse=True) + + best_model, best_provider = items[0] + model_instances[alias] = best_model + provider_map[alias] = [p for _, p in items] + + # Log provider distribution (using top provider for stats) provider_counts: dict[str, int] = {} - for provider in provider_map.values(): - provider_name = getattr(provider, "upstream_name", "unknown") - provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1 + for providers in provider_map.values(): + if providers: + provider = providers[0] + provider_name = getattr(provider, "upstream_name", "unknown") + provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1 logger.debug( f"Updated model mappings with ({len(unique_models)} unique models and {len(model_instances)} aliases)", diff --git a/routstr/auth.py b/routstr/auth.py index b3be04b2..1f869886 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -286,30 +286,55 @@ async def validate_bearer_key( ) +async def get_billing_key(key: ApiKey, session: AsyncSession) -> ApiKey: + """Returns the key that should be charged for the request.""" + if key.parent_key_hash: + parent = await session.get(ApiKey, key.parent_key_hash) + if parent: + # We want to keep the total_requests and total_spent on the child key + # but use the balance and reserved_balance of the parent. + # However, pay_for_request updates reserved_balance and total_requests. + # To stay simple, we charge the parent's balance and update parent's total_requests. + return parent + else: + logger.error( + "Parent key not found for child key", + extra={ + "child_key_hash": key.hashed_key[:8] + "...", + "parent_key_hash": key.parent_key_hash[:8] + "...", + }, + ) + return key + + async def pay_for_request( key: ApiKey, cost_per_request: int, session: AsyncSession ) -> int: """Process payment for a request.""" + billing_key = await get_billing_key(key, session) + logger.info( "Processing payment for request", extra={ "key_hash": key.hashed_key[:8] + "...", - "current_balance": key.balance, + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "current_balance": billing_key.balance, "required_cost": cost_per_request, - "sufficient_balance": key.balance >= cost_per_request, + "sufficient_balance": billing_key.balance >= cost_per_request, }, ) - if key.total_balance < cost_per_request: + if billing_key.total_balance < cost_per_request: logger.warning( "Insufficient balance for request", extra={ "key_hash": key.hashed_key[:8] + "...", - "balance": key.balance, - "reserved_balance": key.reserved_balance, + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "balance": billing_key.balance, + "reserved_balance": billing_key.reserved_balance, "required": cost_per_request, - "shortfall": cost_per_request - key.total_balance, + "shortfall": cost_per_request - billing_key.total_balance, }, ) @@ -317,7 +342,7 @@ async def pay_for_request( status_code=402, detail={ "error": { - "message": f"Insufficient balance: {cost_per_request} mSats required. {key.total_balance} available. (reserved: {key.reserved_balance})", + "message": f"Insufficient balance: {cost_per_request} mSats required. {billing_key.total_balance} available. (reserved: {billing_key.reserved_balance})", "type": "insufficient_quota", "code": "insufficient_balance", } @@ -328,15 +353,16 @@ async def pay_for_request( "Charging base cost for request", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "cost": cost_per_request, - "balance_before": key.balance, + "balance_before": billing_key.balance, }, ) # Charge the base cost for the request atomically to avoid race conditions stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .where(col(ApiKey.balance) - col(ApiKey.reserved_balance) >= cost_per_request) .values( reserved_balance=col(ApiKey.reserved_balance) + cost_per_request, @@ -344,6 +370,16 @@ async def pay_for_request( ) ) result = await session.exec(stmt) # type: ignore[call-overload] + + # Also increment total_requests on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_requests=col(ApiKey.total_requests) + 1) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount == 0: @@ -351,8 +387,9 @@ async def pay_for_request( "Concurrent request depleted balance", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "required_cost": cost_per_request, - "current_balance": key.balance, + "current_balance": billing_key.balance, }, ) @@ -361,23 +398,26 @@ async def pay_for_request( status_code=402, detail={ "error": { - "message": f"Insufficient balance: {cost_per_request} mSats required. {key.balance} available.", + "message": f"Insufficient balance: {cost_per_request} mSats required. {billing_key.balance} available.", "type": "insufficient_quota", "code": "insufficient_balance", } }, ) - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) logger.info( "Payment processed successfully", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "charged_amount": cost_per_request, - "new_balance": key.balance, - "total_spent": key.total_spent, - "total_requests": key.total_requests, + "new_balance": billing_key.balance, + "total_spent": billing_key.total_spent, + "total_requests": billing_key.total_requests, }, ) @@ -387,9 +427,11 @@ async def pay_for_request( async def revert_pay_for_request( key: ApiKey, session: AsyncSession, cost_per_request: int ) -> None: + billing_key = await get_billing_key(key, session) + stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - cost_per_request, total_requests=col(ApiKey.total_requests) - 1, @@ -397,27 +439,40 @@ async def revert_pay_for_request( ) result = await session.exec(stmt) # type: ignore[call-overload] + + # Also decrement total_requests on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_requests=col(ApiKey.total_requests) - 1) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount == 0: logger.error( "Failed to revert payment - insufficient reserved balance", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "cost_to_revert": cost_per_request, - "current_reserved_balance": key.reserved_balance, + "current_reserved_balance": billing_key.reserved_balance, }, ) raise HTTPException( status_code=402, detail={ "error": { - "message": f"failed to revert request payment: {cost_per_request} mSats required. {key.balance} available.", + "message": f"failed to revert request payment: {cost_per_request} mSats required. {billing_key.balance} available.", "type": "payment_error", "code": "payment_error", } }, ) - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) async def adjust_payment_for_tokens( @@ -428,15 +483,17 @@ async def adjust_payment_for_tokens( This is called after the initial payment and the upstream request is complete. Returns cost data to be included in the response. """ + billing_key = await get_billing_key(key, session) model = response_data.get("model", "unknown") logger.debug( "Starting payment adjustment for tokens", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "model": model, "deducted_max_cost": deducted_max_cost, - "current_balance": key.balance, + "current_balance": billing_key.balance, "has_usage": "usage" in response_data, }, ) @@ -446,8 +503,10 @@ async def adjust_payment_for_tokens( try: release_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .values(reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) + .values( + reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost + ) ) await session.exec(release_stmt) # type: ignore[call-overload] await session.commit() @@ -455,13 +514,18 @@ async def adjust_payment_for_tokens( "Released reservation without charging (fallback)", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "deducted_max_cost": deducted_max_cost, }, ) except Exception as e: logger.error( "Failed to release reservation in fallback", - extra={"error": str(e), "key_hash": key.hashed_key[:8] + "..."}, + extra={ + "error": str(e), + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + }, ) match await calculate_cost(response_data, deducted_max_cost, session): @@ -470,6 +534,7 @@ async def adjust_payment_for_tokens( "Using max cost data (no token adjustment)", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "model": model, "max_cost": cost.total_msats, }, @@ -477,7 +542,7 @@ async def adjust_payment_for_tokens( # Finalize by releasing reservation and charging max cost finalize_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, balance=col(ApiKey.balance) - cost.total_msats, @@ -485,27 +550,41 @@ async def adjust_payment_for_tokens( ) ) result = await session.exec(finalize_stmt) # type: ignore[call-overload] + + # Also update total_spent on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_spent=col(ApiKey.total_spent) + cost.total_msats) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount == 0: logger.error( "Failed to finalize max-cost payment - retrying reservation release", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "deducted_max_cost": deducted_max_cost, - "current_reserved_balance": key.reserved_balance, + "current_reserved_balance": billing_key.reserved_balance, "total_cost": cost.total_msats, "model": model, }, ) await release_reservation_only() else: - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) logger.info( "Max cost payment finalized", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "charged_amount": cost.total_msats, - "new_balance": key.balance, + "new_balance": billing_key.balance, "model": model, }, ) @@ -521,6 +600,7 @@ async def adjust_payment_for_tokens( "Calculated token-based cost", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "model": model, "token_cost": cost.total_msats, "deducted_max_cost": deducted_max_cost, @@ -533,11 +613,15 @@ async def adjust_payment_for_tokens( if cost_difference == 0: logger.debug( "Finalizing with exact reserved cost", - extra={"key_hash": key.hashed_key[:8] + "...", "model": model}, + extra={ + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "model": model, + }, ) finalize_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, @@ -546,8 +630,20 @@ async def adjust_payment_for_tokens( ) ) await session.exec(finalize_stmt) # type: ignore[call-overload] + + # Also update total_spent on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_spent=col(ApiKey.total_spent) + total_cost_msats) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) return cost.dict() # this should never happen why do we handle this??? @@ -557,16 +653,17 @@ async def adjust_payment_for_tokens( "Additional charge required for token usage", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "additional_charge": cost_difference, - "current_balance": key.balance, - "sufficient_balance": key.balance >= cost_difference, + "current_balance": billing_key.balance, + "sufficient_balance": billing_key.balance >= cost_difference, "model": model, }, ) finalize_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, @@ -575,18 +672,31 @@ async def adjust_payment_for_tokens( ) ) result = await session.exec(finalize_stmt) # type: ignore[call-overload] + + # Also update total_spent on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_spent=col(ApiKey.total_spent) + total_cost_msats) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount: cost.total_msats = total_cost_msats - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) logger.info( "Finalized payment with additional charge", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "charged_amount": total_cost_msats, - "new_balance": key.balance, + "new_balance": billing_key.balance, "model": model, }, ) @@ -595,6 +705,7 @@ async def adjust_payment_for_tokens( "Failed to finalize additional charge - releasing reservation", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "attempted_charge": total_cost_msats, "model": model, }, @@ -607,15 +718,16 @@ async def adjust_payment_for_tokens( "Refunding excess payment", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "refund_amount": refund, - "current_balance": key.balance, + "current_balance": billing_key.balance, "model": model, }, ) refund_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, @@ -624,6 +736,16 @@ async def adjust_payment_for_tokens( ) ) result = await session.exec(refund_stmt) # type: ignore[call-overload] + + # Also update total_spent on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_spent=col(ApiKey.total_spent) + total_cost_msats) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount == 0: @@ -631,8 +753,9 @@ async def adjust_payment_for_tokens( "Failed to finalize payment - releasing reservation", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "deducted_max_cost": deducted_max_cost, - "current_reserved_balance": key.reserved_balance, + "current_reserved_balance": billing_key.reserved_balance, "total_cost": total_cost_msats, "model": model, }, @@ -640,14 +763,17 @@ async def adjust_payment_for_tokens( await release_reservation_only() else: cost.total_msats = total_cost_msats - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) logger.info( "Refund processed successfully", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "refunded_amount": refund, - "new_balance": key.balance, + "new_balance": billing_key.balance, "final_cost": cost.total_msats, "model": model, }, diff --git a/routstr/balance.py b/routstr/balance.py index 6697a5db..1a1e5850 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -32,16 +32,30 @@ async def get_key_from_header( ) -# TODO: remove this endpoint when frontend is updated -@router.get("/", include_in_schema=False) -async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict: +async def get_balance_info(key: ApiKey, session: AsyncSession) -> dict: + from .auth import get_billing_key + + billing_key = await get_billing_key(key, session) return { "api_key": "sk-" + key.hashed_key, - "balance": key.balance, - "reserved": key.reserved_balance, + "balance": billing_key.balance, + "reserved": billing_key.reserved_balance, + "is_child": key.parent_key_hash is not None, + "parent_key": "sk-" + key.parent_key_hash if key.parent_key_hash else None, + "total_requests": key.total_requests, + "total_spent": key.total_spent, } +# TODO: remove this endpoint when frontend is updated +@router.get("/", include_in_schema=False) +async def account_info( + key: ApiKey = Depends(get_key_from_header), + session: AsyncSession = Depends(get_session), +) -> dict: + return await get_balance_info(key, session) + + # TODO: Implement POST /v1/wallet/create endpoint # This endpoint should accept: # - cashu_token (required): The eCash token to deposit @@ -66,12 +80,11 @@ async def create_balance( @router.get("/info") -async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict: - return { - "api_key": "sk-" + key.hashed_key, - "balance": key.balance, - "reserved": key.reserved_balance, - } +async def wallet_info( + key: ApiKey = Depends(get_key_from_header), + session: AsyncSession = Depends(get_session), +) -> dict: + return await get_balance_info(key, session) class TopupRequest(BaseModel): @@ -85,6 +98,10 @@ async def topup_wallet_endpoint( key: ApiKey = Depends(get_key_from_header), session: AsyncSession = Depends(get_session), ) -> dict[str, int]: + from .auth import get_billing_key + + billing_key = await get_billing_key(key, session) + if topup_request is not None: cashu_token = topup_request.cashu_token if cashu_token is None: @@ -94,7 +111,7 @@ async def topup_wallet_endpoint( if len(cashu_token) < 10 or "cashu" not in cashu_token: raise HTTPException(status_code=400, detail="Invalid token format") try: - amount_msats = await credit_balance(cashu_token, key, session) + amount_msats = await credit_balance(cashu_token, billing_key, session) except ValueError as e: error_msg = str(e) if "already spent" in error_msg.lower(): @@ -155,6 +172,12 @@ async def refund_wallet_endpoint( 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 if key.refund_currency == "sat": @@ -228,6 +251,75 @@ async def donate(token: str, ref: str | None = None) -> str: return "Invalid token." +class ChildKeyRequest(BaseModel): + count: int + + +@router.post("/child-key") +async def create_child_key( + payload: ChildKeyRequest, + key: ApiKey = Depends(get_key_from_header), + session: AsyncSession = Depends(get_session), +) -> dict: + """Creates one or more child API keys that use the parent's balance.""" + # Log incoming request for debugging + logger.debug(f"Child key creation request: count={payload.count}") + + count = payload.count + if count < 1 or count > 50: + raise HTTPException(status_code=400, detail="Count must be between 1 and 50.") + + # 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_per_key = settings.child_key_cost + total_cost = cost_per_key * count + + if key.total_balance < total_cost: + raise HTTPException( + status_code=402, + detail=f"Insufficient balance to create {count} child keys. {total_cost} mSats required.", + ) + + # Deduct cost from parent + key.balance -= total_cost + key.total_spent += total_cost + session.add(key) + + # Generate new keys + import secrets + + new_keys = [] + for _ in range(count): + 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) + new_keys.append("sk-" + new_key_hash) + + await session.commit() + + response_data = { + "api_keys": new_keys, + "count": count, + "cost_msats": total_cost, + "cost_sats": total_cost // 1000, + "parent_balance": key.balance, + "parent_balance_sats": key.balance // 1000, + } + logger.debug(f"Child key creation response: {response_data}") + return response_data + + @router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], diff --git a/routstr/core/admin.py b/routstr/core/admin.py index d4f7ec59..8185591e 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -4,7 +4,6 @@ from datetime import datetime, timezone from pathlib import Path from fastapi import APIRouter, Depends, HTTPException, Query, Request -from fastapi.responses import HTMLResponse, RedirectResponse from pydantic import BaseModel from sqlmodel import select @@ -41,109 +40,6 @@ def require_admin_api(request: Request) -> None: raise HTTPException(status_code=403, detail="Unauthorized") -def is_admin_authenticated(request: Request) -> bool: - auth_header = request.headers.get("Authorization") - if auth_header and auth_header.startswith("Bearer "): - token = auth_header.split(" ", 1)[1] - expiry = admin_sessions.get(token) - if expiry and expiry > int(datetime.now(timezone.utc).timestamp()): - return True - - return False - - -@admin_router.get( - "/partials/balances", - dependencies=[Depends(require_admin_api)], - response_class=HTMLResponse, -) -async def partial_balances(request: Request) -> str: - ( - balance_details, - total_wallet_balance_sats, - total_user_balance_sats, - owner_balance, - ) = await fetch_all_balances() - # Provide JSON for client usage - # Embed a script tag to update balanceDetails and the UI markup - rows = "".join( - [ - f"""
-
{detail["mint_url"].replace("https://", "").replace("http://", "")} • {detail["unit"].upper()}
-
{detail["wallet_balance"] if not detail.get("error") else "error"}
-
{detail["user_balance"] if not detail.get("error") else "-"}
-
0 else ""}">{detail["owner_balance"] if not detail.get("error") else "-"}
-
""" - for detail in balance_details - if detail.get("wallet_balance", 0) > 0 or detail.get("error") - ] - ) - return f""" -

Cashu Wallet Balance

-
- Your Balance (Total) - {owner_balance} sats -
-
- Total Wallet - {total_wallet_balance_sats} sats -
-
- User Balance - {total_user_balance_sats} sats -
-

Your balance = Total wallet - User balance

-
-
-
Mint / Unit
-
Wallet
-
Users
-
Owner
-
- {rows} -
- - """ - - -@admin_router.get( - "/partials/apikeys", - dependencies=[Depends(require_admin_api)], - response_class=HTMLResponse, -) -async def partial_apikeys(request: Request) -> str: - async with create_session() as session: - result = await session.exec(select(ApiKey)) - api_keys = result.all() - - def fmt_time(ts: int | None) -> str: - if ts is None: - return "" - dt = datetime.fromtimestamp(ts, tz=timezone.utc) - return f"{ts} ({dt.strftime('%Y-%m-%d %H:%M:%S')} UTC)" - - rows = "".join( - [ - f"{key.hashed_key}{key.balance}{key.total_spent}{key.total_requests}{key.refund_address}{fmt_time(key.key_expiry_time)}" - for key in api_keys - ] - ) - return f""" -

Temporary Balances

- - - - - - - - - - {rows} -
Hashed KeyBalance (mSats)Total Spent (mSats)Total RequestsRefund AddressRefund Time
- """ - - @admin_router.get("/api/temporary-balances", dependencies=[Depends(require_admin_api)]) async def get_temporary_balances_api(request: Request) -> list[dict[str, object]]: async with create_session() as session: @@ -158,6 +54,7 @@ async def get_temporary_balances_api(request: Request) -> list[dict[str, object] "total_requests": key.total_requests, "refund_address": key.refund_address, "key_expiry_time": key.key_expiry_time, + "parent_key_hash": key.parent_key_hash, } for key in api_keys ] @@ -303,612 +200,6 @@ class WithdrawRequest(BaseModel): unit: str = "sat" -def login_form() -> str: - return """ - - - - - - -
-

🔐 Admin Login

-
- - -
-
- - - """ - - -def setup_form() -> str: - return """ - - - - - - -
-

🔧 Initial Admin Setup

-

Create a secure password for your admin dashboard.

-
- - - -
-
-
- - - """ - - -def info(content: str) -> str: - return f""" - - - - - -
-

{content}

-
- - - """ - - -def admin_auth() -> str: - admin_pw = settings.admin_password - if admin_pw == "": - return setup_form() - else: - return login_form() - - -async def dashboard(request: Request) -> str: - return ( - f""" - - - - - """ - + """ - - - """ - + """ - -

Admin Dashboard

- -
-
Loading balances…
-
- - - - - - - - - - - - - -
- Withdrawal Token: -
- -

Save this token! It represents your withdrawn balance.

-
- -
-

Temporary Balances

-
Loading API keys…
-
- - - """ - ) - - -@admin_router.get("/", response_class=HTMLResponse) -async def admin(request: Request) -> RedirectResponse: - return RedirectResponse("/") - - -@admin_router.get("/logs/{request_id}", response_class=HTMLResponse) -async def view_logs(request: Request, request_id: str) -> str: - if not is_admin_authenticated(request): - return admin_auth() - - logger.info(f"Investigating logs for request_id: {request_id}") - - # Search for log entries with this request_id - log_entries = [] - logs_dir = Path("logs") - - if logs_dir.exists(): - # Get all log files sorted by modification time (most recent first) - log_files = sorted( - logs_dir.glob("*.log"), key=lambda x: x.stat().st_mtime, reverse=True - ) - - for log_file in log_files[:7]: # Check last 7 days of logs - try: - with open(log_file, "r") as f: - for line in f: - if request_id in line: - try: - # Parse JSON log entry - log_data = json.loads(line.strip()) - log_entries.append(log_data) - except json.JSONDecodeError: - # If not JSON, include raw line - log_entries.append({"raw": line.strip()}) - except Exception as e: - logger.error(f"Error reading log file {log_file}: {e}") - - # Sort entries by timestamp if available - log_entries.sort(key=lambda x: x.get("asctime", ""), reverse=False) - - # Format log entries for display - formatted_logs = [] - for entry in log_entries: - if "raw" in entry: - formatted_logs.append(f'
{entry["raw"]}
') - else: - # Format JSON log entry - timestamp = entry.get("asctime", "Unknown time") - level = entry.get("levelname", "INFO") - message = entry.get("message", "") - pathname = entry.get("pathname", "") - lineno = entry.get("lineno", "") - - # Extract additional fields - extra_fields = { - k: v - for k, v in entry.items() - if k - not in [ - "asctime", - "levelname", - "message", - "pathname", - "lineno", - "name", - "version", - "request_id", - ] - } - - level_class = level.lower() - formatted_entry = f""" -
-
- {timestamp} - [{level}] - {pathname}:{lineno} -
-
{message}
- """ - - if extra_fields: - formatted_entry += '
' - for key, value in extra_fields.items(): - formatted_entry += f'
{key}: {json.dumps(value) if isinstance(value, (dict, list)) else value}
' - formatted_entry += "
" - - formatted_entry += "
" - formatted_logs.append(formatted_entry) - - return ( - f""" - - - - - - """ - + f""" - - ← Back to Dashboard -

Log Investigation

-
- Request ID: {request_id} -
-
- {"".join(formatted_logs) if formatted_logs else '
No log entries found for this Request ID
'} -
-

- Found {len(log_entries)} log entries • Searched last 7 days of logs -

- - - """ - ) - - @admin_router.post("/withdraw", dependencies=[Depends(require_admin_api)]) async def withdraw( request: Request, withdraw_request: WithdrawRequest @@ -942,546 +233,6 @@ async def withdraw( return {"token": token} -DASHBOARD_MODELS_JS: str = """ - -""" - - -def models_page() -> str: - return ( - f""" - - - - {DASHBOARD_MODELS_JS} - - """ - + """ - - ← Back to Dashboard -

Models

- -
-

Models Table

-
- - -
- - - - - - - - - - -
ID
Loading…
-
- - -
-
- - - - - - - - - """ - ) - - class ModelCreate(BaseModel): id: str name: str @@ -1498,908 +249,6 @@ class ModelCreate(BaseModel): enabled: bool = True -@admin_router.get("/models", response_class=HTMLResponse) -async def admin_models(request: Request) -> str: - if is_admin_authenticated(request): - return models_page() - return admin_auth() - - -UPSTREAM_PROVIDERS_JS: str = """ - -""" - - -def upstream_providers_page() -> str: - return ( - f""" - - - - {UPSTREAM_PROVIDERS_JS} - - """ - + """ - - ← Back to Dashboard -

Upstream Providers

- -
-

Providers

-
- -
- - - - - - - - - - - - -
TypeBase URLStatusActions
Loading…
-
- - - - - - - - - - - """ - ) - - -@admin_router.get("/upstream-providers", response_class=HTMLResponse) -async def admin_upstream_providers(request: Request) -> str: - if is_admin_authenticated(request): - return upstream_providers_page() - return admin_auth() - - @admin_router.post( "/api/upstream-providers/{provider_id}/models", dependencies=[Depends(require_admin_api)], @@ -2515,7 +364,7 @@ async def get_provider_model(provider_id: int, model_id: str) -> dict[str, objec status_code=404, detail="Model not found for this provider" ) return _row_to_model( - row, apply_provider_fee=True, provider_fee=provider.provider_fee + row, apply_provider_fee=False, provider_fee=provider.provider_fee ).dict() # type: ignore @@ -2553,6 +402,91 @@ async def delete_all_provider_models(provider_id: int) -> dict[str, object]: return {"ok": True, "deleted": len(rows)} +class BatchOverrideRequest(BaseModel): + models: list[ModelCreate] + + +@admin_router.post( + "/api/upstream-providers/{provider_id}/batch-override", + dependencies=[Depends(require_admin_api)], +) +async def batch_override_provider_models( + provider_id: int, payload: BatchOverrideRequest +) -> dict[str, object]: + """Batch override models for a specific provider.""" + logger.info( + f"BATCH_OVERRIDE called: provider_id={provider_id}, count={len(payload.models)}" + ) + + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + + overridden_count = 0 + + for model_data in payload.models: + # Try to get existing model regardless of whether it's enabled or not + existing_row = await session.get(ModelRow, (model_data.id, provider_id)) + + if existing_row: + # Update existing + existing_row.name = model_data.name + existing_row.description = model_data.description + existing_row.created = int(model_data.created) + existing_row.context_length = int(model_data.context_length) + existing_row.architecture = json.dumps(model_data.architecture) + existing_row.pricing = json.dumps(model_data.pricing) + existing_row.sats_pricing = None + existing_row.per_request_limits = ( + json.dumps(model_data.per_request_limits) + if model_data.per_request_limits is not None + else None + ) + existing_row.top_provider = ( + json.dumps(model_data.top_provider) if model_data.top_provider else None + ) + existing_row.canonical_slug = model_data.canonical_slug + existing_row.alias_ids = ( + json.dumps(model_data.alias_ids) if model_data.alias_ids else None + ) + existing_row.enabled = model_data.enabled + session.add(existing_row) + else: + # Create new + row = ModelRow( + id=model_data.id, + name=model_data.name, + description=model_data.description, + created=int(model_data.created), + context_length=int(model_data.context_length), + architecture=json.dumps(model_data.architecture), + pricing=json.dumps(model_data.pricing), + sats_pricing=None, + per_request_limits=( + json.dumps(model_data.per_request_limits) + if model_data.per_request_limits is not None + else None + ), + top_provider=( + json.dumps(model_data.top_provider) if model_data.top_provider else None + ), + canonical_slug=model_data.canonical_slug, + alias_ids=( + json.dumps(model_data.alias_ids) if model_data.alias_ids else None + ), + upstream_provider_id=provider_id, + enabled=model_data.enabled, + ) + session.add(row) + + overridden_count += 1 + + await session.commit() + + await refresh_model_maps() + return {"ok": True, "count": overridden_count, "message": f"Successfully batch overridden {overridden_count} models"} + class UpstreamProviderCreate(BaseModel): provider_type: str base_url: str @@ -2726,7 +660,10 @@ async def get_provider_models(provider_id: int) -> dict[str, object]: raise HTTPException(status_code=404, detail="Provider not found") db_models = await list_models( - session=session, upstream_id=provider_id, include_disabled=True + session=session, + upstream_id=provider_id, + include_disabled=True, + apply_fees=False, ) upstream_models = [] @@ -2734,10 +671,7 @@ async def get_provider_models(provider_id: int) -> dict[str, object]: if upstream_instance: try: raw_models = await upstream_instance.fetch_models() - upstream_models = [ - upstream_instance._apply_provider_fee_to_model(m) - for m in raw_models - ] + upstream_models = raw_models except Exception as e: logger.error( f"Failed to fetch models from {provider.provider_type}: {e}" @@ -2950,72 +884,6 @@ async def get_openrouter_presets() -> list[dict[str, object]]: return models_data -DASHBOARD_CSS: str = """ -* { margin: 0; padding: 0; box-sizing: border-box; } -body { font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', sans-serif; background: #f5f7fa; color: #2c3e50; line-height: 1.6; padding: 2rem; } -h1, h2 { margin-bottom: 1rem; color: #1a202c; } -h1 { font-size: 2rem; } -h2 { font-size: 1.5rem; margin-top: 2rem; } -p { margin-bottom: 0.5rem; color: #4a5568; } -table { width: 100%; border-collapse: collapse; background: white; border-radius: 8px; overflow: hidden; box-shadow: 0 1px 3px rgba(0,0,0,0.1); margin-top: 1rem; } -th { background: #4a5568; color: white; font-weight: 600; padding: 12px; text-align: left; } -td { padding: 12px; border-bottom: 1px solid #e2e8f0; } -tr:hover { background: #f7fafc; } -button { padding: 10px 20px; cursor: pointer; background: #4299e1; color: white; border: none; border-radius: 6px; font-weight: 600; margin-right: 10px; transition: all 0.2s; } -button:hover { background: #3182ce; transform: translateY(-1px); box-shadow: 0 2px 4px rgba(0,0,0,0.1); } -button:disabled { background: #a0aec0; cursor: not-allowed; transform: none; } -.refresh-btn { background: #48bb78; } -.refresh-btn:hover { background: #38a169; } -.investigate-btn { background: #4299e1; } -.balance-card { background: white; padding: 2rem; border-radius: 8px; box-shadow: 0 1px 3px rgba(0,0,0,0.1); margin-bottom: 2rem; } -.balance-item { display: flex; justify-content: space-between; margin-bottom: 1rem; } -.balance-label { color: #718096; } -.balance-value { font-size: 1.5rem; font-weight: 700; color: #2d3748; } -.balance-primary { color: #48bb78; } -.currency-grid { margin-top: 1rem; font-size: 0.9rem; } -.currency-row { display: grid; grid-template-columns: 2fr 1fr 1fr 1fr; gap: 0.5rem; padding: 0.4rem 0; border-bottom: 1px solid #f0f0f0; align-items: center; } -.currency-row:last-child { border-bottom: none; } -.currency-header { font-weight: 600; color: #4a5568; border-bottom: 2px solid #e2e8f0; padding-bottom: 0.5rem; } -.mint-name { color: #2d3748; font-size: 0.85rem; word-break: break-all; } -.balance-num { text-align: right; font-family: monospace; } -.owner-positive { color: #22c55e; } -.error-row { color: #dc2626; font-style: italic; } -#token-result { margin-top: 20px; padding: 20px; background: #e6fffa; border: 1px solid #38b2ac; border-radius: 8px; display: none; } -#token-text { font-family: 'Monaco', monospace; font-size: 13px; background: #2d3748; color: #68d391; padding: 15px; border-radius: 6px; margin: 10px 0; word-break: break-all; } -.copy-btn { background: #38a169; padding: 6px 12px; font-size: 14px; } -.copy-btn:hover { background: #2f855a; } -.modal { display: none; position: fixed; z-index: 1000; left: 0; top: 0; width: 100%; height: 100%; background: rgba(0,0,0,0.5); backdrop-filter: blur(4px); } -.modal-content { background: white; margin: 5% auto; padding: 0.75rem 1rem 2.25rem; width: 90%; max-width: 720px; max-height: 85vh; overflow-y: auto; border-radius: 12px; box-shadow: 0 20px 25px -5px rgba(0,0,0,0.1); animation: slideIn 0.3s ease; } -@keyframes slideIn { from { transform: translateY(-20px); opacity: 0; } to { transform: translateY(0); opacity: 1; } } -.close { color: #a0aec0; float: right; font-size: 28px; font-weight: bold; cursor: pointer; margin: -10px -10px 0 0; } -.close:hover { color: #2d3748; } -input[type="number"], input[type="text"], input[type="password"], select { width: 100%; padding: 10px; margin: 10px 0; border: 2px solid #e2e8f0; border-radius: 6px; font-size: 16px; transition: border 0.2s; } -input[type="number"]:focus, input[type="text"]:focus, input[type="password"]:focus, select:focus { outline: none; border-color: #4299e1; } -.warning { color: #e53e3e; font-weight: 600; margin: 10px 0; padding: 10px; background: #fff5f5; border-radius: 6px; } -""" - - -LOGS_CSS: str = """ -body { font-family: Arial, sans-serif; margin: 20px; background-color: #f5f5f5; } -h1 { color: #333; } -.back-btn { padding: 8px 16px; background-color: #007bff; color: white; border: none; border-radius: 4px; cursor: pointer; text-decoration: none; display: inline-block; margin-bottom: 20px; } -.back-btn:hover { background-color: #0056b3; } -.log-container { background-color: white; border: 1px solid #ddd; border-radius: 8px; padding: 20px; max-height: 80vh; overflow-y: auto; } -.log-entry { margin-bottom: 15px; padding: 10px; border: 1px solid #e0e0e0; border-radius: 4px; font-family: 'Courier New', monospace; font-size: 12px; background-color: #f9f9f9; } -.log-entry.log-error { background-color: #fee; border-color: #fcc; } -.log-entry.log-warning { background-color: #ffc; border-color: #ff9; } -.log-entry.log-debug, .log-entry.log-trace { background-color: #f0f0f0; border-color: #ccc; } -.log-header { margin-bottom: 5px; color: #666; } -.log-timestamp { color: #0066cc; } -.log-level { font-weight: bold; } -.log-message { margin: 5px 0; color: #333; } -.log-extra { margin-top: 5px; padding-top: 5px; border-top: 1px solid #e0e0e0; } -.log-field { margin: 2px 0; color: #666; word-break: break-all; } -.no-logs { text-align: center; color: #666; padding: 40px; } -.request-id-display { background-color: #e9ecef; padding: 10px; border-radius: 4px; margin-bottom: 20px; font-family: monospace; } -""" - - @admin_router.get("/api/usage/metrics", dependencies=[Depends(require_admin_api)]) async def get_usage_metrics( request: Request, diff --git a/routstr/core/db.py b/routstr/core/db.py index 4c236d3a..bbcfc6fe 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -5,6 +5,7 @@ from typing import AsyncGenerator from alembic import command from alembic.config import Config +from sqlalchemy import UniqueConstraint from sqlalchemy.ext.asyncio.engine import create_async_engine from sqlmodel import Field, Relationship, SQLModel, func, select, update from sqlmodel.ext.asyncio.session import AsyncSession @@ -47,6 +48,9 @@ class ApiKey(SQLModel, table=True): # type: ignore default=None, description="Currency of the cashu-token", ) + parent_key_hash: str | None = Field( + default=None, foreign_key="api_keys.hashed_key", index=True + ) @property def total_balance(self) -> int: @@ -108,11 +112,14 @@ class LightningInvoice(SQLModel, table=True): # type: ignore class UpstreamProviderRow(SQLModel, table=True): # type: ignore __tablename__ = "upstream_providers" + __table_args__ = ( + UniqueConstraint("base_url", "api_key", name="uq_upstream_providers_base_url_api_key"), + ) id: int | None = Field(default=None, primary_key=True) provider_type: str = Field( description="Provider type: custom, openai, anthropic, azure, openrouter, etc." ) - base_url: str = Field(unique=True, description="Base URL of the upstream API") + base_url: str = Field(description="Base URL of the upstream API") api_key: str = Field(description="API key for the upstream provider") api_version: str | None = Field( default=None, description="API version for Azure OpenAI" diff --git a/routstr/core/exceptions.py b/routstr/core/exceptions.py index e74d3e05..64111047 100644 --- a/routstr/core/exceptions.py +++ b/routstr/core/exceptions.py @@ -6,6 +6,15 @@ from .logging import get_logger logger = get_logger(__name__) +class UpstreamError(Exception): + """Exception raised when an upstream provider fails.""" + + def __init__(self, message: str, status_code: int = 502): + self.message = message + self.status_code = status_code + super().__init__(message) + + async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse: """Handle HTTP exceptions and include request ID in response.""" request_id = getattr(request.state, "request_id", "unknown") diff --git a/routstr/core/main.py b/routstr/core/main.py index 9a0741ef..f45924b6 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -11,12 +11,9 @@ from fastapi.staticfiles import StaticFiles from starlette.exceptions import HTTPException from ..balance import balance_router, deprecated_wallet_router -from ..discovery import providers_cache_refresher, providers_router -from ..nip91 import announce_provider -from ..payment.models import ( - models_router, - update_sats_pricing, -) +from ..nostr import announce_provider, providers_cache_refresher +from ..nostr.discovery import providers_router +from ..payment.models import models_router, update_sats_pricing from ..payment.price import update_prices_periodically from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically from ..wallet import periodic_payout @@ -33,9 +30,9 @@ setup_logging() logger = get_logger(__name__) if os.getenv("VERSION_SUFFIX") is not None: - __version__ = f"0.2.2-{os.getenv('VERSION_SUFFIX')}" + __version__ = f"0.3.0-{os.getenv('VERSION_SUFFIX')}" else: - __version__ = "0.2.2" + __version__ = "0.3.0" @asynccontextmanager @@ -191,6 +188,7 @@ async def info() -> dict: "mints": global_settings.cashu_mints, "http_url": global_settings.http_url, "onion_url": global_settings.onion_url, + "child_key_cost_msats": global_settings.child_key_cost, } diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 53cd94fa..cd38c0c9 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -52,6 +52,7 @@ class Settings(BaseSettings): exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE") upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE") tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE") + child_key_cost: int = Field(default=1000, env="CHILD_KEY_COST") # Minimum per-request charge in millisatoshis when model pricing is free/zero min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT") reset_reserved_balance_on_startup: bool = Field( @@ -142,7 +143,7 @@ def resolve_bootstrap() -> Settings: pass if not base.onion_url: try: - from ..nip91 import discover_onion_url_from_tor # type: ignore + from ..nostr.listing import discover_onion_url_from_tor # type: ignore discovered = discover_onion_url_from_tor() if discovered: diff --git a/routstr/nostr/__init__.py b/routstr/nostr/__init__.py new file mode 100644 index 00000000..b19039a7 --- /dev/null +++ b/routstr/nostr/__init__.py @@ -0,0 +1,4 @@ +from .discovery import providers_cache_refresher +from .listing import announce_provider + +__all__ = ["providers_cache_refresher", "announce_provider"] diff --git a/routstr/discovery.py b/routstr/nostr/discovery.py similarity index 98% rename from routstr/discovery.py rename to routstr/nostr/discovery.py index 94375b63..268b4c7c 100644 --- a/routstr/discovery.py +++ b/routstr/nostr/discovery.py @@ -8,8 +8,8 @@ import httpx import websockets from fastapi import APIRouter, HTTPException -from .core.logging import get_logger -from .core.settings import settings +from ..core.logging import get_logger +from ..core.settings import settings logger = get_logger(__name__) @@ -72,8 +72,6 @@ async def query_nostr_relay_for_providers( elif data[0] == "NOTICE": try: msg = str(data[1]) - if len(msg) > 200: - msg = msg[:200] + "..." logger.debug(f"Relay notice: {msg}") except Exception: logger.debug("Relay notice received") diff --git a/routstr/nip91.py b/routstr/nostr/listing.py similarity index 91% rename from routstr/nip91.py rename to routstr/nostr/listing.py index d14c2c6e..a512add7 100644 --- a/routstr/nip91.py +++ b/routstr/nostr/listing.py @@ -1,6 +1,6 @@ #!/usr/bin/env python3 """ -NIP-91: Routstr Provider Discoverability Implementation +Listing: Routstr Provider Discoverability Implementation Automatically announces this Routstr proxy instance to Nostr relays. """ @@ -18,15 +18,15 @@ from nostr.key import PrivateKey from nostr.message_type import ClientMessageType from nostr.relay_manager import RelayManager -from .core import get_logger -from .core.settings import settings +from ..core import get_logger +from ..core.settings import settings logger = get_logger(__name__) def get_app_version() -> str | None: try: - from .core.main import __version__ as imported_version + from ..core.main import __version__ as imported_version return imported_version except Exception: @@ -71,7 +71,7 @@ def nsec_to_keypair(nsec: str) -> tuple[str, str] | None: return None -def create_nip91_event( +def create_listing_event( private_key_hex: str, provider_id: str, endpoint_urls: list[str], @@ -80,7 +80,7 @@ def create_nip91_event( metadata: dict[str, Any] | None = None, ) -> dict[str, Any]: """ - Create a NIP-91 compliant provider announcement event (kind:38421). + Create a listing provider announcement event (kind:38421). Args: private_key_hex: 32-byte hex private key for signing @@ -164,14 +164,14 @@ def events_semantically_equal(a: dict[str, Any], b: dict[str, Any]) -> bool: return True -async def query_nip91_events( +async def query_listing_events( relay_url: str, pubkey: str, provider_id: str | None = None, timeout: int = 30, ) -> tuple[list[dict[str, Any]], bool]: """ - Query a Nostr relay for NIP-91 provider announcements (kind:38421) via nostr library. + Query a Nostr relay for listing provider announcements (kind:38421) via nostr library. Returns a tuple of (events, ok) where ok indicates whether the relay interaction succeeded without transport-level errors. @@ -188,7 +188,7 @@ async def query_nip91_events( flt = Filter(kinds=[38421], authors=[pubkey], limit=10) filters = Filters([flt]) - sub_id = f"nip91_{int(time.time())}" + sub_id = f"routstr_listing_{int(time.time())}" rm.add_subscription(sub_id, filters) req: list[Any] = [ClientMessageType.REQUEST, sub_id] req.extend(filters.to_json_array()) @@ -294,7 +294,7 @@ async def _determine_provider_id(public_key_hex: str, relay_urls: list[str]) -> async def query_single_relay(relay_url: str) -> list[dict[str, Any]]: try: - events, _ok = await query_nip91_events(relay_url, public_key_hex, None) + events, _ok = await query_listing_events(relay_url, public_key_hex, None) return events except Exception: return [] @@ -330,7 +330,7 @@ async def publish_to_relay( timeout: int = 30, ) -> bool: """ - Publish a NIP-91 event to a nostr relay via nostr library. + Publish a listing event to a nostr relay via nostr library. """ def _sync_publish() -> bool: @@ -341,7 +341,7 @@ async def publish_to_relay( time.sleep(1.0) # Publish the event as-is via publish_message to preserve signature rm.publish_message(json.dumps(["EVENT", event])) - logger.debug(f"Sent NIP-91 event {event.get('id', '')} to {relay_url}") + logger.debug(f"Sent listing event {event.get('id', '')} to {relay_url}") time.sleep(1.0) return True except Exception as e: @@ -364,13 +364,13 @@ async def announce_provider() -> None: # Check for NSEC in environment (use NSEC only) nsec = settings.nsec if not nsec: - logger.info("Nostr private key not found (NSEC), skipping NIP-91 announcement") + logger.info("Nostr private key not found (NSEC), skipping listing announcement") return # Convert NSEC to keypair keypair = nsec_to_keypair(nsec) if not keypair: - logger.error("Failed to parse NSEC, skipping NIP-91 announcement") + logger.error("Failed to parse NSEC, skipping listing announcement") return private_key_hex, public_key_hex = keypair @@ -409,7 +409,7 @@ async def announce_provider() -> None: if not endpoint_urls: logger.warning( - "No valid endpoints configured (HTTP_URL/ONION_URL). Skipping NIP-91 publish." + "No valid endpoints configured (HTTP_URL/ONION_URL). Skipping listing publish." ) return @@ -434,7 +434,7 @@ async def announce_provider() -> None: # Create the candidate event that we would publish version_str = get_app_version() - candidate_event = create_nip91_event( + candidate_event = create_listing_event( private_key_hex=private_key_hex, provider_id=provider_id, endpoint_urls=endpoint_urls, @@ -474,7 +474,7 @@ async def announce_provider() -> None: if _should_skip(relay_url): logger.debug(f"Skipping {relay_url} due to backoff") continue - events, ok = await query_nip91_events(relay_url, public_key_hex, provider_id) + events, ok = await query_listing_events(relay_url, public_key_hex, provider_id) if ok: _register_success(relay_url) existing_events.extend(events) @@ -489,7 +489,7 @@ async def announce_provider() -> None: if not all_match: logger.debug( - "No matching NIP-91 announcement found or differences detected; publishing update" + "No matching listing announcement found or differences detected; publishing update" ) success_count = 0 for relay_url in relay_urls: @@ -502,11 +502,11 @@ async def announce_provider() -> None: else: _register_failure(relay_url) logger.info( - f"Published NIP-91 announcement to {success_count}/{len(relay_urls)} relays" + f"Published listing announcement to {success_count}/{len(relay_urls)} relays" ) else: logger.debug( - "Matching NIP-91 announcement already present; skipping publish on startup" + "Matching listing announcement already present; skipping publish on startup" ) # Re-announce periodically (every 24 hours) @@ -518,7 +518,7 @@ async def announce_provider() -> None: # Build fresh candidate event for comparison version_str = get_app_version() - candidate_event = create_nip91_event( + candidate_event = create_listing_event( private_key_hex=private_key_hex, provider_id=provider_id, endpoint_urls=endpoint_urls, @@ -533,7 +533,7 @@ async def announce_provider() -> None: if _should_skip(relay_url): logger.debug(f"Skipping {relay_url} due to backoff") continue - events, ok = await query_nip91_events( + events, ok = await query_listing_events( relay_url, public_key_hex, provider_id ) if ok: @@ -549,7 +549,7 @@ async def announce_provider() -> None: if all_match: logger.debug( - "Matching NIP-91 announcement already present; skipping periodic re-announce" + "Matching listing announcement already present; skipping periodic re-announce" ) continue @@ -567,8 +567,8 @@ async def announce_provider() -> None: _register_failure(relay_url) except asyncio.CancelledError: - logger.info("NIP-91 announcement task cancelled") + logger.info("Listing announcement task cancelled") break except Exception as e: - logger.debug(f"Error in NIP-91 announcement loop: {type(e).__name__}") + logger.debug(f"Error in listing announcement loop: {type(e).__name__}") # Continue running despite errors diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 40aed173..1df4caf8 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -15,6 +15,7 @@ class CostData(BaseModel): input_msats: int output_msats: int total_msats: int + total_usd: float = 0.0 class MaxCostData(CostData): @@ -61,6 +62,7 @@ async def calculate_cost( # todo: can be sync input_msats=0, output_msats=0, total_msats=0, + total_usd=0.0, ) usage_data = response_data["usage"] @@ -101,6 +103,7 @@ async def calculate_cost( # todo: can be sync input_msats=-1, # Cost field doesn't break down by token type output_msats=-1, total_msats=cost_in_msats, + total_usd=usd_cost, ) except Exception as e: logger.warning( @@ -210,6 +213,7 @@ async def calculate_cost( # todo: can be sync output_msats = round(output_tokens / 1000 * MSATS_PER_1K_OUTPUT_TOKENS, 3) token_based_cost = math.ceil(input_msats + output_msats) + total_usd = (token_based_cost / 1000.0) * sats_usd_price() logger.info( "Calculated token-based cost", @@ -219,6 +223,7 @@ async def calculate_cost( # todo: can be sync "input_cost_msats": input_msats, "output_cost_msats": output_msats, "total_cost_msats": token_based_cost, + "total_usd": total_usd, "model": response_data.get("model", "unknown"), }, ) @@ -228,4 +233,5 @@ async def calculate_cost( # todo: can be sync input_msats=int(input_msats), output_msats=int(output_msats), total_msats=token_based_cost, + total_usd=total_usd, ) diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py index b625fec1..26cf580d 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -25,82 +25,6 @@ class LNURLError(Exception): """LNURL related errors.""" -def parse_lightning_invoice_amount(invoice: str, currency: str = "sat") -> int: - """Parse Lightning invoice (BOLT-11) to extract amount in specified currency units. - - Args: - invoice: BOLT-11 Lightning invoice string - currency: Target currency unit ("sat" or "msat") - - Returns: - Amount in the specified currency unit - - Raises: - LNURLError: If invoice format is invalid or amount cannot be parsed - """ - invoice = invoice.lower().strip() - - if not invoice.startswith("ln"): - raise LNURLError("Invalid Lightning invoice format") - - # Find the network part (bc, tb, etc.) - network_start = 2 - while network_start < len(invoice) and invoice[network_start] not in "0123456789": - network_start += 1 - - if network_start >= len(invoice): - raise LNURLError("Invalid Lightning invoice format") - - # Parse amount and multiplier - amount_str = "" - multiplier = "" - i = network_start - - # Extract numeric part - while i < len(invoice) and invoice[i].isdigit(): - amount_str += invoice[i] - i += 1 - - # Extract multiplier if present - if i < len(invoice) and invoice[i] in "munp": - multiplier = invoice[i] - i += 1 - - # Check if we have the required "1" separator - if i >= len(invoice) or invoice[i] != "1": - raise LNURLError("Invalid Lightning invoice format") - - if not amount_str: - raise LNURLError("Lightning invoice amount not specified") - - # Convert to base units - try: - amount = int(amount_str) - except ValueError: - raise LNURLError("Invalid Lightning invoice amount") - - # Apply multiplier to get millisatoshis - if multiplier == "m": # milli = 10^-3 - amount_msat = amount * 100_000_000 # amount is in BTC * 10^-3 - elif multiplier == "u": # micro = 10^-6 - amount_msat = amount * 100_000 # amount is in BTC * 10^-6 - elif multiplier == "n": # nano = 10^-9 - amount_msat = amount * 100 # amount is in BTC * 10^-9 - elif multiplier == "p": # pico = 10^-12 - amount_msat = amount // 10 # amount is in BTC * 10^-12 - else: - # No multiplier means the amount is in BTC - amount_msat = amount * 100_000_000_000 # Convert BTC to msat - - # Convert to target currency unit - if currency == "msat": - return amount_msat - elif currency == "sat": - return amount_msat // 1000 - else: - raise LNURLError(f"Unsupported currency for Lightning: {currency}") - - async def decode_lnurl(lnurl: str) -> str: """Decode LNURL to get the actual URL. @@ -291,9 +215,7 @@ async def raw_send_to_lnurl( lnurl_data["callback_url"], final_amount ) - melt_quote_resp = await wallet.melt_quote( - invoice=bolt11_invoice, amount_msat=final_amount - ) + melt_quote_resp = await wallet.melt_quote(invoice=bolt11_invoice) if amount: proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True) diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 4afbf45e..a5cefbc0 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -143,14 +143,6 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis return [] -def is_openrouter_upstream() -> bool: - try: - base = (settings.upstream_base_url or "").strip().rstrip("/") - except Exception: - return False - return base.lower() == "https://openrouter.ai/api/v1" - - def _row_to_model( row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01 ) -> Model: @@ -203,33 +195,11 @@ def _row_to_model( return model -def _model_to_row_payload(model: Model) -> dict[str, str | int | bool | None]: - return { - "id": model.id, - "name": model.name, - "created": model.created, - "description": model.description, - "context_length": model.context_length, - "architecture": json.dumps(model.architecture.dict()), - "pricing": json.dumps(model.pricing.dict()), - "sats_pricing": json.dumps(model.sats_pricing.dict()) - if model.sats_pricing - else None, - "per_request_limits": json.dumps(model.per_request_limits) - if model.per_request_limits is not None - else None, - "top_provider": json.dumps(model.top_provider.dict()) - if model.top_provider is not None - else None, - "enabled": model.enabled, - "upstream_provider_id": model.upstream_provider_id, - } - - async def list_models( session: AsyncSession, upstream_id: int, include_disabled: bool = False, + apply_fees: bool = True, ) -> list[Model]: from sqlmodel import select @@ -247,7 +217,7 @@ async def list_models( return [ _row_to_model( r, - apply_provider_fee=True, + apply_provider_fee=apply_fees, provider_fee=providers_by_id[r.upstream_provider_id].provider_fee if r.upstream_provider_id in providers_by_id else 1.01, @@ -261,21 +231,6 @@ async def list_models( ] -async def get_model_by_id( - model_id: str, provider_id: int, session: AsyncSession -) -> Model | None: - from ..core.db import UpstreamProviderRow - - row = await session.get(ModelRow, (model_id, provider_id)) - if not row or not row.enabled: - return None - provider = await session.get(UpstreamProviderRow, provider_id) - if not provider or not provider.enabled: - return None - provider_fee = provider.provider_fee if provider else 1.01 - return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee) - - def _calculate_usd_max_costs(model: Model) -> tuple[float, float, float]: """Calculate max costs in USD based on model context/token limits. diff --git a/routstr/proxy.py b/routstr/proxy.py index d8dd8cfa..123d86bb 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -16,6 +16,7 @@ from .core.db import ( create_session, get_session, ) +from .core.exceptions import UpstreamError from .payment.helpers import ( calculate_discounted_max_cost, check_token_balance, @@ -31,7 +32,9 @@ proxy_router = APIRouter() _upstreams: list[BaseUpstreamProvider] = [] _model_instances: dict[str, Model] = {} # All aliases -> Model -_provider_map: dict[str, BaseUpstreamProvider] = {} # All aliases -> Provider +_provider_map: dict[ + str, list[BaseUpstreamProvider] +] = {} # All aliases -> List[Provider] _unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates) @@ -68,8 +71,8 @@ def get_model_instance(model_id: str) -> Model | None: return _model_instances.get(model_id.lower()) -def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None: - """Get UpstreamProvider for model ID from global cache.""" +def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None: + """Get UpstreamProvider list for model ID from global cache.""" return _provider_map.get(model_id.lower()) @@ -154,8 +157,8 @@ async def proxy( "invalid_model", f"Model '{model_id}' not found", 400, request=request ) - upstream = get_provider_for_model(model_id) - if not upstream: + upstreams = get_provider_for_model(model_id) + if not upstreams: return create_error_response( "invalid_model", f"No provider found for model '{model_id}'", @@ -163,6 +166,10 @@ async def proxy( request=request, ) + # todo figure out cost calculation since fallback provider is usually not the same price + # Use first provider for initial checks/cost calculation + # primary_upstream = upstreams[0] + _max_cost_for_model = await get_max_cost_for_model( model=model_id, session=session, model_obj=model_obj ) @@ -172,14 +179,31 @@ async def proxy( check_token_balance(headers, request_body_dict, max_cost_for_model) if x_cashu := headers.get("x-cashu", None): - if is_responses_api: - return await upstream.handle_x_cashu_responses( - request, x_cashu, path, max_cost_for_model, model_obj - ) - else: - return await upstream.handle_x_cashu( - request, x_cashu, path, max_cost_for_model, model_obj - ) + last_error = None + for i, upstream in enumerate(upstreams): + try: + if is_responses_api: + return await upstream.handle_x_cashu_responses( + request, x_cashu, path, max_cost_for_model, model_obj + ) + else: + return await upstream.handle_x_cashu( + request, x_cashu, path, max_cost_for_model, model_obj + ) + except UpstreamError as e: + logger.warning( + f"Upstream {upstream.provider_type} failed (x-cashu): {e}" + ) + if i == len(upstreams) - 1: + last_error = e + continue + + return create_error_response( + "upstream_error", + str(last_error) if last_error else "All upstreams failed", + 502, + request=request, + ) elif auth := headers.get("authorization", None): key = await get_bearer_token_key(headers, path, session, auth) @@ -194,70 +218,166 @@ async def proxy( ) logger.debug("Processing unauthenticated GET request", extra={"path": path}) - headers = upstream.prepare_headers(dict(request.headers)) - return await upstream.forward_get_request(request, path, headers) + + last_error_response = None + for i, upstream in enumerate(upstreams): + try: + headers = upstream.prepare_headers(dict(request.headers)) + response = await upstream.forward_get_request(request, path, headers) + + if response.status_code in [502, 429] and i < len(upstreams) - 1: + error_message = "" + try: + if hasattr(response, "body"): + body_bytes = response.body + data = json.loads(body_bytes) + if "error" in data: + error_data = data["error"] + if isinstance(error_data, dict): + error_message = error_data.get("message", "") + elif isinstance(error_data, str): + error_message = error_data + except Exception: + pass + + await upstream.on_upstream_error_redirect( + response.status_code, error_message + ) + + logger.warning( + f"Upstream {upstream.provider_type} returned {response.status_code} (GET), trying next provider", + extra={ + "status_code": response.status_code, + "upstream": upstream.provider_type, + }, + ) + continue + return response + except UpstreamError as e: + logger.warning(f"Upstream {upstream.provider_type} failed (GET): {e}") + if i == len(upstreams) - 1: + last_error_response = create_error_response( + "upstream_error", str(e), 502, request=request + ) + continue + return last_error_response or create_error_response( + "upstream_error", "All upstreams failed", 502, request=request + ) if request_body_dict: await pay_for_request(key, max_cost_for_model, session) - headers = upstream.prepare_headers(dict(request.headers)) + for i, upstream in enumerate(upstreams): + headers = upstream.prepare_headers(dict(request.headers)) - try: - if is_responses_api: - response = await upstream.forward_responses_request( - request, - path, - headers, - request_body, - key, - max_cost_for_model, - session, - model_obj, + try: + try: + if is_responses_api: + response = await upstream.forward_responses_request( + request, + path, + headers, + request_body, + key, + max_cost_for_model, + session, + model_obj, + ) + else: + response = await upstream.forward_request( + request, + path, + headers, + request_body, + key, + max_cost_for_model, + session, + model_obj, + ) + except Exception as e: + logger.error( + "Upstream request failed, ensuring payment is reverted", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "path": path, + "key_hash": key.hashed_key[:8] + "...", + "max_cost_for_model": max_cost_for_model, + }, + ) + await revert_pay_for_request(key, session, max_cost_for_model) + raise + + if response.status_code != 200: + # Check if we should retry (502 Upstream Error or 429 Rate Limit) + should_retry = response.status_code in [502, 429, 400, 401, 403, 404] + if should_retry and i < len(upstreams) - 1: + error_message = "" + try: + if hasattr(response, "body"): + body_bytes = response.body + data = json.loads(body_bytes) + if "error" in data: + error_data = data["error"] + if isinstance(error_data, dict): + error_message = error_data.get("message", "") + elif isinstance(error_data, str): + error_message = error_data + except Exception: + pass + + await upstream.on_upstream_error_redirect( + response.status_code, error_message + ) + + logger.warning( + f"Upstream {upstream.provider_type} returned {response.status_code}, trying next provider", + extra={ + "status_code": response.status_code, + "upstream": upstream.provider_type, + }, + ) + continue + + # 4xx error (user error), or other non-retryable error, or last provider failed + await revert_pay_for_request(key, session, max_cost_for_model) + logger.warning( + "Upstream request failed, revert payment", + extra={ + "status_code": response.status_code, + "path": path, + "key_hash": key.hashed_key[:8] + "...", + "key_balance": key.balance, + "max_cost_for_model": max_cost_for_model, + "upstream_headers": response.headers + if hasattr(response, "headers") + else None, + }, + ) + return response + + return response + + except UpstreamError as e: + logger.warning( + f"Upstream {upstream.provider_type} failed: {e}", + extra={"retry": i < len(upstreams) - 1}, ) - else: - response = await upstream.forward_request( - request, - path, - headers, - request_body, - key, - max_cost_for_model, - session, - model_obj, - ) - except Exception as e: - logger.error( - "Upstream request failed, ensuring payment is reverted", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "path": path, - "key_hash": key.hashed_key[:8] + "...", - "max_cost_for_model": max_cost_for_model, - }, - ) - await revert_pay_for_request(key, session, max_cost_for_model) - raise - if response.status_code != 200: - await revert_pay_for_request(key, session, max_cost_for_model) - logger.warning( - "Upstream request failed, revert payment", - extra={ - "status_code": response.status_code, - "path": path, - "key_hash": key.hashed_key[:8] + "...", - "key_balance": key.balance, - "max_cost_for_model": max_cost_for_model, - "upstream_headers": response.headers - if hasattr(response, "headers") - else None, - }, - ) - # Return the mapped error response generated earlier rather than masking with 502 - return response + # If 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 + ) - return response + # 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( diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 559d05ab..f3ed1dab 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -15,6 +15,7 @@ from pydantic import BaseModel from ..auth import adjust_payment_for_tokens, revert_pay_for_request from ..core import get_logger from ..core.db import ApiKey, AsyncSession, create_session +from ..core.exceptions import UpstreamError if TYPE_CHECKING: from ..core.db import UpstreamProviderRow @@ -340,6 +341,20 @@ class BaseUpstreamProvider: message = preview[:500] return message, upstream_code + async def on_upstream_error_redirect( + self, status_code: int, error_message: str + ) -> None: + """Hook called when the proxy redirects to another provider due to an error. + + Subclasses can implement this to perform actions like disabling the provider + if it's out of balance. + + Args: + status_code: The HTTP status code returned by the upstream + error_message: The error message extracted from the upstream response + """ + pass + async def map_upstream_error_response( self, request: Request, path: str, upstream_response: httpx.Response ) -> Response: @@ -436,164 +451,111 @@ class BaseUpstreamProvider: 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 + usage_chunk_data: dict | None = None + done_seen: bool = False - async def finalize_without_usage() -> bytes | None: + async def finalize_db_only() -> None: nonlocal usage_finalized if usage_finalized: - return None + return async with create_session() as new_session: fresh_key = await new_session.get(key.__class__, key.hashed_key) if not fresh_key: - logger.warning( - "Key not found when finalizing streaming payment", - extra={"key_hash": key.hashed_key[:8] + "..."}, - ) - usage_finalized = True - return None + return 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 + await adjust_payment_for_tokens( + fresh_key, + {"model": last_model_seen or "unknown", "usage": None}, + 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: - async for chunk in response.aiter_bytes(): - stored_chunks.append(chunk) - try: - for part in re.split(b"data: ", chunk): - if not part or part.strip() in (b"[DONE]", b""): - continue - try: - obj = json.loads(part) - if isinstance(obj, dict) and obj.get("model"): - last_model_seen = str(obj.get("model")) - except json.JSONDecodeError: - pass except Exception: pass - yield chunk + try: + async for chunk in response.aiter_bytes(): + # Split chunk into SSE events + parts = re.split(b"data: ", chunk) + for i, part in enumerate(parts): + if not part: + continue - logger.debug( - "Streaming completed, analyzing usage data", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "chunks_count": len(stored_chunks), - }, - ) + stripped_part = part.strip() + if not stripped_part: + continue - 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 - ): - 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 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: + cost_data = await adjust_payment_for_tokens( + fresh_key, + usage_chunk_data, + session, + max_cost_for_model, + ) + # Merge cost into usage + usage_chunk_data["usage"]["cost"] = cost_data.get( + "total_usd", 0.0 + ) + # Keep detailed cost in metadata + usage_chunk_data["metadata"] = usage_chunk_data.get( + "metadata", {} + ) + usage_chunk_data["metadata"]["routstr"] = { + "cost": cost_data + } + yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() + usage_finalized = True + except Exception: + # Fallback: yield original usage chunk if adjustment fails + yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() if not usage_finalized: - maybe_cost_event = await finalize_without_usage() - if maybe_cost_event is not None: - yield maybe_cost_event + await finalize_db_only() + + if done_seen: + yield b"data: [DONE]\n\n" except Exception as stream_error: logger.warning( - "Streaming interrupted; finalizing without usage", + "Streaming interrupted; finalizing in background", extra={ "error": str(stream_error), - "error_type": type(stream_error).__name__, "key_hash": key.hashed_key[:8] + "...", }, ) raise finally: if not usage_finalized: - await finalize_without_usage() + await finalize_db_only() # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) @@ -633,6 +595,7 @@ class BaseUpstreamProvider: }, ) + content: bytes | None = None try: content = await response.aread() response_json = json.loads(content) @@ -649,6 +612,14 @@ class BaseUpstreamProvider: cost_data = await adjust_payment_for_tokens( 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 logger.info( @@ -734,180 +705,135 @@ class BaseUpstreamProvider: async def stream_with_responses_cost( max_cost_for_model: int, ) -> AsyncGenerator[bytes, None]: - stored_chunks: list[bytes] = [] usage_finalized: bool = False last_model_seen: str | None = None reasoning_tokens: int = 0 + usage_chunk_data: dict | None = None + done_seen: bool = False - async def finalize_without_usage() -> bytes | None: + async def finalize_db_only() -> None: nonlocal usage_finalized if usage_finalized: - return None + return async with create_session() as new_session: fresh_key = await new_session.get(key.__class__, key.hashed_key) if not fresh_key: - logger.warning( - "Key not found when finalizing Responses API streaming payment", - extra={"key_hash": key.hashed_key[:8] + "..."}, - ) - usage_finalized = True - return None + return 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 + await adjust_payment_for_tokens( + fresh_key, + {"model": last_model_seen or "unknown", "usage": None}, + 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: - async for chunk in response.aiter_bytes(): - stored_chunks.append(chunk) - try: - for part in re.split(b"data: ", chunk): - if not part or part.strip() in (b"[DONE]", b""): - 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 - ) - except json.JSONDecodeError: - pass except Exception: pass - yield chunk + try: + async for chunk in response.aiter_bytes(): + # Split chunk into SSE events + parts = re.split(b"data: ", chunk) + for i, part in enumerate(parts): + if not part: + continue - 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, - }, - ) + stripped_part = part.strip() + if not stripped_part: + continue - # 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 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 ) - 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] + "...", - }, + + # 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: + cost_data = await adjust_payment_for_tokens( + fresh_key, + usage_chunk_data, + session, + 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 + usage_chunk_data["metadata"] = usage_chunk_data.get( + "metadata", {} + ) + usage_chunk_data["metadata"]["routstr"] = { + "cost": cost_data + } + yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() + usage_finalized = True + except Exception: + # Fallback: yield original usage chunk if adjustment fails + yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() if not usage_finalized: - maybe_cost_event = await finalize_without_usage() - if maybe_cost_event is not None: - yield maybe_cost_event + await finalize_db_only() + + if done_seen: + yield b"data: [DONE]\n\n" except Exception as stream_error: logger.warning( - "Responses API streaming interrupted; finalizing without usage", + "Responses API streaming interrupted; finalizing in background", extra={ "error": str(stream_error), - "error_type": type(stream_error).__name__, "key_hash": key.hashed_key[:8] + "...", }, ) raise finally: if not usage_finalized: - await finalize_without_usage() + await finalize_db_only() # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) @@ -947,6 +873,7 @@ class BaseUpstreamProvider: }, ) + content: bytes | None = None try: content = await response.aread() response_json = json.loads(content) @@ -966,6 +893,14 @@ class BaseUpstreamProvider: cost_data = await adjust_payment_for_tokens( 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 logger.info( @@ -1148,6 +1083,14 @@ class BaseUpstreamProvider: ) if response.status_code != 200: + if response.status_code >= 500: + await response.aclose() + await client.aclose() + raise UpstreamError( + f"Upstream returned status {response.status_code}", + status_code=response.status_code, + ) + try: mapped_error = await self.map_upstream_error_response( request, path, response @@ -1238,6 +1181,9 @@ class BaseUpstreamProvider: background=background_tasks, ) + except UpstreamError: + raise + except httpx.RequestError as exc: await client.aclose() error_type = type(exc).__name__ @@ -1267,9 +1213,7 @@ class BaseUpstreamProvider: else: error_message = f"Error connecting to upstream service: {error_type}" - return create_error_response( - "upstream_error", error_message, 502, request=request - ) + raise UpstreamError(error_message, status_code=502) except Exception as exc: await client.aclose() @@ -1384,6 +1328,14 @@ class BaseUpstreamProvider: ) if response.status_code != 200: + if response.status_code >= 500: + await response.aclose() + await client.aclose() + raise UpstreamError( + f"Upstream returned status {response.status_code}", + status_code=response.status_code, + ) + try: mapped_error = await self.map_upstream_error_response( request, path, response @@ -1451,6 +1403,9 @@ class BaseUpstreamProvider: background=background_tasks, ) + except UpstreamError: + raise + except httpx.RequestError as exc: await client.aclose() error_type = type(exc).__name__ @@ -1480,9 +1435,7 @@ class BaseUpstreamProvider: else: error_message = f"Error connecting to upstream service: {error_type}" - return create_error_response( - "upstream_error", error_message, 502, request=request - ) + raise UpstreamError(error_message, status_code=502) except Exception as exc: await client.aclose() diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index 29e099ca..95b6ca84 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -236,7 +236,7 @@ async def _seed_providers_from_settings( from . import upstream_provider_classes providers_to_add: list[UpstreamProviderRow] = [] - seeded_base_urls: set[str] = set() + seeded_provider_keys: set[tuple[str, str]] = set() provider_classes_by_type = { cls.provider_type: cls @@ -261,7 +261,8 @@ async def _seed_providers_from_settings( base_url = provider_class.default_base_url # type: ignore[attr-defined] result = await session.exec( select(UpstreamProviderRow).where( - UpstreamProviderRow.base_url == base_url + UpstreamProviderRow.base_url == base_url, + UpstreamProviderRow.api_key == api_key, ) ) if not result.first(): @@ -273,13 +274,15 @@ async def _seed_providers_from_settings( enabled=True, ) ) - seeded_base_urls.add(base_url) + seeded_provider_keys.add((base_url, api_key)) ollama_base_url = os.environ.get("OLLAMA_BASE_URL") if ollama_base_url: + ollama_api_key = os.environ.get("OLLAMA_API_KEY", "") result = await session.exec( select(UpstreamProviderRow).where( - UpstreamProviderRow.base_url == ollama_base_url + UpstreamProviderRow.base_url == ollama_base_url, + UpstreamProviderRow.api_key == ollama_api_key, ) ) if not result.first(): @@ -287,18 +290,20 @@ async def _seed_providers_from_settings( UpstreamProviderRow( provider_type="ollama", base_url=ollama_base_url, - api_key=os.environ.get("OLLAMA_API_KEY", ""), + api_key=ollama_api_key, enabled=True, ) ) - seeded_base_urls.add(ollama_base_url) + seeded_provider_keys.add((ollama_base_url, ollama_api_key)) if settings.chat_completions_api_version and settings.upstream_base_url: base_url = settings.upstream_base_url - if base_url not in seeded_base_urls: + api_key = settings.upstream_api_key + if (base_url, api_key) not in seeded_provider_keys: result = await session.exec( select(UpstreamProviderRow).where( - UpstreamProviderRow.base_url == base_url + UpstreamProviderRow.base_url == base_url, + UpstreamProviderRow.api_key == api_key, ) ) if not result.first(): @@ -306,19 +311,21 @@ async def _seed_providers_from_settings( UpstreamProviderRow( provider_type="azure", base_url=base_url, - api_key=settings.upstream_api_key, + api_key=api_key, api_version=settings.chat_completions_api_version, enabled=True, ) ) - seeded_base_urls.add(base_url) + seeded_provider_keys.add((base_url, api_key)) if settings.upstream_base_url and settings.upstream_api_key: base_url = settings.upstream_base_url - if base_url not in seeded_base_urls: + api_key = settings.upstream_api_key + if (base_url, api_key) not in seeded_provider_keys: result = await session.exec( select(UpstreamProviderRow).where( - UpstreamProviderRow.base_url == base_url + UpstreamProviderRow.base_url == base_url, + UpstreamProviderRow.api_key == api_key, ) ) if not result.first(): @@ -326,11 +333,11 @@ async def _seed_providers_from_settings( UpstreamProviderRow( provider_type="custom", base_url=base_url, - api_key=settings.upstream_api_key, + api_key=api_key, enabled=True, ) ) - seeded_base_urls.add(base_url) + seeded_provider_keys.add((base_url, api_key)) for provider in providers_to_add: session.add(provider) diff --git a/routstr/upstream/ppqai.py b/routstr/upstream/ppqai.py index 0b4dfb08..c8e9f6a4 100644 --- a/routstr/upstream/ppqai.py +++ b/routstr/upstream/ppqai.py @@ -196,6 +196,37 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): ) return [] + async def on_upstream_error_redirect( + self, status_code: int, error_message: str + ) -> None: + if "insufficient balance" in error_message.lower(): + logger.warning( + f"Disabling PPQ.AI provider ({self.base_url}) due to insufficient balance", + extra={"error": error_message}, + ) + from sqlmodel import select + + from ..core.db import UpstreamProviderRow, create_session + + async with create_session() as session: + statement = select(UpstreamProviderRow).where( + UpstreamProviderRow.base_url == self.base_url, + UpstreamProviderRow.api_key == self.api_key, + ) + result = await session.exec(statement) + provider = result.first() + + if provider: + provider.enabled = False + session.add(provider) + await session.commit() + + # Trigger re-initialization of providers + # Import here to avoid circular dependency + from ..proxy import reinitialize_upstreams + + await reinitialize_upstreams() + async def create_account(self) -> dict[str, object]: """Create a new PPQ.AI account. diff --git a/routstr/wallet.py b/routstr/wallet.py index ce87b560..71ea18ac 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -319,6 +319,7 @@ async def periodic_payout() -> None: wallet, mint_url, unit, not_reserved=True ) proofs = await slow_filter_spend_proofs(proofs, wallet) + await asyncio.sleep(5) user_balance = await db.balances_for_mint_and_unit( session, mint_url, unit ) @@ -344,8 +345,6 @@ async def periodic_payout() -> None: "amount_received": amount_received, }, ) - - await asyncio.sleep(5) except Exception as e: logger.error( f"Error sending payout: {type(e).__name__}", diff --git a/tests/integration/test_child_keys.py b/tests/integration/test_child_keys.py new file mode 100644 index 00000000..c247ef3d --- /dev/null +++ b/tests/integration/test_child_keys.py @@ -0,0 +1,141 @@ +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 ChildKeyRequest, 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( + ChildKeyRequest(count=1), parent_key, integration_session + ) + + assert "api_keys" in result + assert result["cost_msats"] == 1000 + assert result["parent_balance"] == 9000 + + child_key_raw = result["api_keys"][0][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( + ChildKeyRequest(count=1), 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(ChildKeyRequest(count=1), 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) diff --git a/tests/integration/test_general_info_endpoints.py b/tests/integration/test_general_info_endpoints.py index 4c7af5f1..d9399f09 100644 --- a/tests/integration/test_general_info_endpoints.py +++ b/tests/integration/test_general_info_endpoints.py @@ -267,10 +267,9 @@ async def test_admin_endpoint_unauthenticated( """Test GET /admin/ endpoint redirects to /""" await db_snapshot.capture() - response = await integration_client.get("/admin/") + response = await integration_client.get("/admin/api/settings") - assert response.status_code == 307 - assert response.headers.get("location") == "/" + assert response.status_code == 403 diff = await db_snapshot.diff() assert len(diff["api_keys"]["added"]) == 0 diff --git a/tests/integration/test_provider_management.py b/tests/integration/test_provider_management.py index 584fd8c1..b47ad848 100644 --- a/tests/integration/test_provider_management.py +++ b/tests/integration/test_provider_management.py @@ -9,7 +9,7 @@ from unittest.mock import patch import pytest from httpx import AsyncClient -from routstr.discovery import _PROVIDERS_CACHE +from routstr.nostr.discovery import _PROVIDERS_CACHE from .utils import ResponseValidator @@ -71,9 +71,10 @@ async def test_providers_endpoint_default_response( } with patch( - "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.nostr.discovery.query_nostr_relay_for_providers", + return_value=mock_events, ): - with patch("routstr.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch: # Configure mock to return appropriate responses mock_fetch.side_effect = lambda url: mock_fetch_responses.get( url, {"status_code": 500, "json": {"error": "Unknown provider"}} @@ -135,9 +136,10 @@ async def test_providers_endpoint_with_include_json( } with patch( - "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.nostr.discovery.query_nostr_relay_for_providers", + return_value=mock_events, ): - with patch("routstr.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = { "status_code": 200, "json": mock_provider_response, @@ -209,9 +211,10 @@ async def test_providers_data_structure_validation( } with patch( - "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.nostr.discovery.query_nostr_relay_for_providers", + return_value=mock_events, ): - with patch("routstr.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = mock_health_response response = await integration_client.get("/v1/providers/?include_json=true") @@ -256,7 +259,8 @@ async def test_providers_endpoint_no_providers_found( ] with patch( - "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.nostr.discovery.query_nostr_relay_for_providers", + return_value=mock_events, ): response = await integration_client.get("/v1/providers/") @@ -317,10 +321,11 @@ async def test_providers_endpoint_offline_providers( } with patch( - "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.nostr.discovery.query_nostr_relay_for_providers", + return_value=mock_events, ): with patch( - "routstr.discovery.fetch_provider_health", + "routstr.nostr.discovery.fetch_provider_health", side_effect=mock_fetch_provider_health, ): response = await integration_client.get("/v1/providers/?include_json=true") @@ -386,9 +391,10 @@ async def test_providers_endpoint_duplicate_urls( ] with patch( - "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.nostr.discovery.query_nostr_relay_for_providers", + return_value=mock_events, ): - with patch("routstr.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = { "status_code": 200, "endpoint": "root", @@ -425,7 +431,8 @@ async def test_providers_endpoint_nostr_relay_failures( raise Exception("Connection to relay failed") with patch( - "routstr.discovery.query_nostr_relay_for_providers", side_effect=failing_query + "routstr.nostr.discovery.query_nostr_relay_for_providers", + side_effect=failing_query, ): response = await integration_client.get("/v1/providers/") @@ -463,9 +470,10 @@ async def test_providers_endpoint_malformed_urls( ] with patch( - "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.nostr.discovery.query_nostr_relay_for_providers", + return_value=mock_events, ): - with patch("routstr.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} response = await integration_client.get("/v1/providers/") @@ -495,9 +503,10 @@ async def test_providers_endpoint_response_format( ] with patch( - "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.nostr.discovery.query_nostr_relay_for_providers", + return_value=mock_events, ): - with patch("routstr.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} # Test default format @@ -545,9 +554,10 @@ async def test_providers_endpoint_concurrent_requests( ] with patch( - "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.nostr.discovery.query_nostr_relay_for_providers", + return_value=mock_events, ): - with patch("routstr.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} # Create concurrent requests @@ -587,9 +597,10 @@ async def test_providers_endpoint_parameter_validation( ] with patch( - "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.nostr.discovery.query_nostr_relay_for_providers", + return_value=mock_events, ): - with patch("routstr.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} # Test various parameter values @@ -639,9 +650,10 @@ async def test_no_database_changes_during_provider_operations( ] with patch( - "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events + "routstr.nostr.discovery.query_nostr_relay_for_providers", + return_value=mock_events, ): - with patch("routstr.discovery.fetch_provider_health") as mock_fetch: + with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch: mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} # Make multiple requests with different parameters diff --git a/tests/unit/test_algorithm.py b/tests/unit/test_algorithm.py index d367819f..fd7a837f 100644 --- a/tests/unit/test_algorithm.py +++ b/tests/unit/test_algorithm.py @@ -10,7 +10,6 @@ os.environ["UPSTREAM_API_KEY"] = "test" from routstr.algorithm import ( # noqa: E402 calculate_model_cost_score, get_provider_penalty, - should_prefer_model, ) from routstr.payment.models import Architecture, Model, Pricing # noqa: E402 @@ -100,100 +99,3 @@ def test_get_provider_penalty_openrouter() -> None: provider = create_test_provider("openrouter", "https://openrouter.ai/api/v1") penalty = get_provider_penalty(provider) assert penalty == 1.001 - - -def test_should_prefer_model_cheaper_wins() -> None: - """Test that cheaper model is preferred.""" - cheap_model = create_test_model("cheap", prompt_price=0.001, completion_price=0.002) - expensive_model = create_test_model( - "expensive", prompt_price=0.03, completion_price=0.06 - ) - - provider1 = create_test_provider("provider1") - provider2 = create_test_provider("provider2") - - # Cheaper model should win - assert should_prefer_model( - cheap_model, provider1, expensive_model, provider2, "test-alias" - ) - - # More expensive model should not win - assert not should_prefer_model( - expensive_model, provider2, cheap_model, provider1, "test-alias" - ) - - -def test_should_prefer_model_exact_match_wins() -> None: - """Test that exact alias match beats cheaper price.""" - # Make model IDs match the alias differently - exact_match = create_test_model( - "test-model", prompt_price=0.03, completion_price=0.06 - ) - no_match = create_test_model( - "other-model", prompt_price=0.001, completion_price=0.002 - ) - - provider1 = create_test_provider("provider1") - provider2 = create_test_provider("provider2") - - # Exact match should win even though it's more expensive - assert should_prefer_model( - exact_match, provider1, no_match, provider2, "test-model" - ) - - -def test_should_prefer_model_openrouter_slight_penalty() -> None: - """Test that OpenRouter has slight penalty compared to other providers.""" - model1 = create_test_model("model1", prompt_price=0.001, completion_price=0.002) - model2 = create_test_model("model2", prompt_price=0.001, completion_price=0.002) - - regular_provider = create_test_provider("regular", "http://provider.com") - openrouter_provider = create_test_provider( - "openrouter", "https://openrouter.ai/api/v1" - ) - - # Regular provider should be preferred over OpenRouter at same cost - assert should_prefer_model( - model1, regular_provider, model2, openrouter_provider, "test-alias" - ) - - # OpenRouter should not replace regular provider at same cost - assert not should_prefer_model( - model2, openrouter_provider, model1, regular_provider, "test-alias" - ) - - -def test_should_prefer_model_openrouter_can_win_if_cheaper() -> None: - """Test that OpenRouter can still win if significantly cheaper.""" - cheap_model = create_test_model( - "cheap", prompt_price=0.0001, completion_price=0.0002 - ) - expensive_model = create_test_model( - "expensive", prompt_price=0.03, completion_price=0.06 - ) - - regular_provider = create_test_provider("regular", "http://provider.com") - openrouter_provider = create_test_provider( - "openrouter", "https://openrouter.ai/api/v1" - ) - - # OpenRouter should win if it's much cheaper (even with penalty) - assert should_prefer_model( - cheap_model, - openrouter_provider, - expensive_model, - regular_provider, - "test-alias", - ) - - -def test_should_prefer_model_same_cost_first_wins() -> None: - """Test that when costs are identical, current model is kept.""" - model1 = create_test_model("model1", prompt_price=0.001, completion_price=0.002) - model2 = create_test_model("model2", prompt_price=0.001, completion_price=0.002) - - provider1 = create_test_provider("provider1") - provider2 = create_test_provider("provider2") - - # When costs are equal, should not replace - assert not should_prefer_model(model2, provider2, model1, provider1, "test-alias") diff --git a/tests/unit/test_fee_consistency.py b/tests/unit/test_fee_consistency.py deleted file mode 100644 index 8e72a771..00000000 --- a/tests/unit/test_fee_consistency.py +++ /dev/null @@ -1,155 +0,0 @@ -"""Unit tests for model row payload conversion. - -This module tests that _model_to_row_payload correctly serializes model data -for database storage. Pricing is stored as-is without fee application. -Fees are now applied per-provider when reading from the database. - -Key behaviors tested: -1. Pricing is stored as-is without fee application -2. All model fields are correctly serialized to JSON -3. Optional fields are handled correctly (None values) -4. Pricing structure is preserved -5. Original model objects are not mutated -""" - -import json -import os - -import pytest - -# Set required env vars before importing -os.environ["UPSTREAM_BASE_URL"] = "http://test" -os.environ["UPSTREAM_API_KEY"] = "test" - -from routstr.payment.models import ( # noqa: E402 - Architecture, - Model, - Pricing, - _model_to_row_payload, -) - - -@pytest.fixture -def base_architecture() -> Architecture: - """Provide standard architecture for test models.""" - return Architecture( - modality="text", - input_modalities=["text"], - output_modalities=["text"], - tokenizer="gpt", - instruct_type="chat", - ) - - -@pytest.fixture -def standard_pricing() -> Pricing: - """Provide standard USD pricing with known values for testing.""" - return Pricing( - prompt=0.001, - completion=0.002, - request=0.01, - image=0.05, - web_search=0.03, - internal_reasoning=0.015, - max_prompt_cost=10.0, - max_completion_cost=20.0, - max_cost=30.0, - ) - - -@pytest.fixture -def standard_model(base_architecture: Architecture, standard_pricing: Pricing) -> Model: - """Create a standard test model with known pricing.""" - return Model( - id="test-model-standard", - name="Test Model Standard", - created=1234567890, - description="A standard test model", - context_length=8192, - architecture=base_architecture, - pricing=standard_pricing, - ) - - -def test_pricing_stored_without_fees(standard_model: Model) -> None: - """Verify pricing is stored as-is without any fee application.""" - payload = _model_to_row_payload(standard_model) - pricing_str = payload["pricing"] - assert isinstance(pricing_str, str) - pricing = json.loads(pricing_str) - - assert pricing["prompt"] == pytest.approx(0.001, rel=1e-9) - assert pricing["completion"] == pytest.approx(0.002, rel=1e-9) - assert pricing["request"] == pytest.approx(0.01, rel=1e-9) - assert pricing["image"] == pytest.approx(0.05, rel=1e-9) - assert pricing["web_search"] == pytest.approx(0.03, rel=1e-9) - assert pricing["internal_reasoning"] == pytest.approx(0.015, rel=1e-9) - assert pricing["max_prompt_cost"] == pytest.approx(10.0, rel=1e-9) - assert pricing["max_completion_cost"] == pytest.approx(20.0, rel=1e-9) - assert pricing["max_cost"] == pytest.approx(30.0, rel=1e-9) - - -def test_zero_value_pricing_fields(base_architecture: Architecture) -> None: - """Verify that zero-value pricing fields are stored correctly.""" - zero_pricing = Pricing( - prompt=0.0, - completion=0.0, - request=0.0, - image=0.0, - web_search=0.0, - internal_reasoning=0.0, - max_prompt_cost=0.0, - max_completion_cost=0.0, - max_cost=0.0, - ) - - model = Model( - id="test-model-zero", - name="Test Model Zero", - created=1234567890, - description="A model with zero pricing", - context_length=8192, - architecture=base_architecture, - pricing=zero_pricing, - ) - - payload = _model_to_row_payload(model) - pricing_str = payload["pricing"] - assert isinstance(pricing_str, str) - pricing = json.loads(pricing_str) - - assert pricing["prompt"] == pytest.approx(0.0, rel=1e-9) - assert pricing["completion"] == pytest.approx(0.0, rel=1e-9) - assert pricing["request"] == pytest.approx(0.0, rel=1e-9) - - -def test_payload_structure_unchanged(standard_model: Model) -> None: - """Verify that payload structure matches expectations.""" - payload = _model_to_row_payload(standard_model) - - assert "id" in payload - assert "name" in payload - assert "created" in payload - assert "description" in payload - assert "context_length" in payload - assert "architecture" in payload - assert "pricing" in payload - assert "sats_pricing" in payload - assert "per_request_limits" in payload - assert "top_provider" in payload - assert "enabled" in payload - assert "upstream_provider_id" in payload - - assert isinstance(payload["architecture"], str) - assert isinstance(payload["pricing"], str) - - -def test_original_model_not_mutated(standard_model: Model) -> None: - """Verify that the original model object is not mutated.""" - original_prompt = standard_model.pricing.prompt - original_completion = standard_model.pricing.completion - - _model_to_row_payload(standard_model) - - assert standard_model.pricing.prompt == original_prompt - assert standard_model.pricing.completion == original_completion diff --git a/ui/app/providers/page.tsx b/ui/app/providers/page.tsx index dc2b6e3e..2f377011 100644 --- a/ui/app/providers/page.tsx +++ b/ui/app/providers/page.tsx @@ -21,6 +21,7 @@ import { AdminModel, } from '@/lib/api/services/admin'; import { AddProviderModelDialog } from '@/components/AddProviderModelDialog'; +import { BatchOverrideDialog } from '@/components/BatchOverrideDialog'; import { Skeleton } from '@/components/ui/skeleton'; import { AlertCircle, @@ -286,7 +287,9 @@ function ProviderBalance({
{/* eslint-disable-next-line @next/next/no-img-element */} Lightning Invoice QR Code @@ -403,6 +406,9 @@ export default function ProvidersPage() { mode: 'create', initialData: null, }); + const [batchOverrideProviderId, setBatchOverrideProviderId] = useState< + number | null + >(null); const [formData, setFormData] = useState({ provider_type: 'openrouter', @@ -482,6 +488,25 @@ export default function ProvidersPage() { }, }); + const deleteModelMutation = useMutation({ + mutationFn: ({ + providerId, + modelId, + }: { + providerId: number; + modelId: string; + }) => AdminService.deleteProviderModel(providerId, modelId), + onSuccess: (_, variables) => { + queryClient.invalidateQueries({ + queryKey: ['provider-models', variables.providerId], + }); + toast.success('Model deleted successfully'); + }, + onError: (error: Error) => { + toast.error(`Failed to delete model: ${error.message}`); + }, + }); + const handleCreateAccount = async () => { setIsCreatingAccount(true); try { @@ -560,6 +585,12 @@ export default function ProvidersPage() { } }; + const handleDeleteModel = (providerId: number, modelId: string) => { + if (confirm('Are you sure you want to delete this model?')) { + deleteModelMutation.mutate({ providerId, modelId }); + } + }; + const getDefaultBaseUrl = (type: string) => { const providerType = providerTypes.find((pt) => pt.id === type); return providerType?.default_base_url || ''; @@ -627,6 +658,10 @@ export default function ProvidersPage() { }); }; + const handleBatchOverride = (providerId: number) => { + setBatchOverrideProviderId(providerId); + }; + return ( @@ -933,15 +968,68 @@ export default function ProvidersPage() {
) : providerModels && viewingModels === provider.id ? ( - providerModels.remote_models.length === 0 ? ( - // No provided models - show custom models directly without tabs -
- {providerModels.db_models.length === 0 ? ( -
+ 0 + ? 'provided' + : 'custom' + } + className='w-full' + > + + + + Provided Models + + Provided + + {providerModels.remote_models.length} + + + + + Custom Models + + Custom + + {providerModels.db_models.length} + + + + +
+ {providerModels.db_models.length > 0 && (
- No models configured. Add custom models - to use this provider. + Custom models override or extend the + provider's catalog.
+ )} +
+
+
+ {providerModels.db_models.length === 0 ? ( +
+ No custom models configured +
) : (
{providerModels.db_models.map((model) => ( @@ -1000,103 +1093,48 @@ export default function ProvidersPage() { > +
))}
)} - - ) : ( - // Has provided models - show tabs - + - - - - Provided Models - - - Provided - - - {providerModels.remote_models.length} - - - - - Custom Models - - Custom - - {providerModels.db_models.length} - - - - -
- {providerModels.db_models.length > 0 && ( -
- Custom models override or extend the - provider's catalog. -
- )} - -
- {providerModels.db_models.length === 0 ? ( -
- No custom models configured + {providerModels.remote_models.length > 0 ? ( + <> +
+ Models automatically discovered from the + provider's catalog.
- ) : (
- {providerModels.db_models.map( + {providerModels.remote_models.map( (model) => (
-
- - {model.id} - - - {model.enabled - ? 'Enabled' - : 'Disabled'} - +
+ {model.id}
{model.description || @@ -1109,79 +1147,32 @@ export default function ProvidersPage() { tokens
) )}
- )} - - - {providerModels.remote_models.length > - 0 && ( -
- Models automatically discovered from the - provider's catalog. -
- )} -
- {providerModels.remote_models.map( - (model) => ( -
-
-
- {model.id} -
-
- {model.description || - model.name} -
-
-
-
- {model.context_length?.toLocaleString()}{' '} - tokens -
- -
-
- ) - )} + + ) : ( +
+ No provided models available
- - - ) + )} + + ) : null}
)} @@ -1357,6 +1348,19 @@ export default function ProvidersPage() { mode={modelDialogState.mode} /> )} + + {batchOverrideProviderId && ( + setBatchOverrideProviderId(null)} + onSuccess={() => { + queryClient.invalidateQueries({ + queryKey: ['provider-models', batchOverrideProviderId], + }); + }} + /> + )} ); diff --git a/ui/components/AddProviderModelDialog.tsx b/ui/components/AddProviderModelDialog.tsx index 301682f4..009eb942 100644 --- a/ui/components/AddProviderModelDialog.tsx +++ b/ui/components/AddProviderModelDialog.tsx @@ -248,7 +248,9 @@ export function AddProviderModelDialog({ const pricing = model.pricing as Record; const topProvider = model.top_provider as Record | null; - form.setValue('id', model.id); + if (!isOverride) { + form.setValue('id', model.id); + } form.setValue('name', model.name); form.setValue('description', model.description || ''); form.setValue('context_length', model.context_length); @@ -425,7 +427,7 @@ export function AddProviderModelDialog({ {description} - {!isEdit && !isOverride && ( + {!isEdit && (
Presets
@@ -488,8 +490,9 @@ export function AddProviderModelDialog({
- Prefill fields from a preset model definition, then adjust as - needed. + {isOverride + ? 'Apply pricing and settings from a preset model (keeping the model ID unchanged).' + : 'Prefill fields from a preset model definition, then adjust as needed.'}
)} @@ -805,45 +808,6 @@ export function AddProviderModelDialog({ )} /> - ( - - Max Prompt Cost - - - - - - )} - /> - ( - - Max Completion Cost - - - - - - )} - /> - ( - - Max Total Cost - - - - - - )} - />
diff --git a/ui/components/BatchOverrideDialog.tsx b/ui/components/BatchOverrideDialog.tsx new file mode 100644 index 00000000..799f4223 --- /dev/null +++ b/ui/components/BatchOverrideDialog.tsx @@ -0,0 +1,144 @@ +'use client'; + +import React, { useState } from 'react'; +import { Button } from '@/components/ui/button'; +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle, +} from '@/components/ui/dialog'; +import { Textarea } from '@/components/ui/textarea'; +import { toast } from 'sonner'; +import { AdminService } from '@/lib/api/services/admin'; +import { Loader2, Database } from 'lucide-react'; + +export interface BatchOverrideDialogProps { + providerId: number; + isOpen: boolean; + onClose: () => void; + onSuccess: () => void; +} + +export function BatchOverrideDialog({ + providerId, + isOpen, + onClose, + onSuccess, +}: BatchOverrideDialogProps) { + const [jsonInput, setJsonInput] = useState(''); + const [isSubmitting, setIsSubmitting] = useState(false); + + const sampleJson = { + models: [ + { + id: 'model-id-1', + name: 'Model Name 1', + description: 'Description...', + created: Math.floor(Date.now() / 1000), + context_length: 8192, + architecture: { + modality: 'text', + input_modalities: ['text'], + output_modalities: ['text'], + tokenizer: '', + instruct_type: null, + }, + pricing: { + prompt: 0.0, + completion: 0.0, + request: 0.0, + image: 0.0, + web_search: 0.0, + internal_reasoning: 0.0, + }, + enabled: true, + }, + ], + }; + + const handleBatchOverride = async () => { + if (!jsonInput.trim()) { + toast.error('Please enter JSON content'); + return; + } + + setIsSubmitting(true); + try { + let data; + try { + data = JSON.parse(jsonInput); + } catch { + throw new Error('Invalid JSON format'); + } + + if (!data.models || !Array.isArray(data.models)) { + throw new Error('JSON match follow structure: { "models": [...] }'); + } + + const result = await AdminService.batchOverrideProviderModels( + providerId, + data.models + ); + + if (result.ok) { + toast.success(result.message || 'Batch override successful'); + onSuccess(); + onClose(); + setJsonInput(''); + } else { + throw new Error('Batch override failed'); + } + } catch (error: unknown) { + const message = + error instanceof Error ? error.message : 'Batch override failed'; + toast.error(message); + } finally { + setIsSubmitting(false); + } + }; + + return ( + !open && onClose()}> + + + + + Batch Override Models + + + Paste a JSON object with a "models" array containing model + definitions. Existing models with the same ID will be updated. + + + +
+