diff --git a/migrations/versions/a776ca70e5fe_add_cashu_refunds_table.py b/migrations/versions/a776ca70e5fe_add_cashu_refunds_table.py new file mode 100644 index 00000000..6342f25c --- /dev/null +++ b/migrations/versions/a776ca70e5fe_add_cashu_refunds_table.py @@ -0,0 +1,42 @@ +"""add cashu_transactions table + +Revision ID: a776ca70e5fe +Revises: 614c0a740e68 +Create Date: 2026-03-11 22:00:01.554762 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +# revision identifiers, used by Alembic. +revision = "a776ca70e5fe" +down_revision = "614c0a740e68" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.create_table( + "cashu_transactions", + sa.Column("id", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("token", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("amount", sa.Integer(), nullable=False), + sa.Column("unit", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("mint_url", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column( + "type", + sqlmodel.sql.sqltypes.AutoString(), + nullable=False, + server_default="out", + ), + sa.Column("request_id", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("created_at", sa.Integer(), nullable=False), + sa.Column("collected", sa.Boolean(), nullable=False), + sa.Column("swept", sa.Boolean(), nullable=False), + sa.PrimaryKeyConstraint("id"), + ) + + +def downgrade() -> None: + op.drop_table("cashu_transactions") diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 87770a25..e3efa261 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -126,11 +126,24 @@ def create_model_mappings( if isinstance(db_id, int): providers_by_db_id[db_id] = upstream + # Group upstreams by URL and keep only the one with the lowest fee for each URL + upstreams_by_url: dict[str, list["BaseUpstreamProvider"]] = {} + for upstream in upstreams: + url = getattr(upstream, "base_url", "") + if url not in upstreams_by_url: + upstreams_by_url[url] = [] + upstreams_by_url[url].append(upstream) + + filtered_upstreams: list["BaseUpstreamProvider"] = [] + for providers in upstreams_by_url.values(): + best_provider = min(providers, key=lambda p: p.provider_fee) + filtered_upstreams.append(best_provider) + # Separate OpenRouter from other providers openrouter: "BaseUpstreamProvider" | None = None other_upstreams: list["BaseUpstreamProvider"] = [] - for upstream in upstreams: + for upstream in filtered_upstreams: base_url = getattr(upstream, "base_url", "") if base_url == "https://openrouter.ai/api/v1": openrouter = upstream diff --git a/routstr/auth.py b/routstr/auth.py index 62b6c5b8..f844f1e7 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -372,12 +372,12 @@ async def validate_bearer_key( }, ) + key_preview = bearer_key[:10] + "..." if len(bearer_key) > 10 else bearer_key logger.error( - "Invalid API key format", + f"Invalid API key format: preview={key_preview!r} length={len(bearer_key)} " + f"(expected 'sk-...' or 'cashu...' token)", extra={ - "key_preview": bearer_key[:10] + "..." - if len(bearer_key) > 10 - else bearer_key, + "key_preview": key_preview, "key_length": len(bearer_key), }, ) @@ -386,7 +386,7 @@ async def validate_bearer_key( status_code=401, detail={ "error": { - "message": "Invalid API key", + "message": "Invalid API key format. Expected an 'sk-...' API key or a 'cashu...' token.", "type": "invalid_request_error", "code": "invalid_api_key", } diff --git a/routstr/balance.py b/routstr/balance.py index b738a6c3..7edede9b 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -9,7 +9,7 @@ from pydantic import BaseModel from sqlmodel import select from .auth import get_billing_key, validate_bearer_key -from .core.db import ApiKey, AsyncSession, get_session +from .core.db import ApiKey, AsyncSession, CashuTransaction, get_session from .core.logging import get_logger from .core.settings import settings from .lightning import lightning_router @@ -401,6 +401,27 @@ async def reset_child_key_spent( return {"success": True, "message": "Child key balance reset successfully."} +@router.get("/cashu-refund/{payment_token_hash}") +async def get_cashu_refund( + payment_token_hash: str, + session: AsyncSession = Depends(get_session), +) -> dict: + """Retrieve a stored Cashu refund token by the hash of the original payment token.""" + result = await session.get(CashuTransaction, payment_token_hash) + if result is None: + raise HTTPException(status_code=404, detail="Refund not found") + if result.swept: + raise HTTPException(status_code=410, detail="Refund has been swept") + result.collected = True + session.add(result) + await session.commit() + return { + "refund_token": result.token, + "amount": result.amount, + "unit": result.unit, + } + + @router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], diff --git a/routstr/core/admin.py b/routstr/core/admin.py index b278c00d..46589164 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1,3 +1,4 @@ +import asyncio import json import secrets from datetime import datetime, timezone @@ -16,7 +17,13 @@ from ..wallet import ( send_token, slow_filter_spend_proofs, ) -from .db import ApiKey, ModelRow, UpstreamProviderRow, create_session +from .db import ( + ApiKey, + CashuTransaction, + ModelRow, + UpstreamProviderRow, + create_session, +) from .log_manager import log_manager from .logging import get_logger from .settings import SettingsService, settings @@ -749,7 +756,9 @@ async def get_provider_models(provider_id: int) -> dict[str, object]: ) db_model_ids = {model.id for model in db_models} - filtered_remote_models = [m for m in upstream_models if m.id not in db_model_ids] + filtered_remote_models = [ + m for m in upstream_models if m.id not in db_model_ids + ] return { "provider": { @@ -867,12 +876,6 @@ async def initiate_provider_topup( if not provider: raise HTTPException(status_code=404, detail="Provider not found") - upstream_instance = _instantiate_provider(provider) - if not upstream_instance: - raise HTTPException( - status_code=400, detail="Could not instantiate provider" - ) - try: logger.info( f"Initiating top-up for provider {provider_id}", @@ -886,39 +889,69 @@ async def initiate_provider_topup( async with httpx.AsyncClient() as client: clean_url = provider.base_url.rstrip("/") - # Proxy the request to upstream Routstr - # Use the actual API key from the database - resp = await client.post( - f"{clean_url}/v1/balance/lightning/invoice", - json={ - "amount_sats": int(payload.amount), - "purpose": "topup", - "api_key": provider.api_key, - }, - headers={"Authorization": f"Bearer {provider.api_key}"} if provider.api_key else {}, + request_json = { + "amount_sats": int(payload.amount), + "purpose": "topup", + "api_key": provider.api_key, + } + headers = ( + {"Authorization": f"Bearer {provider.api_key}"} + if provider.api_key + else {} ) - if resp.status_code == 200: - data = resp.json() - return { - "ok": True, - "topup_data": { - "payment_request": data.get("bolt11"), - "invoice_id": data.get("invoice_id"), - "status": "pending", - }, - } - else: - logger.error(f"Upstream topup request failed: {resp.text}") - # Check if it's JSON error - try: - error_detail = resp.json() - except Exception: - error_detail = resp.text - raise HTTPException( - status_code=resp.status_code, detail=error_detail + last_status_code = 500 + last_error_detail: object = "Failed to create top-up invoice" + + # Some upstream Routstr nodes fail the first invoice request after warm-up + # and succeed immediately on retry. Retry once here so the UI stays single-click. + for attempt in range(2): + resp = await client.post( + f"{clean_url}/v1/balance/lightning/invoice", + json=request_json, + headers=headers, ) + if resp.status_code == 200: + data = resp.json() + return { + "ok": True, + "topup_data": { + "payment_request": data.get("bolt11"), + "invoice_id": data.get("invoice_id"), + "status": "pending", + }, + } + + logger.error( + f"Upstream topup request failed: {resp.text}", + extra={ + "provider_id": provider_id, + "attempt": attempt + 1, + "status_code": resp.status_code, + }, + ) + try: + last_error_detail = resp.json() + except Exception: + last_error_detail = resp.text + last_status_code = resp.status_code + + if resp.status_code < 500 or attempt == 1: + break + + await asyncio.sleep(0.2) + + raise HTTPException( + status_code=last_status_code, detail=last_error_detail + ) + + upstream_instance = _instantiate_provider(provider) + if not upstream_instance: + raise HTTPException( + status_code=400, detail="Could not instantiate provider" + ) + topup_data = await upstream_instance.initiate_topup(payload.amount) logger.info( @@ -979,7 +1012,9 @@ async def check_topup_status(provider_id: int, invoice_id: str) -> dict[str, obj clean_url = provider.base_url.rstrip("/") resp = await client.get( f"{clean_url}/v1/balance/lightning/invoice/{invoice_id}/status", - headers={"Authorization": f"Bearer {provider.api_key}"} if provider.api_key else {}, + headers={"Authorization": f"Bearer {provider.api_key}"} + if provider.api_key + else {}, ) if resp.status_code == 200: status_data = resp.json() @@ -1027,15 +1062,46 @@ async def get_provider_balance(provider_id: int) -> dict[str, object]: if provider.provider_type == "routstr": import httpx - async with httpx.AsyncClient() as client: - clean_url = provider.base_url.rstrip("/") - headers = {} - if provider.api_key: - headers["Authorization"] = f"Bearer {provider.api_key}" - resp = await client.get( - f"{clean_url}/v1/balance/info", - headers=headers, - ) + clean_url = provider.base_url.rstrip("/") + headers = {} + if provider.api_key: + headers["Authorization"] = f"Bearer {provider.api_key}" + + async with httpx.AsyncClient(timeout=10.0) as client: + try: + resp = await client.get( + f"{clean_url}/v1/balance/info", + headers=headers, + ) + except httpx.TimeoutException as exc: + logger.error( + "Timed out fetching Routstr provider balance", + extra={ + "provider_id": provider_id, + "base_url": clean_url, + "upstream_url": f"{clean_url}/v1/balance/info", + "error": str(exc), + }, + ) + raise HTTPException( + status_code=504, + detail="Timed out contacting upstream Routstr provider", + ) from exc + except httpx.RequestError as exc: + logger.error( + "Failed to fetch Routstr provider balance", + extra={ + "provider_id": provider_id, + "base_url": clean_url, + "upstream_url": f"{clean_url}/v1/balance/info", + "error": str(exc), + }, + ) + raise HTTPException( + status_code=502, + detail="Failed to contact upstream Routstr provider", + ) from exc + if resp.status_code == 200: data = resp.json() # Return balance in sats @@ -1261,6 +1327,49 @@ async def get_log_dates_api(request: Request) -> dict[str, object]: return {"dates": dates} +@admin_router.get("/api/transactions", dependencies=[Depends(require_admin_api)]) +async def get_transactions_api( + type: str | None = None, + status: str | None = None, + search: str | None = None, + limit: int = 100, +) -> dict: + async with create_session() as session: + from sqlmodel import col + + stmt = select(CashuTransaction) + if type: + stmt = stmt.where(CashuTransaction.type == type) + if status: + if status == "collected": + stmt = stmt.where(CashuTransaction.collected == True) # noqa: E712 + elif status == "swept": + stmt = stmt.where(CashuTransaction.swept == True) # noqa: E712 + elif status == "pending": + stmt = stmt.where( + CashuTransaction.collected == False, # noqa: E712 + CashuTransaction.swept == False, # noqa: E712 + ) + + if search: + search_pattern = f"%{search}%" + stmt = stmt.where( + (col(CashuTransaction.id).like(search_pattern)) + | (col(CashuTransaction.token).like(search_pattern)) + | (col(CashuTransaction.request_id).like(search_pattern)) + ) + + stmt = stmt.order_by(col(CashuTransaction.created_at).desc()).limit(limit) + + results = await session.exec(stmt) + transactions = results.all() + + return { + "transactions": [tx.dict() for tx in transactions], + "total": len(transactions), + } + + @admin_router.post( "/api/upstream-providers/{provider_id}/routstr/refund", dependencies=[Depends(require_admin_api)], diff --git a/routstr/core/db.py b/routstr/core/db.py index 549ac1f6..90258dbb 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -2,11 +2,13 @@ import os import pathlib import sqlite3 import time +import uuid from contextlib import asynccontextmanager from typing import AsyncGenerator from alembic import command from alembic.config import Config +from alembic.util.exc import CommandError from sqlalchemy import UniqueConstraint from sqlalchemy.ext.asyncio.engine import create_async_engine from sqlmodel import Field, Relationship, SQLModel, func, select, update @@ -128,6 +130,59 @@ class LightningInvoice(SQLModel, table=True): # type: ignore paid_at: int | None = Field(default=None, description="Unix timestamp when paid") +class CashuTransaction(SQLModel, table=True): # type: ignore + __tablename__ = "cashu_transactions" + + id: str = Field( + primary_key=True, + default_factory=lambda: uuid.uuid4().hex, + description="Unique transaction identifier", + ) + token: str = Field(description="Serialized Cashu token") + amount: int = Field(description="Amount in the token's unit") + unit: str = Field(description="Token unit (sat or msat)") + mint_url: str | None = Field(default=None, description="Mint URL for the token") + type: str = Field(default="out", description="Transaction type: in or out") + request_id: str | None = Field(default=None, description="Associated request ID") + created_at: int = Field( + default_factory=lambda: int(time.time()), + description="Unix timestamp", + ) + collected: bool = Field(default=False) + swept: bool = Field(default=False) + + +async def store_cashu_transaction( + token: str, + amount: int, + unit: str, + mint_url: str | None = None, + typ: str = "out", + request_id: str | None = None, + collected: bool = False, + created_at: int | None = None, +) -> None: + try: + async with create_session() as session: + tx = CashuTransaction( + token=token, + amount=amount, + unit=unit, + mint_url=mint_url, + type=typ, + request_id=request_id, + collected=collected, + created_at=created_at or int(time.time()), + ) + session.add(tx) + await session.commit() + except Exception as e: + logger.warning( + f"Failed to store cashu transaction: {e} (type={typ})", + extra={"error": str(e), "type": typ}, + ) + + class UpstreamProviderRow(SQLModel, table=True): # type: ignore __tablename__ = "upstream_providers" __table_args__ = ( @@ -227,6 +282,17 @@ def fix_cashu_migrations() -> None: logger.warning(f"Could not check/fix Cashu database {db_file}: {e}") +def _clear_alembic_version() -> None: + """Clear the alembic_version table so stamp/upgrade can proceed.""" + sync_url = DATABASE_URL.replace("+aiosqlite", "") + from sqlalchemy import create_engine, text + + eng = create_engine(sync_url) + with eng.begin() as conn: + conn.execute(text("DELETE FROM alembic_version")) + eng.dispose() + + def run_migrations() -> None: """Run Alembic migrations programmatically.""" try: @@ -248,8 +314,19 @@ def run_migrations() -> None: # Set the database URL in the config alembic_cfg.set_main_option("sqlalchemy.url", DATABASE_URL) - # Run migrations to the latest revision - command.upgrade(alembic_cfg, "head") + try: + command.upgrade(alembic_cfg, "head") + except CommandError as e: + if "Can't locate revision" in str(e): + logger.warning( + "Database stamped with unknown revision (likely from another branch). " + "Re-stamping to current head.", + extra={"error": str(e)}, + ) + _clear_alembic_version() + command.stamp(alembic_cfg, "head") + else: + raise logger.info("Database migrations completed successfully") diff --git a/routstr/core/main.py b/routstr/core/main.py index c906204e..d0365dad 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -22,7 +22,7 @@ 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 ..upstream.auto_topup import periodic_auto_topup -from ..wallet import periodic_payout +from ..wallet import periodic_payout, periodic_refund_sweep from .admin import admin_router from .db import create_session, init_db, run_migrations from .exceptions import general_exception_handler, http_exception_handler @@ -55,6 +55,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: model_maps_refresh_task = None key_reset_task = None auto_topup_task = None + refund_sweep_task = None try: # Run database migrations on startup @@ -113,6 +114,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: providers_task = asyncio.create_task(providers_cache_refresher()) key_reset_task = asyncio.create_task(periodic_key_reset()) auto_topup_task = asyncio.create_task(periodic_auto_topup()) + refund_sweep_task = asyncio.create_task(periodic_refund_sweep()) yield @@ -148,6 +150,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: key_reset_task.cancel() if auto_topup_task is not None: auto_topup_task.cancel() + if refund_sweep_task is not None: + refund_sweep_task.cancel() try: tasks_to_wait = [] @@ -171,6 +175,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: tasks_to_wait.append(key_reset_task) if auto_topup_task is not None: tasks_to_wait.append(auto_topup_task) + if refund_sweep_task is not None: + tasks_to_wait.append(refund_sweep_task) if tasks_to_wait: await asyncio.gather(*tasks_to_wait, return_exceptions=True) @@ -191,7 +197,7 @@ app.add_middleware( allow_credentials=True, allow_methods=["*"], allow_headers=["*"], - expose_headers=["x-routstr-request-id"], + expose_headers=["x-routstr-request-id", "x-cashu"], ) # Add logging middleware diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 20a3053c..027cc0ba 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -74,6 +74,7 @@ class Settings(BaseSettings): enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH") enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH") refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") + refund_sweep_ttl_seconds: int = Field(default=86400, env="REFUND_SWEEP_TTL_SECONDS") # Logging log_level: str = Field(default="INFO", env="LOG_LEVEL") @@ -254,9 +255,10 @@ class SettingsService: db_json_raw = {} db_json = _normalize_settings_data(db_json_raw) + valid_fields = set(env_resolved.dict().keys()) merged_dict: dict[str, Any] = dict(env_resolved.dict()) merged_dict.update( - {k: v for k, v in db_json.items() if v not in (None, "", [], {})} + {k: v for k, v in db_json.items() if v not in (None, "", [], {}) and k in valid_fields} ) merged_dict = Settings(**merged_dict).dict() @@ -322,14 +324,11 @@ class SettingsService: if row is None: raise RuntimeError("Settings row missing") (data_str,) = row - data_raw = ( - json.loads(data_str) if isinstance(data_str, str) else dict(data_str) - ) - if not isinstance(data_raw, dict): - data_raw = {} - data = _normalize_settings_data(data_raw) + data = json.loads(data_str) if isinstance(data_str, str) else dict(data_str) + valid_fields = set(settings.dict().keys()) # Update in-place for k, v in data.items(): - setattr(settings, k, v) + if k in valid_fields: + setattr(settings, k, v) cls._current = settings return settings diff --git a/routstr/proxy.py b/routstr/proxy.py index b9ba6e1d..c884c896 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -455,15 +455,14 @@ async def get_bearer_token_key( ) return key except Exception as e: + key_preview = bearer_key[:20] + "..." if len(bearer_key) > 20 else bearer_key logger.error( - "Bearer token validation failed", + f"Bearer token validation failed: {type(e).__name__}: {e} path={path} key={key_preview!r}", extra={ "error": str(e), "error_type": type(e).__name__, "path": path, - "bearer_key_preview": bearer_key[:20] + "..." - if len(bearer_key) > 20 - else bearer_key, + "bearer_key_preview": key_preview, }, ) raise diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 9b66a480..f7517c4b 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import hashlib import json import re import traceback @@ -15,7 +16,13 @@ from sqlmodel import select from ..auth import adjust_payment_for_tokens from ..core import get_logger -from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow, create_session +from ..core.db import ( + ApiKey, + AsyncSession, + UpstreamProviderRow, + create_session, + store_cashu_transaction, +) from ..core.exceptions import UpstreamError from ..payment.cost_calculation import ( CostData, @@ -203,9 +210,7 @@ class BaseUpstreamProvider: return path.replace("v1/", "", 1) return path - def get_request_base_url( - self, path: str, model_obj: Model | None = None - ) -> str: + def get_request_base_url(self, path: str, model_obj: Model | None = None) -> str: """Get upstream base URL used when building forwarding URL.""" return self.base_url.rstrip("/") @@ -1325,9 +1330,7 @@ class BaseUpstreamProvider: ) # Don't revert here — proxy.py owns payment revert to avoid double-revert - raise UpstreamError( - "An unexpected server error occurred", status_code=500 - ) + raise UpstreamError("An unexpected server error occurred", status_code=500) async def forward_responses_request( self, @@ -1539,9 +1542,7 @@ class BaseUpstreamProvider: ) # Don't revert here — proxy.py owns payment revert to avoid double-revert - raise UpstreamError( - "An unexpected server error occurred", status_code=500 - ) + raise UpstreamError("An unexpected server error occurred", status_code=500) async def forward_get_request( self, @@ -1679,13 +1680,22 @@ class BaseUpstreamProvider: ) return None - async def send_refund(self, amount: int, unit: str, mint: str | None = None) -> str: + async def send_refund( + self, + amount: int, + unit: str, + mint: str | None = None, + payment_token_hash: str | None = None, + request_id: str | None = None, + ) -> str: """Create and send a refund token to the user. Args: amount: Refund amount unit: Unit of the refund (sat or msat) mint: Optional mint URL for the refund token + payment_token_hash: Optional SHA-256 hash of the original payment token for storage + request_id: Optional HTTP request ID for tracking Returns: Refund token string @@ -1715,6 +1725,18 @@ class BaseUpstreamProvider: }, ) + try: + await store_cashu_transaction( + token=refund_token, + amount=amount, + unit=unit, + mint_url=mint, + typ="out", + request_id=request_id, + ) + except Exception: + pass # store_cashu_transaction already logs + return refund_token except Exception as e: last_exception = e @@ -1764,6 +1786,8 @@ class BaseUpstreamProvider: unit: str, max_cost_for_model: int, mint: str | None = None, + payment_token_hash: str | None = None, + request_id: str | None = None, ) -> StreamingResponse: """Handle streaming response for X-Cashu payment, calculating refund if needed. @@ -1773,6 +1797,7 @@ class BaseUpstreamProvider: amount: Payment amount received unit: Payment unit (sat or msat) max_cost_for_model: Maximum cost for the model + payment_token_hash: Optional hash of original payment token for refund storage Returns: StreamingResponse with refund token in header if applicable @@ -1844,7 +1869,10 @@ class BaseUpstreamProvider: }, ) - refund_token = await self.send_refund(refund_amount, unit, mint) + refund_token = await self.send_refund( + refund_amount, unit, mint, payment_token_hash, + request_id=request_id, + ) response_headers["X-Cashu"] = refund_token logger.info( @@ -1897,6 +1925,8 @@ class BaseUpstreamProvider: unit: str, max_cost_for_model: int, mint: str | None = None, + payment_token_hash: str | None = None, + request_id: str | None = None, ) -> Response: """Handle non-streaming response for X-Cashu payment, calculating refund if needed. @@ -1906,6 +1936,7 @@ class BaseUpstreamProvider: amount: Payment amount received unit: Payment unit (sat or msat) max_cost_for_model: Maximum cost for the model + payment_token_hash: Optional hash of original payment token for refund storage Returns: Response with refund token in header if applicable @@ -1967,7 +1998,10 @@ class BaseUpstreamProvider: ) if refund_amount > 0: - refund_token = await self.send_refund(refund_amount, unit, mint) + refund_token = await self.send_refund( + refund_amount, unit, mint, payment_token_hash, + request_id=request_id, + ) response_headers["X-Cashu"] = refund_token logger.info( @@ -2003,6 +2037,17 @@ class BaseUpstreamProvider: emergency_refund = amount refund_token = await send_token(emergency_refund, unit=unit, mint_url=mint) response.headers["X-Cashu"] = refund_token + try: + await store_cashu_transaction( + token=refund_token, + amount=emergency_refund, + unit=unit, + mint_url=mint, + typ="out", + request_id=request_id, + ) + except Exception: + pass logger.warning( "Emergency refund issued due to JSON parse error", @@ -2027,6 +2072,8 @@ class BaseUpstreamProvider: unit: str, max_cost_for_model: int, mint: str | None = None, + payment_token_hash: str | None = None, + request_id: str | None = None, ) -> StreamingResponse | Response: """Handle chat completion response for X-Cashu payment, detecting streaming vs non-streaming. @@ -2063,11 +2110,25 @@ class BaseUpstreamProvider: if is_streaming: return await self.handle_x_cashu_streaming_response( - content_str, response, amount, unit, max_cost_for_model, mint + content_str, + response, + amount, + unit, + max_cost_for_model, + mint, + payment_token_hash, + request_id=request_id, ) else: return await self.handle_x_cashu_non_streaming_response( - content_str, response, amount, unit, max_cost_for_model, mint + content_str, + response, + amount, + unit, + max_cost_for_model, + mint, + payment_token_hash, + request_id=request_id, ) except Exception as e: @@ -2096,6 +2157,7 @@ class BaseUpstreamProvider: max_cost_for_model: int, model_obj: Model, mint: str | None = None, + payment_token_hash: str | None = None, ) -> Response | StreamingResponse: """Forward request paid with X-Cashu token to upstream service. @@ -2166,7 +2228,10 @@ class BaseUpstreamProvider: }, ) - refund_token = await self.send_refund(amount - 60, unit, mint) + refund_token = await self.send_refund( + amount - 60, unit, mint, payment_token_hash, + request_id=getattr(request.state, "request_id", None), + ) logger.info( "Refund processed for failed upstream request", @@ -2204,7 +2269,13 @@ class BaseUpstreamProvider: ) result = await self.handle_x_cashu_chat_completion( - response, amount, unit, max_cost_for_model, mint + response, + amount, + unit, + max_cost_for_model, + mint, + payment_token_hash, + request_id=getattr(request.state, "request_id", None), ) background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) @@ -2279,10 +2350,25 @@ class BaseUpstreamProvider: ) try: + payment_token_hash = hashlib.sha256(x_cashu_token.encode()).hexdigest() headers = dict(request.headers) amount, unit, mint = await recieve_token(x_cashu_token) headers = self.prepare_headers(dict(request.headers)) + request_id = getattr(request.state, "request_id", None) + try: + await store_cashu_transaction( + token=x_cashu_token, + amount=amount, + unit=unit, + mint_url=mint, + typ="in", + request_id=request_id, + collected=True, + ) + except Exception: + pass + logger.info( "X-Cashu token redeemed for Responses API", extra={"amount": amount, "unit": unit, "path": path, "mint": mint}, @@ -2297,6 +2383,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, mint, + payment_token_hash, ) except Exception as e: error_message = str(e) @@ -2356,6 +2443,7 @@ class BaseUpstreamProvider: max_cost_for_model: int, model_obj: Model, mint: str | None = None, + payment_token_hash: str | None = None, ) -> Response | StreamingResponse: """Forward Responses API request paid with X-Cashu token to upstream service. @@ -2427,7 +2515,10 @@ class BaseUpstreamProvider: }, ) - refund_token = await self.send_refund(amount - 60, unit, mint) + refund_token = await self.send_refund( + amount - 60, unit, mint, payment_token_hash, + request_id=getattr(request.state, "request_id", None), + ) logger.info( "Refund processed for failed upstream Responses API request", @@ -2465,7 +2556,13 @@ class BaseUpstreamProvider: ) result = await self.handle_x_cashu_responses_completion( - response, amount, unit, max_cost_for_model, mint + response, + amount, + unit, + max_cost_for_model, + mint, + payment_token_hash, + request_id=getattr(request.state, "request_id", None), ) background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) @@ -2515,6 +2612,8 @@ class BaseUpstreamProvider: unit: str, max_cost_for_model: int, mint: str | None = None, + payment_token_hash: str | None = None, + request_id: str | None = None, ) -> StreamingResponse | Response: """Handle Responses API completion response for X-Cashu payment. @@ -2552,11 +2651,25 @@ class BaseUpstreamProvider: if is_streaming: return await self.handle_x_cashu_streaming_responses_response( - content_str, response, amount, unit, max_cost_for_model, mint + content_str, + response, + amount, + unit, + max_cost_for_model, + mint, + payment_token_hash, + request_id=request_id, ) else: return await self.handle_x_cashu_non_streaming_responses_response( - content_str, response, amount, unit, max_cost_for_model, mint + content_str, + response, + amount, + unit, + max_cost_for_model, + mint, + payment_token_hash, + request_id=request_id, ) except Exception as e: @@ -2583,6 +2696,8 @@ class BaseUpstreamProvider: unit: str, max_cost_for_model: int, mint: str | None = None, + payment_token_hash: str | None = None, + request_id: str | None = None, ) -> StreamingResponse: """Handle streaming Responses API response for X-Cashu payment. @@ -2664,7 +2779,10 @@ class BaseUpstreamProvider: }, ) - refund_token = await self.send_refund(refund_amount, unit, mint) + refund_token = await self.send_refund( + refund_amount, unit, mint, payment_token_hash, + request_id=request_id, + ) response_headers["X-Cashu"] = refund_token logger.info( @@ -2717,6 +2835,8 @@ class BaseUpstreamProvider: unit: str, max_cost_for_model: int, mint: str | None = None, + payment_token_hash: str | None = None, + request_id: str | None = None, ) -> Response: """Handle non-streaming Responses API response for X-Cashu payment.""" logger.debug( @@ -2776,7 +2896,10 @@ class BaseUpstreamProvider: ) if refund_amount > 0: - refund_token = await self.send_refund(refund_amount, unit, mint) + refund_token = await self.send_refund( + refund_amount, unit, mint, payment_token_hash, + request_id=request_id, + ) response_headers["X-Cashu"] = refund_token logger.info( @@ -2812,6 +2935,17 @@ class BaseUpstreamProvider: emergency_refund = amount refund_token = await send_token(emergency_refund, unit=unit, mint_url=mint) response.headers["X-Cashu"] = refund_token + try: + await store_cashu_transaction( + token=refund_token, + amount=emergency_refund, + unit=unit, + mint_url=mint, + typ="out", + request_id=request_id, + ) + except Exception: + pass logger.warning( "Emergency refund issued for Responses API due to JSON parse error", @@ -2861,10 +2995,25 @@ class BaseUpstreamProvider: ) try: + payment_token_hash = hashlib.sha256(x_cashu_token.encode()).hexdigest() headers = dict(request.headers) amount, unit, mint = await recieve_token(x_cashu_token) headers = self.prepare_headers(dict(request.headers)) + request_id = getattr(request.state, "request_id", None) + try: + await store_cashu_transaction( + token=x_cashu_token, + amount=amount, + unit=unit, + mint_url=mint, + typ="in", + request_id=request_id, + collected=True, + ) + except Exception: + pass + logger.info( "X-Cashu token redeemed successfully", extra={"amount": amount, "unit": unit, "path": path, "mint": mint}, @@ -2879,6 +3028,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, mint, + payment_token_hash, ) except Exception as e: error_message = str(e) @@ -3090,7 +3240,7 @@ class BaseUpstreamProvider: async with create_session() as session: stmt = select(UpstreamProviderRow).where( UpstreamProviderRow.base_url == self.base_url, - UpstreamProviderRow.api_key == self.api_key + UpstreamProviderRow.api_key == self.api_key, ) result = await session.exec(stmt) @@ -3111,15 +3261,24 @@ class BaseUpstreamProvider: diff = set(db_model_ids) - set(model_ids) for db_model_id in diff: - found_db_model = next((model_obj for model_obj in db_models if model_obj.id == db_model_id)) + found_db_model = next( + ( + model_obj + for model_obj in db_models + if model_obj.id == db_model_id + ) + ) models.append(found_db_model) - models_with_fees = [self._apply_provider_fee_to_model(m) for m in models] + models_with_fees = [ + self._apply_provider_fee_to_model(m) for m in models + ] try: sats_to_usd = sats_usd_price() self._models_cache = [ - _update_model_sats_pricing(m, sats_to_usd) for m in models_with_fees + _update_model_sats_pricing(m, sats_to_usd) + for m in models_with_fees ] except Exception: self._models_cache = models_with_fees diff --git a/routstr/upstream/routstr.py b/routstr/upstream/routstr.py index abf82a33..ab6dd0bd 100644 --- a/routstr/upstream/routstr.py +++ b/routstr/upstream/routstr.py @@ -83,7 +83,7 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider): Balance in satoshis, or None if failed """ url = f"{self.base_url}/v1/balance/info" - headers = {"Authorization": f"Bearer {self.api_key}"} + headers = {"Authorization": f"Bearer {self.api_key}"} if self.api_key else {} async with httpx.AsyncClient() as client: try: diff --git a/routstr/wallet.py b/routstr/wallet.py index 71ea18ac..74ea5fa8 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -1,11 +1,12 @@ import asyncio import math +import time from typing import TypedDict from cashu.core.base import Proof, Token from cashu.wallet.helpers import deserialize_token_from_string from cashu.wallet.wallet import Wallet -from sqlmodel import col, update +from sqlmodel import col, select, update from .core import db, get_logger from .core.settings import settings @@ -34,6 +35,7 @@ async def recieve_token( wallet.verify_proofs_dleq(token_obj.proofs) await wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True) + return token_obj.amount, token_obj.unit, token_obj.mint @@ -352,6 +354,61 @@ async def periodic_payout() -> None: ) +async def periodic_refund_sweep() -> None: + while True: + await asyncio.sleep(60 * 60) # every hour + try: + cutoff = int(time.time()) - settings.refund_sweep_ttl_seconds + async with db.create_session() as session: + stmt = select(db.CashuTransaction).where( + db.CashuTransaction.type == "out", + db.CashuTransaction.collected == False, # noqa: E712 + db.CashuTransaction.swept == False, # noqa: E712 + db.CashuTransaction.created_at < cutoff, + ) + results = await session.exec(stmt) + refunds = results.all() + + for refund in refunds: + try: + await recieve_token(refund.token) + refund.swept = True + session.add(refund) + logger.info( + "Swept uncollected refund", + extra={ + "id": refund.id, + "amount": refund.amount, + "unit": refund.unit, + }, + ) + except Exception as e: + error_msg = str(e).lower() + if "already spent" in error_msg: + refund.swept = True + session.add(refund) + logger.info( + "Refund already spent (client collected), marking swept", + extra={ + "id": refund.id, + }, + ) + else: + logger.warning( + "Failed to sweep refund", + extra={ + "id": refund.id, + "error": str(e), + }, + ) + await session.commit() + except Exception as e: + logger.error( + "Error in periodic refund sweep", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + + async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int: wallet = await get_wallet(mint, unit) proofs = wallet._get_proofs_per_keyset(wallet.proofs)[wallet.keyset_id] diff --git a/tests/integration/test_admin_provider_balance.py b/tests/integration/test_admin_provider_balance.py new file mode 100644 index 00000000..c5e94fab --- /dev/null +++ b/tests/integration/test_admin_provider_balance.py @@ -0,0 +1,80 @@ +from datetime import datetime, timedelta, timezone +from unittest.mock import AsyncMock, patch + +import httpx +import pytest +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.admin import admin_sessions +from routstr.core.db import UpstreamProviderRow + + +async def _create_routstr_provider() -> UpstreamProviderRow: + return UpstreamProviderRow( + provider_type="routstr", + base_url="https://upstream.example", + api_key="", + enabled=True, + ) + + +def _admin_headers() -> dict[str, str]: + token = "test-admin-token" + admin_sessions[token] = int( + (datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp() + ) + return {"Authorization": f"Bearer {token}"} + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_admin_routstr_provider_balance_timeout_returns_504( + integration_client: httpx.AsyncClient, + integration_session: AsyncSession, +) -> None: + provider = await _create_routstr_provider() + integration_session.add(provider) + await integration_session.commit() + await integration_session.refresh(provider) + + request = httpx.Request("GET", f"{provider.base_url}/v1/balance/info") + timeout_error = httpx.ConnectTimeout("Connect timeout", request=request) + + with patch( + "httpx.AsyncHTTPTransport.handle_async_request", + new=AsyncMock(side_effect=timeout_error), + ): + response = await integration_client.get( + f"/admin/api/upstream-providers/{provider.id}/balance", + headers=_admin_headers(), + ) + + assert response.status_code == 504 + assert response.json()["detail"] == "Timed out contacting upstream Routstr provider" + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_admin_routstr_provider_balance_request_error_returns_502( + integration_client: httpx.AsyncClient, + integration_session: AsyncSession, +) -> None: + provider = await _create_routstr_provider() + integration_session.add(provider) + await integration_session.commit() + await integration_session.refresh(provider) + + request = httpx.Request("GET", f"{provider.base_url}/v1/balance/info") + request_error = httpx.ConnectError("Connection failed", request=request) + + with patch( + "httpx.AsyncHTTPTransport.handle_async_request", + new=AsyncMock(side_effect=request_error), + ): + response = await integration_client.get( + f"/admin/api/upstream-providers/{provider.id}/balance", + headers=_admin_headers(), + ) + + assert response.status_code == 502 + assert response.json()["detail"] == "Failed to contact upstream Routstr provider" diff --git a/tests/integration/test_provider_fee_enforcement.py b/tests/integration/test_provider_fee_enforcement.py new file mode 100644 index 00000000..ff93a231 --- /dev/null +++ b/tests/integration/test_provider_fee_enforcement.py @@ -0,0 +1,143 @@ +from contextlib import asynccontextmanager +from typing import Any, AsyncGenerator, cast +from unittest.mock import patch + +import pytest + +from routstr.core.db import AsyncSession, ModelRow, UpstreamProviderRow +from routstr.payment.models import Architecture, Model, Pricing +from routstr.proxy import refresh_model_maps +from routstr.upstream.base import BaseUpstreamProvider + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_enforce_lowest_provider_fee_for_same_url( + integration_session: Any, +) -> None: + """Test that the algorithm selects the provider with the lowest fee when URLs match.""" + + # 1. Create two providers with the same URL but different fees + url = "https://api.example.com" + p1 = UpstreamProviderRow( + provider_type="custom", + base_url=url, + api_key="key1", + enabled=True, + provider_fee=1.01, + ) + p2 = UpstreamProviderRow( + provider_type="custom", + base_url=url, + api_key="key2", + enabled=True, + provider_fee=1.05, + ) + + integration_session.add(p1) + integration_session.add(p2) + await integration_session.commit() + await integration_session.refresh(p1) + await integration_session.refresh(p2) + + assert p1.id is not None + assert p2.id is not None + + # 2. Add a model for each provider + m1 = ModelRow( + id="model-a", + name="Model A", + created=1, + description="desc", + context_length=100, + architecture='{"modality": "text", "input_modalities": ["text"], "output_modalities": ["text"], "tokenizer": "tiktoken", "instruct_type": "chat"}', + pricing='{"prompt": 1.0, "completion": 1.0}', + upstream_provider_id=p1.id, + enabled=True, + ) + m2 = ModelRow( + id="model-a", + name="Model A", + created=1, + description="desc", + context_length=100, + architecture='{"modality": "text", "input_modalities": ["text"], "output_modalities": ["text"], "tokenizer": "tiktoken", "instruct_type": "chat"}', + pricing='{"prompt": 1.0, "completion": 1.0}', + upstream_provider_id=p2.id, + enabled=True, + ) + + integration_session.add(m1) + integration_session.add(m2) + await integration_session.commit() + + # 3. Create mock provider instances + class MockProvider(BaseUpstreamProvider): + db_id: int + + def __init__(self, db_id: int, base_url: str, api_key: str, fee: float): + super().__init__(base_url, api_key, fee) + self.db_id = db_id + self.provider_type = "custom" + + def get_cached_models(self) -> list[Model]: + return [ + Model( + id="model-a", + name="Model A", + created=1, + description="desc", + context_length=100, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="tiktoken", + instruct_type="chat", + ), + pricing=Pricing(prompt=1.0, completion=1.0), + enabled=True, + upstream_provider_id=self.db_id, + ) + ] + + async def refresh_models_cache(self) -> None: + pass + + def prepare_headers(self, request_headers: dict[str, str]) -> dict[str, str]: + return request_headers + + # 4. Inject mock providers into the proxy + from routstr import proxy + + assert p1.id is not None + assert p2.id is not None + + # Need to patch proxy._upstreams and proxy.create_session + mp1: MockProvider = MockProvider(p1.id, url, "key1", 1.01) + mp2: MockProvider = MockProvider(p2.id, url, "key2", 1.05) + + with ( + patch("routstr.proxy._upstreams", [mp1, mp2]), + patch("routstr.proxy.create_session") as mock_session_factory, + ): + # Configure mock_session_factory to return a session that uses the test engine + @asynccontextmanager + async def mock_create_session() -> AsyncGenerator[AsyncSession, None]: + yield integration_session + + mock_session_factory.return_value = mock_create_session() + + await refresh_model_maps() + + # 5. Check which provider is selected for 'model-a' + provider_map = proxy.get_provider_for_model("model-a") + + # Assertions + assert provider_map is not None + assert len(provider_map) >= 1 + + # Check the first one, cast to MockProvider to access db_id + best_provider = cast(MockProvider, provider_map[0]) + assert best_provider.db_id == p1.id + assert best_provider.provider_fee == 1.01 diff --git a/tests/integration/test_provider_management.py b/tests/integration/test_provider_management.py index b47ad848..b7db0c6c 100644 --- a/tests/integration/test_provider_management.py +++ b/tests/integration/test_provider_management.py @@ -3,12 +3,16 @@ Integration tests for provider management functionality. Tests GET /v1/providers/ endpoint for listing and managing providers. """ +import time +from types import TracebackType from typing import Any, Generator from unittest.mock import patch import pytest from httpx import AsyncClient +from routstr.core.admin import admin_sessions +from routstr.core.db import UpstreamProviderRow from routstr.nostr.discovery import _PROVIDERS_CACHE from .utils import ResponseValidator @@ -678,3 +682,88 @@ async def test_no_database_changes_during_provider_operations( assert final_diff["api_keys"]["added"] == [] assert final_diff["api_keys"]["modified"] == [] assert final_diff["api_keys"]["removed"] == [] + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_admin_routstr_topup_retries_transient_upstream_failure( + integration_client: AsyncClient, + integration_session: Any, +) -> None: + admin_token = "test-admin-token" + admin_sessions[admin_token] = int(time.time()) + 3600 + integration_client.headers["Authorization"] = f"Bearer {admin_token}" + + provider = UpstreamProviderRow( + provider_type="routstr", + base_url="https://node.example", + api_key="sk-upstream-test", + enabled=True, + provider_fee=1.01, + ) + integration_session.add(provider) + await integration_session.commit() + await integration_session.refresh(provider) + + class MockResponse: + def __init__(self, status_code: int, data: dict[str, Any] | None = None): + self.status_code = status_code + self._data = data or {} + self.text = str(self._data) + + def json(self) -> dict[str, Any]: + return self._data + + class MockAsyncClient: + def __init__(self) -> None: + self.calls = 0 + + async def __aenter__(self) -> "MockAsyncClient": + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + tb: TracebackType | None, + ) -> None: + return None + + async def post( + self, url: str, json: dict[str, Any], headers: dict[str, str] + ) -> MockResponse: + self.calls += 1 + assert url == "https://node.example/v1/balance/lightning/invoice" + assert json["amount_sats"] == 10 + assert json["purpose"] == "topup" + assert json["api_key"] == "sk-upstream-test" + assert headers["Authorization"] == "Bearer sk-upstream-test" + + if self.calls == 1: + return MockResponse(500, {"detail": "warmup failure"}) + + return MockResponse( + 200, + { + "bolt11": "lnbc1testinvoice", + "invoice_id": "invoice-123", + }, + ) + + mock_client = MockAsyncClient() + + try: + with patch("httpx.AsyncClient", return_value=mock_client): + response = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/topup", + json={"amount": 10}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["ok"] is True + assert data["topup_data"]["payment_request"] == "lnbc1testinvoice" + assert data["topup_data"]["invoice_id"] == "invoice-123" + assert mock_client.calls == 2 + finally: + admin_sessions.pop(admin_token, None) diff --git a/tests/unit/test_upstream_routstr.py b/tests/unit/test_upstream_routstr.py new file mode 100644 index 00000000..8cb2427d --- /dev/null +++ b/tests/unit/test_upstream_routstr.py @@ -0,0 +1,76 @@ +from types import TracebackType +from unittest.mock import Mock + +import httpx +import pytest + +from routstr.upstream.routstr import RoutstrUpstreamProvider + + +class DummyAsyncClient: + def __init__(self, response: Mock | None = None, error: Exception | None = None): + self.response = response + self.error = error + self.calls: list[dict[str, object]] = [] + + async def __aenter__(self) -> "DummyAsyncClient": + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + tb: TracebackType | None, + ) -> bool: + return False + + async def get( + self, url: str, headers: dict[str, str], timeout: float + ) -> Mock: + self.calls.append({"url": url, "headers": headers, "timeout": timeout}) + if self.error is not None: + raise self.error + assert self.response is not None + return self.response + + +@pytest.mark.asyncio +async def test_get_balance_omits_auth_header_when_api_key_missing( + monkeypatch: pytest.MonkeyPatch, +) -> None: + response = Mock() + response.json.return_value = {"balance_msats": 42000} + response.raise_for_status.return_value = None + + client = DummyAsyncClient(response=response) + monkeypatch.setattr("routstr.upstream.routstr.httpx.AsyncClient", lambda: client) + + provider = RoutstrUpstreamProvider(base_url="https://node.example", api_key="") + + balance = await provider.get_balance() + + assert balance == 42.0 + assert client.calls == [ + { + "url": "https://node.example/v1/balance/info", + "headers": {}, + "timeout": 10.0, + } + ] + + +@pytest.mark.asyncio +async def test_get_balance_returns_none_on_connect_timeout( + monkeypatch: pytest.MonkeyPatch, +) -> None: + client = DummyAsyncClient(error=httpx.ConnectTimeout("timed out")) + monkeypatch.setattr("routstr.upstream.routstr.httpx.AsyncClient", lambda: client) + + provider = RoutstrUpstreamProvider( + base_url="https://node.example", + api_key="secret", + ) + + balance = await provider.get_balance() + + assert balance is None diff --git a/ui/app/transactions/page.tsx b/ui/app/transactions/page.tsx new file mode 100644 index 00000000..07dc0e2a --- /dev/null +++ b/ui/app/transactions/page.tsx @@ -0,0 +1,366 @@ +'use client'; + +import { useState, useEffect } from 'react'; +import { useQuery } from '@tanstack/react-query'; +import { AppPageShell } from '@/components/app-page-shell'; +import { PageHeader } from '@/components/page-header'; +import { + Card, + CardContent, + CardDescription, + CardHeader, + CardTitle, +} from '@/components/ui/card'; +import { Button } from '@/components/ui/button'; +import { Input } from '@/components/ui/input'; +import { Label } from '@/components/ui/label'; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from '@/components/ui/select'; +import { Badge } from '@/components/ui/badge'; +import { + Table, + TableBody, + TableCell, + TableHead, + TableHeader, + TableRow, +} from '@/components/ui/table'; +import { ScrollArea } from '@/components/ui/scroll-area'; +import { Skeleton } from '@/components/ui/skeleton'; +import { + Empty, + EmptyDescription, + EmptyHeader, + EmptyMedia, + EmptyTitle, +} from '@/components/ui/empty'; +import { + RefreshCw, + Search, + ArrowDownLeft, + ArrowUpRight, + Copy, + Check, + Receipt, +} from 'lucide-react'; +import { AdminService, type Transaction } from '@/lib/api/services/admin'; +import { format } from 'date-fns'; +import { toast } from 'sonner'; + +const STORAGE_KEY = 'routstr-transaction-filters'; + +export default function TransactionsPage() { + const [search, setSearch] = useState(''); + const [type, setType] = useState('all'); + const [status, setStatus] = useState('all'); + const [copiedId, setCopiedId] = useState(null); + + // Load filters from localStorage on mount + useEffect(() => { + const saved = localStorage.getItem(STORAGE_KEY); + if (saved) { + try { + const parsed = JSON.parse(saved); + if (parsed.search) setSearch(parsed.search); + if (parsed.type) setType(parsed.type); + if (parsed.status) setStatus(parsed.status); + } catch (e) { + console.error('Failed to load filters from localStorage', e); + } + } + }, []); + + // Save filters to localStorage whenever they change + useEffect(() => { + const filters = { search, type, status }; + localStorage.setItem(STORAGE_KEY, JSON.stringify(filters)); + }, [search, type, status]); + + const { data, isLoading, refetch, isRefetching } = useQuery({ + queryKey: ['transactions', type, status, search], + queryFn: () => + AdminService.getTransactions( + type === 'all' ? undefined : type, + status === 'all' ? undefined : status, + search || undefined, + 100 + ), + }); + + const handleClearFilters = () => { + setSearch(''); + setType('all'); + setStatus('all'); + }; + + const copyToClipboard = (text: string, id: string) => { + navigator.clipboard.writeText(text); + setCopiedId(id); + toast.success('Copied to clipboard'); + setTimeout(() => setCopiedId(null), 2000); + }; + + const getStatusBadge = (tx: Transaction) => { + if (tx.swept) + return ( + + Swept + + ); + if (tx.collected) + return ( + + Collected + + ); + return ( + + Pending + + ); + }; + + const hasActiveFilters = + type !== 'all' || status !== 'all' || Boolean(search); + + const activeFilterDescription = [ + type !== 'all' ? `type ${type === 'in' ? 'incoming' : 'outgoing'}` : null, + status !== 'all' ? `status ${status}` : null, + search ? `search "${search}"` : null, + ] + .filter(Boolean) + .join(' • '); + + return ( + +
+ refetch()} + variant='outline' + size='sm' + disabled={isRefetching} + > + + Refresh + + } + /> + + + + Filters + + Filter transactions by type, status, or search text + + + +
+
+ +
+ + setSearch(e.target.value)} + /> +
+
+
+ + +
+
+ + +
+
+ +
+
+
+
+ + + +
+ Transaction History + {data && ( + + {data.transactions.length} entries + + )} +
+ {hasActiveFilters && ( + + Showing transactions filtered by {activeFilterDescription} + + )} +
+ + {isLoading ? ( +
+ {Array.from({ length: 8 }).map((_, index) => ( + + ))} +
+ ) : data?.transactions && data.transactions.length > 0 ? ( + + + + + Type + Amount + Status + Request ID + Mint + Date + Actions + + + + {data.transactions.map((tx) => ( + + +
+ {tx.type === 'in' ? ( + + ) : ( + + )} + {tx.type} +
+
+ + {tx.amount} {tx.unit} + + {getStatusBadge(tx)} + + {tx.request_id ? ( +
+ + {tx.request_id} + + +
+ ) : ( + + — + + )} +
+ +
+ {tx.mint_url} +
+
+ + {format(tx.created_at * 1000, 'yyyy-MM-dd HH:mm:ss')} + + + + +
+ ))} +
+
+
+ ) : ( + + + + + + No transactions found + + Try adjusting your filters or check back later. + + + + )} +
+
+
+
+ ); +} diff --git a/ui/components/app-page-shell.tsx b/ui/components/app-page-shell.tsx index ad2fdd31..6b613862 100644 --- a/ui/components/app-page-shell.tsx +++ b/ui/components/app-page-shell.tsx @@ -13,6 +13,7 @@ import { ServerIcon, SettingsIcon, WalletIcon, + ArrowRightLeftIcon, } from 'lucide-react'; import Image from 'next/image'; import { toast } from 'sonner'; @@ -41,6 +42,7 @@ const NAV_ITEMS = [ { title: 'Logs', url: '/logs', icon: FileTextIcon }, { title: 'Models', url: '/models', icon: DatabaseIcon }, { title: 'Providers', url: '/providers', icon: ServerIcon }, + { title: 'Transactions', url: '/transactions', icon: ArrowRightLeftIcon }, { title: 'Settings', url: '/settings', icon: SettingsIcon }, ] as const; diff --git a/ui/components/app-sidebar.tsx b/ui/components/app-sidebar.tsx index 76179f1c..7c777240 100644 --- a/ui/components/app-sidebar.tsx +++ b/ui/components/app-sidebar.tsx @@ -9,6 +9,7 @@ import { ServerIcon, SettingsIcon, WalletIcon, + ArrowRightLeftIcon, } from 'lucide-react'; import Image from 'next/image'; import Link from 'next/link'; @@ -45,6 +46,11 @@ const data = { url: '/balances', icon: WalletIcon, }, + { + title: 'Transactions', + url: '/transactions', + icon: ArrowRightLeftIcon, + }, { title: 'Logs', url: '/logs', diff --git a/ui/components/provider-balance.tsx b/ui/components/provider-balance.tsx index b6460d97..6ba57659 100644 --- a/ui/components/provider-balance.tsx +++ b/ui/components/provider-balance.tsx @@ -23,11 +23,15 @@ import { interface ProviderBalanceProps { providerId: number; platformUrl?: string | null; + isRoutstr?: boolean; + nodeUrl?: string; } export function ProviderBalance({ providerId, platformUrl, + isRoutstr = false, + nodeUrl, }: ProviderBalanceProps) { const [isTopupDialogOpen, setIsTopupDialogOpen] = useState(false); const [topupAmount, setTopupAmount] = useState(''); @@ -100,14 +104,23 @@ export function ProviderBalance({ }); const handleTopup = () => { - const amount = parseFloat(topupAmount); + const amount = Number(topupAmount); - if (isNaN(amount)) { - setTopupError('Please enter a valid amount'); + if (Number.isNaN(amount)) { + setTopupError( + isRoutstr + ? 'Please enter a valid amount in sats' + : 'Please enter a valid amount' + ); return; } - if (amount < 1 || amount > 500) { + if (isRoutstr) { + if (!Number.isInteger(amount) || amount < 1) { + setTopupError('Amount must be a whole number of sats'); + return; + } + } else if (amount < 1 || amount > 500) { setTopupError('Amount must be between $1 and $500'); return; } @@ -153,15 +166,24 @@ export function ProviderBalance({ let displayValue = 'N/A'; if (typeof balance === 'number') { - displayValue = `$${balance.toFixed(2)}`; + displayValue = isRoutstr + ? `${balance.toLocaleString()} sats` + : `$${balance.toFixed(2)}`; } else if (balance && typeof balance === 'object') { const b = balance as Record; if (typeof b.balance === 'number') { - displayValue = `$${b.balance.toFixed(2)}`; + displayValue = isRoutstr + ? `${b.balance.toLocaleString()} sats` + : `$${b.balance.toFixed(2)}`; } else if (typeof b.balance === 'string') { displayValue = b.balance; } else if (b.amount !== undefined) { - displayValue = `$${Number(b.amount).toFixed(2)}`; + const amount = Number(b.amount); + if (!Number.isNaN(amount)) { + displayValue = isRoutstr + ? `${amount.toLocaleString()} sats` + : `$${amount.toFixed(2)}`; + } } } @@ -193,7 +215,9 @@ export function ProviderBalance({ ? 'Your account balance has been updated.' : invoiceData ? 'Scan the QR code or copy the Lightning invoice to pay.' - : 'Enter the amount you want to add to your account balance.'} + : isRoutstr + ? `Top up your balance on node ${nodeUrl || ''}`.trim() + : 'Enter the amount you want to add to your account balance.'} @@ -254,20 +278,29 @@ export function ProviderBalance({ ) : (
- + { setTopupAmount(e.target.value); setTopupError(''); }} min='1' - max='500' - step='0.01' + max={isRoutstr ? undefined : '500'} + step={isRoutstr ? '1' : '0.01'} /> + {isRoutstr && ( +

+ The invoice amount will be created in sats. +

+ )} {topupError && (

{topupError}

)} diff --git a/ui/components/provider-card.tsx b/ui/components/provider-card.tsx index d239547e..b48dc64b 100644 --- a/ui/components/provider-card.tsx +++ b/ui/components/provider-card.tsx @@ -19,11 +19,15 @@ import { Pencil, Trash2, Key, + RotateCcw, } from 'lucide-react'; import { ProviderBalance } from '@/components/provider-balance'; import { ProviderModelsPanel } from '@/components/provider-models-panel'; import { RoutstrCreateKeySection } from '@/components/providers/RoutstrCreateKeySection'; +import { RoutstrProviderService } from '@/lib/api/services/routstr-provider'; +import { useMutation, useQueryClient } from '@tanstack/react-query'; import { useState } from 'react'; +import { toast } from 'sonner'; import { cn } from '@/lib/utils'; import { Dialog, @@ -70,10 +74,29 @@ export function ProviderCard({ onDeleteModel, onOverrideModel, onUpdateApiKey, - availableMints, }: ProviderCardProps) { + const queryClient = useQueryClient(); const [isKeyModalOpen, setIsKeyModalOpen] = useState(false); const hasDetails = Boolean(provider.api_version) || isExpanded; + const isRoutstr = provider.provider_type === 'routstr'; + + const refundMutation = useMutation({ + mutationFn: () => RoutstrProviderService.refundBalance(provider.id), + onSuccess: (data) => { + if (data.ok) { + toast.success('Refund successful', { description: data.message }); + queryClient.invalidateQueries({ + queryKey: ['provider-balance', provider.id], + }); + queryClient.invalidateQueries({ queryKey: ['balances'] }); + } else { + toast.error('Refund failed', { description: data.message }); + } + }, + onError: (error: Error) => { + toast.error(`Refund error: ${error.message}`); + }, + }); return ( @@ -107,11 +130,13 @@ export function ProviderCard({
)} - {provider.provider_type === 'routstr' && ( + {isRoutstr && ( + )} +