cache x-cashu tokens

This commit is contained in:
9qeklajc
2026-03-11 22:03:05 +01:00
parent 4857d741de
commit 8a373276ce
7 changed files with 296 additions and 30 deletions
@@ -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')
+22 -1
View File
@@ -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"],
+46
View File
@@ -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__ = (
+7 -1
View File
@@ -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)
+1
View File
@@ -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")
+130 -27
View File
@@ -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
+56 -1
View File
@@ -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]