From 8a373276ce932cf80df44cbe5dbf369eb77f1b63 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 11 Mar 2026 22:03:05 +0100 Subject: [PATCH] cache x-cashu tokens --- .../a776ca70e5fe_add_cashu_refunds_table.py | 34 ++++ routstr/balance.py | 23 ++- routstr/core/db.py | 46 +++++ routstr/core/main.py | 8 +- routstr/core/settings.py | 1 + routstr/upstream/base.py | 157 +++++++++++++++--- routstr/wallet.py | 57 ++++++- 7 files changed, 296 insertions(+), 30 deletions(-) create mode 100644 migrations/versions/a776ca70e5fe_add_cashu_refunds_table.py 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..5b595cb0 --- /dev/null +++ b/migrations/versions/a776ca70e5fe_add_cashu_refunds_table.py @@ -0,0 +1,34 @@ +"""add cashu_refunds 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_refunds', + sa.Column('payment_token_hash', sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column('refund_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('created_at', sa.Integer(), nullable=False), + sa.Column('collected', sa.Boolean(), nullable=False), + sa.Column('swept', sa.Boolean(), nullable=False), + sa.PrimaryKeyConstraint('payment_token_hash'), + ) + + +def downgrade() -> None: + op.drop_table('cashu_refunds') diff --git a/routstr/balance.py b/routstr/balance.py index b738a6c3..fa8c580f 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, CashuRefund, 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(CashuRefund, 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.refund_token, + "amount": result.amount, + "unit": result.unit, + } + + @router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], diff --git a/routstr/core/db.py b/routstr/core/db.py index 549ac1f6..556d5c5e 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -128,6 +128,52 @@ class LightningInvoice(SQLModel, table=True): # type: ignore paid_at: int | None = Field(default=None, description="Unix timestamp when paid") +class CashuRefund(SQLModel, table=True): # type: ignore + __tablename__ = "cashu_refunds" + + payment_token_hash: str = Field( + primary_key=True, + description="SHA-256 hash of the original x-cashu payment token", + ) + refund_token: str = Field(description="Serialized Cashu refund token") + amount: int = Field(description="Refund 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 refund token" + ) + 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_refund( + payment_token_hash: str, + refund_token: str, + amount: int, + unit: str, + mint_url: str | None = None, +) -> None: + try: + async with create_session() as session: + refund = CashuRefund( + payment_token_hash=payment_token_hash, + refund_token=refund_token, + amount=amount, + unit=unit, + mint_url=mint_url, + ) + session.add(refund) + await session.commit() + except Exception as e: + logger.warning( + "Failed to store cashu refund", + extra={"error": str(e), "payment_token_hash": payment_token_hash}, + ) + + class UpstreamProviderRow(SQLModel, table=True): # type: ignore __tablename__ = "upstream_providers" __table_args__ = ( diff --git a/routstr/core/main.py b/routstr/core/main.py index 461f8b6d..b4b6ac8a 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -18,7 +18,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 @@ -50,6 +50,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 @@ -107,6 +108,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 @@ -140,6 +142,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 = [] @@ -161,6 +165,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) diff --git a/routstr/core/settings.py b/routstr/core/settings.py index fa1464ac..f52f8b04 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") diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 9b66a480..faa89871 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_refund, +) 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,20 @@ 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, + ) -> 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 Returns: Refund token string @@ -1715,6 +1723,18 @@ class BaseUpstreamProvider: }, ) + if payment_token_hash: + try: + await store_cashu_refund( + payment_token_hash=payment_token_hash, + refund_token=refund_token, + amount=amount, + unit=unit, + mint_url=mint, + ) + except Exception: + pass # store_cashu_refund already logs + return refund_token except Exception as e: last_exception = e @@ -1764,6 +1784,7 @@ class BaseUpstreamProvider: unit: str, max_cost_for_model: int, mint: str | None = None, + payment_token_hash: str | None = None, ) -> StreamingResponse: """Handle streaming response for X-Cashu payment, calculating refund if needed. @@ -1773,6 +1794,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 +1866,9 @@ 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 + ) response_headers["X-Cashu"] = refund_token logger.info( @@ -1897,6 +1921,7 @@ class BaseUpstreamProvider: unit: str, max_cost_for_model: int, mint: str | None = None, + payment_token_hash: str | None = None, ) -> Response: """Handle non-streaming response for X-Cashu payment, calculating refund if needed. @@ -1906,6 +1931,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 +1993,9 @@ 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 + ) response_headers["X-Cashu"] = refund_token logger.info( @@ -2003,6 +2031,13 @@ class BaseUpstreamProvider: emergency_refund = amount refund_token = await send_token(emergency_refund, unit=unit, mint_url=mint) response.headers["X-Cashu"] = refund_token + if payment_token_hash: + try: + await store_cashu_refund( + payment_token_hash, refund_token, emergency_refund, unit, mint + ) + except Exception: + pass logger.warning( "Emergency refund issued due to JSON parse error", @@ -2027,6 +2062,7 @@ class BaseUpstreamProvider: unit: str, max_cost_for_model: int, mint: str | None = None, + payment_token_hash: str | None = None, ) -> StreamingResponse | Response: """Handle chat completion response for X-Cashu payment, detecting streaming vs non-streaming. @@ -2063,11 +2099,23 @@ 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, ) 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, ) except Exception as e: @@ -2096,6 +2144,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 +2215,9 @@ 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 + ) logger.info( "Refund processed for failed upstream request", @@ -2204,7 +2255,12 @@ 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, ) background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) @@ -2279,6 +2335,7 @@ 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)) @@ -2297,6 +2354,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, mint, + payment_token_hash, ) except Exception as e: error_message = str(e) @@ -2356,6 +2414,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 +2486,9 @@ 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 + ) logger.info( "Refund processed for failed upstream Responses API request", @@ -2465,7 +2526,12 @@ 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, ) background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) @@ -2515,6 +2581,7 @@ class BaseUpstreamProvider: unit: str, max_cost_for_model: int, mint: str | None = None, + payment_token_hash: str | None = None, ) -> StreamingResponse | Response: """Handle Responses API completion response for X-Cashu payment. @@ -2552,11 +2619,23 @@ 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, ) 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, ) except Exception as e: @@ -2583,6 +2662,7 @@ class BaseUpstreamProvider: unit: str, max_cost_for_model: int, mint: str | None = None, + payment_token_hash: str | None = None, ) -> StreamingResponse: """Handle streaming Responses API response for X-Cashu payment. @@ -2664,7 +2744,9 @@ 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 + ) response_headers["X-Cashu"] = refund_token logger.info( @@ -2717,6 +2799,7 @@ class BaseUpstreamProvider: unit: str, max_cost_for_model: int, mint: str | None = None, + payment_token_hash: str | None = None, ) -> Response: """Handle non-streaming Responses API response for X-Cashu payment.""" logger.debug( @@ -2776,7 +2859,9 @@ 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 + ) response_headers["X-Cashu"] = refund_token logger.info( @@ -2812,6 +2897,13 @@ class BaseUpstreamProvider: emergency_refund = amount refund_token = await send_token(emergency_refund, unit=unit, mint_url=mint) response.headers["X-Cashu"] = refund_token + if payment_token_hash: + try: + await store_cashu_refund( + payment_token_hash, refund_token, emergency_refund, unit, mint + ) + except Exception: + pass logger.warning( "Emergency refund issued for Responses API due to JSON parse error", @@ -2861,6 +2953,7 @@ 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)) @@ -2879,6 +2972,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, mint, + payment_token_hash, ) except Exception as e: error_message = str(e) @@ -3090,7 +3184,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 +3205,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/wallet.py b/routstr/wallet.py index 71ea18ac..f7bbcd73 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 @@ -352,6 +353,60 @@ 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.CashuRefund).where( + db.CashuRefund.collected == False, # noqa: E712 + db.CashuRefund.swept == False, # noqa: E712 + db.CashuRefund.created_at < cutoff, + ) + results = await session.exec(stmt) + refunds = results.all() + + for refund in refunds: + try: + await recieve_token(refund.refund_token) + refund.swept = True + session.add(refund) + logger.info( + "Swept uncollected refund", + extra={ + "payment_token_hash": refund.payment_token_hash, + "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={ + "payment_token_hash": refund.payment_token_hash, + }, + ) + else: + logger.warning( + "Failed to sweep refund", + extra={ + "payment_token_hash": refund.payment_token_hash, + "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]