Merge remote-tracking branch 'origin/main' into feat/lightning-error-envelope

# Conflicts:
#	routstr/core/main.py
This commit is contained in:
9qeklajc
2026-09-18 21:23:17 +02:00
181 changed files with 24044 additions and 8621 deletions
+2
View File
@@ -33,6 +33,8 @@ ROUTSTR_SECRET_KEY=
# DATABASE_POOL_PRE_PING=false
# Warn when a checkout is held this many seconds.
# DATABASE_POOL_HOLD_WARN_SECONDS=10
# SQLite write-lock timeout, in seconds.
# DATABASE_BUSY_TIMEOUT=30
# SQLite serialises writes; increasing its pool can trade pool timeouts for
# "database is locked" errors rather than increasing write throughput.
+5 -5
View File
@@ -11,7 +11,7 @@ jobs:
runs-on: ubuntu-latest
strategy:
matrix:
python-version: ["3.11", "3.12"]
python-version: ["3.11", "3.12", "3.14"]
steps:
- name: Checkout code
@@ -25,22 +25,22 @@ jobs:
- name: Install dependencies
run: |
uv sync --dev
uv sync --python ${{ matrix.python-version }} --dev
- name: Run linting with ruff
run: |
uv run ruff check .
uv run --python ${{ matrix.python-version }} ruff check .
- name: Run type checking with mypy
run: |
uv run mypy .
uv run --python ${{ matrix.python-version }} mypy .
- name: Run tests with pytest
env:
UPSTREAM_BASE_URL: "http://test"
UPSTREAM_API_KEY: "test"
run: |
uv run pytest --verbose --tb=short
uv run --python ${{ matrix.python-version }} pytest --verbose --tb=short
- name: Upload test results
if: always()
+3
View File
@@ -11,6 +11,9 @@ dist/
*.egg
.mypy_cache/**
# MkDocs build output
site/
# Development
.notes
.*keys.db
+1 -1
View File
@@ -1 +1 @@
3.11
3.14
+3 -1
View File
@@ -1,10 +1,12 @@
FROM ghcr.io/astral-sh/uv:python3.11-bookworm-slim
ARG PYTHON_VERSION=3.14
FROM ghcr.io/astral-sh/uv:python${PYTHON_VERSION}-bookworm-slim
RUN apt-get update \
&& apt-get install -y --no-install-recommends \
git \
build-essential \
pkg-config \
libffi-dev \
libsecp256k1-dev \
autoconf \
automake \
+4 -1
View File
@@ -1,4 +1,6 @@
# Multi-stage Dockerfile for Routstr (includes UI build)
ARG PYTHON_VERSION=3.14
# Stage 1: Build the UI
FROM node:23-alpine AS ui-builder
WORKDIR /app/ui
@@ -16,13 +18,14 @@ ENV NEXT_TELEMETRY_DISABLED=1
RUN pnpm run build
# Stage 2: Build the Routstr Node
FROM ghcr.io/astral-sh/uv:python3.11-bookworm-slim AS runner
FROM ghcr.io/astral-sh/uv:python${PYTHON_VERSION}-bookworm-slim AS runner
RUN apt-get update \
&& apt-get install -y --no-install-recommends \
git \
build-essential \
pkg-config \
libffi-dev \
libsecp256k1-dev \
autoconf \
automake \
-33
View File
@@ -307,23 +307,6 @@ ANALYTICS_KEY = os.getenv("ROUTSTR_ANALYTICS_KEY")
api_key = PROD_KEY if is_production() else DEV_KEY
```
### Delegated Authentication
Create sub-keys with limited permissions:
```bash
POST /v1/wallet/create/subkey
Authorization: Bearer sk-parent-key
Content-Type: application/json
{
"name": "Limited Subkey",
"balance_limit": 1000,
"allowed_models": ["gpt-3.5-turbo"],
"expires_in_hours": 24
}
```
## Rate Limiting
Rate limits are applied per API key:
@@ -388,22 +371,6 @@ Content-Type: application/json
## Monitoring
### Usage Alerts
Set up usage notifications:
```bash
POST /v1/wallet/alerts
Authorization: Bearer sk-...
Content-Type: application/json
{
"low_balance_threshold": 1000,
"daily_spend_limit": 5000,
"webhook_url": "https://your-app.com/webhook"
}
```
### Audit Logging
All API key usage is logged:
+64 -57
View File
@@ -418,7 +418,7 @@ POST /v1/wallet/create
### Get Key Information
Get current balance, consumption data, and child keys for an API key.
Get current balance and consumption data for an API key.
```http
GET /v1/balance/info
@@ -432,23 +432,9 @@ Authorization: Bearer sk-...
"api_key": "sk-abc...",
"balance": 8500000,
"reserved": 0,
"is_child": false,
"parent_key": null,
"total_requests": 42,
"total_spent": 1500000,
"balance_limit": null,
"balance_limit_reset": null,
"validity_date": null,
"child_keys": [
{
"api_key": "sk-child1...",
"total_requests": 10,
"total_spent": 500000,
"balance_limit": 1000000,
"balance_limit_reset": "daily",
"validity_date": 1738000000
}
]
"validity_date": null
}
```
@@ -500,48 +486,23 @@ Authorization: Bearer sk-...
}
```
### Withdraw Funds
### Refund Balance
Withdraw balance as eCash.
Pay out the current balance. The key remains valid at zero balance and can be topped up again. The payout goes to a Lightning address when one is given (in the request or stored on the key), otherwise a Cashu token is returned.
```http
POST /v1/wallet/withdraw
POST /v1/balance/refund
Authorization: Bearer sk-...
Content-Type: application/json
```
**Request Body:**
`/v1/wallet/refund` is a deprecated alias.
**Request Body** (optional):
```json
{
"amount": 5000,
"mint": "https://mint.example.com"
}
```
**Response:**
```json
{
"cashu_token": "cashuAeyJ0...",
"amount": 5000,
"mint": "https://mint.example.com"
}
```
### Create Child Key
Creates one or more child API keys that share the parent's balance. Each child key creation costs a fixed amount (configurable).
```http
POST /v1/balance/child-key
Authorization: Bearer sk-...
```
**Request Body:**
```json
{
"count": 1
"lightning_address": "user@getalby.com"
}
```
@@ -549,21 +510,67 @@ Authorization: Bearer sk-...
| Parameter | Type | Required | Default | Description |
|-----------|------|----------|---------|-------------|
| `count` | integer | Yes | - | Number of child keys to create (1-50) |
| `lightning_address` | string | No | Key's stored refund address | Lightning address or LNURL to pay. Overrides the stored address for this request. The effective address (request or stored) is resolved only for a request that can open a new claim, before any balance is debited. |
**Response:**
**Response (Lightning):**
```json
{
"api_keys": ["sk-abc...", "sk-def..."],
"count": 2,
"cost_msats": 2000,
"cost_sats": 2,
"parent_balance": 98000,
"parent_balance_sats": 98
"refund_id": "3f9c1e2d8b7a4c6e9f0a1b2c3d4e5f60",
"status": "paid",
"recipient": "user@getalby.com",
"sats": "4500"
}
```
**Response (Cashu):**
```json
{
"refund_id": "3f9c1e2d8b7a4c6e9f0a1b2c3d4e5f60",
"status": "paid",
"token": "cashuAeyJ0...",
"sats": "4500"
}
```
The amount field is `sats` or `msats` depending on the key's refund currency. It reports the gross balance debited by the claim. For Lightning refunds, mint and input fees can reduce the amount actually delivered to the recipient.
**Behaviour:**
- The balance is debited and a refund claim is recorded before the payout is attempted. A key has at most one open claim at a time.
- If the payout fails cleanly, the claim is closed and the balance is restored. Retry the request.
- Once a melt quote has been recorded or a Cashu token has been issued, the payout may already have happened, so any later failure returns `502` and withholds the balance rather than restoring it. The exception is the mint answering the melt itself with `unpaid`: that is proof nothing was sent, so the balance is restored at once and the request returns `503`.
- If the Lightning payment is dispatched but the mint cannot confirm the outcome, the request returns `502`, the balance stays withheld, and a background reconciler asks the mint until it answers. The balance is restored if the mint reports the payment unpaid.
- An unresolved claim is reported before any replay: a request on a key with an open claim returns `409` with that claim's `refund_id` and `status`.
- Calling again on a zero-balance key with no open claim returns the last paid Lightning refund, or the Cashu token issued by the last paid claim while it remains uncollected.
**Errors:**
| Status | Meaning |
|--------|---------|
| `400` | Invalid Lightning destination, no balance, or balance too small for the refund unit |
| `400` | Ongoing requests are still reserving balance on this key |
| `401` | Unknown key |
| `409` | Balance changed concurrently. Retry. |
| `409` | `refund_in_progress`: another refund claim for this key is still open. The body carries its `refund_id` and `status` |
| `409` | `refund_unresolved`: a claim for this key is `stuck` and needs operator reconciliation |
| `410` | Previously issued Cashu refund token has been swept |
| `500` | Payout failed before anything was dispatched. Balance restored. Retry. |
| `502` | Payment dispatched, outcome unconfirmed. Balance withheld pending reconciliation. Do not retry. |
| `503` | Mint unavailable, or the mint reported the Lightning payment unpaid. Balance restored. Retry later. |
**X-Cashu refunds:**
Requests paid per-call with an `X-Cashu` header get their change from this endpoint by sending the same header instead of `Authorization`:
```http
POST /v1/balance/refund
X-Cashu: cashuAeyJ0...
```
Returns the change token in the body and in an `X-Cashu` response header. `404` if no matching request exists, `425` while the change is still being minted, `410` if it was swept.
## Provider Discovery
## Admin Settings
+12 -8
View File
@@ -153,28 +153,32 @@ granularity) on any of them.
|--------|--------|--------|-----------|---------|
| `token_already_spent` | 400 | `cashu_token_already_spent` | No | The token was already redeemed. |
| `invalid_token` | 400 | `invalid_cashu_token` | No | The token is malformed or cannot be decoded. |
| `mint_error` | 422 | `cashu_token_swap_fees_exceed_amount` | No | Token value is too small to cover the mint's swap/melt fees. |
| `mint_error` | 422 | `cashu_foreign_mint_swap_failed` | No | Swapping the token from a foreign mint to the primary mint failed. |
| `mint_error` | 422 | `cashu_token_swap_fees_exceed_amount` | No | Token value is too small to cover the mint's NUT-02 input fees. |
| `untrusted_mint` | 400 | `cashu_untrusted_source_mint` | No | The token was issued by a mint this node does not accept. Only the node's configured mints (`PRIMARY_MINT_URL` / `CASHU_MINTS`) are redeemable. |
| `mint_unreachable` | 503 | `cashu_source_mint_unreachable` | **Yes** | The mint that issued the token could not be reached; it cannot be redeemed at another mint. |
| `mint_rate_limited` | 503 | `cashu_mint_rate_limited` | **Yes** | The mint rate-limited the request; retry after the cooldown. |
| `mint_unreachable` | 503 | `cashu_mint_unreachable` | **Yes** | The mint could not be reached (DNS failure, refused/reset connection, timeout). The token is fine — retry once the mint recovers. |
| `mint_timeout` | 503 | `cashu_mint_timeout` | **Yes** | The mint did not respond in time; retry later. |
| `mint_unreachable` | 503 | `cashu_mint_unreachable` | **Yes** | The mint could not be reached (DNS failure, refused/reset connection). The token is fine — retry once the mint recovers. |
| `cashu_error` | 400 | `cashu_token_redemption_failed` | No | The token could not be redeemed for another expected reason. |
| `cashu_error` | 400 | `cashu_token_zero_value` | No | The token redeemed to zero (empty/dust token, or value fully consumed by fees). |
| `token_consumed` | 500 | `cashu_token_consumed` | No | The token was **spent** (melted/redeemed) but crediting it then failed. Do not retry — the token is gone; contact support to reconcile. |
| `api_error` | 500 | `internal_error` | Maybe | Unexpected server-side fault during redemption. |
!!! important "Retry only transient mint failures"
Only `mint_unreachable` and `mint_rate_limited` (503) are retryable — the
same token may work again later. Everything else is a permanent property of
the token and must not be blindly retried. Use exponential backoff for the
Only `mint_unreachable`, `mint_rate_limited` and `mint_timeout` (503) are
retryable — the same token may work again later. Everything else is a
permanent property of the token and must not be blindly retried.
`untrusted_mint` is permanent: the node will never accept that mint until
an operator adds it to `CASHU_MINTS`. Use exponential backoff for the
503 responses, and honor the mint's cooldown for `mint_rate_limited`. In
particular, a `token_consumed` 500 means the mint already spent the token,
so a retry would fail as `token_already_spent`.
#### Mint failures (retryable)
`mint_unreachable` and `mint_rate_limited` are retryable redemption errors. For
`mint_rate_limited`, honor the mint's cooldown before retrying.
`mint_unreachable`, `mint_rate_limited` and `mint_timeout` are retryable
redemption errors. For `mint_rate_limited`, honor the mint's cooldown before
retrying.
```json
{
+3 -2
View File
@@ -40,8 +40,9 @@ Every time you make a request to `/v1/chat/completions` (or others), the cost is
`Cost = (Input_Tokens * Price_Input) + (Output_Tokens * Price_Output) + Request_Fee`
- Prices are defined per model (see `/v1/models`).
- If you stream the response, the balance is deducted incrementally or finalized at the end of the stream.
- If your balance hits 0 mid-stream, the connection is closed.
- Routstr reserves an authorization ceiling before forwarding, then finalizes the request at measured token cost.
- If a successful upstream omits usage, Routstr estimates input tokens from the provider-bound request and output tokens from the returned body or streamed deltas, then applies normal model pricing.
- A reservation is only a temporary hold. Missing usage or unusable prices must never turn the full reservation into the charge; if no auditable estimate can be priced, the reservation is released without charge.
### Headers
+1 -1
View File
@@ -300,7 +300,7 @@ Project metadata and dependencies:
name = "routstr"
version = "0.2.2"
dependencies = [
"fastapi[standard]>=0.115",
"fastapi[standard-no-fastapi-cloud-cli]>=0.141",
"sqlmodel>=0.0.24",
"cashu",
# ...
+3 -3
View File
@@ -58,10 +58,10 @@ Contains the shared opaque EHBP transport and billing helpers:
- `EHBPForwardingTarget` — provider-specific target URL plus extra headers
- `forward_ehbp_request()` — forwards the encrypted body, captures Tinfoil
usage from a response header or streaming HTTP trailer, and finalizes bearer
billing at actual cost (falling back to max cost when usage is unavailable)
billing at actual cost (releasing the reservation when usage is unavailable)
- `forward_ehbp_x_cashu_request()` — redeems the Cashu token, refunds the full
token on upstream failure, and refunds the difference between the redeemed
amount and actual cost (or max cost when usage is unavailable)
amount and actual cost (or the full amount when usage is unavailable)
### Provider support
@@ -84,7 +84,7 @@ The proxy is a **blind relay** for EHBP requests. It cannot decrypt the body
Cost tracking happens at the proxy level. Routstr reserves or redeems up to
`max_cost_for_model`, then Tinfoil's out-of-band usage header/trailer allows it
to finalize at actual token cost. If trusted usage is missing or invalid, the
proxy safely falls back to max-cost billing.
proxy releases/refunds rather than treating the authorization ceiling as usage.
## End-to-end flow
+14
View File
@@ -26,6 +26,20 @@ If you want to run a node, resell API access, or monetize hardware.
---
## 👥 For Teams (Remote Nodes)
If you want one shared Routstr endpoint for a whole team, with per-member identities and per-member usage tracking.
- **[Overview](teams/index.md)**: What a remote node is, and when to use one.
- **[Deploy on Cloudron](teams/deploy-cloudron.md)**: The packaged, supported deployment.
- **[Deploy with Docker](teams/deploy-docker.md)**: Run it on any host behind your own TLS.
- **[Team Members](teams/team-members.md)**: Bootstrap the first admin and invite people.
- **[Connecting Clients](teams/clients.md)**: Wire up Claude Code, Pi, OpenCode, and API keys.
- **[Usage and Model Policy](teams/usage-and-policy.md)**: Per-member spend and model allowlists.
- **[Security Model](teams/security.md)**: Auth rules and endpoint scoping.
---
## 🔌 API Reference
- **[Overview](api/overview.md)**: Base URL, headers, and standards.
+13
View File
@@ -9,6 +9,19 @@ Routstr uses **Nostr** as a decentralized directory for service discovery. Your
1. **Provider Advertisement (Kind 38421)**: Your node periodically publishes an event with its URL, models, and pricing
2. **Client Discovery**: Clients query relays for these events to find suitable providers
### When announcements are published
Your node publishes an advertisement as soon as it has both a **Nsec** and at least one
reachable endpoint (a public `HTTP_URL`, or an `.onion` address). Saving the Nsec in the
dashboard is enough — the announcement follows within a minute, and **no restart is
required**. After the first publish it re-announces every 24 hours, and immediately
whenever the Nsec, endpoints, mints or relays change.
A node that has no Nsec yet simply waits, and starts announcing the moment one is
configured. Note that `HTTP_URL` defaults to `http://localhost:8000`, which is not a
reachable endpoint: a node with the default value and no onion address has nothing to
advertise and will not publish until one is set.
---
## Configuration
+153
View File
@@ -0,0 +1,153 @@
# Connecting Clients
A **client** is one agent or application talking to the node. Each client gets its own ID and its own API key (`sk-...`). Clients are how a team node attributes usage: every request is billed against the client that made it.
Members create and manage **their own** clients. They cannot see or touch anyone else's.
---
## Two credentials, two jobs
The node accepts two entirely different kinds of credential, and confusing them is the most common source of `403`s.
| Credential | Header | Purpose | Can do |
|---|---|---|---|
| **API key** | `Authorization: Bearer sk-...` | Inference | Send chat/completion requests. Cannot touch wallets, clients, or npubs. |
| **NIP-98** | `Authorization: Nostr <base64-event>` | Management | Manage clients, npubs, wallet, node control. Signed per-request by the member's `nsec`. |
Your agents use the **API key**. The `routstrd` CLI uses **NIP-98** automatically, which is why it needs your `nsec` in `~/.routstrd/config.json`.
---
## Add a client
From the member's own machine:
```bash
# Name it explicitly
routstrd clients add --name "My Laptop"
# Or use a one-shot integration setup
routstrd clients add --claude-code
routstrd clients add --pi-agent
routstrd clients add --opencode
routstrd clients add --openclaw
routstrd clients add --hermes
```
The integration flags configure the agent's own config file as well as registering the client, so you do not have to hand-edit anything. Several can be combined in one call.
On success the CLI prints the credentials and the endpoint to point at:
```text
Client created.
ID: my-laptop
Name: My Laptop
API Key: sk-9f2a...
Access Routstr at: https://team.example.com/v1
```
!!! warning "The API key is a secret"
Treat `sk-...` like a password. It bills inference to the team wallet. Do not commit it, and do not paste it into a chat — unlike an npub, it is not safe to share.
### Adding is idempotent
Running `clients add` with a name that already exists does not create a duplicate. It looks the client up and prints the existing record — including its API key — so re-running is a safe way to recover a key you lost:
```text
Client 'my-laptop' already exists.
ID: my-laptop
Name: My Laptop
API Key: sk-9f2a...
```
### List and delete
```bash
routstrd clients list
routstrd clients delete my-laptop
```
`clients list` shows **only your own** clients.
---
## Going beyond the CLI
The agent integrations cover the common tools, but any OpenAI-compatible client works — point it at the node and use the API key:
```bash
curl https://team.example.com/v1/chat/completions \
-H "Authorization: Bearer sk-9f2a..." \
-H "Content-Type: application/json" \
-d '{
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "hello"}]
}'
```
The base URL is always the node host plus `/v1`. Discover available models without any credential at all:
```bash
curl https://team.example.com/v1/models
```
---
## How client ownership works
This is the part that explains the odd-looking IDs on the node.
### IDs are derived from the name
A client's ID is the name lowercased with internal whitespace collapsed to hyphens, then stripped of anything that is not alphanumeric or a hyphen. `"My Laptop!"` becomes `my-laptop`.
### The node appends an owner suffix
So that two members can both have a client called `my-laptop` without colliding, the auth proxy appends the **last 7 characters of the owner's npub** to the ID before it reaches the daemon:
| Where you look | Client ID |
|---|---|
| Member's `routstrd clients list` | `my-laptop` |
| On the node (`cloudron exec`, then `routstrd clients list`) | `my-laptop-4f2x9k7` |
The suffix is stripped again on the way back, so members always see the clean ID. On the node you deliberately see the suffixed form — **the trailing characters are what tell you which member owns a client.**
### Ownership is recorded explicitly
Newly created clients store the owner's npub in an `ownerNpub` field, and the proxy authorises against that field. Clients created before that field existed fall back to matching the ID suffix, so older installs keep working until those clients are recreated.
### Consequence: admins are not automatically superusers here
`/clients`, `/clients/add`, and `/clients/delete` are owner-scoped by the calling npub. Even an `admin` cannot list or delete a colleague's clients through these endpoints. Cross-member visibility comes from running the CLI **on the node itself**, where the daemon is unauthenticated on loopback:
```bash
cloudron exec --app routstr.example.com
routstrd clients list # all clients, all owners, suffixed IDs
```
That is also the only practical way to clean up a departing member's keys — see [Team Members](team-members.md#what-revocation-does-and-does-not-do).
---
## Refreshing models and integrations
The `clients` command carries options for the daemon's scheduled refresh job, which updates the Routstr 21 model list and re-syncs client integrations:
```bash
routstrd clients --manual-refresh # refresh now, once
routstrd clients --disable-automatic-refresh # stop the scheduled job
routstrd clients --enable-automatic-refresh # start it again
```
The model list matters because it is also what the [model allowlist](usage-and-policy.md#model-allowlist) is enforced against.
---
## Next steps
- [Usage and Model Policy](usage-and-policy.md) — watch what those clients are spending.
- [Security Model](security.md) — the exact rules applied to each credential.
+169
View File
@@ -0,0 +1,169 @@
# Deploy on Cloudron
[Cloudron](https://www.cloudron.io/) is the supported deployment target for a team node. The packaged image already contains **both** processes — the `routstrd` daemon and the `routstrd-auth` proxy — supervised inside a single container, with `/app/data` handled as persistent storage and TLS terminated by the platform.
| | |
|---|---|
| **App ID** | `io.routstr.routstrd-auth` |
| **Public port** | `8008` (Cloudron proxies it over HTTPS on 443) |
| **Health check** | `GET /health` |
| **Memory limit** | 512 MB |
| **Minimum box version** | Cloudron 9.1.0 |
---
## Prerequisites
- A running Cloudron box with a domain that can get a certificate.
- The [`cloudron` CLI](https://docs.cloudron.io/cli/) installed and logged in, **only if** you are building the image yourself:
```bash
npm install -g cloudron
cloudron login my.example.com
```
## Install
### Option A — from the published version list
The app is published as a custom Cloudron app with a version list (`CloudronVersions.json`), currently at `0.1.26`. Once that app store entry is registered on your Cloudron instance, install it from the dashboard, or:
```bash
cloudron install --appstore-id io.routstr.routstrd-auth --location routstr.example.com
```
### Option B — build the image yourself
Use this when you want to run a local modification:
```bash
git clone https://github.com/routstr/routstrd-remote
cd routstrd-remote
cloudron build # builds the Dockerfile and pushes it to your registry
cloudron install --image <registry>/routstrd-remote:<tag> --location routstr.example.com
```
!!! note "The Dockerfile is the Cloudron image"
The repository's `Dockerfile` is built `FROM cloudron/base:5.0.0` and its `CMD` is `cloudron/start.sh`, which prepares `/app/data` and starts `supervisord`. It expects Cloudron's filesystem conventions and should not be confused with a generic Docker image. See [Deploy with Docker](deploy-docker.md) for what that means in practice.
---
## Bootstrap the first admin
The moment the app is healthy, the npub table is **empty**, and nothing except the public endpoints can be reached. Claim it before anything else — while the table is empty, `POST /npubs` is accepted without authentication, so this is the only window in which an unauthenticated registration succeeds.
On the machine of whoever will be the first admin:
```bash
bun i -g routstrd
routstrd remote https://routstr.example.com
routstrd npubs register --name "Alice"
```
`routstrd remote` generates a fresh Nostr identity if you do not have one, stores it in `~/.routstrd/config.json`, and prints your npub. `routstrd npubs register` then posts that npub and, because no npubs exist yet, receives `admin`.
!!! warning "Register immediately after install"
Until the first admin registers, anyone who knows the URL can claim the node. Do this as part of the install, not later.
Verify:
```bash
routstrd npubs list
```
---
## Configuration
Cloudron defaults are set by `cloudron/start.sh` and the two supervisor programs. Everything below can be overridden through the Cloudron **Environment Variables** tab.
| Variable | Default | Purpose |
|---|---|---|
| `ROUTSTRD_AUTH_PORT` | `8008` | Public port served by the auth proxy. Must match the manifest's `httpPort`. |
| `ROUTSTRD_AUTH_HOST` | `0.0.0.0` | Bind address of the auth proxy. |
| `ROUTSTRD_UPSTREAM` | `http://localhost:8009` | Where the daemon listens. Keep this on loopback. |
| `ROUTSTRD_PORT` | `8009` | Port the daemon binds. |
| `ROUTSTRD_DIR` | `/app/data/routstrd` | Config directory shared by the daemon and the proxy. |
| `ROUTSTRD_DB_PATH` | `/app/data/routstrd/routstr.db` | Shared SQLite database. |
| `ROUTSTRD_CONFIG_FILE` | `$ROUTSTRD_DIR/config.json` | Daemon config file. |
| `ROUTSTRD_AUTH_MODEL_ALLOWLIST` | `false` | Set to `true` to restrict the team to the Routstr 21 model list. See [Usage and Model Policy](usage-and-policy.md). |
| `ROUTSTRD_AUTH_ADMIN_NPUBS` | *(unset)* | Optional bootstrap admins. See below. |
### Bootstrapping admins from the environment
Instead of the interactive `npubs register` step you can seed admins declaratively. Three variables are accepted and merged: `ROUTSTRD_AUTH_ADMIN_NPUBS`, `ROUTSTRD_AUTH_ADMIN_PUBKEYS`, and `ROUTSTRD_AUTH_BOOTSTRAP_NPUB`. Values are comma- or whitespace-separated and may be either `npub1...` or 64-character hex.
Rows created this way are tagged `source = 'env'`. At every startup the proxy **reconciles** them: an env-sourced row whose pubkey is no longer present in the environment is **deleted**. This means the environment variables are the source of truth for those rows — removing someone from the variable revokes their access on the next restart.
!!! tip "Prefer `npubs register` for the first admin"
There is deliberately **no** hardcoded default admin npub in the image. An image with a baked-in admin pubkey would hand control of every deployment to the same key.
### Filesystem layout
| Path | Lifetime | Contents |
|---|---|---|
| `/app/code` | replaced on update | Auth proxy source and the `start.sh` / `run-auth.sh` scripts. |
| `/app/data` | persistent, backed up | `routstrd/config.json`, `routstrd/routstr.db`, `logs/`, and a `.initialized` marker. |
| `/run` | ephemeral | `supervisord` socket and pid file. |
The startup script writes `authUrl` into `routstrd/config.json` pointing at the local proxy, and generates the container's Nostr identity (`nsec`) on first boot if one is missing. That generated identity is what authorises the daemon's own calls back through the proxy.
---
## Backups
Cloudron's `localstorage` addon makes `/app/data` persistent and includes it in regular backups. The entire data directory is covered, which means the SQLite database at `/app/data/routstrd/routstr.db` travels with it.
To restore, restore the app from a Cloudron backup — the wallet configuration, npub table, and client records all come back together. Do not hand-copy files between hosts: the daemon's `config.json` holds the container's `nsec`, and losing it breaks the node's ability to authenticate to its own proxy.
!!! warning "Use `cloudron exec` carefully"
The database uses SQLite WAL mode. If you need a manual snapshot, stop the app first (`cloudron stop`) or use `sqlite3 .backup` inside the container — copying the `.db` file while the daemon is writing can produce an inconsistent file.
---
## Updates
```bash
cloudron update --app routstr.example.com
```
Updates are the primary lifecycle event on Cloudron and are designed to preserve `/app/data`. After the app comes back, confirm health and that your admin npub is still recognised:
```bash
curl https://routstr.example.com/health
routstrd npubs list
```
## Day-to-day operations
```bash
cloudron logs -f --app routstr.example.com # follow both processes' output
cloudron exec --app routstr.example.com # shell into the container
cloudron stop --app routstr.example.com
cloudron start --app routstr.example.com
cloudron debug --app routstr.example.com # read-write filesystem, app paused
cloudron debug --disable --app routstr.example.com
```
Inside the container you are on the machine that holds the wallet and the shared database, so `routstrd` commands there operate on the node itself rather than as a remote member:
```bash
routstrd npubs list # everyone with access, and their roles
routstrd clients list # every client on the node, not just your own
routstrd top # interactive usage TUI across all members
```
Both processes log to stdout/stderr and are collected by Cloudron; there are no log files to rotate inside `/app/data`.
### Failure handling
`supervisord` runs both programs with `autorestart=true` and a start priority that brings the **daemon up first** (priority 10) and the **proxy second** (priority 20). The proxy's launcher additionally waits — up to 120 seconds — for the database file to exist *and* for the daemon's `/health` to answer before it starts serving. A crash-looping proxy therefore usually means the daemon never became healthy.
---
## Next steps
- [Team Members](team-members.md) — invite the rest of your team.
- [Security Model](security.md) — what is exposed and what is not.
- [Troubleshooting](troubleshooting.md) — when the proxy will not start.
+170
View File
@@ -0,0 +1,170 @@
# Deploy with Docker
The team node is a single container running two supervised processes. You can run that same image on any Docker host and terminate TLS with your own reverse proxy.
!!! note "Read this first"
The image is built `FROM cloudron/base:5.0.0` and its entrypoint is the Cloudron startup script. It works outside Cloudron — it only requires a writable `/app/data` — but its filesystem conventions are Cloudron's, and the repository's `docker-compose.yml` predates this image. The commands below are the ones that match the current image.
---
## Build
```bash
git clone https://github.com/routstr/routstrd-remote
cd routstrd-remote
docker build -t routstr-remote:0.1.26 .
```
The Dockerfile installs Bun (the x64 **baseline** build, chosen because some hosts do not expose AVX/AVX2 and the default binary crashes with `SIGILL`), installs the `routstrd` daemon globally, installs the proxy's dependencies, and copies in the supervisor configuration for both processes.
## Run
```bash
mkdir -p "$HOME/routstr-remote-data"
docker run -d \
--name routstr-remote \
--restart unless-stopped \
-p 127.0.0.1:8008:8008 \
-v "$HOME/routstr-remote-data:/app/data" \
--memory 1g \
routstr-remote:0.1.26
```
**Why each flag matters:**
| Flag | Reason |
|---|---|
| `-v ...:/app/data` | The only persistent path. It holds `routstrd/config.json` (including the container's `nsec`), `routstrd/routstr.db`, and `logs/`. **Without it the node loses its identity and every npub on restart.** |
| `-p 127.0.0.1:8008:8008` | Publish on loopback only and let a TLS reverse proxy in front of it expose the service. Binding `0.0.0.0:8008` on a public host sends API keys over plaintext HTTP. |
| `--memory 1g` | Two Bun processes plus the daemon's model and usage state. The Cloudron manifest requests 512 MB; give a bare Docker host at least that, and prefer more. |
The container listens on exactly two ports: `8008` (public auth proxy) and `8009` (daemon, bound to loopback **inside** the container and deliberately not published).
Check it came up:
```bash
curl http://127.0.0.1:8008/health
docker logs -f routstr-remote
```
### Startup ordering
The entrypoint starts `supervisord`, which brings up the daemon first and the proxy second. The proxy's launcher then waits for the database file to appear and for `http://localhost:8009/health` to answer, retrying up to 120 times at one-second intervals. If that window expires, the proxy exits with `Timed out waiting for routstrd to become ready.` and restarts.
## Put TLS in front
Anything that terminates TLS and forwards to `127.0.0.1:8008` works. Two properties are worth configuring explicitly:
- **Disable response buffering.** The proxy already sends `X-Accel-Buffering: no` upstream, but your own proxy should also be configured not to buffer, otherwise streamed LLM responses appear to truncate.
- **Raise idle timeouts.** Model responses can be silent for a long time while reasoning or waiting on tools. The proxy disables Bun's per-request idle timeout for exactly this reason, but nginx's default 60-second `proxy_read_timeout` will still cut streams. Set it to something generous.
```nginx
location / {
proxy_pass http://127.0.0.1:8008;
proxy_http_version 1.1;
proxy_set_header Host $host;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
proxy_set_header X-Forwarded-Proto $scheme;
proxy_buffering off;
proxy_read_timeout 3600s;
proxy_send_timeout 3600s;
}
```
## Bootstrap the first admin
Identical to Cloudron — run this from the admin's own machine:
```bash
bun i -g routstrd
routstrd remote https://routstr.example.com
routstrd npubs register --name "Alice"
```
While the npub table is empty, `POST /npubs` needs no authentication, so the first registration wins. See [Team Members](team-members.md).
## Configuration
Override any of the variables from [Deploy on Cloudron](deploy-cloudron.md#configuration) with `-e`, for example:
```bash
docker run -d \
--name routstr-remote \
--restart unless-stopped \
-p 127.0.0.1:8008:8008 \
-v "$HOME/routstr-remote-data:/app/data" \
-e ROUTSTRD_AUTH_MODEL_ALLOWLIST=true \
routstr-remote:0.1.26
```
Inside the container the defaults are `ROUTSTRD_DIR=/app/data/routstrd`, `ROUTSTRD_DB_PATH=/app/data/routstrd/routstr.db`, `ROUTSTRD_UPSTREAM=http://localhost:8009`, `ROUTSTRD_AUTH_HOST=0.0.0.0`, `ROUTSTRD_AUTH_PORT=8008`.
!!! warning "Use a bind mount, not an anonymous volume"
If you recreate the container (`docker rm` then `docker run`), an unnamed volume is orphaned and the node comes back with a **new** Nostr identity and an **empty** npub table — meaning the next person to hit `POST /npubs` becomes admin. Always mount a known host directory.
## Backups
The database is SQLite in WAL mode, so stop the container before copying the data directory, or take a proper snapshot from inside:
```bash
docker stop routstr-remote
tar czf routstr-remote-$(date +%F).tar.gz -C "$HOME" routstr-remote-data
docker start routstr-remote
```
Or, without downtime:
```bash
docker exec routstr-remote sqlite3 /app/data/routstrd/routstr.db ".backup /app/data/routstrd/backup.db"
docker cp routstr-remote:/app/data/routstrd/backup.db .
```
## Updates
```bash
docker stop routstr-remote && docker rm routstr-remote
git pull
docker build -t routstr-remote:new .
docker run -d ... routstr-remote:new # same -v and -p flags as before
```
Because `/app/data` is a bind mount, the node's identity, npub table, and client records survive the swap. That is the entire reason the mount is not optional.
---
## Running from source (evaluation and development)
You do not need Docker to try the proxy. Point it at any routstrd daemon's database:
```bash
bun install
bun run src/index.ts validate # checks config and opens the DB
bun run src/index.ts start # binds 0.0.0.0:8008
```
Useful flags: `--port`, `--host`, `--upstream`, `--db-path`. The `validate` subcommand prints the effective configuration and reports how many npubs are registered, split by role — it is the fastest way to confirm the proxy can see the right database:
```text
Configuration:
Port: 8008
Host: 0.0.0.0
Upstream: http://localhost:8009
DB path: /app/data/routstrd/routstr.db
Bootstrap admin npubs/pubkeys from env: 0
Model allowlist: disabled
✅ DB accessible. 3 npub(s) registered (1 admin, 2 user).
```
!!! warning "`validate` fails before the daemon has run once"
If the database does not exist, validation stops with `Database not found at ... Make sure routstrd has been initialized`. The proxy shares the daemon's database; it never creates the schema itself.
---
## Next steps
- [Team Members](team-members.md) — invite the rest of your team.
- [Security Model](security.md) — which endpoints require what.
- [Troubleshooting](troubleshooting.md) — startup and streaming failures.
+112
View File
@@ -0,0 +1,112 @@
# Teams and Remote Nodes
A normal `routstrd` install is a **single-user daemon** running on your own machine. A **remote node** turns that same daemon into a **shared instance for a team**: one server runs the daemon behind an authentication proxy, and every team member gets their own Nostr identity, their own API keys, and their own usage accounting.
This is the product sometimes called **Routstrd Remote** — the repository is [`Routstr/routstrd-remote`](https://github.com/routstr/routstrd-remote) and the running service is the package `routstrd-auth`.
---
## Who this is for
- **Small teams and orgs** that want one funded Routstr endpoint instead of one daemon per laptop.
- **Anyone who wants per-person spend visibility** without building a billing system.
- **Self-hosters** who already run [Cloudron](https://www.cloudron.io/) and want a one-click install.
## What it gives you
| Capability | Detail |
|---|---|
| **One endpoint, many people** | Members point their coding agents at a single HTTPS URL. No per-machine setup beyond one CLI command. |
| **Per-member identity** | Each person has their own Nostr keypair (npub). Access is granted and revoked by adding or deleting that npub. |
| **Per-member attribution** | Every client registration gets a unique ID, and usage is reported per client, so you can see who is spending what. |
| **Two levels of privilege** | `admin` (manage people, move funds, control the node) and `user` (run inference, manage only their own clients). |
| **Scoped API keys** | Agent API keys can buy inference but cannot touch the wallet or other members' clients. |
| **Model policy** | Optional allowlist restricts the team to approved models. |
| **One wallet to fund** | The team tops up a single node wallet rather than N personal wallets. |
## What it is not
- It is **not** multi-tenant SaaS. Everyone shares the node's wallet and upstream provider set.
- It is **not** a billing or chargeback system. Usage is *attributed* per client; invoicing your teammates is up to you.
- It does **not** give each member a separate balance. See [Usage and Model Policy](usage-and-policy.md) for exactly what is tracked.
---
## Architecture
Two processes run together inside one container. Only one of them is reachable from outside.
```mermaid
flowchart TD
CLI["routstrd CLI<br/>on a member laptop"]
Agent["Coding agents<br/>Claude Code, Pi, OpenCode"]
App["App holding an sk- API key"]
TLS["Reverse proxy<br/>TLS termination on 443"]
Proxy["routstrd-auth<br/>0.0.0.0:8008 public"]
Daemon["routstrd daemon<br/>localhost:8009 no auth"]
DB[("routstr.db<br/>shared SQLite")]
Providers["Upstream model providers"]
CLI -->|https| TLS
Agent -->|https| TLS
App -->|https| TLS
TLS -->|http| Proxy
Proxy -->|forward| Daemon
Proxy -->|npubs and clients| DB
Daemon -->|usage and models| DB
Daemon -->|inference| Providers
```
**The security property that matters:** the daemon runs with **no authentication at all** because it is bound to `localhost` and never published. The auth proxy is the only public surface. If you expose port `8009`, you have removed the entire security model.
### Components
| Component | Role | Bind |
|---|---|---|
| `routstrd-auth` | Public auth proxy. Validates credentials, enforces roles and model policy, forwards to the daemon. | `0.0.0.0:8008` |
| `routstrd` | The inference daemon. Owns the wallet, providers, clients, and usage records. | `localhost:8009` |
| `routstr.db` | Shared SQLite database. Holds `routstr_auth_npubs`, `clients`, usage rows, and `sdk_storage` (including the Routstr 21 model list). | on disk |
| Reverse proxy | TLS termination and the public hostname. On Cloudron this is managed by the platform. | `443` |
---
## Roles
Registration lives in a single table, `routstr_auth_npubs`, where each row has a `role` of `admin` or `user`.
| Capability | `admin` | `user` |
|---|---|---|
| Run inference with own API keys | yes | yes |
| Create and delete **own** clients | yes | yes |
| Read **own** usage | yes | yes |
| List all registered npubs | yes | yes |
| Add / update / delete npubs | yes | no |
| Send funds from the node wallet | yes | no |
| Node control (providers, refunds, stop) | yes | yes |
| Read wallet balance / status | yes | yes |
`user` is the default role when an admin adds someone. Promote with `routstrd npubs update <npub> --role admin`.
---
## Bootstrap order
A fresh node has an empty npub table. That produces exactly one unaudited window, and only one:
1. **The first person** runs `routstrd npubs register` against the new node. Because the table is empty, `POST /npubs` is accepted **without authentication**, and the caller becomes `admin`.
2. From that moment on, **every** npub operation requires NIP-98 auth from an existing admin. A second unauthenticated registration is refused with `409` / "already configured".
Nobody else can self-register. Team members must be added by an admin. See [Team Members](team-members.md).
---
## Next steps
- **[Deploy on Cloudron](deploy-cloudron.md)** — the supported, packaged deployment with TLS and backups handled for you.
- **[Deploy with Docker](deploy-docker.md)** — run the same image anywhere, behind your own reverse proxy.
- **[Team Members](team-members.md)** — bootstrap the first admin and invite people.
- **[Connecting Clients](clients.md)** — wire up Claude Code, Pi, OpenCode, and raw API keys.
- **[Usage and Model Policy](usage-and-policy.md)** — per-member spend tracking and the model allowlist.
- **[Security Model](security.md)** — the exact auth rules, public paths, and restricted endpoints.
- **[Troubleshooting](troubleshooting.md)** — diagnosing the failures people actually hit.
+172
View File
@@ -0,0 +1,172 @@
# Security Model
The whole design rests on one property: **the daemon has no authentication because it is never reachable.** `routstrd` binds loopback-only on port `8009`; the auth proxy on `8008` is the single public surface, and it is the only component that makes authorisation decisions.
!!! danger "Never publish port 8009"
The daemon is unauthenticated by design. Exposing it — or port-forwarding it for debugging, or forgetting to restrict a Docker port mapping to loopback — removes the entire security model at once. Anyone who reaches it can read the wallet, list every member's clients, and spend the team's funds.
---
## Request decision flow
The proxy is **default-deny**. A request is only forwarded if some rule explicitly allows it.
```mermaid
flowchart TD
A["request arrives on 8008"] --> B{"management path?<br/>npubs, clients, usage"}
B -->|yes| C["own handler<br/>NIP-98 required"]
B -->|no| D{"GET or HEAD<br/>on a public path?"}
D -->|yes| E["forward, no auth"]
D -->|no| F{"Authorization header?"}
F -->|missing| G["401"]
F -->|"Bearer sk-..."| H{"key found in clients?"}
H -->|no| I["401"]
H -->|yes| J{"restricted path?<br/>wallet, node control"}
J -->|yes| K["403"]
J -->|no| L["forward with header intact"]
F -->|"Nostr event"| M{"valid NIP-98?<br/>url, method, body hash, sig"}
M -->|no| N["401"]
M -->|yes| O{"pubkey registered?<br/>and role sufficient?"}
O -->|no| P["403"]
O -->|yes| Q["forward, header stripped"]
```
Two details are easy to miss:
- **Public means `GET`/`HEAD` only.** The proxy applies the public-path rule only to read methods, because the daemon routes a `POST` to those same paths as a *paid* request. `GET /v1/models` needs no credential; `POST /v1/models` does.
- **Management paths are matched before anything else.** `/npubs`, `/clients`, `/clients/add`, `/clients/delete`, `/usage`, and `/usage/summary` are handled by the proxy's own handlers and never forwarded to the daemon wholesale.
---
## Public paths
Reachable with no credential at all (`GET`/`HEAD` only):
| Path | Purpose |
|---|---|
| `/health` | Liveness. Used by Cloudron's health check and by the proxy's own upstream probe. |
| `/ping` | Lightweight reachability. |
| `/models` | Model directory. |
| `/v1/models` | OpenAI-compatible model list. |
| `/models/*`, `/v1/models/*` | Prefixes covering per-model detail paths. |
This is intentional — an agent needs to discover models before it has a key, and provider discovery is public information in Routstr.
---
## Credential one: API keys (`Bearer sk-...`)
An API key is looked up in the client records. If no client carries it, the request is rejected with `401 Invalid API key.`
Keys are **deliberately narrow**. A valid key is refused with `403` on every restricted path:
| Restricted endpoint | Why |
|---|---|
| `/wallet/status`, `/wallet/unlock`, `/wallet/balance` | An inference key must not read wallet state. |
| `/wallet/receive/cashu`, `/wallet/receive/bolt11` | No minting funds with a key. |
| `/wallet/send/cashu`, `/wallet/send/bolt11` | Admin-only in any case. |
| `/wallet/mints`, `/wallet/mints/info` | No mint inspection. |
| `/stop`, `/refund`, `/refund/xcashu` | No node control. |
| `/providers`, `/providers/enable`, `/providers/disable` | `?refresh=true` rewrites the stored provider list. |
| `/nwc/*` | No payment-channel access. |
| `/npubs`, `/clients/add`, `/clients/delete`, `/usage` | No management surface at all. |
The rule of thumb: **an API key buys inference, and nothing else.**
Key handling on the way through: the `Authorization` header is **preserved** so the daemon can validate the key itself, and its own accounting stays authoritative.
---
## Credential two: NIP-98 (`Nostr <base64-event>`)
Management and wallet operations require a signed [NIP-98](https://github.com/nostr-protocol/nips/blob/master/98.md) event. This is not a bearer token — it is a signature over the specific request, so it cannot be replayed against a different endpoint.
The proxy enforces, in order:
| Check | Rule |
|---|---|
| Event kind | must be `27235` |
| Timestamp | within **±60 seconds** of now |
| `u` tag | must equal the **absolute request URL**, including scheme and host |
| `method` tag | must match the HTTP method (case-insensitive) |
| `payload` tag | **required when the body is non-empty**; must equal the SHA-256 hex digest of the raw body, compared in constant time |
| Signature | verified with `verifyEvent` |
The proxy then looks the pubkey up in `routstr_auth_npubs`:
- **Not registered** → `403`. The error message is context-aware: on a node with no npubs at all it tells you to run `routstrd npubs register`; otherwise it says registered auth is required.
- **Registered but role insufficient** → `403 Admin access required.`
- **Registered and sufficient** → forwarded, with the `Authorization` header **stripped** so it does not reach the daemon or the upstream provider.
!!! warning "Behind a reverse proxy, forwarded headers are not optional"
The `u` tag is checked against the **public** URL the client signed. The proxy reconstructs that URL from `X-Forwarded-Proto` and `X-Forwarded-Host` (falling back to `Host`). If your reverse proxy does not set them, the comparison fails and every NIP-98 request is rejected with `NIP-98 URL tag does not match this request.` See the nginx snippet in [Deploy with Docker](deploy-docker.md#put-tls-in-front).
!!! note "The ±60 second window means clock skew matters"
A client whose clock is more than a minute off will produce events that are rejected as `outside the allowed window`. If one machine alone fails to authenticate, check its clock before suspecting the node.
---
## Role requirements by endpoint
| Endpoint group | Required |
|---|---|
| `/wallet/send/cashu`, `/wallet/send/bolt11` | `admin` |
| `/wallet/status`, `/wallet/unlock`, `/wallet/balance`, `/wallet/receive/*`, `/wallet/mints*`, `/stop`, `/refund*`, `/providers*`, `/nwc/*` | any registered npub (`admin` or `user`) |
| `/clients`, `/clients/add`, `/clients/delete` | any registered npub, **scoped to own clients** |
| `/usage`, `/usage/summary` | any registered npub, **scoped to own usage** |
| `/npubs` read | any registered npub |
| `/npubs` create / update / delete | `admin` |
| Everything else | valid API key **or** registered npub |
---
## Bootstrap window
While `routstr_auth_npubs` is empty, `POST /npubs` is accepted **without authentication**. This exists solely so a fresh node can be claimed, and it closes permanently after the first registration.
The practical implication: a node that is deployed and healthy but has not had its first admin register is **unclaimed**. Treat deployment and bootstrap as one operation.
There is deliberately no hardcoded default admin in the image. The absence of one means an image cannot be shipped with a known admin key — but it also means a half-finished deployment is claimable by whoever finds it first.
As an alternative to the interactive step, admins can be seeded with `ROUTSTRD_AUTH_ADMIN_NPUBS`, `ROUTSTRD_AUTH_ADMIN_PUBKEYS`, or `ROUTSTRD_AUTH_BOOTSTRAP_NPUB`. Rows created this way are tagged `source = 'env'` and **reconciled at every startup** — remove the value from the environment and the row is deleted, which is a clean way to authorise a node declaratively.
---
## CORS
The proxy answers with:
```text
Access-Control-Allow-Origin: *
Access-Control-Allow-Methods: GET, POST, PATCH, DELETE, OPTIONS
Access-Control-Allow-Headers: Authorization, Content-Type, X-Cashu, X-Routstr-Model
Access-Control-Expose-Headers: X-Cashu, X-Routstr-Request-Id, X-Routstr-Cost-Msats,
X-Routstr-Cost-Usd, X-Routstr-Input-Cost-Msats,
X-Routstr-Output-Cost-Msats
```
A wildcard origin is safe **here specifically** because the app uses no cookies and no sessions. There is no ambient browser identity for a cross-origin page to borrow — a malicious page cannot make an authenticated request on a visitor's behalf, because every non-public request still needs its own API key or signature.
!!! warning "If you ever add cookie or session auth, revisit this"
The wildcard is only correct while authentication is entirely credential-based. Adding session cookies would turn this into a real vulnerability.
---
## Hardening checklist
- [ ] **Daemon is loopback-only.** Verify `8009` is not published (`docker port routstr-remote`, or check the Cloudron app's port config).
- [ ] **TLS everywhere.** No member or agent should ever send an `sk-...` key over plaintext HTTP.
- [ ] **First admin registered** immediately after install.
- [ ] **Reverse proxy sets `X-Forwarded-Proto` / `X-Forwarded-Host`**, or NIP-98 fails.
- [ ] **Streaming hangs are fixed with timeouts, not by buffering.** Disable `proxy_buffering` and raise `proxy_read_timeout`; do not "fix" a truncated stream by publishing the daemon directly.
- [ ] **`ROUTSTRD_AUTH_ADMIN_NPUBS` reflects reality** if you use env bootstrapping — those rows are deleted on restart when the variable changes.
- [ ] **Departed members have their clients deleted**, not just their npub. Deleting an npub does **not** revoke existing API keys.
- [ ] **Backups cover `/app/data`** and are tested, including the `nsec` in `routstrd/config.json`.
- [ ] **Model allowlist verified** if you rely on it — it fails open when the model list has not been populated.
---
## Next steps
- [Troubleshooting](troubleshooting.md) — diagnosing `401` and `403` responses.
- [Team Members](team-members.md#what-revocation-does-and-does-not-do) — the two-step offboarding that revocation alone does not cover.
+177
View File
@@ -0,0 +1,177 @@
# Team Members
Access to a team node is an entry in one table: `routstr_auth_npubs`. Each row holds a Nostr pubkey, an optional display name, and a `role` of `admin` or `user`. Adding someone grants access; deleting their row revokes it.
There are no passwords, no invite links, and no email addresses. Identity is a Nostr keypair that each person generates on their own machine.
---
## The invite loop
A new member does the first two steps themselves; an existing admin does the third.
```mermaid
sequenceDiagram
participant M as New member
participant A as Existing admin
participant N as Team node
M->>M: install the routstrd CLI
M->>N: set the remote URL
N-->>M: generates keypair and prints npub
M->>A: send npub out of band
A->>N: add npub with role user
N-->>A: access confirmed
M->>N: add a client integration
N-->>M: API key issued
```
### 1. The member installs the CLI and connects
```bash
bun i -g routstrd
routstrd remote https://team.example.com
```
`routstrd remote` writes the node URL into `~/.routstrd/config.json` and, **only if no Nostr identity exists yet**, generates one and prints the npub:
```text
Remote daemon URL set to: https://team.example.com
A new Nostr identity has been generated for remote authentication.
Your npub: npub1abc...xyz
You can view it in the config file at: /home/bob/.routstrd/config.json
```
If you already had an identity, it is reused and no npub is printed — run `routstrd remote` with no arguments to display the current node and identity.
!!! tip "Your npub is not a secret"
It is a public key and safe to paste into a team chat. The corresponding `nsec` lives in `~/.routstrd/config.json` and **is** a secret: it signs every management request. Never share it, and treat any host that has it as holding that member's credentials.
### 2. The member sends their npub to an admin
Out of band — chat, ticket, whatever. There is no self-service join.
### 3. An admin adds them
```bash
routstrd npubs add npub1abc...xyz --name "Bob"
```
The role defaults to `user`. The new member can now run inference and manage their own clients. To make them an admin, pass `--role admin` (or promote later).
---
## Bootstrap: the first admin
A brand-new node has an empty table, which is a special case: `POST /npubs` is accepted **without authentication** so that someone can claim the node.
```bash
routstrd remote https://team.example.com
routstrd npubs register --name "Alice"
```
`npubs register` is deliberately narrow — it refuses to do anything if any npub already exists:
```text
Admin npubs already configured (3). Ask your admin to add your npub.
Your npub: npub1...
```
So `register` only ever works once per node. After that, `npubs add` is the command, and it requires admin NIP-98 auth.
!!! warning "Claim the node during install"
Between the app becoming healthy and the first `npubs register`, the node is unclaimed — anybody who reaches the URL can become admin. See [Deploy on Cloudron](deploy-cloudron.md#bootstrap-the-first-admin).
---
## Command reference
All of these talk to the auth proxy over NIP-98, so the caller must be a registered npub, and the mutating ones require the `admin` role.
| Command | Role needed | Notes |
|---|---|---|
| `routstrd npubs list` | any registered | Shows role and name for everyone, and marks your own row with `→ you`. |
| `routstrd npubs register` | none, **once** | Only succeeds while the table is empty. |
| `routstrd npubs add <npub>` | admin | Accepts `npub1...` or a 64-char hex pubkey. `--role admin\|user` (default `user`), `--name`. |
| `routstrd npubs update <npub>` | admin | `--role` and/or `--name`. Passing `--name ""` clears the name. |
| `routstrd npubs delete <npub>` | admin | Revokes access. |
Names are trimmed and capped at 64 characters.
`routstrd npubs list` output looks like this:
```text
Npubs (3):
- npub1qqq...4f2 [admin] "Alice" → you
- npub1xxx...9k7 [user] "Bob"
- npub1zzz...3md [user]
```
If your own npub is missing from the list, the CLI tells you whom to send it to:
```text
Your npub is not in the npub list. Ask an admin to add your npub:
npub1yyy...0pl
```
## Underlying HTTP API
The CLI is a thin wrapper over four endpoints on the auth proxy. Useful for scripting or a custom onboarding form.
| Method | Path | Auth | Body |
|---|---|---|---|
| `GET` | `/npubs` | any registered npub | — |
| `POST` | `/npubs` | none **if the table is empty**, otherwise admin | `{ "npub": "npub1..." }` or `{ "pubkey": "<hex>" }`, plus optional `role` and `name` |
| `PATCH` | `/npubs` | admin | `{ "npub": "npub1..." }` plus `role` and/or `name` (`name: null` clears it) |
| `DELETE` | `/npubs/<npub-or-pubkey>` | admin | — (also accepts `/npubs?npub=...`) |
Responses:
- `GET /npubs` returns `{ "npubs": [ { "npub": "...", "name": "...", "role": "admin" } ] }`.
- Adding a pubkey that is **already registered** returns `409` rather than silently succeeding; use `PATCH` to change an existing entry.
- Every mutation is performed by, and recorded against, the requesting admin.
!!! note "Revocation is immediate"
Roles and removals are read from the database on **every request** with no caching layer. Deleting an npub stops that member's management access on their next request. Their **API keys are a separate matter** — see below.
---
## What revocation does and does not do
Both halves of offboarding are separate, and only the first is available through the normal member-facing API.
**1. Management access — revoked by deleting the npub.**
```bash
routstrd npubs delete npub1abc...xyz
```
Roles and rows are read from the database on **every request** with no caching layer, so the next NIP-98 request from that key is rejected immediately.
**2. Inference keys — a separate, manually managed thing.**
The Bearer path validates an `sk-...` key by looking it up in the client records and nothing else. Deleting an npub does **not** touch those rows, so an offboarded member's existing API keys keep working for inference until the client itself is deleted.
Here the proxy's scoping matters: `/clients`, `/clients/add`, and `/clients/delete` are **strictly owner-scoped** — the proxy filters and authorises by the calling npub. Being an `admin` does not grant access to a *colleague's* client records through those endpoints. So the realistic offboarding paths are:
- **Have the member delete their own clients** before you remove their npub: `routstrd clients list`, then `routstrd clients delete <id>`.
- **Or do it on the node itself.** Inside the container the CLI talks to the daemon directly on loopback, where there is no auth layer and therefore no ownership filter:
```bash
cloudron exec --app routstr.example.com
routstrd clients list # every client on the node, all owners
routstrd clients delete <id>
```
On the node, client IDs appear **with** their owner suffix (`my-laptop-4f2x9k7`) — see [Connecting Clients](clients.md). That suffix is exactly what tells you which member a client belongs to.
!!! danger "Removing an npub is not the same as revoking access"
If you skip step 2, a departed member's agents continue to consume the team's wallet. Always pair `npubs delete` with deleting their client records.
---
## Next steps
- [Connecting Clients](clients.md) — get each member's agents talking to the node.
- [Usage and Model Policy](usage-and-policy.md) — see what each person is spending.
+161
View File
@@ -0,0 +1,161 @@
# Troubleshooting
Most team-node problems are one of five things: the daemon never came up, the wrong database, a credential problem, a clock or proxy-header problem, or a networking/timeout issue around streaming.
---
## First triage
Run these in order. Together they answer "is the node up, does it have people, and is the proxy seeing the right database".
```bash
# 1. Is the public surface alive?
curl -sS https://team.example.com/health
# 2. Do both processes run, and what did they log on boot?
cloudron logs --app routstr.example.com | tail -50 # or: docker logs routstr-remote
# 3. Can the proxy see the database, and does it know your people?
cloudron exec --app routstr.example.com
routstrd-auth validate
```
Step 3 is the most informative. Its output ends with a line like `✅ DB accessible. 3 npub(s) registered (1 admin, 2 user).` — if that count is wrong, or the path is wrong, you have found your problem.
On startup the proxy also logs a summary, and warns loudly if the node is unclaimed:
```text
routstrd-auth proxy listening on http://0.0.0.0:8008
Upstream: http://localhost:8009
DB path: /app/data/routstrd/routstr.db
Registered npubs: 0
Model allowlist: disabled
Warning: no registered npub/pubkey. The first admin can be registered without auth using POST /npubs.
```
!!! warning "`Registered npubs: 0` on an established node is an emergency"
Either the database was lost (usually a missing volume mount) or the proxy is pointed at the wrong file. While the table is empty the node is claimable by anyone who reaches it. Fix it before anything else — see [Lost identity or empty npub list](#lost-identity-or-empty-npub-list).
---
## Error reference
### Startup
| Symptom | Cause | Fix |
|---|---|---|
| `Timed out waiting for routstrd to become ready.` | The proxy's launcher waited 120 seconds for the database file to exist **and** for `http://localhost:8009/health` to answer. The daemon is unhealthy or crashing. | Read the daemon's log lines above this one. It is a daemon problem, not a proxy problem. |
| `Database not found at /app/data/routstrd/routstr.db. Make sure routstrd has been initialized (routstrd onboard).` | The proxy shares the daemon's database and never creates the schema. It ran before the daemon had ever initialised. | Start the daemon first (`routstrd start`, or let `supervisord` do it — priority 10 before 20). |
| `Invalid admin Nostr pubkey(s): <value>. Use npub or 64-char hex pubkeys.` | A bootstrap admin variable contains something unparseable. | Fix `ROUTSTRD_AUTH_ADMIN_NPUBS` / `_PUBKEYS` / `_BOOTSTRAP_NPUB`. Use `npub1...` or 64-char hex. |
| Container exits with `Illegal instruction` / `SIGILL` | Bun's default x64 build needs AVX/AVX2, which some hosts and VMs do not expose. | The shipped image already uses the **baseline** build for this reason. If you build your own, keep `BUN_TARGET=bun-linux-x64-baseline`. |
| Proxy crash-loops immediately after a config change | Validation fails, so `start` exits non-zero and `supervisord` restarts it. | Run `routstrd-auth validate` to see the message. |
### Authentication
| Response | Meaning | Fix |
|---|---|---|
| `401 Missing Authorization header. Use 'Authorization: Bearer sk-...' or 'Authorization: Nostr <base64-event>'.` | No credential sent, and the path is not public. Remember public paths are `GET`/`HEAD` only. | Send a credential, or use a read method on a public path. |
| `401 Invalid API key.` | The key is not in the client records. | Recover it with `routstrd clients add --name "<existing name>"`, which prints the existing key. |
| `403 API keys cannot access this endpoint. Use NIP-98 auth from a registered npub/pubkey.` | You used an `sk-...` key on a wallet, node-control, or management endpoint. | Use the CLI, which signs with NIP-98 automatically. |
| `403 Admin access required. Only admin npubs can perform this action.` | The npub is registered but holds the `user` role. | An admin promotes them: `routstrd npubs update <npub> --role admin`. |
| `403 This endpoint requires a registered npub/pubkey, but none is configured. Register the first admin with 'routstrd npubs register'.` | The npub table is empty. | Bootstrap the first admin. |
| `403 This endpoint requires NIP-98 auth from a registered npub/pubkey.` | The signature was valid but the pubkey is not in the table. | An admin adds it: `routstrd npubs add <npub>`. |
| `401 Invalid Authorization format. Expected 'Bearer sk-...' or 'Nostr <base64-event>'.` | The header used a different scheme or a typo'd prefix. | Fix the prefix. |
A useful diagnostic: the token **type** determines the error. A `403` naming a specific capability means you authenticated successfully and were then refused by policy. A `401` means you did not authenticate at all.
### NIP-98 signature rejections
These all return `401` with a precise reason.
| Message | Cause | Fix |
|---|---|---|
| `NIP-98 URL tag does not match this request.` | The most common one. The `u` tag holds the URL the client signed, and the proxy reconstructs the public URL from `X-Forwarded-Proto` / `X-Forwarded-Host` (falling back to `Host`). Behind TLS termination, a missing forwarded header makes the reconstructed URL `http://...` while the client signed `https://...`. | Configure the reverse proxy to set both headers. See [Deploy with Docker](deploy-docker.md#put-tls-in-front). |
| `NIP-98 event timestamp is outside the allowed window.` | Events must be within **±60 seconds**. Clock skew. | Sync the client's clock (`timedatectl`, NTP). If only one machine fails, it is that machine. |
| `NIP-98 payload tag is required for requests with a body.` \| `NIP-98 payload tag does not match the request body hash.` | The signed SHA-256 does not match the body that arrived — something rewrote the body in transit. | Check for a middleware, WAF, or forward proxy that re-encodes request bodies. |
| `NIP-98 method tag does not match this request.` | The event was signed for a different method (often a `GET` signature reused on a `POST`). | Sign per request; do not reuse events. |
| `Invalid NIP-98 event signature.` \| `Invalid NIP-98 event kind.` | Corrupted token, or not a kind `27235` event. | Regenerate the request with the CLI. |
| `Invalid NIP-98 token encoding.` \| `Invalid NIP-98 event JSON.` | The base64 payload is truncated — common when a long `Authorization` header is split or truncated by a client. | Check for a header-size limit in the client or proxy. |
### Client-side CLI messages
| Message | Meaning | Fix |
|---|---|---|
| `The daemon at <url> rejected this account.` then `Register/authorize this npub on the remote daemon first: <npub>` | The node does not recognise your npub. | Send the printed npub to an admin. |
| `No remote node is set up.` | No `daemonUrl` in `~/.routstrd/config.json`. | `routstrd remote https://team.example.com`. |
| `Your npub is not in the npub list. Ask an admin to add your npub:` | `npubs list` worked (so you *are* registered) but shows you as absent — typically a stale local identity after a config reset. | Re-run `routstrd remote` to display your current npub, and confirm with an admin which one is registered. |
| `Daemon is not running` | The local daemon is unreachable. | Only relevant when running from source; `routstrd start`. |
---
## Common scenarios
### Streaming responses get cut off mid-answer
The symptom is a reply that stops abruptly — often after exactly the same number of seconds — with no error from the model.
This is **not** a node problem. It is response buffering or an idle timeout in a proxy in front of the node. Model turns can be silent for a long time while reasoning or waiting on tools, and intermediaries treat that silence as a dead connection.
Fix, in order of what actually bites:
1. **Disable response buffering** where the proxy talks to `8008`. The node already sends `X-Accel-Buffering: no`, but the fronting proxy must also be told (`proxy_buffering off` in nginx).
2. **Raise read timeouts** well above the default (`proxy_read_timeout 3600s`). nginx's 60-second default is the usual culprit.
3. **Do not work around it by exposing the daemon.** Publishing `8009` removes authentication entirely.
The proxy itself disables Bun's per-request idle timeout (`server.timeout(req, 0)`) precisely because valid streams can be quiet for a long time; anything still timing out is outside the node.
### Lost identity or empty npub list
If `routstrd npubs list` is empty after a restart, or `routstrd-auth validate` reports `0 npub(s)`, the node lost its data directory. In Docker this is almost always an **unmounted or changed volume**: `/app/data` is the only persistent path.
The damage is two-fold:
- The `nsec` in `routstrd/config.json` is gone, so the node has a **new** identity.
- The npub table is empty, so `POST /npubs` is **unauthenticated again** — anyone who reaches the URL can claim the node.
Recover in this order:
1. **Stop the app** so nobody claims it.
2. Restore `/app/data` from a backup, or re-mount the correct volume and restart.
3. If data is unrecoverable, accept the new identity and re-register the first admin — then have everyone whose npub was lost send theirs again, and recreate their clients (keys do not survive either).
### Model requests return `403` unexpectedly
Either the sender is genuinely outside the allowlist, or the allowlist is enforcing a **stale** list.
```bash
cloudron exec --app routstr.example.com
sqlite3 /app/data/routstrd/routstr.db \
"SELECT substr(value,1,120) FROM sdk_storage WHERE key = 'routstr21Models';"
```
- **No output:** the daemon never bootstrapped the list, so the proxy is failing **open** and allowing everything. The `403` is coming from somewhere else.
- **Output present but missing the model you want:** the list is stale. Refresh it with `routstrd clients --manual-refresh`, and check the scheduled job has not been disabled with `--disable-automatic-refresh`.
Remember the check applies to `POST`/`PUT`/`PATCH` bodies containing a `model` field only, and comparisons are case-sensitive.
### A member cannot authenticate at all
Work down this list:
1. **Are they registered?** `routstrd npubs list` as an admin.
2. **Is their npub the one you registered?** They run `routstrd remote` with no arguments to print the identity actually in their config. A machine with an old or regenerated `nsec` presents a different npub.
3. **Is their clock correct?** Off by more than 60 seconds means every NIP-98 event is rejected.
4. **Are forwarded headers set?** If *everyone* is failing with a URL tag mismatch, it is the reverse proxy, not the people.
5. **Do they have admin-requiring needs?** A `403 Admin access required` is a role problem, not an authentication problem.
### The node works but nobody can see usage
Usage is attributed per client, and `/usage` is scoped to the caller's npub. A member with no clients has no usage to show. Aggregate figures require running the CLI **on the node**, where the daemon is unauthenticated and unfiltered:
```bash
cloudron exec --app routstr.example.com
routstrd top
```
---
## Next steps
- [Security Model](security.md) — the rules behind the `401`s and `403`s.
- [Usage and Model Policy](usage-and-policy.md) — allowlist behaviour and its fail-open case.
+144
View File
@@ -0,0 +1,144 @@
# Usage and Model Policy
Two things a team admin cares about: **who is spending what**, and **which models the team may use**. The first is always tracked; the second is an opt-in policy.
---
## Usage tracking
Every request is attributed to the **client** that made it, and every client is owned by exactly one member. That chain is what produces per-person reporting.
### As a member
```bash
routstrd usage # your own usage summary
routstrd top # interactive TUI, alias for 'monitor'
routstrd balance # wallet balance and status
```
### As the operator, on the node
The most complete view is the TUI, run inside the container:
```bash
cloudron exec --app routstr.example.com
routstrd top
```
The **clients** tab shows individual usage, and because each client ID carries the **last 7 characters of its owner's npub**, you can tell teammates apart at a glance even when several of them named their client `my-laptop`:
```text
my-laptop-4f2x9k7 1,204 req $3.18
my-laptop-9k7qp2 881 req $2.07
ci-runner-3md1zz 412 req $0.94
```
Run from the node, this view is **not** filtered to any one member — that is the difference between the CLI on the box and the same CLI on a laptop.
### Over HTTP
| Method | Path | Auth | Scope |
|---|---|---|---|
| `GET` | `/usage` | NIP-98 from a registered npub | The **caller's** usage only |
| `GET` | `/usage/summary` | NIP-98 from a registered npub | The **caller's** usage only |
The proxy forces the scope: it takes the authenticated npub and sets the `npub` query parameter before forwarding to the daemon, so the daemon returns that member's records. There is no query parameter you can pass to widen it — a member cannot read a colleague's usage through the API. Aggregate visibility requires access to the node itself.
---
## The wallet
The team shares **one node wallet**. Funding it is a single action rather than N personal top-ups, and that is one of the main reasons to run a team node.
| Endpoint | Required role |
|---|---|
| `/wallet/status`, `/wallet/balance`, `/wallet/mints`, `/wallet/mints/info` | any registered npub |
| `/wallet/receive/cashu`, `/wallet/receive/bolt11` | any registered npub |
| `/wallet/send/cashu`, `/wallet/send/bolt11` | **admin only** |
Reading the balance and receiving funds are open to every registered member; **moving funds out is admin-only.** Note the consequence: members can see the team's total balance but cannot withdraw from it. API keys can do none of this.
Funding, mint management, and payment semantics are the daemon's domain and are documented in the provider and client guides rather than repeated here.
!!! warning "Usage is attributed, not isolated"
There is no per-member balance and no spending cap on this layer. A member's clients spend from the same wallet everyone else does. If you need hard limits, they are not enforced here — track the usage view and manage access accordingly.
---
## Model allowlist
The proxy can restrict the team to the **Routstr 21 model list**. This is **disabled by default**; enable it with:
```bash
ROUTSTRD_AUTH_MODEL_ALLOWLIST=true
```
When enabled, a request naming a model outside the list is rejected with `403` **before it reaches the daemon**, so no tokens are spent. When disabled, every model passes through untouched.
### How it works
```mermaid
flowchart TD
A["routstrd CLI or agent"] --> B["routstrd-auth<br/>checks auth, then model"]
B -->|"allowed"| C["routstrd daemon"]
B -->|"403 not allowed"| A
C --> D["upstream provider"]
E["Nostr kind 38423<br/>Routstr 21 list"] -->|fetched by daemon| F[("sdk_storage<br/>routstr21Models")]
F -->|read on every request| B
```
1. The **daemon's** SDK fetches the Routstr 21 list from Nostr (kind `38423` events) and stores it in the shared SQLite database under the `sdk_storage` table, key `routstr21Models`.
2. The **proxy** reads that key from the same database. It has **zero Nostr dependency** — no relay connections, no keys of its own for this purpose.
3. The value is read fresh on every request, with no caching, so a list update takes effect immediately.
Because the proxy shares the daemon's database, this costs a single indexed key lookup — typically under a millisecond.
### What is and is not checked
| Request | Checked |
|---|---|
| `POST` / `PUT` / `PATCH` with a JSON body containing `model` | yes |
| `GET` requests, including `/models` and `/v1/models` | no — public paths, forwarded immediately |
| Management endpoints (`/npubs`, `/clients`, `/usage`) | no — routed to their own handlers before this check |
| Non-JSON bodies | no — no `model` field can be extracted |
| Requests with no `model` field | no — forwarded, and the upstream produces the error |
Model IDs are compared **case-sensitively**, matching how they are stored in the list.
### Fail-open behaviour
If `routstr21Models` is absent — typically because the daemon has not bootstrapped it yet — the proxy **fails open and allows every model**. This is deliberate: the alternative is that a fresh node blocks all traffic until Nostr bootstrapping completes. The trade-off is that an allowlist can be silently ineffective early in a node's life, so verify it after enabling:
```bash
# confirm the daemon has populated the list
cloudron exec --app routstr.example.com
sqlite3 /app/data/routstrd/routstr.db \
"SELECT substr(value,1,120) FROM sdk_storage WHERE key = 'routstr21Models';"
```
If that returns nothing, the allowlist is not yet meaningful.
### Performance note
When enforcement is enabled, the proxy buffers `POST`/`PUT`/`PATCH` request bodies so it can inspect the `model` field. That is a small latency cost on request upload. **Response streaming is unaffected** — SSE and LLM token streams pass through unbuffered, and the proxy explicitly tells intermediaries not to buffer. When the allowlist is disabled, the body is not buffered at all.
---
## Keeping the model list current
The list is updated by the daemon, not the proxy:
```bash
routstrd clients --manual-refresh # refresh now
routstrd clients --disable-automatic-refresh # stop the scheduled job
routstrd clients --enable-automatic-refresh # resume it
```
If the scheduled refresh job is disabled and nobody runs a manual refresh, the allowlist enforces a **stale** list — and would eventually stop matching newly approved models.
---
## Next steps
- [Security Model](security.md) — the full endpoint-by-endpoint auth matrix.
- [Troubleshooting](troubleshooting.md) — including "why am I getting a 403".
+52 -17
View File
@@ -197,9 +197,9 @@ Properties:
This is the only architecture that preserves end-to-end encryption from the user to the PPQ/Tinfoil enclave while still letting Routstr mediate payment. The key requirement is that usage/cost metadata must be returned outside the encrypted body, ideally as a response header available before body streaming begins.
## Current Routstr problem
## Original Routstr problem
The current EHBP implementation charges successful EHBP requests at `max_cost_for_model` because Routstr cannot decrypt the response body:
The original EHBP implementation charged successful EHBP requests at `max_cost_for_model` because Routstr could not decrypt the response body:
```text
successful EHBP request -> charge full reserved max cost
@@ -260,7 +260,13 @@ During local PPQ testing, PPQ responses included this CORS exposure header:
Access-Control-Expose-Headers: Ehbp-Response-Nonce, X-Private-Usage-Metrics, X-Encrypted-Usage-Metrics, X-Tinfoil-Usage-Metrics
```
However, the actual tested non-streaming response did not include any of these usage headers, even when `X-Tinfoil-Request-Usage-Metrics: true` was sent.
However, when this was tested against PPQ's `/private/` endpoint
(`private/gpt-oss-120b`) on 2026-06-21, the non-streaming response did not
include any of these usage headers, even when
`X-Tinfoil-Request-Usage-Metrics: true` was sent. That observation does *not*
hold for the direct Tinfoil enclave upstream that Routstr ships: see
[Usage metrics header format](#usage-metrics-header-format) below, where the
response header and the streaming trailer are both verified present.
The decrypted body did include normal OpenAI usage, but only the decrypting Tinfoil client can see that body.
@@ -369,7 +375,7 @@ Possible approaches:
- PPQ private models are billed per actual input/output tokens.
- Private model rates are available from `GET /v1/models?type=all`.
- Current Routstr EHBP billing at max cost is wrong for PPQ private models.
- Max-cost EHBP fallback is wrong for PPQ private models; current code releases/refunds when trusted usage metadata is absent.
- Direct Tinfoil integration inside Routstr would enable exact usage billing but would make Routstr see plaintext.
- A blind EHBP relay preserves privacy but requires PPQ/Tinfoil to expose usage/cost in plaintext headers/trailers.
- The preferred solution is to keep Routstr blind and have PPQ return billing metadata outside the encrypted body.
@@ -385,7 +391,9 @@ and `routstr/upstream/ehbp.py`.
- Base URL: `https://inference.tinfoil.sh`
- Fetches models from the public `GET /v1/models` endpoint (no auth needed).
- Parses Tinfoil's pricing (`inputTokenPricePer1M`, `outputTokenPricePer1M`,
`requestPrice`) into the standard `Model`/`Pricing` schema.
`cachedInputTokenPricePer1M`, `requestPrice`) into the standard
`Model`/`Pricing` schema. Cached reads use the cached rate when present,
otherwise the full input rate; cache writes always use the full input rate.
- `supports_ehbp = True` — acts as a blind EHBP relay.
- `get_ehbp_forwarding_target()` returns a target that includes
`X-Tinfoil-Request-Usage-Metrics: true`.
@@ -396,9 +404,12 @@ and `routstr/upstream/ehbp.py`.
- `routstr/upstream/ehbp.py`:
- `parse_tinfoil_usage_metrics()` parses
`prompt=N,completion=N[,total=N][,model=<name>]` into an OpenAI-style
usage dict. The `model` field (added in tinfoilsh/confidential-model-router
PR #385) is extracted as a string.
`prompt=N,completion=N[,total=N][,cached_prompt_tokens=N,
uncached_prompt_tokens=N][,model=<name>][,cost_usd=<usd>]` into an
OpenAI-style usage dict. Cache reads map to ``cache_read_input_tokens``
so ``calculate_cost`` can bill them at the cached rate. The ``model``
field (added in tinfoilsh/confidential-model-router PR #385) is extracted
as a string; ``cost_usd`` is parsed for logging only.
- `_resolve_ehbp_target_url()` overrides the forwarding URL with
`X-Tinfoil-Enclave-Url` when the SDK sends it.
- `_strip_proxy_headers()` removes `X-Routstr-Model`,
@@ -410,8 +421,9 @@ and `routstr/upstream/ehbp.py`.
actual served model's pricing is used for cost calculation.
- `forward_ehbp_request()` (bearer auth): if `X-Tinfoil-Usage-Metrics` is
present in the response header, finalizes with `adjust_payment_for_tokens()`
for exact billing; otherwise falls back to max-cost. Billing uses the
actual served model when it differs from the requested one.
for exact billing; otherwise releases the reservation. The encrypted body
cannot be estimated locally, and the authorization ceiling is not billed.
Billing uses the actual served model when it differs from the requested one.
- `forward_ehbp_x_cashu_request()`: if usage is available, computes the
refund from actual cost instead of max cost, using the actual served
model's pricing when applicable.
@@ -425,10 +437,10 @@ and `routstr/upstream/ehbp.py`.
|---|---|---|
| Bearer, non-streaming | `X-Tinfoil-Usage-Metrics` response header | Exact token cost via `adjust_payment_for_tokens` |
| Bearer, streaming | `X-Tinfoil-Usage-Metrics` HTTP trailer | Exact token cost (h11 captures trailers) |
| Bearer, no usage header/trailer | N/A | Max-cost fallback |
| Bearer, no usage header/trailer | N/A | Release reservation; zero charge |
| X-Cashu, non-streaming | `X-Tinfoil-Usage-Metrics` response header | Refund = `redeemed - actual_cost` |
| X-Cashu, streaming | `X-Tinfoil-Usage-Metrics` HTTP trailer | Refund = `redeemed - actual_cost` (h11 captures trailers) |
| X-Cashu, no usage header/trailer | N/A | Refund = `redeemed - max_cost` |
| X-Cashu, no usage header/trailer | N/A | Full refund |
### Cost response headers
@@ -438,13 +450,20 @@ Routstr returns cost info as response headers:
| Header | Auth | Description |
|---|---|---|
| `X-Routstr-Cost-Msats` | Bearer, X-Cashu | Total msats charged for this request |
| `X-Routstr-Cost-Usd` | Bearer | USD equivalent of the charge |
| `X-Routstr-Cost-Msats` | Bearer, X-Cashu | Settled msats debited for this request |
| `X-Routstr-Computed-Cost-Msats` | Bearer, X-Cashu | Computed usage cost, emitted when it differs from the settled debit |
| `X-Routstr-Cost-Usd` | Bearer | USD equivalent of the computed usage |
| `X-Routstr-Input-Cost-Msats` | Bearer, X-Cashu | msats attributed to input tokens |
| `X-Routstr-Output-Cost-Msats` | Bearer, X-Cashu | msats attributed to output tokens |
| `X-Routstr-Cache-Read-Msats` | Bearer, X-Cashu | msats attributed to cached input (cache reads) |
| `X-Routstr-Cache-Creation-Msats` | Bearer, X-Cashu | msats attributed to cache creation (0 for Tinfoil today) |
The client/Tinfoil SDK can read these headers from the HTTP response without
needing to decrypt the body.
needing to decrypt the body. A duplicate or rejected finalization can therefore
report a zero settled debit while preserving the non-zero computed usage cost.
Normal JSON responses use the same contract: `total_msats` and `charged_msats`
are settled values, while `computed_msats` retains the usage calculation when
it differs.
### Setup
@@ -461,9 +480,23 @@ Tinfoil returns usage metrics in the `X-Tinfoil-Usage-Metrics` response header
true` is sent. As of tinfoilsh/confidential-model-router PR #385, the format is:
```
prompt=<prompt_tokens>,completion=<completion_tokens>,total=<total_tokens>,model=<served_model>
prompt=<prompt_tokens>,completion=<completion_tokens>,total=<total_tokens>[,cached_prompt_tokens=<n>,uncached_prompt_tokens=<n>][,model=<served_model>][,cost_usd=<usd>]
```
`prompt` is the inclusive prompt total; `cached_prompt_tokens` is the portion
already in Tinfoil's prefix cache and is billed at the model's
`cachedInputTokenPricePer1M` rate (or the full input rate when the model has
no cached rate). `cost_usd` is Tinfoil's own computed request cost and is
currently parsed for observability only — Routstr bills from token counts.
Note that the header/trailer value is not always a single occurrence: for
streaming responses the trailer is emitted twice, so a client that reads the
trailer directly may see the same `prompt=...,completion=...,...` string twice
in one field, comma-joined. Parsers must be tolerant of the duplicate rather
than assuming exactly one occurrence. `parse_tinfoil_usage_metrics()` is
unaffected: it assigns each `key=value` part as it walks the comma-separated
value, and both occurrences carry identical numbers.
The `model` field carries the actual model name served by the enclave.
Routstr uses this to:
@@ -490,4 +523,6 @@ back to the requested model's pricing.
finalizers for bearer and X-Cashu requests. This provides actual-cost billing
today, at the cost of full time-to-last-byte latency for streaming responses.
- Whether Tinfoil's `/v1/responses` endpoint also returns usage metrics
headers or trailers.
headers or trailers. Verified: yes — `/v1/responses` returns
`X-Tinfoil-Usage-Metrics` as a plaintext response header, with the same field
set as `/v1/chat/completions` (including `cost_usd`).
-45
View File
@@ -1,45 +0,0 @@
import json
import sys
import httpx
def create_child_keys(base_url: str, api_key: str, count: int = 3) -> list[str]:
headers = {"Authorization": f"Bearer {api_key}"}
print(f"Requesting {count} child keys from {base_url}...")
child_keys = []
for i in range(count):
try:
response = httpx.post(f"{base_url}/v1/balance/child-key", headers=headers)
if response.status_code == 200:
data = response.json()
child_keys.append(data["api_key"])
print(
f" [{i + 1}] Created: {data['api_key']} (Cost: {data['cost_msats']} msats)"
)
else:
print(f" [{i + 1}] Failed: {response.status_code} - {response.text}")
except Exception as e:
print(f" [{i + 1}] Error: {str(e)}")
return child_keys
if __name__ == "__main__":
if len(sys.argv) < 2:
print("Usage: python create_child_keys.py <api_key_or_cashu_token> [base_url]")
sys.exit(1)
auth_key = sys.argv[1]
base_url = sys.argv[2] if len(sys.argv) > 2 else "http://localhost:8000"
keys = create_child_keys(base_url, auth_key)
if keys:
print("\nSuccessfully created child keys:")
print(json.dumps(keys, indent=2))
else:
print("\nNo child keys were created.")
+1 -1
View File
@@ -7,7 +7,7 @@ from openai import OpenAI
client = OpenAI(
api_key=os.environ.get("TOKEN"),
base_url=os.environ.get("ONION_URL", "http://roustrjfsdgfiueghsklchg.onion/v1"),
http_client=httpx.Client(proxies="socks5://localhost:9050"),
http_client=httpx.Client(proxy="socks5://localhost:9050"),
)
print(
@@ -0,0 +1,58 @@
"""add refunds table
Revision ID: 3a0fbd387f10
Revises: a3f1b6c204de
Create Date: 2026-09-16
"""
import sqlalchemy as sa
import sqlmodel
from alembic import op
revision = "3a0fbd387f10"
down_revision = "a3f1b6c204de"
branch_labels = None
depends_on = None
OPEN_STATUSES = "status IN ('pending', 'ambiguous')"
def upgrade() -> None:
op.create_table(
"refunds",
sa.Column("id", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
sa.Column(
"api_key_hashed_key", sqlmodel.sql.sqltypes.AutoString(), nullable=False
),
sa.Column("method", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
sa.Column("destination", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
sa.Column("amount_msats", sa.Integer(), nullable=False),
sa.Column("unit", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
sa.Column("mint_url", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
sa.Column("status", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
sa.Column("quote_id", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
sa.Column("token", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
sa.Column("claimed_at", sa.Integer(), nullable=True),
sa.Column("created_at", sa.Integer(), nullable=False),
sa.Column("updated_at", sa.Integer(), nullable=False),
sa.ForeignKeyConstraint(["api_key_hashed_key"], ["api_keys.hashed_key"]),
sa.PrimaryKeyConstraint("id"),
)
op.create_index("ix_refunds_api_key_hashed_key", "refunds", ["api_key_hashed_key"])
op.create_index("ix_refunds_status", "refunds", ["status"])
op.create_index(
"ux_refunds_open_per_key",
"refunds",
["api_key_hashed_key"],
unique=True,
sqlite_where=sa.text(OPEN_STATUSES),
postgresql_where=sa.text(OPEN_STATUSES),
)
def downgrade() -> None:
op.drop_index("ux_refunds_open_per_key", table_name="refunds")
op.drop_index("ix_refunds_status", table_name="refunds")
op.drop_index("ix_refunds_api_key_hashed_key", table_name="refunds")
op.drop_table("refunds")
@@ -0,0 +1,27 @@
"""add model_metadata to model_paths
Revision ID: a3f1b6c204de
Revises: e5a6b7c8d9f0
Create Date: 2026-09-15 21:50:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "a3f1b6c204de"
down_revision = "e5a6b7c8d9f0"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"model_paths",
sa.Column("model_metadata", sa.Text(), nullable=False, server_default="{}"),
)
def downgrade() -> None:
op.drop_column("model_paths", "model_metadata")
@@ -0,0 +1,76 @@
"""Remove child keys and balance limits.
Removes the child-key feature (parent_key_hash) and the balance-limit
machinery (balance_limit, balance_limit_reset, balance_limit_reset_date)
from api_keys, plus the balance_limit/balance_limit_reset pass-through on
lightning_invoices.
Data preservation: before dropping the columns, every child key is
converted into a standalone key by clearing parent_key_hash. Child keys
never hold their own balance (they always spent from their parent), so no
funds are lost: the parent keeps its full balance, and the former child
rows are preserved with their total_spent/total_requests history intact.
"""
import sqlalchemy as sa
from alembic import op
revision = "e5a6b7c8d9f0"
down_revision = "b4f7a1c9d2e3"
branch_labels = None
depends_on = None
def upgrade() -> None:
# Convert child keys into standalone keys before dropping the link.
# Their balance is always 0 (they spent from the parent), so this
# cannot strand any funds.
op.execute("UPDATE api_keys SET parent_key_hash = NULL")
with op.batch_alter_table("api_keys") as batch_op:
batch_op.drop_index("ix_api_keys_parent_key_hash")
batch_op.drop_column("parent_key_hash")
batch_op.drop_column("balance_limit")
batch_op.drop_column("balance_limit_reset")
batch_op.drop_column("balance_limit_reset_date")
with op.batch_alter_table("lightning_invoices") as batch_op:
batch_op.drop_column("balance_limit")
batch_op.drop_column("balance_limit_reset")
def downgrade() -> None:
with op.batch_alter_table("lightning_invoices") as batch_op:
batch_op.add_column(sa.Column("balance_limit", sa.Integer(), nullable=True))
batch_op.add_column(
sa.Column("balance_limit_reset", sa.String(), nullable=True)
)
with op.batch_alter_table("api_keys") as batch_op:
batch_op.add_column(
sa.Column("balance_limit_reset_date", sa.Integer(), nullable=True)
)
batch_op.add_column(
sa.Column(
"balance_limit_reset",
sa.String(),
nullable=True,
)
)
batch_op.add_column(sa.Column("balance_limit", sa.Integer(), nullable=True))
batch_op.add_column(
sa.Column(
"parent_key_hash",
sa.String(),
nullable=True,
)
)
batch_op.create_foreign_key(
"fk_api_keys_parent_key_hash",
"api_keys",
["parent_key_hash"],
["hashed_key"],
)
batch_op.create_index(
"ix_api_keys_parent_key_hash", ["parent_key_hash"], unique=False
)
+9
View File
@@ -91,6 +91,15 @@ nav:
- Advanced Pricing: provider/advanced-pricing.md
- Discovery: provider/discovery.md
- Tor Support: provider/tor.md
- Teams (Remote):
- Overview: teams/index.md
- Deploy on Cloudron: teams/deploy-cloudron.md
- Deploy with Docker: teams/deploy-docker.md
- Team Members: teams/team-members.md
- Connecting Clients: teams/clients.md
- Usage and Model Policy: teams/usage-and-policy.md
- Security Model: teams/security.md
- Troubleshooting: teams/troubleshooting.md
- API Reference:
- Overview: api/overview.md
- Authentication: api/authentication.md
+36 -7
View File
@@ -1,27 +1,27 @@
[project]
name = "routstr"
version = "0.4.5"
version = "0.4.7"
description = "Payment proxy for your LLM endpoint using cashu and nostr."
readme = "README.md"
requires-python = ">=3.11"
dependencies = [
"fastapi[standard]>=0.115",
"fastapi[standard-no-fastapi-cloud-cli]>=0.141",
"aiosqlite>=0.20",
"sqlmodel>=0.0.24",
"httpx[socks]>=0.25.2",
"h11>=0.14",
"sqlmodel>=0.0.42", # Python 3.14 deferred-annotation support
"httpx[socks]>=0.28.1",
"h11>=0.16",
"greenlet>=3.2.1",
"alembic>=1.13",
"python-json-logger>=2.0.0",
"cashu>=0.20",
"marshmallow>=3.13,<4.0",
"websockets>=12.0",
"nostr>=0.0.2",
"nostr-sdk>=0.45.1,<0.46",
"mdurl==0.1.2",
"pillow>=10",
"openai>=1.98.0",
"litellm>=1.55.0",
"litellm>=1.93.0,<1.94", # 1.93 is the first line supporting Python 3.14
]
[dependency-groups]
@@ -87,3 +87,32 @@ disallow_untyped_decorators = true
[tool.uv.sources]
routstr = { workspace = true }
# Security floors. cashu 0.20.x caps httpx, h11, fastapi, cryptography,
# setuptools and wheel below their patched versions. routstr uses only cashu's
# wallet modules, not its server paths, so the caps are safe to lift here.
[tool.uv]
# coincurve 20's build config uses cmake.verbose, removed in
# scikit-build-core 0.10. This keeps its source build working on Python 3.14.
build-constraint-dependencies = ["scikit-build-core<0.10"]
override-dependencies = [
# httpx 0.28 (required by litellm 1.93) removed the `proxies` kwarg cashu
# passes on every mint call; routstr/cashu_compat.py restores it for cashu.
"httpx[socks]>=0.28.1,<1.0",
"importlib-metadata>=8.0.0,<9.0",
"h11>=0.16.0",
"fastapi[standard-no-fastapi-cloud-cli]>=0.141",
"cryptography>=49.0.0",
"setuptools>=83.0.0",
"wheel>=0.46.2",
]
# Transitive deps whose dependents allow the patched version but don't require
# it. Constraints raise the floor without bypassing any upstream pin.
constraint-dependencies = [
"starlette>=1.3.1",
"httpcore>=1.0.9", # 1.0.8 caps h11<0.15
# 1.76 is the first grpcio-tools release with CPython 3.14 wheels.
"grpcio>=1.76.0,<2.0.0",
"grpcio-tools>=1.76.0,<2.0.0",
]
+79 -19
View File
@@ -59,6 +59,23 @@ def calculate_model_cost_score(model: "Model") -> float:
return total_cost
def calculate_model_reservation_score(model: "Model") -> float:
"""Context-based reservation ceiling for ranking same-model candidates.
The balance gate reserves on ``sats_pricing.max_cost``, which scales with
``context_length``. Ranking on the same ceiling keeps the advertised,
routed, and reserved candidate consistent. Lower is better.
"""
from .payment.models import _calculate_usd_max_costs
try:
_, _, max_cost = _calculate_usd_max_costs(model)
return float(max_cost)
except Exception:
# Pricing shape missing/invalid; per-token score keeps order deterministic
return calculate_model_cost_score(model)
def get_provider_penalty(provider: "BaseUpstreamProvider") -> float:
"""Calculate a penalty multiplier for certain providers.
@@ -118,9 +135,21 @@ def create_model_mappings(
Returns:
Tuple of (model_instances, provider_map, unique_models)
"""
from .payment.models import _row_to_model
from .payment.models import _row_to_model, has_usable_pricing
from .upstream.helpers import resolve_model_alias
def _unusable_price(model: "Model") -> bool:
"""A candidate may only route on rates a request can be billed against.
Mirrors the served-catalog backstop in ``list_models``: a negative or
non-finite rate is not a price, and the cost calculation cannot bill on
one, so every request on the model would be charged the full maximum
reservation instead. Applies to provider-discovered models as well as
persisted overrides — no override row need exist for a malformed price
to be built into the candidate map.
"""
return not has_usable_pricing(model.pricing)
candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {}
unique_models: dict[str, "Model"] = {}
unique_model_keys: dict[str, str] = {}
@@ -208,12 +237,33 @@ def create_model_mappings(
# Apply overrides only for this provider's model row.
if model_key is not None and model_key in overrides_by_key:
override_row, provider_fee = overrides_by_key[model_key]
model_to_use = _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee
)
try:
model_to_use = _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee
)
except Exception as exc:
# Stored pricing is JSON from whatever wrote the row, so
# converting it can raise. Doing that inside this loop let
# one such row unwind the whole map build: at boot the node
# came up routing nothing, and on a later refresh the map it
# already had went permanently stale. The sibling loop over
# override-only rows already skips and logs such a row.
logger.warning(
"Skipping invalid model override while building model mappings",
extra={
"model_id": model.id,
"upstream_provider_id": upstream_db_id,
"error": str(exc),
"error_type": type(exc).__name__,
},
)
continue
else:
model_to_use = model
if _unusable_price(model_to_use):
continue
forwarded_model_id = get_effective_forwarded_model_id(model_to_use)
# Get all aliases for this model
@@ -280,6 +330,8 @@ def create_model_mappings(
continue
if not model_to_use.enabled:
continue
if _unusable_price(model_to_use):
continue
forwarded_model_id = get_effective_forwarded_model_id(model_to_use)
@@ -323,20 +375,24 @@ def create_model_mappings(
def alias_priority(model: "Model", alias: str) -> int:
"""Rank how strong the mapping of alias->model is.
An exact model ID is authoritative and must be cost-ranked against the
other exact matches before considering forwarded aliases. This keeps a
provider-specific forwarded ID from shadowing a directly available,
A provider that serves the requested ID directly is authoritative and
must be cost-ranked before providers that only reach it through a
forwarded alias, so a forwarded ID cannot shadow a directly available,
cheaper model with the requested ID.
"Directly served" covers both the exact model ID and the same ID behind
a provider prefix (e.g. ``gpt-oss-120b`` on Tinfoil vs
``openai/gpt-oss-120b`` on OpenRouter). Both name the same model, so
they share the top tier and cost decides between them; otherwise the
provider whose catalog omits the org prefix would always win on ID
spelling regardless of price.
"""
if model.id and model.id.lower() == alias:
return 5
model_base = get_base_model_id(model.id)
if (model.id and model.id.lower() == alias) or model_base.lower() == alias:
return 4
forwarded_model_id = get_effective_forwarded_model_id(model)
if forwarded_model_id and forwarded_model_id.lower() == alias:
return 4
model_base = get_base_model_id(model.id)
if model_base == alias:
return 3
if model.canonical_slug:
canonical_base = get_base_model_id(model.canonical_slug)
@@ -345,15 +401,19 @@ def create_model_mappings(
return 1
for alias, items in candidates.items():
# Sort key: (priority DESC, cost ASC)
# Using negative cost for DESC sort overall to keep high priority first
def sort_key(item: tuple["Model", "BaseUpstreamProvider"]) -> tuple[int, float]:
# Sort key: (priority DESC, reservation ASC, cost ASC)
# Using negative costs for DESC sort overall to keep high priority first
def sort_key(
item: tuple["Model", "BaseUpstreamProvider"],
) -> tuple[int, float, float]:
model, provider = item
priority = alias_priority(model, alias)
cost = calculate_model_cost_score(model)
penalty = get_provider_penalty(provider)
adjusted_cost = cost * penalty
return (priority, -adjusted_cost)
# Rank on the reservation ceiling the balance gate enforces, with
# per-token typical-usage cost as tiebreaker
adjusted_reservation = calculate_model_reservation_score(model) * penalty
adjusted_cost = calculate_model_cost_score(model) * penalty
return (priority, -adjusted_reservation, -adjusted_cost)
items.sort(key=sort_key, reverse=True)
+421 -509
View File
File diff suppressed because it is too large Load Diff
+50 -407
View File
@@ -1,16 +1,13 @@
import asyncio
import hashlib
import time
from time import monotonic
from typing import Annotated, NoReturn
from fastapi import APIRouter, Depends, Header, HTTPException
from fastapi.responses import JSONResponse
from pydantic import BaseModel
from sqlmodel import col, select, update
from sqlmodel import col, select
from . import refund
from .auth import (
get_billing_key,
redemption_error_to_http_exception,
validate_bearer_key,
)
@@ -21,20 +18,13 @@ from .core.db import (
get_session,
release_stale_reservations,
)
from .core.db import (
store_cashu_transaction_with_retry as store_cashu_transaction,
)
from .core.logging import get_logger
from .core.settings import settings
from .lightning import lightning_router
from .payment.lnurl import MeltOutcomeAmbiguousError
from .wallet import (
classify_redemption_error,
credit_balance,
is_mint_connection_error,
recieve_token,
send_to_lnurl,
send_token,
token_mint_url,
)
@@ -58,39 +48,15 @@ async def get_key_from_header(
async def get_balance_info(key: ApiKey, session: AsyncSession) -> dict:
billing_key = await get_billing_key(key, session)
info = {
"api_key": "sk-" + key.hashed_key,
"balance": billing_key.total_balance,
"reserved": billing_key.reserved_balance,
"is_child": key.parent_key_hash is not None,
"balance": key.total_balance,
"reserved": key.reserved_balance,
"total_requests": key.total_requests,
"total_spent": key.total_spent,
"balance_limit": key.balance_limit,
"balance_limit_reset": key.balance_limit_reset,
"validity_date": key.validity_date,
}
if key.parent_key_hash:
info["parent_key_preview"] = key.parent_key_hash[:8] + "..."
else:
# Fetch child keys if this is a parent key
statement = select(ApiKey).where(ApiKey.parent_key_hash == key.hashed_key)
results = await session.exec(statement)
child_keys = results.all()
if child_keys:
info["child_keys"] = [
{
"api_key": "sk-" + ck.hashed_key,
"total_requests": ck.total_requests,
"total_spent": ck.total_spent,
"balance_limit": ck.balance_limit,
"balance_limit_reset": ck.balance_limit_reset,
"validity_date": ck.validity_date,
}
for ck in child_keys
]
return info
@@ -117,26 +83,18 @@ async def account_info(
class BalanceCreateRequest(BaseModel):
initial_balance_token: str
balance_limit: int | None = None
balance_limit_reset: str | None = None
validity_date: int | None = None
async def _create_balance(
initial_balance_token: str,
balance_limit: int | None,
balance_limit_reset: str | None,
validity_date: int | None,
session: AsyncSession,
) -> dict:
key = await validate_bearer_key(initial_balance_token, session)
if balance_limit is not None or balance_limit_reset or validity_date:
key.balance_limit = balance_limit
key.balance_limit_reset = balance_limit_reset
if validity_date is not None:
key.validity_date = validity_date
if balance_limit_reset:
key.balance_limit_reset_date = int(time.time())
session.add(key)
await session.commit()
await session.refresh(key)
@@ -154,8 +112,6 @@ async def create_balance_from_body(
) -> dict:
return await _create_balance(
payload.initial_balance_token,
payload.balance_limit,
payload.balance_limit_reset,
payload.validity_date,
session,
)
@@ -164,15 +120,11 @@ async def create_balance_from_body(
@router.get("/create")
async def create_balance(
initial_balance_token: str,
balance_limit: int | None = None,
balance_limit_reset: str | None = None,
validity_date: int | None = None,
session: AsyncSession = Depends(get_session),
) -> dict:
return await _create_balance(
initial_balance_token,
balance_limit,
balance_limit_reset,
validity_date,
session,
)
@@ -208,7 +160,7 @@ async def topup_wallet_endpoint(
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
) -> dict[str, int]:
billing_key = await get_billing_key(key, session)
billing_key = key
if topup_request is not None:
cashu_token = topup_request.cashu_token
@@ -276,35 +228,6 @@ async def topup_wallet_endpoint(
return {"msats": amount_msats}
_REFUND_CACHE_TTL_SECONDS: int = settings.refund_cache_ttl_seconds
_refund_cache_lock: asyncio.Lock = asyncio.Lock()
_refund_cache: dict[str, tuple[float, dict[str, str]]] = {}
def _cache_key_for_authorization(authorization: str) -> str:
return hashlib.sha256(authorization.strip().encode()).hexdigest()
async def _refund_cache_get(authorization: str) -> dict[str, str] | None:
key = _cache_key_for_authorization(authorization)
async with _refund_cache_lock:
item = _refund_cache.get(key)
if item is None:
return None
expires_at, value = item
if expires_at <= monotonic():
del _refund_cache[key]
return None
return value
async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None:
key = _cache_key_for_authorization(authorization)
expiry = monotonic() + _REFUND_CACHE_TTL_SECONDS
async with _refund_cache_lock:
_refund_cache[key] = (expiry, value)
async def _lookup_key_no_create(
bearer_value: str, session: AsyncSession
) -> ApiKey | None:
@@ -318,17 +241,16 @@ async def _lookup_key_no_create(
async def _get_persisted_api_key_refund(
key: ApiKey, session: AsyncSession
key: ApiKey, session: AsyncSession, token: str | None = None
) -> dict[str, str] | None:
result = await session.exec(
select(CashuTransaction)
.where(
CashuTransaction.api_key_hashed_key == key.hashed_key,
CashuTransaction.type == "out",
CashuTransaction.source == "apikey",
)
.order_by(col(CashuTransaction.created_at).desc())
query = select(CashuTransaction).where(
CashuTransaction.api_key_hashed_key == key.hashed_key,
CashuTransaction.type == "out",
CashuTransaction.source == "apikey",
)
if token is not None:
query = query.where(CashuTransaction.token == token)
result = await session.exec(query.order_by(col(CashuTransaction.created_at).desc()))
refund = result.first()
if refund is None:
return None
@@ -347,36 +269,13 @@ async def _get_persisted_api_key_refund(
return persisted
async def _restore_balance(
session: AsyncSession,
hashed_key: str,
balance: int,
reserved_balance: int,
mint_url: str,
) -> None:
"""Restore balance after a failed refund mint attempt."""
restore_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == hashed_key)
.values(
balance=col(ApiKey.balance) + balance,
reserved_balance=col(ApiKey.reserved_balance) + reserved_balance,
)
)
await session.exec(restore_stmt) # type: ignore[call-overload]
await session.commit()
logger.info(
"refund_wallet_endpoint: balance restored after mint failure",
extra={
"hashed_key": hashed_key,
"restored_balance": balance,
"mint_url": mint_url,
},
)
class RefundRequest(BaseModel):
lightning_address: str | None = None
@router.post("/refund", response_model=None)
async def refund_wallet_endpoint(
refund_request: RefundRequest | None = None,
authorization: Annotated[str | None, Header()] = None,
x_cashu: Annotated[str | None, Header()] = None,
session: AsyncSession = Depends(get_session),
@@ -452,17 +351,28 @@ async def refund_wallet_endpoint(
},
)
# Check for an open claim before any replay or destination lookup.
if open_claim := await refund.latest_open(session, key):
raise refund.refund_in_progress_error(open_claim)
if key.total_balance <= 0:
if cached := await _refund_cache_get(bearer_value):
return cached
paid = await refund.latest_terminal(session, key)
if paid and paid.method == "lightning":
return refund.describe(paid)
if paid and paid.token:
# Match the ledger row to this claim's token, not the latest one.
if persisted := await _get_persisted_api_key_refund(
key, session, paid.token
):
return persisted
return refund.describe(paid)
# Legacy payouts predate the claim row, so fall back to the ledger.
if persisted := await _get_persisted_api_key_refund(key, session):
return persisted
if key.parent_key_hash:
raise HTTPException(
status_code=400,
detail="Cannot refund child key. Please refund the parent key instead.",
)
if paid:
return refund.describe(paid)
if stuck := await refund.latest_stuck(session, key):
raise refund.refund_in_progress_error(stuck)
if key.reserved_balance > 0:
# Release only durable reservations old enough to be stale. A newer
@@ -481,168 +391,33 @@ async def refund_wallet_endpoint(
logger.warning(
"refund_wallet_endpoint: released stale reservation before refund",
extra={
"hashed_key": key.hashed_key,
"key_hash": key.hashed_key[:8],
"stale_timeout_seconds": settings.stale_reservation_timeout_seconds,
},
)
remaining_balance_msats: int = key.total_balance
if key.refund_currency == "sat":
remaining_balance = remaining_balance_msats // 1000
else:
remaining_balance = remaining_balance_msats
unit = refund.refund_unit(key)
remaining_balance = refund.amount_in_unit(remaining_balance_msats, unit)
if remaining_balance_msats > 0 and remaining_balance <= 0:
raise HTTPException(status_code=400, detail="Balance too small to refund")
elif remaining_balance <= 0:
raise HTTPException(status_code=400, detail="No balance to refund")
# Capture values before debit — the session may refresh key after commit
pre_debit_balance = key.balance
pre_debit_reserved = key.reserved_balance
requested = refund_request.lightning_address if refund_request else None
destination = requested or key.refund_address
if destination:
# Stored addresses can rot too; reject before any balance is debited.
await refund.validate_lightning_destination(destination)
# --- DEBIT FIRST: atomically zero the balance before minting tokens ---
# This prevents the race where a concurrent topup/spend happens between
# reading the balance and minting the refund token (double-spend).
debit_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.balance) == pre_debit_balance)
.where(col(ApiKey.reserved_balance) == pre_debit_reserved)
.values(balance=0, reserved_balance=0, reserved_at=None)
claim = await refund.open_claim(
session,
key,
method="lightning" if destination else "cashu",
destination=destination,
)
debit_result = await session.exec(debit_stmt) # type: ignore[call-overload]
await session.commit()
if debit_result.rowcount == 0:
# Balance changed between read and debit — another request is active
raise HTTPException(
status_code=409,
detail="Balance changed concurrently. Please retry the refund.",
)
# The balance is locked at zero, so it is safe to create the refund token.
effective_refund_mint = (
key.refund_mint_url
if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints
else settings.primary_mint
)
try:
refund_currency = key.refund_currency or "sat"
if key.refund_address:
await send_to_lnurl(
remaining_balance,
key.refund_currency or "sat",
effective_refund_mint,
key.refund_address,
)
result = {"recipient": key.refund_address}
else:
token = await send_token(
remaining_balance, refund_currency, effective_refund_mint
)
effective_refund_mint = token_mint_url(token, effective_refund_mint)
result = {"token": token}
if key.refund_currency == "sat":
result["sats"] = str(remaining_balance_msats // 1000)
else:
result["msats"] = str(remaining_balance_msats)
if "token" in result:
logger.info(
"refund_wallet_endpoint: cashu token issued",
extra={
"path": "/v1/wallet/refund",
"token": result["token"],
"amount": remaining_balance,
"currency": key.refund_currency or "sat",
},
)
except MeltOutcomeAmbiguousError as e:
# The melt was dispatched and may still settle. Restoring the balance
# here would let the same debit be paid out twice; keep the debit and
# leave the outcome to reconciliation.
logger.error(
"refund_wallet_endpoint: melt outcome ambiguous; balance withheld "
"pending reconciliation",
extra={
"error": str(e),
"hashed_key": key.hashed_key,
"remaining_balance": remaining_balance,
"refund_currency": key.refund_currency,
"refund_mint_url": key.refund_mint_url,
},
)
raise HTTPException(
status_code=502,
detail=(
"Refund was dispatched but its outcome is unconfirmed; the "
"balance is withheld until reconciliation completes"
),
)
except HTTPException:
# Minting failed — restore the debited balance
await _restore_balance(
session,
key.hashed_key,
pre_debit_balance,
pre_debit_reserved,
key.refund_mint_url or "",
)
raise
except Exception as e:
# Minting failed — restore the debited balance
await _restore_balance(
session,
key.hashed_key,
pre_debit_balance,
pre_debit_reserved,
key.refund_mint_url or "",
)
error_msg = str(e)
logger.error(
"refund_wallet_endpoint: mint/send failed",
extra={
"error": error_msg,
"error_type": type(e).__name__,
"hashed_key": key.hashed_key,
"remaining_balance": remaining_balance,
"refund_currency": key.refund_currency,
"refund_mint_url": key.refund_mint_url,
"has_refund_address": bool(key.refund_address),
},
)
if is_mint_connection_error(e):
raise HTTPException(status_code=503, detail="Mint service unavailable")
else:
raise HTTPException(status_code=500, detail="Refund failed")
await _refund_cache_set(bearer_value, result)
if "token" in result:
await store_cashu_transaction(
token=result["token"],
amount=remaining_balance,
unit=key.refund_currency or "sat",
mint_url=effective_refund_mint,
typ="out",
collected=False,
source="apikey",
api_key_hashed_key=key.hashed_key,
)
logger.info(
"refund_wallet_endpoint: refund successful",
extra={
"refunded_msats": remaining_balance_msats,
"previous_reserved_balance": key.reserved_balance,
},
)
return result
return await refund.execute(session, claim)
@router.get("/history")
@@ -650,12 +425,6 @@ async def wallet_history(
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
) -> dict[str, list[dict[str, str | int | bool | None]]]:
if key.parent_key_hash:
raise HTTPException(
status_code=400,
detail="Cannot view child key history. Please use the parent key instead.",
)
result = await session.exec(
select(CashuTransaction)
.where(CashuTransaction.api_key_hashed_key == key.hashed_key)
@@ -693,132 +462,6 @@ async def donate(token: str, ref: str | None = None) -> str:
return "Invalid token."
class ChildKeyRequest(BaseModel):
count: int
balance_limit: int | None = None
balance_limit_reset: str | None = None
validity_date: int | None = None
@router.post("/child-key")
async def create_child_key(
payload: ChildKeyRequest,
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
) -> dict:
"""Creates one or more child API keys that use the parent's balance."""
# Log incoming request for debugging
logger.debug(f"Child key creation request: count={payload.count}")
count = payload.count
if count < 1 or count > 50:
raise HTTPException(status_code=400, detail="Count must be between 1 and 50.")
# Check if this is already a child key
if key.parent_key_hash:
raise HTTPException(
status_code=400,
detail="Cannot create a child key for another child key.",
)
cost_per_key = settings.child_key_cost
total_cost = cost_per_key * count
if key.total_balance < total_cost:
raise HTTPException(
status_code=402,
detail=f"Insufficient balance to create {count} child keys. {total_cost} mSats required.",
)
# Deduct cost from parent atomically — guards against concurrent requests
# that both pass the balance check above on stale in-memory state.
deduct_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.balance) - col(ApiKey.reserved_balance) >= total_cost)
.values(
balance=col(ApiKey.balance) - total_cost,
total_spent=col(ApiKey.total_spent) + total_cost,
)
)
result = await session.exec(deduct_stmt) # type: ignore[call-overload]
if result.rowcount == 0:
raise HTTPException(
status_code=402,
detail=f"Insufficient balance to create {count} child keys. {total_cost} mSats required.",
)
# Generate new keys
import secrets
new_keys = []
for _ in range(count):
new_key_raw = secrets.token_hex(32)
new_key_hash = new_key_raw # We use the raw key as the hash for sk- keys
child_key = ApiKey(
hashed_key=new_key_hash,
balance=0,
parent_key_hash=key.hashed_key,
balance_limit=payload.balance_limit,
balance_limit_reset=payload.balance_limit_reset,
balance_limit_reset_date=int(time.time())
if payload.balance_limit_reset
else None,
validity_date=payload.validity_date,
)
session.add(child_key)
new_keys.append("sk-" + new_key_hash)
await session.commit()
await session.refresh(key)
response_data = {
"api_keys": new_keys,
"count": count,
"cost_msats": total_cost,
"cost_sats": total_cost // 1000,
"parent_balance": key.balance,
"parent_balance_sats": key.balance // 1000,
}
logger.debug(f"Child key creation response: {response_data}")
return response_data
class ChildKeyResetRequest(BaseModel):
child_key: str
@router.post("/child-key/reset")
async def reset_child_key_spent(
payload: ChildKeyResetRequest,
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
) -> dict:
"""Resets the total_spent of a child key. Must be called by the parent."""
child_key_raw = payload.child_key
if child_key_raw.startswith("sk-"):
child_key_raw = child_key_raw[3:]
child_key = await session.get(ApiKey, child_key_raw)
if not child_key:
raise HTTPException(status_code=404, detail="Child key not found.")
if child_key.parent_key_hash != key.hashed_key:
raise HTTPException(
status_code=403, detail="Unauthorized. You are not the parent of this key."
)
child_key.total_spent = 0
if child_key.balance_limit_reset:
child_key.balance_limit_reset_date = int(time.time())
session.add(child_key)
await session.commit()
return {"success": True, "message": "Child key balance reset successfully."}
@router.api_route(
"/{path:path}",
methods=["GET", "POST", "PUT", "DELETE"],
+108
View File
@@ -0,0 +1,108 @@
"""Compatibility shim that keeps cashu 0.20.x working on httpx>=0.28.
cashu's ``async_set_httpx_client`` decorator builds the client for *every* mint
call as ``httpx.AsyncClient(proxies=proxies_dict, ...)``. httpx deprecated
``proxies`` in 0.26 and removed it in 0.28, so on httpx>=0.28 every wallet
operation routstr performs -- ``load_mint_keysets``, ``mint_quote``,
``melt_quote``, token redeem -- raises::
TypeError: AsyncClient.__init__() got an unexpected keyword argument 'proxies'
Upstream cashu (0.20.3, the latest release) still caps ``httpx<0.26`` and has no
release that fixes this, while litellm>=1.84 requires ``httpx>=0.28``. We bridge
the gap by translating the keyword *inside cashu's module namespace only* --
global httpx behaviour is untouched, and the shim is a no-op on httpx<0.28.
Delete this module once cashu ships a release that passes ``proxy=``.
"""
from typing import Any
import httpx
__all__ = ["install_cashu_httpx_shim"]
_ALL_SCHEMES = "all://"
def _single_proxy(proxies: Any) -> str | None:
"""Collapse an httpx<0.28 ``proxies`` mapping into a single ``proxy`` URL.
cashu only ever builds ``{}`` or ``{"all://": url}``, so a mapping with one
distinct URL is all we need to support. Anything richer is unrepresentable
as httpx 0.28's scalar ``proxy=``; we fail closed and raise rather than
return ``None``, because dropping the entry would silently send mint
traffic direct instead of through the configured Tor/SOCKS proxy.
"""
if not proxies:
return None
if isinstance(proxies, str):
return proxies
if isinstance(proxies, dict):
if _ALL_SCHEMES in proxies:
value = proxies[_ALL_SCHEMES]
return str(value) if value is not None else None
distinct = {str(v) for v in proxies.values() if v is not None}
if not distinct:
return None
if len(distinct) == 1:
return distinct.pop()
raise ValueError(
f"cannot represent proxies={proxies!r} as httpx 0.28 proxy=; "
"refusing to send proxied traffic direct"
)
class _ProxiesCompatAsyncClient(httpx.AsyncClient):
"""``httpx.AsyncClient`` that still accepts the removed ``proxies`` kwarg."""
def __init__(self, *args: Any, **kwargs: Any) -> None:
if "proxies" in kwargs:
proxies = kwargs.pop("proxies")
proxy = _single_proxy(proxies)
if proxy is not None and kwargs.get("proxy") is None:
kwargs["proxy"] = proxy
elif isinstance(proxies, dict) and not proxies:
# In httpx<0.28, proxies={} disabled environment proxy discovery.
kwargs.setdefault("trust_env", False)
super().__init__(*args, **kwargs)
class _HttpxNamespace:
"""Stand-in for the ``httpx`` module inside cashu's ``v1_api``.
Every attribute resolves against the real module except ``AsyncClient``,
so cashu keeps using genuine httpx types everywhere else.
"""
AsyncClient = _ProxiesCompatAsyncClient
def __getattr__(self, name: str) -> Any:
return getattr(httpx, name)
def _httpx_accepts_proxies() -> bool:
import inspect
try:
return "proxies" in inspect.signature(httpx.AsyncClient.__init__).parameters
except (TypeError, ValueError): # pragma: no cover - defensive
return False
def install_cashu_httpx_shim() -> bool:
"""Patch cashu's mint client to survive httpx>=0.28.
Returns True when the shim was installed, False when it wasn't needed.
Safe to call repeatedly.
"""
if _httpx_accepts_proxies():
return False
from cashu.wallet import v1_api
if isinstance(getattr(v1_api, "httpx", None), _HttpxNamespace):
return True
v1_api.httpx = _HttpxNamespace() # type: ignore[assignment]
return True
+279 -113
View File
@@ -1,4 +1,3 @@
import asyncio
import json
import re
import secrets
@@ -6,12 +5,17 @@ from datetime import datetime, timezone
from pathlib import Path
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from pydantic import BaseModel, RootModel
from pydantic import BaseModel, RootModel, field_validator
from pydantic.v1 import ValidationError as PydanticValidationError
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from ..payment.models import _row_to_model, list_models
from ..payment.models import (
REQUIRED_PRICING_FIELDS,
_row_to_model,
list_models,
)
from ..payment.rates import BILLABLE_PRICING_FIELDS, coerce_rate
from ..proxy import refresh_model_maps, reinitialize_upstreams
from ..wallet import fetch_all_balances, send_token, token_mint_url
from . import vault
@@ -30,6 +34,7 @@ from .db import (
from .db import (
store_cashu_transaction_with_retry as store_cashu_transaction,
)
from .exceptions import json_compliant
from .log_manager import log_manager
from .logging import get_logger
from .provider_slugs import allocate_unique_provider_slug
@@ -69,7 +74,9 @@ async def require_admin_api(request: Request) -> None:
async with create_session() as session:
result = await session.exec(select(CliToken).where(CliToken.token == token))
cli_token = result.first()
if cli_token and (cli_token.expires_at is None or cli_token.expires_at > now_ts):
if cli_token and (
cli_token.expires_at is None or cli_token.expires_at > now_ts
):
cli_token.last_used_at = now_ts
session.add(cli_token)
await session.commit()
@@ -104,25 +111,28 @@ async def get_temporary_balances_api(
)
total = count_result.one()
# Aggregate totals across the whole (search-filtered) set, not just the
# current page. Balance counts only parent (non-child) keys to avoid
# double-counting, since child keys draw from their parent's balance.
totals_result = await session.exec(
# Aggregate totals across the whole search-filtered set, not just this page.
balance_totals_result = await session.exec(
select(
func.coalesce(func.sum(ApiKey.balance), 0),
func.coalesce(func.sum(ApiKey.reserved_balance), 0),
func.coalesce(
func.sum(
case(
(col(ApiKey.parent_key_hash).is_(None), ApiKey.balance),
else_=0,
)
),
0,
func.sum(col(ApiKey.balance) - col(ApiKey.reserved_balance)), 0
),
).where(*filters)
)
(
total_balance,
total_reserved_balance,
total_available_balance,
) = balance_totals_result.one()
usage_totals_result = await session.exec(
select(
func.coalesce(func.sum(ApiKey.total_spent), 0),
func.coalesce(func.sum(ApiKey.total_requests), 0),
).where(*filters)
)
total_balance, total_spent, total_requests = totals_result.one()
total_spent, total_requests = usage_totals_result.one()
# Latest created first; keys with no created_at (legacy rows) sort last.
# Use an explicit CASE rather than relying on dialect NULL-ordering so
@@ -143,13 +153,12 @@ async def get_temporary_balances_api(
{
"hashed_key": key.hashed_key,
"balance": key.balance,
"reserved_balance": key.reserved_balance,
"available_balance": key.total_balance,
"total_spent": key.total_spent,
"total_requests": key.total_requests,
"refund_address": key.refund_address,
"key_expiry_time": key.key_expiry_time,
"parent_key_hash": key.parent_key_hash,
"balance_limit": key.balance_limit,
"balance_limit_reset": key.balance_limit_reset,
"validity_date": key.validity_date,
"created_at": key.created_at,
}
@@ -158,6 +167,8 @@ async def get_temporary_balances_api(
"total": total,
"totals": {
"total_balance": total_balance,
"total_reserved_balance": total_reserved_balance,
"total_available_balance": total_available_balance,
"total_spent": total_spent,
"total_requests": total_requests,
},
@@ -165,8 +176,6 @@ async def get_temporary_balances_api(
class ApiKeyUpdate(BaseModel):
balance_limit: int | None = None
balance_limit_reset: str | None = None
validity_date: int | None = None
@@ -181,10 +190,6 @@ async def update_apikey(
if not key:
raise HTTPException(status_code=404, detail="API key not found")
if update.balance_limit is not None:
key.balance_limit = update.balance_limit
if update.balance_limit_reset is not None:
key.balance_limit_reset = update.balance_limit_reset
if update.validity_date is not None:
key.validity_date = update.validity_date
@@ -194,8 +199,6 @@ async def update_apikey(
return {
"hashed_key": key.hashed_key,
"balance_limit": key.balance_limit,
"balance_limit_reset": key.balance_limit_reset,
"validity_date": key.validity_date,
}
@@ -256,16 +259,12 @@ async def update_password(request: Request, password_update: PasswordUpdate) ->
secret = await get_secret(session)
if not secret.admin_password_hash:
raise HTTPException(
status_code=500, detail="Admin password not configured"
)
raise HTTPException(status_code=500, detail="Admin password not configured")
if not vault.verify_password(
password_update.current_password, secret.admin_password_hash
):
raise HTTPException(
status_code=401, detail="Current password is incorrect"
)
raise HTTPException(status_code=401, detail="Current password is incorrect")
# Validate new password
new_password = password_update.new_password.strip()
@@ -483,6 +482,38 @@ class ModelCreate(BaseModel):
enabled: bool = True
forwarded_model_id: str | None = None
@field_validator("pricing")
@classmethod
def _validate_pricing(cls, value: dict[str, object]) -> dict[str, object]:
"""Reject a rate that is malformed, non-finite, negative or not there.
A present-but-invalid rate would otherwise slip through: a non-numeric
string coerces to $0 on the read path (an unpriced-looking row), while a
negative or ``NaN``/``inf`` value is truthy and reads back as a real
price, so the model could be enabled and bill a nonsensical amount.
Surfacing a 422 reports the client bug as a client bug instead of
persisting it. Numeric strings (``"0.000005"``) stay valid, and so does
an omitted auxiliary rate — the stored JSON accepts both.
"""
for field in BILLABLE_PRICING_FIELDS:
if field not in value:
# ``dict.get`` cannot tell this from an explicit ``null``, so
# both were skipped and a row that ``Pricing`` cannot parse was
# written — and then raised out of the response that reads it
# back, after the row had been committed.
if field in REQUIRED_PRICING_FIELDS:
raise ValueError(f"{field} is required")
continue
# The shared coercion also absorbs the OverflowError an oversized
# integer raises, which pydantic does not convert into a validation
# error — unhandled it escaped as a 500 for a bad client value.
if coerce_rate(value[field]) is None:
raise ValueError(
f"{field} must be a finite, non-negative number, "
f"got {value[field]!r}"
)
return value
def _normalize_forwarded_model_id(value: str | None) -> str | None:
if value is None:
@@ -609,9 +640,13 @@ async def get_provider_model(provider_id: str, model_id: str) -> dict[str, objec
raise HTTPException(
status_code=404, detail="Model not found for this provider"
)
return _row_to_model(
row, apply_provider_fee=False, provider_fee=provider.provider_fee
).dict() # type: ignore
# Same duty as the listing this view is opened from: a stored rate that
# is not a usable number must be shown as it is, not encoded as `null`.
return json_compliant( # type: ignore[return-value]
_row_to_model(
row, apply_provider_fee=False, provider_fee=provider.provider_fee
).dict()
)
@admin_router.delete(
@@ -870,29 +905,43 @@ class UpstreamProviderUpdateBySlug(BaseModel):
provider_settings: dict | None = None
async def _active_ppq_claim_in_session(session: AsyncSession, provider_id: int) -> bool:
async def _active_auto_topup_claim_in_session(
session: AsyncSession, provider_id: int, provider_type: str
) -> bool:
"""Check for an active claim inside the caller's transaction.
Must share the transaction of whatever destructive write it is guarding —
a check in its own session leaves a window for a worker to create the
claim between the check and the commit.
"""
from ..upstream.auto_topup import _ppq_state_id_for_provider
from ..upstream.auto_topup import (
_ppq_state_id_for_provider,
_routstr_state_id_for_provider,
)
claim = await session.get(CashuTransaction, _ppq_state_id_for_provider(provider_id))
state_id = (
_ppq_state_id_for_provider(provider_id)
if provider_type == "ppqai"
else _routstr_state_id_for_provider(provider_id)
)
claim = await session.get(CashuTransaction, state_id)
return claim is not None and not claim.collected and not claim.swept
def _require_valid_ppq_auto_topup(
provider_type: str, settings: dict | None
) -> None:
"""Reject PPQ auto top-up settings the worker would later refuse."""
if provider_type != "ppqai":
def _require_valid_auto_topup(provider_type: str, settings: dict | None) -> None:
"""Reject auto top-up settings the worker would later refuse."""
from ..upstream.auto_topup import (
validate_ppq_auto_topup_settings,
validate_routstr_auto_topup_settings,
)
if provider_type == "ppqai":
problem = validate_ppq_auto_topup_settings(settings)
elif provider_type == "routstr":
problem = validate_routstr_auto_topup_settings(settings)
else:
return
from ..upstream.auto_topup import validate_ppq_auto_topup_settings
problem = validate_ppq_auto_topup_settings(settings)
if problem is not None:
raise HTTPException(status_code=400, detail=problem)
@@ -917,16 +966,18 @@ async def _apply_provider_update(
)
if (
provider_type_changed
and provider.provider_type == "ppqai"
and provider.provider_type in ("ppqai", "routstr")
and provider.id is not None
and await _active_ppq_claim_in_session(session, provider.id)
and await _active_auto_topup_claim_in_session(
session, provider.id, provider.provider_type
)
):
# Changing the type would orphan the claim: the PPQ endpoints refuse
# non-ppqai providers, so nobody could ever inspect or release it.
# Changing the type would orphan the claim: the claim endpoints refuse
# providers of the wrong type, so nobody could inspect or release it.
raise HTTPException(
status_code=409,
detail=(
"This provider has an active PPQ auto top-up claim. Release "
"This provider has an active auto top-up claim. Release "
"it before changing the provider type"
),
)
@@ -977,7 +1028,7 @@ async def _apply_provider_update(
except (json.JSONDecodeError, TypeError):
effective_settings = None
if effective_settings is not None:
_require_valid_ppq_auto_topup(provider.provider_type, effective_settings)
_require_valid_auto_topup(provider.provider_type, effective_settings)
if payload.provider_settings is not None:
provider.provider_settings = json.dumps(payload.provider_settings)
@@ -1017,9 +1068,7 @@ async def create_upstream_provider(
else:
slug = await allocate_unique_provider_slug(session, payload.provider_type)
_require_valid_ppq_auto_topup(
payload.provider_type, payload.provider_settings
)
_require_valid_auto_topup(payload.provider_type, payload.provider_settings)
provider = UpstreamProviderRow(
slug=slug,
@@ -1078,9 +1127,7 @@ async def update_upstream_provider_by_slug(
lookup = _validate_slug(payload.slug)
async with create_session() as session:
result = await session.exec(
select(UpstreamProviderRow).where(
UpstreamProviderRow.slug == lookup
)
select(UpstreamProviderRow).where(UpstreamProviderRow.slug == lookup)
)
provider = result.first()
if not provider:
@@ -1117,16 +1164,19 @@ async def delete_upstream_provider(provider_id: str) -> dict[str, object]:
# re-reads the provider inside its own transaction, so these two
# writes serialise — either the claim lands first and this 409s, or
# the delete lands first and the worker refuses to claim.
if provider.provider_type == "ppqai" and await _active_ppq_claim_in_session(
session, deleted_id
if provider.provider_type in (
"ppqai",
"routstr",
) and await _active_auto_topup_claim_in_session(
session, deleted_id, provider.provider_type
):
# Deleting now would orphan the claim and any funds it tracks:
# the PPQ endpoints 404 without the provider row, so the claim
# the claim endpoints 404 without the provider row, so the claim
# could never again be inspected or released.
raise HTTPException(
status_code=409,
detail=(
"This provider has an active PPQ auto top-up claim. "
"This provider has an active auto top-up claim. "
"Resolve and release it before deleting the provider"
),
)
@@ -1186,8 +1236,13 @@ async def get_provider_models(provider_id: str) -> dict[str, object]:
"provider_type": provider.provider_type,
"base_url": provider.base_url,
},
"db_models": [m.dict() for m in db_models],
"remote_models": [m.dict() for m in filtered_remote_models],
# This listing includes disabled models, so it is the one view that
# still carries a row the served-catalog backstop holds back —
# including one whose stored rate is not a usable number. The
# encoder would report that rate as `null`, indistinguishable from a
# missing one; show the operator the value that needs fixing.
"db_models": [json_compliant(m.dict()) for m in db_models],
"remote_models": [json_compliant(m.dict()) for m in filtered_remote_models],
}
@@ -1316,50 +1371,37 @@ async def initiate_provider_topup(
else {}
)
last_status_code = 500
last_error_detail: object = "Failed to create top-up invoice"
# Quote creation is unsafe to retry without idempotency.
resp = await client.post(
f"{clean_url}/v1/balance/lightning/invoice",
json=request_json,
headers=headers,
)
# 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,
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",
},
)
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)
}
logger.error(
f"Upstream topup request failed: {resp.text}",
extra={
"provider_id": provider_id,
"status_code": resp.status_code,
},
)
try:
error_detail: object = resp.json()
except Exception:
error_detail = resp.text
raise HTTPException(
status_code=last_status_code, detail=last_error_detail
status_code=resp.status_code, detail=error_detail
)
upstream_instance = _instantiate_provider(provider)
@@ -1718,6 +1760,34 @@ async def get_logs_api(
}
@admin_router.get(
"/api/logs/request/{request_id}", dependencies=[Depends(require_admin_api)]
)
async def get_logs_by_request_id_api(
request: Request,
request_id: str,
date: str | None = None,
limit: int = Query(default=200, ge=1, le=1000),
) -> dict[str, object]:
"""
Get every log entry belonging to a single request ID, oldest first.
"""
log_entries = log_manager.search_logs(
date=date,
request_id=request_id,
limit=limit,
)
log_entries.sort(key=lambda entry: str(entry.get("asctime", "")))
return {
"logs": log_entries,
"total": len(log_entries),
"request_id": request_id,
"date": date,
"limit": limit,
}
@admin_router.get("/api/logs/dates", dependencies=[Depends(require_admin_api)])
async def get_log_dates_api(request: Request) -> dict[str, object]:
logs_dir = Path("logs")
@@ -1811,6 +1881,87 @@ async def release_ppq_auto_topup_api(
return {"ok": True, "released": True}
_ROUTSTR_RELEASE_ERRORS = {
"no_active_claim": "No active Routstr auto top-up claim to release",
"stale_state": "The claim changed since it was reviewed; reload and check again",
"claim_changed": (
"The claim changed while the release was being applied; reload and check again"
),
}
class ReleaseRoutstrAutoTopupRequest(BaseModel):
confirmed_peer_reconciled: bool
state_token: str | None = None
async def _require_routstr_provider(provider_id: int) -> UpstreamProviderRow:
async with create_session() as session:
provider = await session.get(UpstreamProviderRow, provider_id)
if provider is None:
raise HTTPException(status_code=404, detail="Provider not found")
if provider.provider_type != "routstr":
raise HTTPException(status_code=400, detail="Provider is not a Routstr node")
return provider
@admin_router.get(
"/api/upstream-providers/{provider_id}/routstr-auto-topup",
dependencies=[Depends(require_admin_api)],
)
async def get_routstr_auto_topup_api(provider_id: int) -> dict[str, object]:
await _require_routstr_provider(provider_id)
from ..upstream.auto_topup import get_routstr_auto_topup_state
return {"ok": True, **await get_routstr_auto_topup_state(provider_id)}
@admin_router.post(
"/api/upstream-providers/{provider_id}/routstr-auto-topup/release",
dependencies=[Depends(require_admin_api)],
)
async def release_routstr_auto_topup_api(
provider_id: int, payload: ReleaseRoutstrAutoTopupRequest
) -> dict[str, object]:
await _require_routstr_provider(provider_id)
if not payload.confirmed_peer_reconciled:
raise HTTPException(
status_code=400,
detail="Confirm the peer credited or returned the token before releasing",
)
from ..upstream.auto_topup import release_routstr_auto_topup_state
outcome = await release_routstr_auto_topup_state(
provider_id, state_token=payload.state_token
)
if not outcome.released:
raise HTTPException(
status_code=409, detail=_ROUTSTR_RELEASE_ERRORS[outcome.reason]
)
logger.warning(
"Admin released Routstr auto top-up claim after manual reconciliation",
extra={"provider_id": provider_id, "state_token": payload.state_token},
)
return {"ok": True, "released": True}
def _transaction_status(tx: CashuTransaction) -> str:
"""An outgoing admin withdrawal ends at "issued": the node hands the bearer
token over and never learns whether it was redeemed, so its flags stay false
and it would otherwise read pending forever. Incoming admin rows are redeemed
by the node itself and keep the normal collected/swept lifecycle.
"""
if tx.swept:
return "swept"
if tx.collected:
return "collected"
if tx.source == "admin" and tx.type == "out":
return "issued"
return "pending"
@admin_router.get("/api/transactions", dependencies=[Depends(require_admin_api)])
async def get_transactions_api(
type: str | None = None,
@@ -1823,10 +1974,11 @@ async def get_transactions_api(
async with create_session() as session:
from sqlmodel import col, func
# Hide only the deterministic PPQ claim-lock rows. Append-only PPQ
# payment rows remain visible as the audit trail for irreversible melts.
# Hide only the deterministic claim-lock rows. Append-only PPQ payment
# rows and auto-topup token rows remain visible as the audit trail.
base = select(CashuTransaction).where(
~col(CashuTransaction.id).like("ppq-auto-topup-%")
~col(CashuTransaction.id).like("ppq-auto-topup-%"),
~col(CashuTransaction.id).like("routstr-auto-topup-%"),
)
if type:
base = base.where(CashuTransaction.type == type)
@@ -1840,13 +1992,25 @@ async def get_transactions_api(
base = base.where(CashuTransaction.source == source)
if status:
if status == "collected":
base = base.where(CashuTransaction.collected == True) # noqa: E712
base = base.where(
CashuTransaction.collected == True, # noqa: E712
CashuTransaction.swept == False, # noqa: E712
)
elif status == "swept":
base = base.where(CashuTransaction.swept == True) # noqa: E712
elif status == "issued":
base = base.where(
CashuTransaction.source == "admin",
CashuTransaction.type == "out",
CashuTransaction.collected == False, # noqa: E712
CashuTransaction.swept == False, # noqa: E712
)
elif status == "pending":
base = base.where(
CashuTransaction.collected == False, # noqa: E712
CashuTransaction.swept == False, # noqa: E712
(CashuTransaction.source != "admin")
| (CashuTransaction.type != "out"),
)
if search:
@@ -1873,15 +2037,17 @@ async def get_transactions_api(
return {
"transactions": [
tx.dict(exclude={"sweep_started_at"}) for tx in transactions
{
**tx.dict(exclude={"sweep_started_at"}),
"status": _transaction_status(tx),
}
for tx in transactions
],
"total": total,
}
@admin_router.get(
"/api/lightning-invoices", dependencies=[Depends(require_admin_api)]
)
@admin_router.get("/api/lightning-invoices", dependencies=[Depends(require_admin_api)])
async def get_lightning_invoices_api(
status: str | None = None,
purpose: str | None = None,
+181 -76
View File
@@ -12,12 +12,11 @@ from typing import AsyncGenerator
from alembic import command
from alembic.config import Config
from alembic.util.exc import CommandError
from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_
from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_, text
from sqlalchemy.engine import make_url
from sqlalchemy.exc import IntegrityError, OperationalError
from sqlalchemy.ext.asyncio import AsyncEngine
from sqlalchemy.ext.asyncio.engine import create_async_engine
from sqlalchemy.orm import aliased
from sqlmodel import Field, Relationship, SQLModel, col, func, select, update
from sqlmodel.ext.asyncio.session import AsyncSession
@@ -37,6 +36,9 @@ def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine:
is_memory_sqlite = is_sqlite and url.database in {None, "", ":memory:"}
pool_pre_ping = settings.database_pool_pre_ping or not is_sqlite
options: dict[str, int | float | bool] = {"pool_pre_ping": pool_pre_ping}
connect_args: dict[str, object] = {}
if is_sqlite and not is_memory_sqlite:
connect_args["timeout"] = settings.database_busy_timeout
if not is_memory_sqlite:
options.update(
pool_size=settings.database_pool_size,
@@ -51,9 +53,12 @@ def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine:
"database_url_backend": backend,
"in_memory_sqlite": is_memory_sqlite,
**options,
"connect_args": connect_args,
},
)
created_engine = create_async_engine(database_url, echo=False, **options)
created_engine = create_async_engine(
database_url, echo=False, connect_args=connect_args, **options
)
hold_warn_seconds = settings.database_pool_hold_warn_seconds
def record_pool_checkout(
@@ -132,21 +137,6 @@ class ApiKey(SQLModel, table=True): # type: ignore
default=None,
description="Currency of the cashu-token",
)
parent_key_hash: str | None = Field(
default=None, foreign_key="api_keys.hashed_key", index=True
)
balance_limit: int | None = Field(
default=None,
description="Max spendable balance in msats for this key (mostly for child keys)",
)
balance_limit_reset: str | None = Field(
default=None,
description="Reset policy for balance limit (manual, daily, monthly, etc.)",
)
balance_limit_reset_date: int | None = Field(
default=None,
description="Unix timestamp of the last time the balance limit was reset",
)
validity_date: int | None = Field(
default=None,
description="Unix timestamp after which the key is no longer valid",
@@ -171,6 +161,55 @@ async def reset_all_reserved_balances(session: AsyncSession) -> None:
logger.info("Reset reserved balances on startup")
async def _transition_stale_reservation(
session: AsyncSession, reservation_id: str, cutoff: int
) -> bool:
"""Mark one reservation released iff its lease is still older than cutoff.
``created_at`` doubles as the heartbeat lease timestamp, so the guard must
be part of this update: a reservation renewed between the sweeper's select
and this transition is in flight and must survive.
"""
transition = await session.exec( # type: ignore[call-overload]
update(ReservationRelease)
.where(col(ReservationRelease.id) == reservation_id)
.where(col(ReservationRelease.status) == "active")
.where(col(ReservationRelease.created_at) < cutoff)
.values(status="released")
)
return bool(transition.rowcount == 1)
async def _release_legacy_aggregate(
session: AsyncSession,
key_hash: str,
observed_reserved: int,
observed_reserved_at: int | None,
) -> bool:
"""Zero one legacy aggregate reservation iff it is exactly as observed.
A new reservation committing between the sweeper's read and this update
changes ``reserved_balance``/``reserved_at`` in the same transaction that
creates its durable row, so this compare-and-swap fails instead of erasing
the newcomer's reserved funds.
"""
if observed_reserved <= 0:
return False
reserved_at_guard = (
col(ApiKey.reserved_at).is_(None)
if observed_reserved_at is None
else col(ApiKey.reserved_at) == observed_reserved_at
)
result = await session.exec( # type: ignore[call-overload]
update(ApiKey)
.where(col(ApiKey.hashed_key) == key_hash)
.where(col(ApiKey.reserved_balance) == observed_reserved)
.where(reserved_at_guard)
.values(reserved_balance=0, reserved_at=None)
)
return bool(result.rowcount == 1)
async def release_stale_reservations(
session: AsyncSession,
max_age_seconds: int,
@@ -191,25 +230,22 @@ async def release_stale_reservations(
col(ReservationRelease.billing_key_hash) == key_hash,
)
)
reservations = (await session.exec(query)).all()
# Capture primitives: a repair rollback below would expire ORM instances.
reservation_rows = [
(r.id, r.key_hash, r.billing_key_hash, r.reserved_msats)
for r in (await session.exec(query)).all()
]
released = 0
for reservation in reservations:
transition = await session.exec( # type: ignore[call-overload]
update(ReservationRelease)
.where(col(ReservationRelease.id) == reservation.id)
.where(col(ReservationRelease.status) == "active")
.values(status="released")
)
if transition.rowcount != 1:
for res_id, res_key_hash, res_billing_hash, res_msats in reservation_rows:
if not await _transition_stale_reservation(session, res_id, cutoff):
continue
values = {
"reserved_balance": col(ApiKey.reserved_balance)
- reservation.reserved_msats,
"reserved_balance": col(ApiKey.reserved_balance) - res_msats,
"reserved_at": case(
(
col(ApiKey.reserved_balance) - reservation.reserved_msats > 0,
col(ApiKey.reserved_balance) - res_msats > 0,
col(ApiKey.reserved_at),
),
else_=None,
@@ -217,24 +253,42 @@ async def release_stale_reservations(
}
parent_result = await session.exec( # type: ignore[call-overload]
update(ApiKey)
.where(col(ApiKey.hashed_key) == reservation.billing_key_hash)
.where(col(ApiKey.reserved_balance) >= reservation.reserved_msats)
.where(col(ApiKey.hashed_key) == res_billing_hash)
.where(col(ApiKey.reserved_balance) >= res_msats)
.values(**values)
)
if parent_result.rowcount != 1:
await session.rollback()
return 0
if reservation.billing_key_hash != reservation.key_hash:
aggregates_ok = parent_result.rowcount == 1
if aggregates_ok and res_billing_hash != res_key_hash:
child_result = await session.exec( # type: ignore[call-overload]
update(ApiKey)
.where(col(ApiKey.hashed_key) == reservation.key_hash)
.where(col(ApiKey.reserved_balance) >= reservation.reserved_msats)
.where(col(ApiKey.hashed_key) == res_key_hash)
.where(col(ApiKey.reserved_balance) >= res_msats)
.values(**values)
)
if child_result.rowcount != 1:
await session.rollback()
return 0
aggregates_ok = child_result.rowcount == 1
if not aggregates_ok:
# The aggregates no longer hold this reservation's msats — the
# durable row is corrupt. Repair by terminalizing it WITHOUT
# subtracting uncertain aggregates (legacy cleanup below reconciles
# any stale remainder) and keep sweeping the rest of the batch:
# one corrupt row must not poison all stale cleanup.
await session.rollback()
if await _transition_stale_reservation(session, res_id, cutoff):
await session.commit()
released += 1
logger.error(
"Released corrupt stale reservation without aggregate subtraction",
extra={
"reservation_id": res_id,
"billing_key_hash": res_billing_hash[:8] + "...",
"reserved_msats": res_msats,
},
)
continue
# Commit each release on its own so a later corrupt record's rollback
# cannot discard the healthy releases already processed in this batch.
await session.commit()
released += 1
# Rolling upgrades can leave aggregate reservations created before durable
@@ -246,16 +300,13 @@ async def release_stale_reservations(
col(ApiKey.reserved_at) < cutoff
)
else:
legacy_query = legacy_query.where(
or_(
col(ApiKey.hashed_key) == key_hash,
col(ApiKey.parent_key_hash) == key_hash,
)
).where(
legacy_query = legacy_query.where(col(ApiKey.hashed_key) == key_hash).where(
or_(col(ApiKey.reserved_at).is_(None), col(ApiKey.reserved_at) < cutoff)
)
for legacy_key in (await session.exec(legacy_query)).all():
observed_reserved = legacy_key.reserved_balance
observed_reserved_at = legacy_key.reserved_at
active_owner = (
await session.exec(
select(ReservationRelease.id)
@@ -272,10 +323,10 @@ async def release_stale_reservations(
).first()
if active_owner is not None:
continue
legacy_key.reserved_balance = 0
legacy_key.reserved_at = None
session.add(legacy_key)
released += 1
if await _release_legacy_aggregate(
session, legacy_key.hashed_key, observed_reserved, observed_reserved_at
):
released += 1
await session.commit()
if released:
@@ -290,21 +341,15 @@ async def release_stale_reservations(
async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> int:
"""Delete dead parentless API keys; return the count removed.
"""Delete dead API keys; return the count removed.
Dead = 0 balance/reservation/spend/requests, older than the grace period,
no parent, no children, no invoice that could still settle. Cashu rows are
Dead = 0 balance/reservation/spend/requests, older than the grace
period, no invoice that could still settle. Cashu rows are
unlinked (not deleted) first to keep the audit trail.
"""
now = int(time.time())
cutoff = now - min_age_seconds
child = aliased(ApiKey)
has_children = (
select(child.hashed_key).where(
col(child.parent_key_hash) == col(ApiKey.hashed_key)
)
).exists()
# An expired invoice stays creditable for the grace window, and crediting it
# after its target key is gone strands the payment at the mint.
settleable_invoice = (
@@ -322,16 +367,22 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in
)
).exists()
has_refund_claim = (
select(Refund.id).where(
col(Refund.api_key_hashed_key) == col(ApiKey.hashed_key)
)
).exists()
eligible_hashes = (
select(ApiKey.hashed_key)
.where(col(ApiKey.balance) == 0)
.where(col(ApiKey.reserved_balance) == 0)
.where(col(ApiKey.total_spent) == 0)
.where(col(ApiKey.total_requests) == 0)
.where(col(ApiKey.parent_key_hash).is_(None))
.where((col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff))
.where(~settleable_invoice)
.where(~has_children)
# refunds holds a non-null FK to the key.
.where(~has_refund_claim)
)
# Unlink transactions rather than cascade-deleting them, so the financial
@@ -386,9 +437,10 @@ class ModelRow(SQLModel, table=True): # type: ignore
class ModelPathRow(SQLModel, table=True): # type: ignore
"""Upstream provider path a model is reachable through.
Discovery/visibility data only. ``model_id`` is intentionally NOT globally
unique: it is the client-visible ``/v1/models`` id (``forwarded_model_id or
id``) grouped across every provider that exposes the model. A single model
Discovery data plus provider-specific model metadata. ``model_id`` is
intentionally NOT globally unique: it is the client-visible ``/v1/models``
id (``forwarded_model_id or id``) grouped across every provider that exposes
the model. A single model
can therefore have several rows — one per direct provider path plus one per
OpenRouter sub-provider endpoint.
"""
@@ -425,6 +477,10 @@ class ModelPathRow(SQLModel, table=True): # type: ignore
endpoint_name: str | None = Field(
default=None, description="Human-readable endpoint display name"
)
model_metadata: str = Field(
default="{}",
description="JSON model metadata specific to this provider path",
)
upstream_provider_id: int = Field(
index=True,
foreign_key="upstream_providers.id",
@@ -470,14 +526,6 @@ class LightningInvoice(SQLModel, table=True): # type: ignore
)
expires_at: int = Field(description="Unix timestamp when invoice expires")
paid_at: int | None = Field(default=None, description="Unix timestamp when paid")
balance_limit: int | None = Field(
default=None,
description="Max spendable msats for the created key",
)
balance_limit_reset: str | None = Field(
default=None,
description="Reset policy for balance limit (daily, weekly, monthly)",
)
validity_date: int | None = Field(
default=None,
description="Unix timestamp after which the created key expires",
@@ -520,6 +568,53 @@ class CashuTransaction(SQLModel, table=True): # type: ignore
)
REFUND_OPEN_STATUSES = ("pending", "ambiguous")
# Debited from the key but neither paid out nor restored, so still owed.
REFUND_UNRESOLVED_STATUSES = ("pending", "ambiguous", "stuck")
_REFUND_OPEN_PREDICATE = "status IN ('pending', 'ambiguous')"
class Refund(SQLModel, table=True): # type: ignore
"""One payout claim; the partial unique index allows one open claim per key."""
__tablename__ = "refunds"
__table_args__ = (
Index(
"ux_refunds_open_per_key",
"api_key_hashed_key",
unique=True,
sqlite_where=text(_REFUND_OPEN_PREDICATE),
postgresql_where=text(_REFUND_OPEN_PREDICATE),
),
)
id: str = Field(primary_key=True, default_factory=lambda: uuid.uuid4().hex)
api_key_hashed_key: str = Field(foreign_key="api_keys.hashed_key", index=True)
method: str = Field(description="Payout method: lightning or cashu")
destination: str | None = Field(
default=None, description="Lightning address or LNURL, NULL for cashu"
)
amount_msats: int = Field(description="Balance debited when the claim opened")
unit: str = Field(description="Mint unit the payout is denominated in")
mint_url: str = Field(description="Mint the payout is drawn from")
status: str = Field(
default="pending",
index=True,
description="pending, paid, failed, ambiguous, or stuck",
)
quote_id: str | None = Field(
default=None, description="Melt quote id, for reconciling an ambiguous payout"
)
token: str | None = Field(default=None, description="Issued cashu token")
claimed_at: int | None = Field(
default=None, description="Reconciler lease timestamp"
)
created_at: int = Field(default_factory=lambda: int(time.time()))
updated_at: int = Field(default_factory=lambda: int(time.time()))
async def store_cashu_transaction(
token: str,
amount: int,
@@ -911,8 +1006,18 @@ async def complete_routstr_fee_payout(
async def total_user_liability(db_session: AsyncSession) -> int:
"""Return all outstanding API-key balances in millisatoshis."""
result = await db_session.exec(select(func.sum(ApiKey.balance)))
"""Return all outstanding user funds in millisatoshis.
Key balances and unresolved refunds are summed in one statement so a
claim opened between two reads cannot be missed by both.
"""
key_balances = select(func.coalesce(func.sum(ApiKey.balance), 0)).scalar_subquery()
unresolved_refunds = (
select(func.coalesce(func.sum(Refund.amount_msats), 0))
.where(col(Refund.status).in_(REFUND_UNRESOLVED_STATUSES))
.scalar_subquery()
)
result = await db_session.exec(select(key_balances + unresolved_refunds))
return int(result.one() or 0)
+75 -2
View File
@@ -1,4 +1,8 @@
import math
from fastapi import Request
from fastapi.encoders import jsonable_encoder
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse
from .logging import get_logger
@@ -30,6 +34,26 @@ class UpstreamError(Exception):
super().__init__(message)
class EhbpTimeoutError(UpstreamError):
"""Raised when an EHBP upstream times out waiting for a response.
Distinct from a generic :class:`UpstreamError` so callers can map the
failure to a ``504 Gateway Timeout`` with a stable ``UPSTREAM_TIMEOUT``
code instead of a misleading ``500`` internal server error.
``details`` carries optional structured, redaction-safe context and is
forwarded to the client by ``create_upstream_error_response``.
"""
def __init__(self, message: str, details: dict[str, object] | None = None):
super().__init__(
message,
status_code=504,
code="UPSTREAM_TIMEOUT",
details=details,
)
async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse:
"""Handle HTTP exceptions and include request ID in response."""
request_id = getattr(request.state, "request_id", "unknown")
@@ -40,15 +64,25 @@ async def http_exception_handler(request: Request, exc: Exception) -> JSONRespon
path = request.url.path
# 4xx is client behaviour; the uvicorn access log already records it.
# Only 5xx warrants a server-side warning/error log here.
if status_code >= 500:
logger.error(
error_type = None
if isinstance(detail, dict):
error = detail.get("error")
if isinstance(error, dict):
error_type = error.get("type")
log = (
logger.warning
if error_type in {"mint_unreachable", "mint_rate_limited"}
else logger.error
)
log(
f"HTTP {status_code} on {path}: {detail}",
extra={
"request_id": request_id,
"status_code": status_code,
"detail": detail,
"path": path,
"error_type": error_type,
},
)
@@ -61,6 +95,45 @@ async def http_exception_handler(request: Request, exc: Exception) -> JSONRespon
return JSONResponse(status_code=status_code, content=content)
def json_compliant(value: object) -> object:
"""Render non-finite floats as text so a reply carrying them can serialize.
``json`` parses the bare ``NaN``/``Infinity``/``-Infinity`` literals into
real floats, so a request body — and a stored row written from one — may
hold one anywhere. ``JSONResponse`` encodes with ``allow_nan=False`` and
raises on them, which would turn a reply that merely *quotes* the offending
value into a 500.
"""
if isinstance(value, float) and not math.isfinite(value):
return repr(value)
if isinstance(value, dict):
return {key: json_compliant(item) for key, item in value.items()}
if isinstance(value, (list, tuple)):
return [json_compliant(item) for item in value]
return value
async def validation_exception_handler(
request: Request, exc: Exception
) -> JSONResponse:
"""Answer a request-validation failure with a 422 that always serializes.
Pydantic echoes the rejected value back in each error's ``input`` field. A
non-finite float there breaks the encoder, so the 422 escapes as a 500 and
reports a client's bad rate as a server fault.
"""
request_id = getattr(request.state, "request_id", "unknown")
errors = exc.errors() if isinstance(exc, RequestValidationError) else []
return JSONResponse(
status_code=422,
content={
"detail": json_compliant(jsonable_encoder(errors)),
"request_id": request_id,
},
)
async def general_exception_handler(request: Request, exc: Exception) -> JSONResponse:
"""Handle general exceptions and include request ID in response."""
request_id = getattr(request.state, "request_id", "unknown")
+30 -12
View File
@@ -16,7 +16,7 @@ DO NOT modify or remove these messages without updating the usage tracking logic
- The 'token_cost', 'model', 'input_tokens', and 'output_tokens' fields are extracted for dashboard metrics
3. "Max cost payment finalized" (INFO) - routstr/auth.py
- Used as the successful completion fallback when token usage is unavailable
- Used for explicit flat-price/MaxCostData settlements; missing usage alone must not create this charge
- The 'charged_amount', 'model', 'input_tokens', and 'output_tokens' fields are extracted for dashboard metrics
4. "Payment processed successfully" (INFO) - routstr/auth.py
@@ -51,7 +51,7 @@ from pythonjsonlogger import jsonlogger
from rich.console import Console
from rich.logging import RichHandler
from .redaction import redact_obj, redact_org_ids
from .redaction import redact_field, redact_org_ids
# Only use RichHandler when stdout is a real TTY. In non-TTY contexts
# (docker logs, pipes, CI) Rich pads every line to width and wraps long
@@ -100,8 +100,10 @@ class DailyRotatingFileHandler(logging.handlers.TimedRotatingFileHandler):
self.baseFilename = new_filename
self.current_date = new_date
# FIX ME: not sure if we need this
# self._cleanup_old_files()
# `backupCount` alone never prunes these files: the base filename moves
# with the date, so the inherited rollover finds no siblings to expire
# and every day of logged credentials is retained indefinitely.
self._cleanup_old_files()
if not self.delay:
self.stream = self._open()
@@ -182,6 +184,18 @@ class RequestIdFilter(logging.Filter):
return True
class ClientAppFilter(logging.Filter):
"""Attach request-local app attribution to log records."""
def filter(self, record: logging.LogRecord) -> bool:
# Import here to avoid circular imports
from .middleware import UNKNOWN_CLIENT_APP, client_app_context
client_app = client_app_context.get(None)
record.client_app = client_app if client_app else UNKNOWN_CLIENT_APP
return True
# Standard ``LogRecord`` attributes that are never user-supplied ``extra``
# fields; skipped when redacting structured extras (``msg``/``message`` are
# handled separately above).
@@ -260,13 +274,11 @@ class SecurityFilter(logging.Filter):
# Structured `extra={...}` fields are emitted by the JSON formatter
# straight from the record dict and never pass through the message
# formatting above. Redact organization IDs from any string-valued
# extra so they cannot leak via structured logs.
# formatting above, so they need their own recursive pass.
for attr, value in list(record.__dict__.items()):
if attr in _NON_EXTRA_RECORD_ATTRS:
continue
if isinstance(value, (str, dict, list, tuple)):
record.__dict__[attr] = redact_obj(value)
record.__dict__[attr] = redact_field(attr, value)
except Exception:
pass
@@ -323,7 +335,7 @@ def setup_logging() -> None:
"rich_tracebacks": True,
"markup": True,
"console": _console,
"filters": ["request_id_filter", "security_filter"],
"filters": ["request_id_filter", "client_app_filter", "security_filter"],
}
else:
console_handler = {
@@ -331,7 +343,7 @@ def setup_logging() -> None:
"level": log_level,
"formatter": "plain",
"stream": "ext://sys.stdout",
"filters": ["request_id_filter", "security_filter"],
"filters": ["request_id_filter", "client_app_filter", "security_filter"],
}
LOGGING_CONFIG = {
@@ -340,7 +352,7 @@ def setup_logging() -> None:
"formatters": {
"json": {
"()": jsonlogger.JsonFormatter,
"format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s %(request_id)s",
"format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s %(request_id)s %(client_app)s",
"datefmt": "%Y-%m-%d %H:%M:%S",
},
"plain": {
@@ -351,6 +363,7 @@ def setup_logging() -> None:
"filters": {
"version_filter": {"()": VersionFilter},
"request_id_filter": {"()": RequestIdFilter},
"client_app_filter": {"()": ClientAppFilter},
"security_filter": {"()": SecurityFilter},
},
"handlers": {
@@ -364,7 +377,12 @@ def setup_logging() -> None:
"interval": 1, # Every 1 day
"backupCount": 30, # Keep 30 days of logs
"atTime": None, # Rotate at midnight (00:00)
"filters": ["version_filter", "request_id_filter", "security_filter"],
"filters": [
"version_filter",
"request_id_filter",
"client_app_filter",
"security_filter",
],
},
},
"loggers": {
+25 -11
View File
@@ -4,6 +4,7 @@ from pathlib import Path
from typing import AsyncGenerator
from fastapi import FastAPI
from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse, RedirectResponse
from fastapi.staticfiles import StaticFiles
@@ -13,10 +14,10 @@ from starlette.types import Scope
from ..auth import (
periodic_dead_key_prune,
periodic_key_reset,
periodic_stale_reservation_sweep,
)
from ..balance import balance_router, deprecated_wallet_router
from ..cashu_compat import install_cashu_httpx_shim
from ..lightning import (
lightning_router,
periodic_invoice_watcher,
@@ -31,13 +32,18 @@ from ..nostr.discovery import providers_router
from ..payment.models import models_router, update_sats_pricing
from ..payment.price import update_prices_periodically
from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically
from ..refund import periodic_refund_reconcile
from ..upstream.auto_topup import periodic_auto_topup
from ..upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing
from ..upstream.litellm_routing import configure_litellm
from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout
from .admin import admin_router
from .db import create_session, init_db, run_migrations
from .exceptions import general_exception_handler, http_exception_handler
from .exceptions import (
general_exception_handler,
http_exception_handler,
validation_exception_handler,
)
from .logging import get_logger, setup_logging
from .middleware import LoggingMiddleware
from .not_found import _NOT_FOUND_HTML, not_found_catch_all # noqa: F401
@@ -63,15 +69,21 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
models_refresh_task = None
model_maps_refresh_task = None
model_paths_refresh_task = None
key_reset_task = None
stale_reservation_task = None
dead_key_prune_task = None
auto_topup_task = None
refund_sweep_task = None
refund_reconcile_task = None
routstr_fee_task = None
invoice_watcher_task = None
try:
# cashu 0.20.x passes the `proxies` kwarg httpx removed in 0.28.
# routstr.wallet and routstr.payment.lnurl also install this at import;
# repeating it here keeps startup correct for any future module that
# reaches cashu's mint client without going through those two.
install_cashu_httpx_shim()
# Apply litellm-wide settings (drop_params, chat-completions URL,
# debug logging) before any upstream provider dispatches a request.
configure_litellm()
@@ -143,16 +155,18 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
refresh_model_paths_periodically(get_upstreams)
)
payout_task = asyncio.create_task(periodic_payout())
if global_settings.nsec:
nip91_task = asyncio.create_task(announce_provider())
# Always started: the loop idles until an NSEC is configured and re-reads
# it every iteration, so a key saved (or cleared) through the admin UI
# takes effect without a restart.
nip91_task = asyncio.create_task(announce_provider())
analytics_task = asyncio.create_task(publish_usage_analytics())
if global_settings.providers_refresh_interval_seconds > 0:
providers_task = asyncio.create_task(providers_cache_refresher())
key_reset_task = asyncio.create_task(periodic_key_reset())
stale_reservation_task = asyncio.create_task(periodic_stale_reservation_sweep())
dead_key_prune_task = asyncio.create_task(periodic_dead_key_prune())
auto_topup_task = asyncio.create_task(periodic_auto_topup())
refund_sweep_task = asyncio.create_task(periodic_refund_sweep())
refund_reconcile_task = asyncio.create_task(periodic_refund_reconcile())
routstr_fee_task = asyncio.create_task(periodic_routstr_fee_payout())
invoice_watcher_task = asyncio.create_task(periodic_invoice_watcher())
@@ -188,8 +202,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
model_maps_refresh_task.cancel()
if model_paths_refresh_task is not None:
model_paths_refresh_task.cancel()
if key_reset_task is not None:
key_reset_task.cancel()
if stale_reservation_task is not None:
stale_reservation_task.cancel()
if dead_key_prune_task is not None:
@@ -198,6 +210,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
auto_topup_task.cancel()
if refund_sweep_task is not None:
refund_sweep_task.cancel()
if refund_reconcile_task is not None:
refund_reconcile_task.cancel()
if routstr_fee_task is not None:
routstr_fee_task.cancel()
if invoice_watcher_task is not None:
@@ -223,8 +237,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
tasks_to_wait.append(model_maps_refresh_task)
if model_paths_refresh_task is not None:
tasks_to_wait.append(model_paths_refresh_task)
if key_reset_task is not None:
tasks_to_wait.append(key_reset_task)
if stale_reservation_task is not None:
tasks_to_wait.append(stale_reservation_task)
if dead_key_prune_task is not None:
@@ -233,6 +245,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
tasks_to_wait.append(auto_topup_task)
if refund_sweep_task is not None:
tasks_to_wait.append(refund_sweep_task)
if refund_reconcile_task is not None:
tasks_to_wait.append(refund_reconcile_task)
if routstr_fee_task is not None:
tasks_to_wait.append(routstr_fee_task)
if invoice_watcher_task is not None:
@@ -293,6 +307,7 @@ app.add_middleware(LoggingMiddleware)
# Add exception handlers
app.add_exception_handler(HTTPException, http_exception_handler) # type: ignore
app.add_exception_handler(RequestValidationError, validation_exception_handler)
app.add_exception_handler(Exception, general_exception_handler)
@@ -306,7 +321,6 @@ async def info() -> dict:
"mints": global_settings.cashu_mints,
"http_url": global_settings.http_url,
"onion_url": global_settings.onion_url,
"child_key_cost_msats": global_settings.child_key_cost,
}
+51 -2
View File
@@ -2,8 +2,10 @@ import time
import uuid
from contextvars import ContextVar
from typing import Callable
from urllib.parse import urlsplit
from fastapi import Request, Response
from starlette.datastructures import Headers
from starlette.middleware.base import BaseHTTPMiddleware
from .logging import get_logger
@@ -13,6 +15,41 @@ logger = get_logger(__name__)
# Context variable to store request ID across async context
request_id_context: ContextVar[str | None] = ContextVar("request_id")
client_app_context: ContextVar[str | None] = ContextVar("client_app")
UNKNOWN_CLIENT_APP = "unknown"
# Prefer OpenRouter app headers, then browser and SDK fallbacks.
_CLIENT_APP_HEADERS: tuple[str, ...] = (
"x-title",
"http-referer",
"referer",
"user-agent",
)
# Limit untrusted header data repeated in every log record.
_CLIENT_APP_MAX_LENGTH = 120
def client_app_from_headers(headers: Headers) -> str:
for header in _CLIENT_APP_HEADERS:
raw = headers.get(header)
if raw is None:
continue
cleaned = "".join(ch for ch in raw if ch.isprintable()).strip()
if header in ("http-referer", "referer"):
try:
url = urlsplit(cleaned)
if url.scheme not in ("http", "https") or not url.hostname:
continue
except ValueError:
continue
# Attribution needs the origin, not credentials or private page URLs.
cleaned = f"{url.scheme}://{url.netloc.rsplit('@', 1)[-1]}"
if cleaned:
return cleaned[:_CLIENT_APP_MAX_LENGTH]
return UNKNOWN_CLIENT_APP
# Methods that are never logged: HEAD requests are health probes from
# monitoring/load balancers, OPTIONS are CORS preflights — both are framework
@@ -71,6 +108,10 @@ class LoggingMiddleware(BaseHTTPMiddleware):
# Set request ID in context for logging
token = request_id_context.set(request_id)
client_app_token = client_app_context.set(
client_app_from_headers(request.headers)
)
path = request.url.path
should_log = _should_log(request.method, path)
@@ -84,7 +125,9 @@ class LoggingMiddleware(BaseHTTPMiddleware):
"request_id": request_id,
"method": request.method,
"path": path,
"query_params": dict(request.query_params),
# Names only: query values carry API keys and refund
# tokens on the wallet routes.
"query_param_names": sorted(request.query_params.keys()),
},
)
@@ -128,6 +171,12 @@ class LoggingMiddleware(BaseHTTPMiddleware):
finally:
# Reset context
request_id_context.reset(token)
client_app_context.reset(client_app_token)
__all__ = ["LoggingMiddleware", "request_id_context"]
__all__ = [
"LoggingMiddleware",
"UNKNOWN_CLIENT_APP",
"client_app_context",
"request_id_context",
]
+99 -16
View File
@@ -1,8 +1,9 @@
"""Redaction helpers for sensitive provider identifiers.
"""Redaction helpers for sensitive provider identifiers and credentials.
Single source of truth for stripping account-scoped identifiers (e.g. OpenAI
organization IDs) from any text before it is logged, returned to a caller, or
written to an audit entry.
organization IDs) and spendable credentials (Cashu tokens, bearer keys, key
hashes) from any text before it is logged, returned to a caller, or written to
an audit entry.
"""
from __future__ import annotations
@@ -33,19 +34,101 @@ def redact_org_ids(text: str) -> str:
return _ORG_ID_PATTERN.sub(ORG_ID_PLACEHOLDER, text)
def redact_obj(obj: Any) -> Any:
"""Recursively redact organization IDs in arbitrary nested structures.
SECRET_PLACEHOLDER = "[REDACTED]"
Strings are redacted in place; dicts and lists/tuples are walked so that
identifiers nested inside structured payloads (e.g. log ``extra`` fields or
error ``details``) are also stripped. Other types are returned unchanged.
"""
# Field names whose value is spendable or authenticating on its own. Matched as
# substrings of the lowercased key, so ``hashed_key`` (a live ``sk-`` credential)
# is stripped while the truncated ``key_hash`` prefix used for correlation is
# not. Numeric values are never stripped, which keeps ``input_tokens`` and the
# other usage-analytics fields intact.
_SECRET_KEY_HINTS = (
"authorization",
"api_key",
"apikey",
"bearer",
"cashu",
"cookie",
"credential",
"hashed_key",
"mnemonic",
"nsec",
"passphrase",
"password",
"private_key",
"privkey",
"secret",
"token",
)
# Value shapes that are spendable wherever they appear, including inside URLs,
# query strings and free-form error text. Every alternative is anchored on a
# literal prefix and uses a single bounded character class, so matching stays
# linear on the logging hot path.
_SECRET_VALUE_PATTERNS: tuple[re.Pattern[str], ...] = (
re.compile(r"cashu[A-Z][A-Za-z0-9_\-=/+]{20,}"),
re.compile(r"\bnsec1[a-z0-9]{20,}"),
re.compile(r"\bBearer\s+[A-Za-z0-9_\-.=]{10,}", re.IGNORECASE),
re.compile(r"\bsk-[A-Za-z0-9]{16,}"),
re.compile(r"\b[0-9a-f]{64}\b"),
re.compile(
r"(?<=[?&])([^=&\s]*(?:token|key|secret|password|auth|sig)[^=&\s]*=)[^&\s\"']+",
re.IGNORECASE,
),
)
_MAX_REDACTION_DEPTH = 12
def _redact_secret_text(text: str) -> str:
for pattern in _SECRET_VALUE_PATTERNS:
text = pattern.sub(
lambda match: (match.group(1) if match.groups() else "")
+ SECRET_PLACEHOLDER,
text,
)
return redact_org_ids(text)
def _is_secret_key(key: object) -> bool:
if not isinstance(key, str):
return False
lowered = key.lower()
return any(hint in lowered for hint in _SECRET_KEY_HINTS)
def _redact_secrets(obj: Any, depth: int, seen: frozenset[int]) -> Any:
if isinstance(obj, str):
return redact_org_ids(obj)
return _redact_secret_text(obj)
if not isinstance(obj, (dict, list, tuple)):
return obj
if depth >= _MAX_REDACTION_DEPTH or id(obj) in seen:
return SECRET_PLACEHOLDER
nested = seen | {id(obj)}
if isinstance(obj, dict):
return {key: redact_obj(value) for key, value in obj.items()}
if isinstance(obj, list):
return [redact_obj(value) for value in obj]
if isinstance(obj, tuple):
return tuple(redact_obj(value) for value in obj)
return obj
return {
key: _redact_value(key, value, depth + 1, nested)
for key, value in obj.items()
}
redacted = [_redact_secrets(value, depth + 1, nested) for value in obj]
return redacted if isinstance(obj, list) else tuple(redacted)
def _redact_value(key: object, value: Any, depth: int, seen: frozenset[int]) -> Any:
# Containers keep being walked even under a secret-shaped key so that the
# surrounding structure stays readable for operators.
if isinstance(value, (bool, int, float, dict, list, tuple)) or value is None:
return _redact_secrets(value, depth, seen)
if _is_secret_key(key):
return SECRET_PLACEHOLDER
return _redact_secrets(value, depth, seen)
def redact_field(key: str, value: Any) -> Any:
"""Strip credentials from one named field, e.g. a log ``extra`` entry.
Both secret-shaped keys and secret-shaped values are stripped, so a leak
survives neither a renamed field nor a credential embedded in free text.
The walk is depth-limited and cycle-safe: a malformed payload degrades to
``[REDACTED]`` rather than taking the logging call down with it.
"""
return _redact_value(key, value, 0, frozenset())
+25 -4
View File
@@ -77,7 +77,6 @@ class Settings(BaseSettings):
exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE")
upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE")
tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE")
child_key_cost: int = Field(default=0, env="CHILD_KEY_COST")
# Minimum per-request charge in millisatoshis when model pricing is free/zero
min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT")
reset_reserved_balance_on_startup: bool = Field(
@@ -101,6 +100,12 @@ class Settings(BaseSettings):
# Network
cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS")
# Comma-separated METHOD:path pairs adding to the proxy's canonical
# endpoint allowlist, e.g. "POST:v1/rerank,GET:batches". Only for upstreams
# exposing an endpoint outside the OpenAI-compatible set; each addition
# widens what the provider credential can be spent against, so wildcards
# and prefixes are not supported.
proxy_extra_allowed_paths: str = Field(default="", env="PROXY_EXTRA_ALLOWED_PATHS")
tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL")
providers_refresh_interval_seconds: int = Field(
default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS"
@@ -119,7 +124,6 @@ class Settings(BaseSettings):
enable_model_paths_refresh: bool = Field(
default=True, env="ENABLE_MODEL_PATHS_REFRESH"
)
refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS")
# Uncollected refund tokens are swept after ~6 months (180 days).
# Fixed for now: not configurable via env or the settings DB/admin API
# (empty env list disables env binding; see FIXED_FIELDS).
@@ -127,6 +131,14 @@ class Settings(BaseSettings):
refund_sweep_claim_timeout_seconds: int = Field(
default=900, gt=0, env="REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS"
)
# How long an open refund claim may sit before the reconciler asks the mint
# what became of it. Doubles as the reconciler's per-row lease.
refund_claim_timeout_seconds: int = Field(
default=300, gt=0, env="REFUND_CLAIM_TIMEOUT_SECONDS"
)
refund_reconcile_interval_seconds: int = Field(
default=60, gt=0, env="REFUND_RECONCILE_INTERVAL_SECONDS"
)
# Database connection-pool controls (advanced). Capacity defaults provide
# headroom for Routstr's concurrent request and background-payment workload.
@@ -142,6 +154,9 @@ class Settings(BaseSettings):
database_pool_hold_warn_seconds: float = Field(
default=10.0, gt=0, env="DATABASE_POOL_HOLD_WARN_SECONDS"
)
database_busy_timeout: float = Field(
default=30.0, gt=0, env="DATABASE_BUSY_TIMEOUT"
)
# Logging
log_level: str = Field(default="INFO", env="LOG_LEVEL")
@@ -192,12 +207,18 @@ SECRET_FIELDS = frozenset({"admin_password", "nsec"})
# neither store nor shadow them; env is always authoritative.
ENV_ONLY_FIELDS = frozenset(
{
# Widening the proxy's reachable upstream surface is a deployment
# decision, not a runtime toggle: it changes what the provider
# credential can be spent against. Keeping it env-only also lets the
# proxy parse it once at import without going stale.
"proxy_extra_allowed_paths",
"database_pool_size",
"database_max_overflow",
"database_pool_timeout",
"database_pool_recycle",
"database_pool_pre_ping",
"database_pool_hold_warn_seconds",
"database_busy_timeout",
}
)
@@ -248,7 +269,7 @@ def derive_npub_from_nsec(nsec: str) -> str | None:
boot.
"""
try:
from nostr.key import PublicKey # type: ignore
from nostr_sdk import PublicKey
from ..nostr.listing import nsec_to_keypair
except ImportError:
@@ -260,7 +281,7 @@ def derive_npub_from_nsec(nsec: str) -> str | None:
_privkey_hex, pubkey_hex = keypair
try:
return PublicKey(bytes.fromhex(pubkey_hex)).bech32()
return PublicKey.parse(pubkey_hex).to_bech32()
except (ValueError, AttributeError):
return None
+1 -1
View File
@@ -20,7 +20,7 @@ import subprocess
from functools import lru_cache
from pathlib import Path
BASE_VERSION = "0.4.5"
BASE_VERSION = "0.4.7"
_REPO_ROOT = Path(__file__).resolve().parents[2]
_GIT_TIMEOUT_SECONDS = 2.0
+9 -12
View File
@@ -80,8 +80,6 @@ class _InvoiceSettlement:
purpose: str
api_key_hash: str | None
mint_url: str | None
balance_limit: int | None
balance_limit_reset: str | None
validity_date: int | None
@classmethod
@@ -93,8 +91,6 @@ class _InvoiceSettlement:
purpose=invoice.purpose,
api_key_hash=invoice.api_key_hash,
mint_url=invoice.mint_url,
balance_limit=invoice.balance_limit,
balance_limit_reset=invoice.balance_limit_reset,
validity_date=invoice.validity_date,
)
@@ -118,8 +114,6 @@ class InvoiceCreateRequest(BaseModel):
default=None,
description="Deprecated: legacy field for topup. Prefer Authorization header.",
)
balance_limit: int | None = Field(default=None)
balance_limit_reset: str | None = Field(default=None)
validity_date: int | None = Field(default=None)
@@ -220,11 +214,18 @@ async def _request_mint_with_fallback(
)
continue
try:
wallet = await get_wallet(mint_url, "sat", retry_on_rate_limit=False)
wallet = await get_wallet(
mint_url,
"sat",
retry_on_rate_limit=False,
load_proofs=False,
)
quote = await run_mint_operation(
lambda: wallet.request_mint(amount_sats),
op_name="request_mint_invoice",
mint_url=mint_url,
# Quote creation is unsafe to retry without idempotency.
retry_timeouts=False,
retry_on_rate_limit=False,
)
return quote.request, quote.quote, mint_url
@@ -400,8 +401,6 @@ async def create_invoice(
api_key_hash=api_key_token[3:] if api_key_token else None,
purpose=request.purpose,
mint_url=mint_url,
balance_limit=request.balance_limit,
balance_limit_reset=request.balance_limit_reset,
validity_date=request.validity_date,
expires_at=expires_at,
)
@@ -582,7 +581,7 @@ async def check_invoice_payment(
await session.commit()
mint_url = settlement.mint_url or settings.primary_mint
wallet = await get_wallet(mint_url, "sat")
wallet = await get_wallet(mint_url, "sat", load_proofs=False)
try:
mint_status = await run_mint_operation(
lambda: wallet.get_mint_quote(settlement.payment_hash),
@@ -797,8 +796,6 @@ async def _create_api_key_record(
balance=invoice.amount_sats * 1000,
refund_currency="sat",
refund_mint_url=mint_url,
balance_limit=invoice.balance_limit,
balance_limit_reset=invoice.balance_limit_reset,
validity_date=invoice.validity_date,
)
session.add(api_key)
+24 -1
View File
@@ -177,6 +177,8 @@ class MintRateGuard:
if isinstance(error, httpx.HTTPStatusError):
retry_after = parse_retry_after(error.response.headers)
self.apply_rate_limit_cooldown(retry_after)
elif is_mint_transport_error(error):
self.apply_cooldown(MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport")
else:
self.apply_cooldown(1.0)
logger.warning(
@@ -236,6 +238,17 @@ def mint_cooldown_reason(mint_url: str) -> str | None:
return MintRateGuard.get(mint_url).cooldown_reason()
def is_mint_transport_error(error: BaseException) -> bool:
current: BaseException | None = error
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
if isinstance(current, MINT_TRANSPORT_EXCEPTIONS):
return True
current = current.__cause__ or current.__context__
return False
def is_mint_rate_limited(error: BaseException) -> bool:
"""Return whether an exception chain represents HTTP 429/cooldown."""
@@ -269,6 +282,7 @@ async def run_mint_operation(
mint_url: str = "",
retry_timeouts: bool = True,
retry_on_rate_limit: bool = True,
allow_during_cooldown: bool = False,
) -> Any:
"""Run one mint operation with bounded concurrency and adaptive cooldown."""
@@ -282,7 +296,7 @@ async def run_mint_operation(
return await factory()
async def invoke() -> Any:
if guard is not None:
if guard is not None and not allow_during_cooldown:
return await guard.run(timed_factory)
return await timed_factory()
@@ -293,6 +307,7 @@ async def run_mint_operation(
raise
except (asyncio.TimeoutError, httpx.TimeoutException) as exc:
if retry_timeouts and attempt < max_attempts - 1:
# Cooldown opens only after retries; earlier would stretch each backoff to a full cooldown wait.
backoff = (2**attempt) + (time.monotonic() % 1.0)
logger.warning(
"Mint operation timed out, retrying",
@@ -305,11 +320,19 @@ async def run_mint_operation(
)
await asyncio.sleep(backoff)
continue
if guard is not None:
guard.apply_cooldown(
MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport"
)
raise httpx.TimeoutException(
f"{op_name} timed out (attempts: {attempt + 1})"
) from exc
except Exception as exc:
if not is_mint_rate_limited(exc):
if guard is not None and is_mint_transport_error(exc):
guard.apply_cooldown(
MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport"
)
raise
backoff = (2**attempt) + (time.monotonic() % 1.0)
+4 -21
View File
@@ -12,13 +12,11 @@ import json
import time
from typing import Any
from nostr.event import Event
from nostr.key import PrivateKey
from ..core import get_logger
from ..core.log_manager import log_manager
from ..core.settings import settings
from .listing import nsec_to_keypair, publish_to_relay
from .sdk import create_signed_event
logger = get_logger(__name__)
@@ -44,18 +42,6 @@ WINDOW_DEFINITIONS: tuple[tuple[str, int, int], ...] = (
)
def _event_to_dict(ev: Event) -> dict[str, Any]:
return {
"id": ev.id,
"pubkey": ev.public_key,
"created_at": ev.created_at,
"kind": int(ev.kind) if not isinstance(ev.kind, int) else ev.kind,
"tags": ev.tags,
"content": ev.content,
"sig": ev.signature,
}
def _resolve_provider_id(public_key_hex: str) -> str:
explicit_provider_id = (settings.provider_id or "").strip()
if explicit_provider_id:
@@ -293,21 +279,18 @@ def create_stats_snapshot_event(
*,
d_tag: str,
) -> dict[str, Any]:
private_key = PrivateKey(bytes.fromhex(private_key_hex))
tags = [
["d", d_tag],
["provider", provider_id],
["schema", ANALYTICS_SCHEMA],
]
event = Event(
public_key=private_key.public_key.hex(),
content=payload_json,
return create_signed_event(
private_key_hex,
kind=ANALYTICS_KIND,
content=payload_json,
tags=tags,
)
private_key.sign_event(event)
return _event_to_dict(event)
def _fingerprint_payload(payload: dict[str, Any]) -> str:
+3 -3
View File
@@ -320,18 +320,18 @@ async def fetch_provider_health(endpoint_url: str) -> dict[str, Any]:
is_onion = ".onion" in endpoint_url
# Set up client arguments conditionally
proxies = None
proxy: str | None = None
if is_onion:
try:
tor_proxy = settings.tor_proxy_url
except Exception:
tor_proxy = "socks5://127.0.0.1:9050"
proxies = {"http://": tor_proxy, "https://": tor_proxy} # type: ignore[assignment]
proxy = tor_proxy
async with httpx.AsyncClient(
timeout=httpx.Timeout(30.0),
follow_redirects=True,
proxies=proxies, # type: ignore[arg-type]
proxy=proxy,
) as client:
# Prefer provider's /v1/info for full details
info_url = f"{endpoint_url.rstrip('/')}/v1/info"
+191 -262
View File
@@ -8,18 +8,12 @@ import asyncio
import json
import os
import random
import ssl
import time
from typing import Any, cast
from nostr.event import Event
from nostr.filter import Filter, Filters
from nostr.key import PrivateKey
from nostr.message_type import ClientMessageType
from nostr.relay_manager import RelayManager
from ..core import get_logger
from ..core.settings import settings
from .sdk import create_signed_event, fetch_events, parse_keypair, send_event
logger = get_logger(__name__)
@@ -33,18 +27,6 @@ def get_app_version() -> str | None:
return None
def _event_to_dict(ev: Event) -> dict[str, Any]:
return {
"id": ev.id,
"pubkey": ev.public_key,
"created_at": ev.created_at,
"kind": int(ev.kind) if not isinstance(ev.kind, int) else ev.kind,
"tags": ev.tags,
"content": ev.content,
"sig": ev.signature,
}
def nsec_to_keypair(nsec: str) -> tuple[str, str] | None:
"""
Convert a Nostr private key (nsec) to a keypair (privkey_hex, pubkey_hex).
@@ -56,16 +38,10 @@ def nsec_to_keypair(nsec: str) -> tuple[str, str] | None:
Tuple of (private_key_hex, public_key_hex) or None if invalid
"""
try:
if nsec.startswith("nsec"):
pk = PrivateKey.from_nsec(nsec)
return (pk.hex(), pk.public_key.hex())
if len(nsec) == 64:
pk = PrivateKey(bytes.fromhex(nsec))
return (pk.hex(), pk.public_key.hex())
logger.error(f"Invalid private key format/length: {len(nsec)}")
return None
if not (nsec.startswith("nsec") or len(nsec) == 64):
logger.error(f"Invalid private key format/length: {len(nsec)}")
return None
return parse_keypair(nsec)
except Exception as e:
logger.error(f"Failed to convert nsec to keypair: {e}")
return None
@@ -93,8 +69,6 @@ def create_listing_event(
Returns:
Complete signed nostr event as a dict ready for publishing
"""
pk = PrivateKey(bytes.fromhex(private_key_hex))
tags = [["d", provider_id]]
for url in endpoint_urls:
tags.append(["u", url])
@@ -107,9 +81,12 @@ def create_listing_event(
content = json.dumps(metadata, separators=(",", ":")) if metadata else ""
ev = Event(pk.public_key.hex(), content, kind=38421, tags=tags)
pk.sign_event(ev)
return _event_to_dict(ev)
return create_signed_event(
private_key_hex,
kind=38421,
content=content,
tags=tags,
)
def _get_tag_values(event: dict[str, Any], key: str) -> list[str]:
@@ -177,75 +154,25 @@ async def query_listing_events(
succeeded without transport-level errors.
"""
def _sync_query() -> tuple[list[dict[str, Any]], bool]:
rm = RelayManager()
rm.add_relay(relay_url)
events_out: list[dict[str, Any]] = []
ok = True
try:
rm.open_connections({"cert_reqs": ssl.CERT_NONE})
time.sleep(1.0)
try:
events_out = await fetch_events(
relay_url,
kind=38421,
author=pubkey,
limit=10,
timeout=timeout,
)
except Exception as e:
logger.debug(f"Failed to query relay {relay_url}: {type(e).__name__}")
return [], False
flt = Filter(kinds=[38421], authors=[pubkey], limit=10)
filters = Filters([flt])
sub_id = f"routstr_listing_{int(time.time())}"
rm.add_subscription(sub_id, filters)
req: list[Any] = [ClientMessageType.REQUEST, sub_id]
req.extend(filters.to_json_array())
rm.publish_message(json.dumps(req))
start = time.time()
last_event_ts = start
while time.time() - start < timeout:
drained = False
while rm.message_pool.has_events():
drained = True
ev_msg = rm.message_pool.get_event()
ev = ev_msg.event
ev_dict = _event_to_dict(ev)
if provider_id is not None:
tags = ev_dict.get("tags", [])
if not any(
isinstance(t, list)
and len(t) >= 2
and t[0] == "d"
and t[1] == provider_id
for t in tags
):
continue
events_out.append(ev_dict)
logger.debug(
f"Found listing event: {ev_dict.get('id', '')[:6]}...{ev_dict.get('id', '')[-6:]}"
)
if drained:
last_event_ts = time.time()
while rm.message_pool.has_notices():
notice = rm.message_pool.get_notice()
try:
content = getattr(notice, "content", notice)
s = str(content)
if len(s) > 200:
s = s[:200] + "..."
logger.debug(f"Relay notice: {s}")
except Exception:
pass
if time.time() - last_event_ts > 2.5:
break
time.sleep(0.1)
except Exception as e:
ok = False
logger.debug(f"Failed to query relay {relay_url}: {type(e).__name__}")
finally:
try:
rm.close_connections()
except Exception:
pass
return events_out, ok
return await asyncio.to_thread(_sync_query)
if provider_id is not None:
events_out = [
event
for event in events_out
if _get_single_tag_value(event, "d") == provider_id
]
return events_out, True
def discover_onion_url_from_tor(base_dir: str = "/var/lib/tor") -> str | None:
@@ -333,117 +260,102 @@ async def publish_to_relay(
Publish a listing event to a nostr relay via nostr library.
"""
def _sync_publish() -> bool:
rm = RelayManager()
rm.add_relay(relay_url)
try:
rm.open_connections({"cert_reqs": ssl.CERT_NONE})
time.sleep(1.0)
# Publish the event as-is via publish_message to preserve signature
rm.publish_message(json.dumps(["EVENT", event]))
logger.debug(f"Sent listing event {event.get('id', '')} to {relay_url}")
time.sleep(1.0)
return True
except Exception as e:
logger.debug(f"Failed to publish to {relay_url}: {type(e).__name__}")
return False
finally:
try:
rm.close_connections()
except Exception:
pass
return await asyncio.to_thread(_sync_publish)
async def announce_provider() -> None:
"""
Background task to announce this Routstr provider to Nostr relays.
Checks for existing announcements and creates new ones if needed.
"""
# Check for NSEC in environment (use NSEC only)
nsec = settings.nsec
if not nsec:
logger.info("Nostr private key not found (NSEC), skipping listing announcement")
return
# Convert NSEC to keypair
keypair = nsec_to_keypair(nsec)
if not keypair:
logger.error("Failed to parse NSEC, skipping listing announcement")
return
private_key_hex, public_key_hex = keypair
logger.info(f"Using Nostr pubkey: {public_key_hex}")
# Resolve settings and determine if we can publish BEFORE touching relays
try:
base_url: str | None = settings.http_url
onion_url: str | None = settings.onion_url
provider_name = settings.name or "Routstr Proxy"
provider_about = settings.description or "Privacy-preserving AI proxy via Nostr"
cashu_mints = [m.strip() for m in settings.cashu_mints if m.strip()]
except Exception:
base_url = settings.http_url or None
onion_url = settings.onion_url or None
provider_name = settings.name or "Routstr Proxy"
provider_about = settings.description or "Privacy-preserving AI proxy via Nostr"
cashu_mints = [m.strip() for m in settings.cashu_mints if m.strip()]
await send_event(relay_url, event, timeout=timeout)
logger.debug(f"Sent listing event {event.get('id', '')} to {relay_url}")
return True
except Exception as e:
logger.debug(f"Failed to publish to {relay_url}: {type(e).__name__}")
return False
# Re-announce cadence once a provider is listed.
ANNOUNCEMENT_INTERVAL_SECONDS = 24 * 60 * 60
# Poll cadence while there is nothing to announce (no NSEC, no endpoint, ...).
DISABLED_POLL_SECONDS = 60
# How often the long re-announce sleep re-checks the configured NSEC, so a
# newly saved identity is announced promptly instead of up to 24h later.
IDENTITY_POLL_SECONDS = 30
DEFAULT_RELAY_URLS = [
"wss://relay.nostr.band",
"wss://relay.damus.io",
"wss://relay.routstr.com",
"wss://nos.lol",
]
def _resolve_endpoint_urls() -> list[str]:
"""Endpoints to advertise: a public HTTP URL and/or an onion URL."""
endpoint_urls: list[str] = []
base_url = (settings.http_url or "").strip()
if base_url and base_url != "http://localhost:8000":
endpoint_urls.append(base_url)
onion_url = (settings.onion_url or "").strip()
if not onion_url:
discovered = discover_onion_url_from_tor()
if discovered:
onion_url = discovered
logger.info(f"Discovered onion URL via Tor volume: {onion_url}")
mint_urls = cashu_mints if cashu_mints else None
endpoint_urls: list[str] = []
if base_url and base_url.strip() and base_url.strip() != "http://localhost:8000":
endpoint_urls.append(base_url.strip())
if onion_url and onion_url.strip():
ou = onion_url.strip()
if ou.endswith(".onion") and not (
ou.startswith("http://") or ou.startswith("https://")
if onion_url:
if onion_url.endswith(".onion") and not (
onion_url.startswith("http://") or onion_url.startswith("https://")
):
ou = f"http://{ou}"
endpoint_urls.append(ou)
onion_url = f"http://{onion_url}"
endpoint_urls.append(onion_url)
if not endpoint_urls:
logger.warning(
"No valid endpoints configured (HTTP_URL/ONION_URL). Skipping listing publish."
)
return
return endpoint_urls
# Only now configure relays and determine provider_id (may query relays)
def _resolve_relay_urls() -> list[str]:
relay_urls = [u.strip() for u in getattr(settings, "relays", []) if u.strip()]
if not relay_urls:
relay_urls = [
"wss://relay.nostr.band",
"wss://relay.damus.io",
"wss://relay.routstr.com",
"wss://nos.lol",
]
return relay_urls or list(DEFAULT_RELAY_URLS)
provider_id = await _determine_provider_id(public_key_hex, relay_urls)
logger.info(f"Using provider_id: {provider_id}")
# Build metadata
metadata = {
"name": provider_name,
"about": provider_about,
}
def _resolve_mint_urls() -> list[str] | None:
mints = [m.strip() for m in (settings.cashu_mints or []) if m.strip()]
return mints or None
# Create the candidate event that we would publish
version_str = get_app_version()
candidate_event = create_listing_event(
private_key_hex=private_key_hex,
provider_id=provider_id,
endpoint_urls=endpoint_urls,
mint_urls=mint_urls,
version=version_str,
metadata=metadata,
)
# Backoff configuration and state (sensible defaults)
async def _sleep_until_next_announcement(
seconds: float, parsed_nsec: str | None
) -> None:
"""Sleep up to ``seconds``, returning early if the configured NSEC changes.
Without the early wake, a node runner who replaces the NSEC in the admin UI
would wait out the whole re-announce interval before the new identity (and,
with it, the new ``d`` tag and npub) is announced.
"""
remaining = float(seconds)
while remaining > 0:
if (settings.nsec or "").strip() != (parsed_nsec or ""):
return
chunk = min(float(IDENTITY_POLL_SECONDS), remaining)
await asyncio.sleep(chunk)
remaining -= chunk
async def announce_provider() -> None:
"""Background task announcing this Routstr provider to Nostr relays.
Started unconditionally at boot: while the node has no NSEC the task idles
and re-checks, so an identity configured later through the admin UI is
picked up (and announced) without a restart. The identity, endpoints, mints
and relays are all re-read every iteration, mirroring
``publish_usage_analytics``.
"""
parsed_nsec: str | None = None
private_key_hex: str | None = None
public_key_hex: str | None = None
provider_id: str | None = None
warned_missing_nsec = False
# Backoff state is deliberately long-lived: it has to survive an idle poll,
# a full re-announce cycle and an identity change, otherwise a failing relay
# would be retried at full rate on every pass.
backoff_base = 5.0
backoff_max = 900.0
backoff_jitter_ratio = 0.2
@@ -468,67 +380,74 @@ async def announce_provider() -> None:
f"Backoff: {relay} delay={delay:.1f}s jitter={jitter:.1f}s next={int(scheduled)}"
)
# Fetch existing events for this provider_id
existing_events: list[dict[str, Any]] = []
for relay_url in relay_urls:
if _should_skip(relay_url):
logger.debug(f"Skipping {relay_url} due to backoff")
continue
events, ok = await query_listing_events(relay_url, public_key_hex, provider_id)
if ok:
_register_success(relay_url)
existing_events.extend(events)
else:
_register_failure(relay_url)
# Decide whether to publish: publish if none exist or any differ from candidate
found_any = len(existing_events) > 0
all_match = found_any and all(
events_semantically_equal(ev, candidate_event) for ev in existing_events
)
if not all_match:
logger.debug(
"No matching listing announcement found or differences detected; publishing update"
)
success_count = 0
for relay_url in relay_urls:
if _should_skip(relay_url):
logger.debug(f"Skipping publish to {relay_url} due to backoff")
continue
if await publish_to_relay(relay_url, candidate_event):
_register_success(relay_url)
success_count += 1
else:
_register_failure(relay_url)
logger.info(
f"Published listing announcement to {success_count}/{len(relay_urls)} relays"
)
else:
logger.debug(
"Matching listing announcement already present; skipping publish on startup"
)
# Re-announce periodically (every 24 hours)
announcement_interval = 24 * 60 * 60
while True:
try:
await asyncio.sleep(announcement_interval)
nsec = (settings.nsec or "").strip()
if not nsec:
if not warned_missing_nsec:
logger.info(
"Nostr private key not configured (NSEC); waiting for one "
"to be set before announcing this provider"
)
warned_missing_nsec = True
parsed_nsec = None
await asyncio.sleep(DISABLED_POLL_SECONDS)
continue
# Re-derive the identity whenever the configured NSEC changes, so a
# key saved (or replaced) through the admin UI takes effect live.
if nsec != parsed_nsec:
keypair = nsec_to_keypair(nsec)
if not keypair:
logger.error(
"Invalid NSEC; waiting for a valid one before announcing"
)
parsed_nsec = None
await asyncio.sleep(DISABLED_POLL_SECONDS)
continue
private_key_hex, public_key_hex = keypair
parsed_nsec = nsec
provider_id = None
warned_missing_nsec = False
logger.info(f"Using Nostr pubkey: {public_key_hex}")
if private_key_hex is None or public_key_hex is None:
await asyncio.sleep(DISABLED_POLL_SECONDS)
continue
endpoint_urls = _resolve_endpoint_urls()
if not endpoint_urls:
logger.warning(
"No valid endpoints configured (HTTP_URL/ONION_URL). "
"Skipping listing publish until one is set."
)
await asyncio.sleep(DISABLED_POLL_SECONDS)
continue
relay_urls = _resolve_relay_urls()
if provider_id is None:
provider_id = await _determine_provider_id(public_key_hex, relay_urls)
logger.info(f"Using provider_id: {provider_id}")
metadata = {
"name": settings.name or "Routstr Proxy",
"about": settings.description
or "Privacy-preserving AI proxy via Nostr",
}
# Build fresh candidate event for comparison
version_str = get_app_version()
candidate_event = create_listing_event(
private_key_hex=private_key_hex,
provider_id=provider_id,
endpoint_urls=endpoint_urls,
mint_urls=mint_urls,
version=version_str,
mint_urls=_resolve_mint_urls(),
version=get_app_version(),
metadata=metadata,
)
# Fetch existing events for this provider_id
existing_events = []
existing_events: list[dict[str, Any]] = []
for relay_url in relay_urls:
if _should_skip(relay_url):
logger.debug(f"Skipping {relay_url} due to backoff")
@@ -549,26 +468,36 @@ async def announce_provider() -> None:
if all_match:
logger.debug(
"Matching listing announcement already present; skipping periodic re-announce"
"Matching listing announcement already present; skipping publish"
)
else:
logger.debug(
"No matching listing announcement found or differences "
"detected; publishing update"
)
success_count = 0
for relay_url in relay_urls:
if _should_skip(relay_url):
logger.debug(f"Skipping publish to {relay_url} due to backoff")
continue
if await publish_to_relay(relay_url, candidate_event):
_register_success(relay_url)
success_count += 1
else:
_register_failure(relay_url)
logger.info(
"Published listing announcement to "
f"{success_count}/{len(relay_urls)} relays"
)
continue
logger.debug(
f"Re-announcing provider due to differences or absence: {candidate_event['id']}"
# Re-announce periodically; wakes early if the NSEC changes.
await _sleep_until_next_announcement(
ANNOUNCEMENT_INTERVAL_SECONDS, parsed_nsec
)
for relay_url in relay_urls:
if _should_skip(relay_url):
logger.debug(f"Skipping publish to {relay_url} due to backoff")
continue
ok = await publish_to_relay(relay_url, candidate_event)
if ok:
_register_success(relay_url)
else:
_register_failure(relay_url)
except asyncio.CancelledError:
logger.info("Listing announcement task cancelled")
break
except Exception as e:
logger.debug(f"Error in listing announcement loop: {type(e).__name__}")
# Continue running despite errors
await asyncio.sleep(DISABLED_POLL_SECONDS)
+94
View File
@@ -0,0 +1,94 @@
"""Small adapter around the maintained ``nostr-sdk`` package."""
from __future__ import annotations
import json
from datetime import timedelta
from typing import Any, cast
from nostr_sdk import (
AckPolicy,
Client,
Event,
EventBuilder,
Filter,
Keys,
Kind,
PublicKey,
RelayUrl,
ReqExitPolicy,
ReqTarget,
SendEventTarget,
Tag,
)
def parse_keypair(secret_key: str) -> tuple[str, str]:
keys = Keys.parse(secret_key)
return keys.secret_key().to_hex(), keys.public_key().to_hex()
def create_signed_event(
secret_key_hex: str,
*,
kind: int,
content: str,
tags: list[list[str]],
) -> dict[str, Any]:
keys = Keys.parse(secret_key_hex)
event = (
EventBuilder(Kind(kind), content)
.tags([Tag.parse(tag) for tag in tags])
.finalize(keys)
)
return cast(dict[str, Any], json.loads(event.as_json()))
async def fetch_events(
relay_url: str,
*,
kind: int,
author: str,
limit: int,
timeout: int,
) -> list[dict[str, Any]]:
relay = RelayUrl.parse(relay_url)
client = Client()
await client.add_relay(relay)
try:
await client.connect()
event_filter = (
Filter().kinds([Kind(kind)]).authors([PublicKey.parse(author)]).limit(limit)
)
events = await client.fetch_events(
ReqTarget.single(relay, [event_filter]),
timeout=timedelta(seconds=timeout),
policy=ReqExitPolicy.WAIT_DURATION_AFTER_EOSE(timedelta(seconds=2.5)),
max_events=limit,
)
return [cast(dict[str, Any], json.loads(event.as_json())) for event in events]
finally:
await client.shutdown()
async def send_event(relay_url: str, event: dict[str, Any], *, timeout: int) -> None:
relay = RelayUrl.parse(relay_url)
client = Client()
await client.add_relay(relay)
try:
await client.connect()
output = await client.send_event(
Event.from_json(json.dumps(event)),
target=SendEventTarget.to([relay]),
ack_policy=AckPolicy.all(),
ok_timeout=timedelta(seconds=timeout),
)
if output.failed or relay not in output.success:
reasons = ", ".join(
f"{url}: {reason}" for url, reason in output.failed.items()
)
raise RuntimeError(
f"Relay did not accept event: {reasons or 'no OK from relay'}"
)
finally:
await client.shutdown()
+58 -45
View File
@@ -1,11 +1,12 @@
import math
from typing import TYPE_CHECKING
from pydantic.v1 import BaseModel
from pydantic.v1 import BaseModel, Field
from ..core import get_logger
from ..core.settings import settings
from .price import sats_usd_price
from .rates import coerce_rate, is_usable_rate
from .usage import normalize_usage, parse_token_count
if TYPE_CHECKING:
@@ -34,6 +35,9 @@ class CostData(BaseModel):
cache_creation_input_tokens: int = 0
cache_read_msats: int = 0
cache_creation_msats: int = 0
# Actual debit after finalization; None means settlement has not run yet.
charged_msats: int | None = None
upstream_usd: float = Field(default=0.0, exclude=True)
class MaxCostData(CostData):
@@ -107,11 +111,11 @@ async def calculate_cost(
if usage is None:
logger.warning(
"No usage data in response — billing at MaxCostData with zero "
"tokens. Dashboard will show this request as `(0+0)`. Most "
"common cause: upstream stream did not include a final usage "
"chunk (OpenAI-compat backends require "
"`stream_options.include_usage=true`).",
"No usage data or local estimate in response — releasing the "
"reservation without charging it as usage. Dashboard will show "
"this request as `(0+0)` tokens. Most common cause: upstream "
"stream did not include a final usage chunk (OpenAI-compat "
"backends require `stream_options.include_usage=true`).",
extra={
"max_cost_msats": max_cost,
"model": response_data.get("model", "unknown"),
@@ -175,13 +179,14 @@ async def calculate_cost(
cost_details = usage_data.get("cost_details", {})
if not isinstance(cost_details, dict):
cost_details = {}
input_usd = _coerce_usd(
cost_details.get("input_cost")
or cost_details.get("upstream_inference_prompt_cost")
# Coerce each spelling before choosing between them: `inf` and `NaN`
# are truthy, so a malformed first field would otherwise win the
# fallback and the usable figure beside it would never be read.
input_usd = _coerce_usd(cost_details.get("input_cost")) or _coerce_usd(
cost_details.get("upstream_inference_prompt_cost")
)
output_usd = _coerce_usd(
cost_details.get("output_cost")
or cost_details.get("upstream_inference_completions_cost")
output_usd = _coerce_usd(cost_details.get("output_cost")) or _coerce_usd(
cost_details.get("upstream_inference_completions_cost")
)
cache_pricing_rates: tuple[float, float, float, float] | None = None
if cache_read_tokens > 0 or cache_creation_tokens > 0:
@@ -241,27 +246,32 @@ async def calculate_cost(
else:
input_rate, output_rate, cache_read_rate, cache_creation_rate = pricing_rates
if not (input_rate and output_rate):
# Truthiness is not the question: `NaN` and a negative rate are both truthy
# and sailed past this gate into the token math, while a rate of zero is a
# price — free — and reading it as a missing one charged the whole
# reservation for a request the model serves for nothing.
rates = (input_rate, output_rate, cache_read_rate, cache_creation_rate)
if not all(is_usable_rate(rate) for rate in rates):
logger.warning(
"No token pricing configured — billing at flat MaxCostData. "
"Token counts %s in the upstream response but cannot be "
"priced; the request will appear in dashboards with the "
"raw counts and a fixed max-cost charge.",
"are present"
if (input_tokens > 0 or output_tokens > 0)
else "are zero",
"No usable token pricing — releasing the reservation instead of "
"treating its ceiling as the charge. Token counts %s in the "
"upstream response but cannot be converted to money; the request "
"will appear in dashboards with raw counts and a zero charge.",
"are present" if (input_tokens > 0 or output_tokens > 0) else "are zero",
extra={
"base_cost_msats": max_cost,
"model": response_data.get("model", "unknown"),
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"input_rate": input_rate,
"output_rate": output_rate,
},
)
return MaxCostData(
base_msats=max_cost,
base_msats=0,
input_msats=0,
output_msats=0,
total_msats=max_cost,
total_msats=0,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_read_input_tokens=cache_read_tokens,
@@ -289,15 +299,25 @@ async def calculate_cost(
def _coerce_usd(value: object) -> float:
"""Coerce a value to USD float, handling various formats safely."""
if value is None or isinstance(value, bool):
return 0.0
if not isinstance(value, (int, float, str)):
return 0.0
try:
return max(0.0, float(value))
except (TypeError, ValueError):
return 0.0
"""Coerce an upstream-reported USD figure to a usable amount, else ``0.0``.
These values come straight off the upstream response, where ``json.loads``
accepts the bare ``NaN``/``Infinity`` literals and overflows ``1e999`` to
``inf``. A non-finite figure is not a cost, and letting one through poisoned
the proportional split in ``_calculate_from_usd_cost`` (``inf / inf`` is
``NaN``): the resulting exception was absorbed by the broad handler around
the USD path, so a request whose *total* cost was perfectly valid fell
through to token-estimated pricing and was billed a fraction of what the
upstream charged.
``0.0`` means "no usable figure" to every caller, which is the same thing an
absent field means, so the caller's existing ``> 0`` checks handle it.
"""
# A cost figure is coerced exactly like a rate; only the way an unusable one
# is reported differs. A negative is rejected here, where the previous
# `max(0.0, …)` clamped it.
amount = coerce_rate(value)
return amount if amount is not None else 0.0
def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
@@ -326,9 +346,7 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
# actually deducts from the balance. For non-BYOK providers (e.g.
# OpenRouter) usage.cost already equals upstream_inference_cost, so we
# fall through to the normal ``cost`` lookup below.
upstream_cost = _coerce_usd(
cost_details.get("upstream_inference_cost")
)
upstream_cost = _coerce_usd(cost_details.get("upstream_inference_cost"))
if upstream_cost > 0 and usage_data.get("is_byok"):
byok_fee = _coerce_usd(usage_data.get("cost"))
return upstream_cost + byok_fee
@@ -359,8 +377,7 @@ def _get_pricing_rates(
``None`` means configured fixed pricing should be used by the caller.
"""
if settings.fixed_pricing and (
settings.fixed_per_1k_input_tokens
or settings.fixed_per_1k_output_tokens
settings.fixed_per_1k_input_tokens or settings.fixed_per_1k_output_tokens
):
return None
@@ -416,12 +433,8 @@ def _get_pricing_rates(
usd_per_sat = sats_usd_price()
mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat
mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat
cache_read_usd = _coerce_usd(
pricing.get("cache_read_input_token_cost")
)
cache_write_usd = _coerce_usd(
pricing.get("cache_creation_input_token_cost")
)
cache_read_usd = _coerce_usd(pricing.get("cache_read_input_token_cost"))
cache_write_usd = _coerce_usd(pricing.get("cache_creation_input_token_cost"))
mscr_1k = (
cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat
if cache_read_usd > 0
@@ -479,6 +492,7 @@ def _calculate_from_usd_cost(
"""Calculate cost from USD figures, deriving input/output split from tokens."""
if provider_fee is None:
provider_fee = _resolve_provider_fee(response_data.get("model", ""))
reported_usd = usd_cost
usd_cost = usd_cost * provider_fee
input_usd = input_usd * provider_fee
output_usd = output_usd * provider_fee
@@ -525,9 +539,7 @@ def _calculate_from_usd_cost(
regular_weight = input_tokens * input_rate
cache_read_weight = cache_read_tokens * cache_read_rate
cache_creation_weight = cache_creation_tokens * cache_creation_rate
total_input_weight = (
regular_weight + cache_read_weight + cache_creation_weight
)
total_input_weight = regular_weight + cache_read_weight + cache_creation_weight
if total_input_weight > 0:
cache_read_msats = int(
round(
@@ -566,6 +578,7 @@ def _calculate_from_usd_cost(
cache_creation_input_tokens=cache_creation_tokens,
cache_read_msats=cache_read_msats,
cache_creation_msats=cache_creation_msats,
upstream_usd=reported_usd,
)
+303 -70
View File
@@ -1,8 +1,12 @@
import asyncio
import base64
import ipaddress
import json
import math
import socket
from io import BytesIO
from typing import Any
from urllib.parse import urlsplit, urlunsplit
import httpx
from fastapi import HTTPException, Response
@@ -14,10 +18,22 @@ from ..core import get_logger
from ..core.exceptions import UpstreamError
from ..core.redaction import redact_org_ids
from ..core.settings import settings
from ..wallet import deserialize_token_from_string
from ..wallet import (
UntrustedSourceMintError,
classify_redemption_error,
deserialize_token_from_string,
is_trusted_source_mint,
)
from .responses_input import (
FILE_ID_URL_PREFIX,
count_input_images,
input_image_part_to_image_url,
responses_input_to_messages,
)
logger = get_logger(__name__)
def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> None:
if x_cashu := headers.get("x-cashu", None):
cashu_token = x_cashu
@@ -68,6 +84,19 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N
detail="Invalid authentication token format",
)
if not is_trusted_source_mint(token_obj.mint):
classified = classify_redemption_error(
UntrustedSourceMintError(f"Untrusted source mint: {token_obj.mint}")
)
assert classified is not None
error_type, status_code, message, error_code = classified
raise HTTPException(
status_code=status_code,
detail={
"error": {"message": message, "type": error_type, "code": error_code}
},
)
amount_msat = (
token_obj.amount if token_obj.unit == "msat" else token_obj.amount * 1000
)
@@ -156,7 +185,12 @@ async def calculate_discounted_max_cost(
body: dict,
model_obj: Any | None = None,
) -> int:
"""Calculate the discounted max cost for a request using model pricing when available."""
"""Calculate the discounted max cost for a request using model pricing when available.
Completion discounts are trimmed from the largest declared cap among
``max_tokens`` and ``max_completion_tokens`` (chat/completions) or
``max_output_tokens`` (responses).
"""
if settings.fixed_pricing:
return max_cost_for_model
@@ -196,41 +230,75 @@ async def calculate_discounted_max_cost(
adjusted = max_cost_for_model
if messages := body.get("messages"):
prompt_tokens = estimate_tokens(messages)
messages = body.get("messages")
# Estimated over the whole body: a discount driven by message text alone lets
# a caller hide prompt weight elsewhere, shrink the reservation, and be billed
# for work the reservation never covered.
prompt_tokens = estimate_prompt_tokens(body)
image_tokens = await estimate_image_tokens_in_messages(messages)
if image_tokens > 0:
logger.debug(
"Found images in request",
extra={
"model": model,
"image_tokens": image_tokens,
},
)
prompt_tokens += image_tokens
# Images are billed as tokens by the upstream but carry no text for
# ``estimate_prompt_tokens`` to count, so they are estimated separately and
# added on both the chat (``messages``) and Responses (``input``) paths.
image_tokens = 0
if isinstance(messages, list):
image_tokens += await estimate_image_tokens_in_messages(messages)
input_data = body.get("input")
if isinstance(input_data, list):
converted = responses_input_to_messages(input_data)
if converted is None:
image_tokens += count_input_images(input_data) * _MAX_ORIGINAL_IMAGE_TOKENS
else:
image_tokens += await estimate_image_tokens_in_messages(converted)
if image_tokens > 0:
logger.debug(
"Found images in request",
extra={
"model": model,
"image_tokens": image_tokens,
},
)
prompt_tokens += image_tokens
if prompt_tokens > 0:
estimated_prompt_delta_sats = (
max_prompt_allowed_sats - prompt_tokens * model_pricing.prompt
)
if estimated_prompt_delta_sats > 0:
adjusted = adjusted - math.floor(estimated_prompt_delta_sats * 1000)
max_tokens_raw = body.get("max_tokens", None)
if max_tokens_raw is not None:
# Completion caps arrive under several names: ``max_tokens`` (legacy
# chat), ``max_completion_tokens`` (modern chat) and ``max_output_tokens``
# (Responses API). When a request declares more than one, reserve against
# the largest: upstream precedence between the fields varies by provider,
# so the smaller cap may not be honored and the reservation must never
# under-cover what the upstream could bill.
max_tokens_int: int | None = None
for cap_field in ("max_tokens", "max_completion_tokens", "max_output_tokens"):
cap_raw = body.get(cap_field)
if cap_raw is None:
continue
try:
max_tokens_int = int(max_tokens_raw)
cap_int = int(cap_raw)
except (TypeError, ValueError):
logger.warning(
"Invalid max_tokens; ignoring in cost adjustment",
extra={"max_tokens": str(max_tokens_raw)[:64], "model": model},
"Invalid completion token cap; ignoring in cost adjustment",
extra={
"field": cap_field,
"value": str(cap_raw)[:64],
"model": model,
},
)
else:
estimated_completion_delta_sats = (
max_completion_allowed_sats - max_tokens_int * model_pricing.completion
)
if estimated_completion_delta_sats > 0:
adjusted = adjusted - math.floor(estimated_completion_delta_sats * 1000)
continue
max_tokens_int = (
cap_int if max_tokens_int is None else max(max_tokens_int, cap_int)
)
if max_tokens_int is not None:
estimated_completion_delta_sats = (
max_completion_allowed_sats - max_tokens_int * model_pricing.completion
)
if estimated_completion_delta_sats > 0:
adjusted = adjusted - math.floor(estimated_completion_delta_sats * 1000)
logger.debug(
"Discounted max cost computed",
@@ -262,6 +330,51 @@ def estimate_tokens(messages: list) -> int:
return total // 3
def _sum_string_chars(node: Any) -> int:
"""Recursively sum the length of every string in the tree, keys included.
Nothing is excluded. Keys count because JSON-schema property names are
forwarded to the provider, and no exclusion rule can be trusted here: every
part of the body is caller-controlled, so any carve-out (by key name or by
value shape) is a place to hide prompt weight for free. Inline image data is
therefore counted as text too, which only makes the discount smaller.
"""
if isinstance(node, str):
return len(node)
if isinstance(node, dict):
return sum(
len(str(key)) + _sum_string_chars(value) for key, value in node.items()
)
if isinstance(node, list):
return sum(_sum_string_chars(item) for item in node)
return 0
def _count_prompt_token_ids(node: Any) -> int:
if isinstance(node, int) and not isinstance(node, bool):
return 1
if isinstance(node, list):
return sum(_count_prompt_token_ids(item) for item in node)
return 0
def estimate_prompt_tokens(body: dict) -> int:
"""Conservatively estimate prompt tokens for the whole provider-bound body.
Every string counts, as do token IDs in legacy ``prompt`` arrays, so no
forwarded field can hide prompt weight and shrink its reservation.
"""
return _sum_string_chars(body) // 3 + _count_prompt_token_ids(body.get("prompt"))
IMAGE_FETCH_TIMEOUT_SECONDS = 10.0
# Dimensions live in the header, so a prefix suffices and an endless body cannot
# pin memory.
IMAGE_FETCH_MAX_BYTES = 512 * 1024
# Fetches are sequential, so an unbounded URL list is a request-time amplifier.
IMAGE_FETCH_MAX_PER_REQUEST = 8
def _get_image_dimensions(image_data: bytes) -> tuple[int, int]:
"""Extract image dimensions from image bytes."""
try:
@@ -275,13 +388,81 @@ def _get_image_dimensions(image_data: bytes) -> tuple[int, int]:
return (512, 512)
async def _fetch_image_from_url(url: str) -> bytes | None:
"""Fetch image from URL."""
def _is_blocked_address(address: str) -> bool:
"""Allow only globally reachable addresses (RFC 6890)."""
try:
async with httpx.AsyncClient(timeout=10.0) as client:
response = await client.get(url)
response.raise_for_status()
return response.content
ip = ipaddress.ip_address(address)
except ValueError:
return True
if isinstance(ip, ipaddress.IPv6Address):
# An embedded v4 address would otherwise smuggle a rejected target past
# the v6 checks.
for embedded in (ip.ipv4_mapped, ip.sixtofour):
if embedded is not None:
return _is_blocked_address(str(embedded))
return not ip.is_global or ip.is_multicast
async def _validated_fetch_target(url: str) -> tuple[str, str]:
"""Return the URL to request and its ``Host`` header.
Cost estimation runs on the unauthenticated request body, so a caller can
otherwise aim the node at internal hosts. HTTP is rewritten to the resolved
address so the name cannot rebind between check and connect; HTTPS keeps its
hostname because certificate validation already binds the connection.
"""
parts = urlsplit(url)
if parts.scheme not in ("http", "https"):
raise ValueError(f"unsupported scheme: {parts.scheme or 'none'}")
host = parts.hostname
if not host:
raise ValueError("missing host")
default_port = 443 if parts.scheme == "https" else 80
port = parts.port or default_port
host_header = f"[{host}]" if ":" in host else host
if parts.port is not None:
host_header = f"{host_header}:{parts.port}"
infos = await asyncio.get_running_loop().getaddrinfo(
host, port, proto=socket.IPPROTO_TCP
)
if not infos:
raise ValueError("host did not resolve")
for info in infos:
if _is_blocked_address(str(info[4][0])):
raise ValueError("host resolves to a blocked address")
if parts.scheme == "https":
return url, host_header
family, _, _, _, sockaddr = infos[0]
address = str(sockaddr[0])
pinned = f"[{address}]" if family == socket.AF_INET6 else address
if parts.port is not None:
pinned = f"{pinned}:{parts.port}"
return urlunsplit((parts.scheme, pinned, parts.path, parts.query, "")), host_header
async def _fetch_image_from_url(url: str) -> bytes | None:
"""Fetch the leading bytes of an image, enough to read its dimensions."""
try:
target, host_header = await _validated_fetch_target(url)
async with httpx.AsyncClient(
timeout=IMAGE_FETCH_TIMEOUT_SECONDS, follow_redirects=False
) as client:
async with client.stream(
"GET", target, headers={"Host": host_header}
) as response:
response.raise_for_status()
chunks: list[bytes] = []
downloaded = 0
async for chunk in response.aiter_bytes():
chunks.append(chunk)
downloaded += len(chunk)
if downloaded >= IMAGE_FETCH_MAX_BYTES:
break
return b"".join(chunks)[:IMAGE_FETCH_MAX_BYTES]
except Exception as e:
logger.warning(
"Failed to fetch image from URL",
@@ -290,15 +471,46 @@ async def _fetch_image_from_url(url: str) -> bytes | None:
return None
# Patch-based image pricing (OpenAI ``detail: "original"``): the image is
# covered with 32x32px patches and billed as ceil(patches * multiplier)
# tokens, with no 512px-tile downscaling. The API rejects images above
# 30,000 patches, so at the 1.2x multiplier documented for the
# original-capable model families (gpt-5.4/5.5/5.6) the worst case a
# single image can bill is 36,000 tokens.
_IMAGE_PATCH_PX = 32
_MAX_IMAGE_PATCHES = 30_000
_MAX_ORIGINAL_IMAGE_TOKENS = (_MAX_IMAGE_PATCHES * 6 + 4) // 5 # 36,000
def _calculate_original_image_tokens(width: int, height: int) -> int:
"""Estimate tokens for an image billed at ``detail: "original"``.
Patch-based models cover the image with 32x32px patches and bill
``ceil(patches * 1.2)`` tokens. The estimate is bounded by the
30,000-patch rejection limit, which is more conservative than the
per-model resizing patch budgets (e.g. 10,000 patches on gpt-5.4/5.5)
so it never under-reserves.
"""
patches = ((width + _IMAGE_PATCH_PX - 1) // _IMAGE_PATCH_PX) * (
(height + _IMAGE_PATCH_PX - 1) // _IMAGE_PATCH_PX
)
bounded = min(patches, _MAX_IMAGE_PATCHES)
return (bounded * 6 + 4) // 5 # ceil(bounded * 1.2) in exact integer math
def _calculate_image_tokens(width: int, height: int, detail: str = "auto") -> int:
"""Calculate image tokens based on OpenAI's vision pricing.
For low detail: 85 tokens
For high detail/auto: 85 base tokens + 170 tokens per 512px tile
For original detail: patch-based pricing at the original resolution
"""
if detail == "low":
return 85
if detail == "original":
return _calculate_original_image_tokens(width, height)
if width > 2048 or height > 2048:
aspect_ratio = width / height
if width > height:
@@ -330,6 +542,7 @@ async def estimate_image_tokens_in_messages(messages: list) -> int:
Supports both base64 encoded images and image URLs.
"""
total_image_tokens = 0
fetches = 0
for message in messages:
if not isinstance(message, dict):
@@ -350,7 +563,9 @@ async def estimate_image_tokens_in_messages(messages: list) -> int:
continue
content_type = content_item.get("type")
if content_type not in ("image_url", "input_image"):
if content_type == "input_image":
content_item = input_image_part_to_image_url(content_item)
elif content_type != "image_url":
continue
image_url_data = content_item.get("image_url")
@@ -362,7 +577,7 @@ async def estimate_image_tokens_in_messages(messages: list) -> int:
detail = "auto"
elif isinstance(image_url_data, dict):
url = image_url_data.get("url", "")
detail = image_url_data.get("detail", "auto")
detail = image_url_data.get("detail") or "auto"
else:
continue
@@ -370,49 +585,67 @@ async def estimate_image_tokens_in_messages(messages: list) -> int:
continue
if url.startswith("data:image/"):
try:
header, base64_data = url.split(",", 1)
image_bytes = base64.b64decode(base64_data)
width, height = _get_image_dimensions(image_bytes)
tokens = _calculate_image_tokens(width, height, detail)
total_image_tokens += tokens
logger.debug(
"Calculated tokens for base64 image",
extra={
"width": width,
"height": height,
"detail": detail,
"tokens": tokens,
},
)
except Exception as e:
logger.warning(
"Failed to process base64 image",
extra={"error": str(e)},
)
total_image_tokens += 85
total_image_tokens += _data_url_image_tokens(url, detail)
elif url.startswith(FILE_ID_URL_PREFIX):
total_image_tokens += _worst_case_image_tokens(detail)
elif fetches >= IMAGE_FETCH_MAX_PER_REQUEST:
logger.warning(
"Skipping image URL fetch above per-request limit",
extra={"url": url[:100], "limit": IMAGE_FETCH_MAX_PER_REQUEST},
)
total_image_tokens += _worst_case_image_tokens(detail)
else:
fetches += 1
image_bytes_or_none = await _fetch_image_from_url(url)
if image_bytes_or_none:
width, height = _get_image_dimensions(image_bytes_or_none)
tokens = _calculate_image_tokens(width, height, detail)
total_image_tokens += tokens
logger.debug(
"Calculated tokens for URL image",
extra={
"url": url[:100],
"width": width,
"height": height,
"detail": detail,
"tokens": tokens,
},
)
else:
total_image_tokens += 85
total_image_tokens += _image_bytes_tokens(
image_bytes_or_none, detail, source=url[:100]
)
return total_image_tokens
def _worst_case_image_tokens(detail: str) -> int:
"""Dimensions unknown: reserve the most ``detail`` can bill."""
if detail == "original":
return _MAX_ORIGINAL_IMAGE_TOKENS
return _calculate_image_tokens(2048, 2048, detail)
def _data_url_image_tokens(url: str, detail: str) -> int:
try:
_, base64_data = url.split(",", 1)
image_bytes = base64.b64decode(base64_data, validate=True)
except Exception as e:
logger.warning("Failed to decode base64 image", extra={"error": str(e)})
return _worst_case_image_tokens(detail)
return _image_bytes_tokens(image_bytes, detail, source="data-url")
def _image_bytes_tokens(image_bytes: bytes | None, detail: str, source: str) -> int:
if not image_bytes:
return _worst_case_image_tokens(detail)
try:
width, height = Image.open(BytesIO(image_bytes)).size
except Exception as e:
logger.warning(
"Failed to read image dimensions",
extra={"error": str(e), "source": source},
)
return _worst_case_image_tokens(detail)
tokens = _calculate_image_tokens(width, height, detail)
logger.debug(
"Calculated image tokens",
extra={
"source": source,
"width": width,
"height": height,
"detail": detail,
"tokens": tokens,
},
)
return tokens
def create_error_response(
error_type: str,
message: str,
+254 -51
View File
@@ -1,19 +1,28 @@
from __future__ import annotations
import math
import asyncio
import ipaddress
import json
import socket
from collections.abc import Awaitable, Callable
from typing import TypedDict
from typing import Any, TypedDict
import httpx
from cashu.core.base import MeltQuoteState
from cashu.core.settings import settings as cashu_settings
from cashu.wallet.wallet import Proof, Wallet
from ..cashu_compat import install_cashu_httpx_shim
from ..mint import (
MINT_TRANSPORT_EXCEPTIONS,
is_mint_rate_limited,
is_mint_transport_error,
run_mint_operation,
)
# cashu 0.20.x passes the `proxies` kwarg httpx removed in 0.28; see the module
# docstring. Installed at import so no mint call can run before the patch.
install_cashu_httpx_shim()
try:
from bech32 import bech32_decode, convertbits # type: ignore
except ModuleNotFoundError: # pragma: no cover – allow runtime miss
@@ -42,6 +51,114 @@ class MeltOutcomeAmbiguousError(LNURLError):
"""
class MeltUnpaidError(LNURLError):
"""The mint answered the melt request itself with ``unpaid``.
Unlike :class:`MeltOutcomeAmbiguousError` this is proof that no Lightning
payment was made, so callers may restore what they debited.
"""
_MAX_LNURL_REDIRECTS = 3
_MAX_LNURL_RESPONSE_BYTES = 64 * 1024
_NON_PUBLIC_HOST_SUFFIXES = (".localhost", ".local", ".internal")
async def _require_public_https_destination(url: httpx.URL) -> None:
"""Reject anything that is not a public HTTPS endpoint.
LNURL destinations and their redirect targets are attacker-influenced, so
every hop has to be re-checked: a single ``https://`` origin says nothing
about where a 302 points. A bare hostname check is not enough either: a
public-looking name can resolve to a loopback/link-local/private address
(SSRF), so DNS is resolved here and every resulting address must be global.
"""
if url.scheme != "https":
raise LNURLError("LNURL destination must be an HTTPS URL")
host = (url.host or "").rstrip(".").lower()
if not host:
raise LNURLError("LNURL destination has no host")
try:
literal = ipaddress.ip_address(host)
except ValueError:
literal = None
if literal is not None:
if not literal.is_global:
raise LNURLError("LNURL destination is not a public host")
return
if host == "localhost" or host.endswith(_NON_PUBLIC_HOST_SUFFIXES):
raise LNURLError("LNURL destination is not a public host")
port = url.port or 443
try:
infos = await asyncio.get_running_loop().getaddrinfo(
host, port, proto=socket.IPPROTO_TCP
)
except socket.gaierror as e:
raise LNURLError("LNURL destination could not be resolved") from e
if not infos:
raise LNURLError("LNURL destination could not be resolved")
for info in infos:
try:
resolved = ipaddress.ip_address(info[4][0])
except ValueError as e:
raise LNURLError("LNURL destination resolved to an invalid address") from e
if not resolved.is_global:
raise LNURLError("LNURL destination is not a public host")
async def _fetch_lnurl_json(
url: str, params: dict[str, int] | None = None
) -> dict[str, Any]:
"""GET an LNURL endpoint, validating the destination at every redirect.
Response bodies are never echoed: an LNURL service is untrusted, and its
payload would otherwise reach operator logs through raised errors.
"""
try:
target = httpx.URL(url, params=params) if params else httpx.URL(url)
except httpx.InvalidURL as e:
raise LNURLError("LNURL destination is not a usable URL") from e
await _require_public_https_destination(target)
raw: bytes | None = None
async with httpx.AsyncClient() as client:
for _ in range(_MAX_LNURL_REDIRECTS + 1):
async with client.stream(
"GET", target, follow_redirects=False, timeout=10
) as response:
if response.is_redirect:
target = target.join(response.headers.get("location", ""))
await _require_public_https_destination(target)
continue
response.raise_for_status()
chunks = bytearray()
async for chunk in response.aiter_bytes():
chunks.extend(chunk)
if len(chunks) > _MAX_LNURL_RESPONSE_BYTES:
raise LNURLError("LNURL response exceeded the size limit")
raw = bytes(chunks)
break
else:
raise LNURLError("LNURL destination exceeded the redirect limit")
if raw is None:
raise LNURLError("LNURL destination exceeded the redirect limit")
try:
data = json.loads(raw)
except ValueError as e:
raise LNURLError("LNURL response was not valid JSON") from e
if not isinstance(data, dict):
raise LNURLError("LNURL response was not a JSON object")
return data
async def decode_lnurl(lnurl: str) -> str:
"""Decode LNURL to get the actual URL.
@@ -111,26 +228,30 @@ async def get_lnurl_data(lnurl: str) -> LNURLData:
httpx.HTTPError: If the HTTP request fails
"""
url = await decode_lnurl(lnurl)
async with httpx.AsyncClient() as client:
response = await client.get(url, follow_redirects=True, timeout=10)
response.raise_for_status()
lnurl_data = response.json()
lnurl_data = await _fetch_lnurl_json(url)
# Validate payRequest data
if lnurl_data.get("tag") != "payRequest":
raise LNURLError(
f"Invalid LNURL tag: expected 'payRequest', got '{lnurl_data.get('tag')}'"
)
raise LNURLError("Invalid LNURL tag: expected 'payRequest'")
if not isinstance(lnurl_data.get("callback"), str):
callback_url = lnurl_data.get("callback")
if not isinstance(callback_url, str):
raise LNURLError("Invalid LNURL payRequest: missing callback URL")
try:
callback_target = httpx.URL(callback_url)
except httpx.InvalidURL as e:
raise LNURLError("Invalid LNURL callback URL") from e
await _require_public_https_destination(callback_target)
min_sendable = lnurl_data.get("minSendable", 1000) # Default 1 sat
max_sendable = lnurl_data.get("maxSendable", 1000000000) # Default 1000 BTC
if not isinstance(min_sendable, int) or not isinstance(max_sendable, int):
raise LNURLError("Invalid LNURL payRequest: non-integer sendable limits")
return LNURLData(
callback_url=lnurl_data["callback"],
min_sendable=lnurl_data.get("minSendable", 1000), # Default 1 sat
max_sendable=lnurl_data.get("maxSendable", 1000000000), # Default 1000 BTC
callback_url=callback_url,
min_sendable=min_sendable,
max_sendable=max_sendable,
)
@@ -150,26 +271,53 @@ async def get_lnurl_invoice(
LNURLError: If the response is invalid
httpx.HTTPError: If the HTTP request fails
"""
async with httpx.AsyncClient() as client:
response = await client.get(
callback_url,
params={"amount": amount_msat},
follow_redirects=True,
timeout=10,
)
response.raise_for_status()
invoice_data = await _fetch_lnurl_json(callback_url, params={"amount": amount_msat})
invoice_data = response.json()
if "pr" not in invoice_data:
# Check if there's an error in the response
if "reason" in invoice_data:
raise LNURLError(f"LNURL error: {invoice_data['reason']}")
raise LNURLError(f"Invalid LNURL invoice response: {invoice_data}")
if not isinstance(invoice_data.get("pr"), str):
raise LNURLError("LNURL callback returned no invoice")
return invoice_data["pr"], invoice_data
def _select_melt_proofs(
wallet: Wallet,
proofs: list[Proof],
*,
quote_amount: int,
fee_reserve: int,
gross_budget: int,
) -> tuple[list[Proof] | None, int]:
"""Select proofs that cover the quote and exact NUT-02 input fees.
Cashu 0.20's ``select_to_send`` may recursively swap when asked to spend a
wallet's full balance. Melts accept overpayment and return change, so a
bounded, largest-first selection is both safer and minimizes input fees.
Mints reject a melt carrying more than ``mint_max_request_length`` inputs,
so a dust-heavy wallet can only pay what its largest inputs cover; the
caller lowers the amount and the rest goes out on later payouts.
"""
selected: list[Proof] = []
selected_amount = 0
required = quote_amount + fee_reserve
spendable = [
proof
for proof in sorted(proofs, key=lambda item: item.amount, reverse=True)
if getattr(proof, "reserved", False) is not True
]
for proof in spendable[: cashu_settings.mint_max_request_length]:
selected.append(proof)
selected_amount += proof.amount
input_fees = int(wallet.get_fees_for_proofs(selected))
required = quote_amount + fee_reserve + input_fees
if selected_amount >= required:
if required <= gross_budget:
return selected, 0
# Covered but over budget; more proofs only raise input fees.
break
return None, max(1, required - min(selected_amount, gross_budget))
async def raw_send_to_lnurl(
wallet: Wallet,
proofs: list[Proof],
@@ -201,12 +349,11 @@ async def raw_send_to_lnurl(
# Send USD to Lightning Address
paid = await wallet.send_to_lnurl("user@getalby.com", 50, unit="usd")
"""
total_balance = sum(proof.amount for proof in proofs)
if amount and total_balance < amount:
if not isinstance(amount, int) or isinstance(amount, bool) or amount <= 0:
raise ValueError("A positive integer amount is required to send to an LNURL.")
if sum(proof.amount for proof in proofs) < amount:
raise ValueError("Amount to send is higher than available proofs.")
else:
assert isinstance(amount, int)
total_balance = amount
total_balance = amount
lnurl_data = await get_lnurl_data(lnurl)
if unit == "sat":
@@ -226,25 +373,51 @@ async def raw_send_to_lnurl(
f"({min_sendable_sat} - {max_sendable_sat} {unit})"
)
estimated_fees_sat = int(max(math.ceil((amount_msat / 1000) * 0.01), 2)) + 1
estimated_fees_msat = estimated_fees_sat * 1000
final_amount = amount_msat - estimated_fees_msat
final_amount = amount_msat
bolt11_invoice, _ = await get_lnurl_invoice(
lnurl_data["callback_url"], final_amount
)
selected_proofs: list[Proof] | None = None
# Find the largest amount covered by the budget after reserve and input fees.
for _ in range(8):
if final_amount < lnurl_data["min_sendable"]:
raise LNURLError("Cashu melt fees leave no payable LNURL amount")
bolt11_invoice, _ = await get_lnurl_invoice(
lnurl_data["callback_url"], final_amount
)
melt_quote_resp = await run_mint_operation(
lambda: wallet.melt_quote(invoice=bolt11_invoice),
op_name="lnurl_melt_quote",
mint_url=str(wallet.url),
# Quote creation is unsafe to retry without idempotency.
retry_timeouts=False,
)
melt_quote_resp = await run_mint_operation(
lambda: wallet.melt_quote(invoice=bolt11_invoice),
op_name="lnurl_melt_quote",
mint_url=str(wallet.url),
)
quoted_amount = int(melt_quote_resp.amount)
expected_amount = final_amount // 1000 if unit == "sat" else final_amount
if quoted_amount != expected_amount:
raise LNURLError(
f"LNURL invoice amount does not match the requested amount "
f"(quoted {quoted_amount} {unit}, expected {expected_amount} {unit})"
)
selected_proofs, shortfall = _select_melt_proofs(
wallet,
proofs,
quote_amount=quoted_amount,
fee_reserve=int(melt_quote_resp.fee_reserve),
gross_budget=amount,
)
if selected_proofs is not None:
break
final_amount -= shortfall * (1000 if unit == "sat" else 1)
else:
raise LNURLError("Cashu melt fees exceed the requested gross amount")
if on_melt_quote is not None:
await on_melt_quote(melt_quote_resp.quote)
if amount:
proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True)
assert selected_proofs is not None
proofs = selected_proofs
await wallet.set_reserved_for_send(proofs, reserved=True)
try:
melt_response = await run_mint_operation(
@@ -265,15 +438,29 @@ async def raw_send_to_lnurl(
# reserved as though a Lightning payment could still settle.
await wallet.set_reserved_for_send(proofs, reserved=False)
raise
if not isinstance(error, MINT_TRANSPORT_EXCEPTIONS):
if not is_mint_transport_error(error):
raise
# Cashu clears reservations on transport errors despite an unknown outcome.
try:
await wallet.set_reserved_for_melt(
proofs, reserved=True, quote_id=melt_quote_resp.quote
)
except Exception as reservation_error:
raise MeltOutcomeAmbiguousError(
"Melt outcome is ambiguous and its proof reservation could not "
"be restored; proofs must not be retried"
) from reservation_error
melt_response = None
melt_error: BaseException | None = error
else:
melt_error = None
if getattr(melt_response, "state", None) == MeltQuoteState.paid:
melt_state = getattr(melt_response, "state", None)
if melt_state == MeltQuoteState.paid:
return final_amount
if melt_state == MeltQuoteState.unpaid:
await wallet.set_reserved_for_send(proofs, reserved=False)
raise MeltUnpaidError("Cashu mint confirmed that the melt was unpaid")
try:
quote = await run_mint_operation(
@@ -281,6 +468,8 @@ async def raw_send_to_lnurl(
op_name="reconcile_lnurl_melt_quote",
mint_url=str(wallet.url),
retry_timeouts=False,
# Reconciliation must bypass the cooldown opened by this failure.
allow_during_cooldown=True,
)
except Exception as reconciliation_error:
raise MeltOutcomeAmbiguousError(
@@ -290,6 +479,20 @@ async def raw_send_to_lnurl(
if quote is not None and quote.state == MeltQuoteState.paid:
return final_amount
if quote is not None and quote.state == MeltQuoteState.unpaid:
# A just-dispatched quote can briefly report unpaid before transitioning.
try:
await wallet.set_reserved_for_melt(
proofs, reserved=True, quote_id=melt_quote_resp.quote
)
except Exception as reservation_error:
raise MeltOutcomeAmbiguousError(
"Melt outcome is ambiguous and its proof reservation could not "
"be restored; proofs must not be retried"
) from reservation_error
raise MeltOutcomeAmbiguousError(
"Melt outcome is ambiguous; an immediate unpaid state is not final"
) from melt_error
state = getattr(getattr(quote, "state", None), "value", "unknown")
raise MeltOutcomeAmbiguousError(
+146 -40
View File
@@ -5,13 +5,14 @@ import random
import httpx
from fastapi import APIRouter, Depends, HTTPException, Request
from pydantic import BaseModel as V2BaseModel
from pydantic.v1 import BaseModel
from pydantic.v1 import BaseModel, validator
from sqlmodel.ext.asyncio.session import AsyncSession
from ..core.db import ModelRow, UpstreamProviderRow, get_session
from ..core.logging import get_logger
from ..core.settings import settings
from .price import sats_usd_price
from .rates import BILLABLE_PRICING_FIELDS, coerce_rate, is_usable_rate
logger = get_logger(__name__)
@@ -58,12 +59,56 @@ class Pricing(BaseModel):
max_cost: float = 0.0 # in sats not msats
# The rates ``Pricing`` declares without a default, derived from the model so the
# two cannot drift. A payload that omits one writes a row that will not parse.
REQUIRED_PRICING_FIELDS = tuple(
name for name, field in Pricing.__fields__.items() if field.required
)
def has_usable_pricing(pricing: Pricing) -> bool:
"""True if every billable rate is a number a request could be billed on.
Free is usable — a rate of zero is a real price. This asks only whether the
price is well-formed. One unusable rate disqualifies the whole price even
alongside a valid one, since a request can bill on the bad field: a positive
``completion`` must not hide a negative ``prompt``.
"""
return all(
is_usable_rate(getattr(pricing, field)) for field in BILLABLE_PRICING_FIELDS
)
class TopProvider(BaseModel):
context_length: int | None = None
max_completion_tokens: int | None = None
is_moderated: bool | None = None
class Reasoning(BaseModel):
"""Per-model reasoning-effort metadata, matching OpenRouter's shape."""
mandatory: bool | None = None
default_enabled: bool | None = None
supported_efforts: list[str] | None = None
default_effort: str | None = None
supports_max_tokens: bool | None = None
class Config:
extra = "ignore"
def is_empty(self) -> bool:
return not any(
(
self.mandatory is not None,
self.default_enabled is not None,
self.supported_efforts,
self.default_effort,
self.supports_max_tokens is not None,
)
)
class Model(BaseModel):
id: str
name: str
@@ -80,10 +125,43 @@ class Model(BaseModel):
canonical_slug: str | None = None
alias_ids: list[str] | None = None
forwarded_model_id: str | None = None
reasoning: Reasoning | None = None
class Config:
extra = "ignore"
def __hash__(self) -> int:
return hash(self.id)
@validator("reasoning", pre=True)
def _coerce_reasoning(cls, value: object) -> object:
if value is None or value is False:
return None
if isinstance(value, Reasoning):
return None if value.is_empty() else value
if not isinstance(value, dict) or not value:
return None
try:
parsed = Reasoning.parse_obj(value)
except Exception:
return None
return None if parsed.is_empty() else parsed
def dict(self, **kwargs: object) -> dict:
# Non-reasoning models omit the field entirely so the catalog stays
# additive: existing clients never see a new null key.
data = super().dict(**kwargs) # type: ignore[arg-type]
reasoning = data.get("reasoning")
if not reasoning:
data.pop("reasoning", None)
elif isinstance(reasoning, dict):
cleaned = {k: v for k, v in reasoning.items() if v is not None}
if cleaned:
data["reasoning"] = cleaned
else:
data.pop("reasoning", None)
return data
def litellm_cost_entry(model_id: str) -> dict | None:
"""Look up ``model_id`` in litellm's bundled cost map.
@@ -144,18 +222,17 @@ def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing:
def _has_valid_pricing(model: dict) -> bool:
"""Check if model has valid pricing (not free, no negative values)."""
"""Check if model has valid pricing (usable rates, and not free)."""
pricing = model.get("pricing", {})
if not pricing:
return False
try:
prompt = float(pricing.get("prompt", 0))
completion = float(pricing.get("completion", 0))
except (ValueError, TypeError):
return False
if prompt < 0 or completion < 0:
# Coercion runs before the both-zero test below, which `NaN` would defeat
# on its own — and one entry the coercion chokes on must not unwind the
# whole fetch, which once cost the node an entire upstream catalog.
prompt = coerce_rate(pricing.get("prompt", 0))
completion = coerce_rate(pricing.get("completion", 0))
if prompt is None or completion is None:
return False
if prompt == 0 and completion == 0:
@@ -221,9 +298,10 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis
return []
def _row_to_model(
def _build_model_from_row(
row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01
) -> Model:
"""The deterministic USD view of a stored model row, before the sats conversion."""
architecture = json.loads(row.architecture)
pricing = json.loads(row.pricing)
per_request_limits = (
@@ -281,6 +359,14 @@ def _row_to_model(
parsed_pricing.max_cost,
) = _calculate_usd_max_costs(model)
return model
def _row_to_model(
row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01
) -> Model:
model = _build_model_from_row(row, apply_provider_fee, provider_fee)
try:
sats_to_usd = sats_usd_price()
model = _update_model_sats_pricing(model, sats_to_usd)
@@ -309,21 +395,57 @@ async def list_models(
rows = (await session.exec(query)).all() # type: ignore
provider_result = await session.exec(select(UpstreamProviderRow))
providers_by_id = {p.id: p for p in provider_result.all()}
return [
_row_to_model(
r,
apply_provider_fee=apply_fees,
provider_fee=providers_by_id[r.upstream_provider_id].provider_fee
if r.upstream_provider_id in providers_by_id
else 1.01,
)
for r in rows
if include_disabled
or (
models: list[Model] = []
for r in rows:
if not include_disabled and not (
r.upstream_provider_id in providers_by_id
and providers_by_id[r.upstream_provider_id].enabled
)
]
):
continue
try:
model = _row_to_model(
r,
apply_provider_fee=apply_fees,
provider_fee=providers_by_id[r.upstream_provider_id].provider_fee
if r.upstream_provider_id in providers_by_id
else 1.01,
)
except Exception as e:
# Stored pricing/architecture is JSON from whatever wrote the row, so
# a legacy import or foreign writer can leave a field that will not
# parse. Converting inside this loop meant one such row raised out of
# the whole listing and the node advertised nothing at all. Drop the
# row we cannot read — it is unservable either way — and keep serving
# the rest.
logger.warning(
"Skipping model row that could not be read",
extra={
"model_id": r.id,
"upstream_provider_id": r.upstream_provider_id,
"error": str(e),
"error_type": type(e).__name__,
},
)
continue
# Served-map backstop for legacy rows and writers that bypass the admin
# edge: a negative or non-finite rate is not a price. Serving one
# advertises a rate the cost calculation cannot bill on, so the request
# falls through to the flat maximum reservation — or, if the rate is
# negative, bills an amount settlement credits back to the caller.
# ``include_disabled`` is the operator's listing, which must keep showing
# the row so it can be repaired.
if not include_disabled and not has_usable_pricing(model.pricing):
logger.warning(
"Withholding model with an unusable stored rate from the catalog",
extra={
"model_id": r.id,
"upstream_provider_id": r.upstream_provider_id,
},
)
continue
models.append(model)
return models
def _calculate_usd_max_costs(model: Model) -> tuple[float, float, float]:
@@ -409,23 +531,7 @@ def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model:
if (sats.max_cost or 0.0) < min_req_sats:
sats.max_cost = min_req_sats
return Model(
id=model.id,
name=model.name,
created=model.created,
description=model.description,
context_length=model.context_length,
architecture=model.architecture,
pricing=model.pricing,
sats_pricing=sats,
per_request_limits=model.per_request_limits,
top_provider=model.top_provider,
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
alias_ids=model.alias_ids,
forwarded_model_id=model.forwarded_model_id,
)
return model.copy(update={"sats_pricing": sats})
except Exception as e:
logger.error(
"Failed to update sats pricing for model",
+35 -13
View File
@@ -5,12 +5,37 @@ import httpx
from ..core import get_logger
from ..core.settings import settings
from .rates import coerce_rate
logger = get_logger(__name__)
BTC_USD_PRICE: float | None = None
SATS_USD_PRICE: float | None = None
SATS_PER_BTC = 100_000_000
def _parse_quote(raw: object, exchange: str) -> float | None:
"""Coerce an exchange quote to a price, or ``None`` if it is not one.
Every quote passes through here because the aggregator takes the ``min()``
of what it collects: an unusable quote does not merely join the sample, it
*wins* it, and the result is the rate every model and every request on the
node is priced at. A quote is stricter than a billable rate — it must be
positive, and positive *after* the sats conversion the node prices in: a
subnormal quote survives every guard here and still underflows to a zero
sats price, which then divides by zero on every model's rate.
"""
price = coerce_rate(raw)
if price is None or price <= 0 or price / SATS_PER_BTC <= 0:
logger.warning(
"Unusable price quote — ignoring this exchange",
extra={"exchange": exchange, "quote": repr(raw)},
)
return None
return price
async def _kraken_btc_usd(client: httpx.AsyncClient) -> float | None:
"""Fetch BTC/USD price from Kraken API."""
@@ -18,10 +43,11 @@ async def _kraken_btc_usd(client: httpx.AsyncClient) -> float | None:
try:
response = await client.get(api)
price_data = response.json()
price = float(price_data["result"]["XXBTZUSD"]["c"][0])
return price
except (httpx.RequestError, KeyError) as e:
return _parse_quote(price_data["result"]["XXBTZUSD"]["c"][0], "kraken")
except (httpx.RequestError, KeyError, IndexError, TypeError, ValueError) as e:
# A payload whose *shape* changed raises IndexError/TypeError, and a
# non-JSON body raises ValueError; unhandled, one exchange's bad day
# aborted the whole aggregation instead of dropping a single quote.
logger.warning(
"Kraken API error",
extra={
@@ -39,10 +65,8 @@ async def _coinbase_btc_usd(client: httpx.AsyncClient) -> float | None:
try:
response = await client.get(api)
price_data = response.json()
price = float(price_data["data"]["amount"])
return price
except (httpx.RequestError, KeyError) as e:
return _parse_quote(price_data["data"]["amount"], "coinbase")
except (httpx.RequestError, KeyError, IndexError, TypeError, ValueError) as e:
logger.warning(
"Coinbase API error",
extra={
@@ -60,10 +84,8 @@ async def _binance_btc_usdt(client: httpx.AsyncClient) -> float | None:
try:
response = await client.get(api)
price_data = response.json()
price = float(price_data["price"])
return price
except (httpx.RequestError, KeyError) as e:
return _parse_quote(price_data["price"], "binance")
except (httpx.RequestError, KeyError, IndexError, TypeError, ValueError) as e:
logger.warning(
"Binance API error",
extra={
@@ -123,7 +145,7 @@ async def _update_prices() -> None:
)
return
BTC_USD_PRICE = btc_price
SATS_USD_PRICE = btc_price / 100_000_000
SATS_USD_PRICE = btc_price / SATS_PER_BTC
def btc_usd_price() -> float:
+71
View File
@@ -0,0 +1,71 @@
"""The one definition of a billable rate, with no dependencies of its own.
A rate reaches the node from an upstream catalog, the LiteLLM cost map, an
operator's admin edit, a legacy database row and the BTC/USD feed. Each of those
readers needs the same two questions answered — is this value a rate at all, and
is it a rate a request can be billed on — so both answers live here, in a module
that imports nothing from the package. Every guard then shares one definition
instead of drifting, and no caller needs a deferred import to reach it.
"""
import math
# The rates a request can bill on. Derived fields (``max_*_cost``) are excluded
# — they are computed carriers, not charged rates. One definition, shared by the
# admin write edge and the served/routed guards, so they all cover the same set.
BILLABLE_PRICING_FIELDS = (
"prompt",
"completion",
"request",
"image",
"web_search",
"internal_reasoning",
"input_cache_read",
"input_cache_write",
)
def is_usable_rate(rate: float) -> bool:
"""True if a single billable rate is a number a request could be billed on.
The one definition of a usable rate, so every guard that asks the question
answers it identically. A rate qualifies only when it is finite and
non-negative; zero is usable (it means "free", which is a real price) but
``NaN``, ``±inf`` and negatives are not prices at all.
Non-finite: ``inf > 0`` is True, so an infinite rate reads as chargeable and
would be served, routed and billed as ``inf``; ``NaN`` poisons every total it
enters and defeats ordinary comparisons, since ``NaN > 0``, ``NaN < 0`` and
``NaN == 0`` are all False. Negative: a negative rate produces a negative
cost, which the settlement path subtracts from the balance — it pays the
caller to make requests. Both reach a stored row from upstream catalogs as
well as the admin edge (``json.loads`` accepts the bare ``NaN``/``Infinity``
literals and overflows ``1e999`` to ``inf``).
This is the rationale for every guard that calls it; the call sites say what
they do with the answer, not why the answer matters.
"""
return math.isfinite(rate) and rate >= 0.0
def coerce_rate(value: object) -> float | None:
"""Coerce a value from outside the node to a usable rate, or ``None``.
The one coercion, shared by every reader of a rate the node did not compute
itself: an upstream catalog, the LiteLLM cost map, the exchange feed and the
admin write edge. Each of them was parsing for itself, and they disagreed —
which is how a boolean became a price on some paths and not others.
A boolean is rejected outright: it is a change of shape, not a rate, and
Python would make ``True`` a finite, positive ``1.0`` that passes every
numeric guard downstream — a dollar per token. A numeric string is accepted,
because feeds report prices as strings. An oversized integer raises
``OverflowError`` rather than ``ValueError``, so that is caught too.
"""
if isinstance(value, bool) or not isinstance(value, (int, float, str)):
return None
try:
rate = float(value)
except (TypeError, ValueError, OverflowError):
return None
return rate if is_usable_rate(rate) else None
+79
View File
@@ -0,0 +1,79 @@
"""Convert a Responses API ``input`` into chat ``messages`` via litellm.
litellm drops ``file_id`` (emits ``url: ""``) and nests a dict-form ``image_url``
as-is, so ``input_image`` parts are flattened to ``{image_url: str, detail}`` first.
``file_id`` becomes a sentinel URL the image walker treats as unfetchable.
"""
from typing import Any
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
from ..core import get_logger
logger = get_logger(__name__)
FILE_ID_URL_PREFIX = "file-id:"
def _flatten_input_image(part: dict[str, Any]) -> tuple[str, str]:
raw = part.get("image_url")
url = raw.get("url", "") if isinstance(raw, dict) else raw
detail = part.get("detail") or (
raw.get("detail") if isinstance(raw, dict) else None
)
if not url and part.get("file_id"):
url = f"{FILE_ID_URL_PREFIX}{part['file_id']}"
return (url if isinstance(url, str) else ""), (detail or "auto")
def _normalize_item(item: Any) -> Any:
if not isinstance(item, dict):
return item
if item.get("type") == "input_image":
url, detail = _flatten_input_image(item)
return {**item, "image_url": url, "detail": detail}
content = item.get("content")
if isinstance(content, list):
return {**item, "content": [_normalize_item(part) for part in content]}
return item
def input_image_part_to_image_url(part: dict[str, Any]) -> dict[str, Any]:
"""Reshape an ``input_image`` part found inside chat ``messages``."""
url, detail = _flatten_input_image(part)
return {"type": "image_url", "image_url": {"url": url, "detail": detail}}
def count_input_images(input_data: Any) -> int:
if isinstance(input_data, dict):
own = 1 if input_data.get("type") == "input_image" else 0
return own + count_input_images(input_data.get("content"))
if isinstance(input_data, list):
return sum(count_input_images(item) for item in input_data)
return 0
def responses_input_to_messages(input_data: Any) -> list[dict[str, Any]] | None:
"""Returns ``None`` when the transform fails so the caller can worst-case."""
if isinstance(input_data, str):
return [{"role": "user", "content": input_data}]
if not isinstance(input_data, list):
return []
try:
normalized = [_normalize_item(item) for item in input_data]
converted = (
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=normalized, # type: ignore[arg-type]
responses_api_request={},
)
)
return [dict(message) for message in converted]
except Exception as e:
logger.warning(
"Responses input transform failed; using conservative image fallback",
extra={"error": str(e)},
)
return None
+276 -26
View File
@@ -26,6 +26,7 @@ from .core.db import (
)
from .core.exceptions import UpstreamError
from .core.not_found import build_not_found_response
from .core.settings import settings
from .payment.helpers import (
calculate_discounted_max_cost,
check_token_balance,
@@ -37,9 +38,18 @@ from .payment.models import Model
from .upstream import BaseUpstreamProvider
from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request
from .upstream.helpers import init_upstreams
from .upstream.model_paths import (
ModelPathSelector,
decode_model_path,
is_openrouter_base_url,
public_model_id,
public_provider_url,
)
from .upstream.request_correction import correct_request, extract_error_message
logger = get_logger(__name__)
MODEL_PATH_HEADER = "x-routstr-model-path"
proxy_router = APIRouter()
_upstreams: list[BaseUpstreamProvider] = []
@@ -112,6 +122,25 @@ def get_candidates(
return None
def _model_ids_match(requested: str, selected: str) -> bool:
if requested.lower() == selected.lower():
return True
return public_model_id(requested).lower() == public_model_id(selected).lower()
def _candidate_for_selector(
selector: ModelPathSelector,
candidates: list[tuple[Model, BaseUpstreamProvider]],
) -> tuple[Model, BaseUpstreamProvider] | None:
for model_obj, upstream in candidates:
if (
upstream.db_id == selector.provider_id
and public_provider_url(upstream.base_url) == selector.base_url
):
return model_obj, upstream
return None
def get_model_instance(model_id: str) -> Model | None:
"""Get the best-ranked Model instance for a model ID."""
candidates = get_candidates(model_id)
@@ -221,22 +250,141 @@ async def refresh_model_maps_periodically() -> None:
)
_API_PATH_PREFIXES = (
"v1/",
"responses",
"chat/",
"completions",
"models",
"embeddings",
"audio/",
"images/",
"moderations",
"providers",
"tee/",
"attestation",
# Canonical endpoints this proxy will forward, keyed by the path with any
# leading "v1/" and trailing slash removed, mapped to the methods allowed on
# each. The provider credential is attached during forwarding, so endpoint
# permission has to come from this table rather than from the client-supplied
# path: an upstream's key-management, organization, or billing routes live
# under the same origin and must never be reachable through the proxy.
_ALLOWED_ENDPOINTS: dict[str, frozenset[str]] = {
"chat/completions": frozenset({"POST"}),
"completions": frozenset({"POST"}),
"responses": frozenset({"POST"}),
"messages": frozenset({"POST"}),
"embeddings": frozenset({"POST"}),
"models": frozenset({"GET"}),
"attestation": frozenset({"GET"}),
"tee/attestation": frozenset({"GET"}),
}
_ALLOWED_METHODS = frozenset({"GET", "POST"})
def _canonical_api_path(path: str) -> str:
"""Reduce a request path to its allowlist key.
OpenAI-style clients reach the same endpoint with or without the ``v1/``
prefix and with or without a trailing slash, so both spellings collapse to
one key. Callers must screen the path with
:func:`_is_ambiguously_spelled_path` first — this function assumes the path
has no dot segments, empty segments, or encoded separators left to resolve.
"""
core = path[:-1] if path.endswith("/") else path
if core.startswith("v1/"):
core = core[len("v1/") :]
return core
def _parse_extra_allowed_endpoints(raw: str) -> dict[str, frozenset[str]]:
"""Parse operator-configured additions to the endpoint allowlist.
Deployments whose provider exposes an endpoint outside the canonical set
opt in explicitly with ``PROXY_EXTRA_ALLOWED_PATHS``, a comma-separated
list of ``METHOD:path`` pairs (e.g. ``POST:v1/rerank,GET:batches``). Every
entry must name one concrete method and one unambiguous path; wildcards
and bare prefixes are deliberately unsupported, so widening the proxy's
reach is always a per-endpoint decision. Malformed entries are dropped
with a warning rather than silently widening or narrowing the surface.
"""
extra: dict[str, frozenset[str]] = {}
for entry in raw.split(","):
entry = entry.strip()
if not entry:
continue
method, separator, endpoint = entry.partition(":")
method = method.strip().upper()
endpoint = endpoint.strip()
if not separator or method not in _ALLOWED_METHODS or not endpoint:
logger.warning(
"Ignoring malformed PROXY_EXTRA_ALLOWED_PATHS entry",
extra={"entry": entry},
)
continue
if _is_ambiguously_spelled_path(endpoint):
logger.warning(
"Ignoring ambiguously spelled PROXY_EXTRA_ALLOWED_PATHS entry",
extra={"entry": entry},
)
continue
if any(character in endpoint for character in "*?["):
# Refuse glob syntax outright. Kept as a literal endpoint name it
# would never match a real request, so the operator would think
# they had widened the proxy when they had not.
logger.warning(
"Ignoring wildcard PROXY_EXTRA_ALLOWED_PATHS entry; "
"list each endpoint explicitly",
extra={"entry": entry},
)
continue
key = _canonical_api_path(endpoint)
extra[key] = extra.get(key, frozenset()) | {method}
return extra
def _is_ambiguously_spelled_path(path: str) -> bool:
"""Reject paths whose spelling could resolve somewhere the allowlist did not.
``{path:path}`` arrives percent-decoded, so a client that sent ``%2e%2e`` or
``%2f`` shows up here as ``..`` / ``/``. Dot segments, backslashes, duplicate
or leading separators, NUL bytes, and any residual encoded separator are
treated as unsafe: they let a caller walk off the canonical API surface (and
onto a sensitive upstream endpoint) even though the literal prefix check
would pass. Reject rather than trying to rewrite the path.
"""
if not path or path != path.strip() or path.startswith("/"):
return True
if "\x00" in path or "\\" in path:
return True
# A single trailing slash is canonical (e.g. "attestation/"); ignore it,
# then no remaining segment may be empty (covers "//") or a dot segment.
core = path[:-1] if path.endswith("/") else path
if any(segment in ("", ".", "..") for segment in core.split("/")):
return True
lowered = path.lower()
return "%2e" in lowered or "%2f" in lowered or "%5c" in lowered
_EXTRA_ALLOWED_ENDPOINTS = _parse_extra_allowed_endpoints(
settings.proxy_extra_allowed_paths
)
def _allowed_methods_for(endpoint: str) -> frozenset[str]:
"""Return the methods allowed on a canonical endpoint, empty if unknown."""
methods = _ALLOWED_ENDPOINTS.get(endpoint, frozenset())
methods |= _EXTRA_ALLOWED_ENDPOINTS.get(endpoint, frozenset())
return methods
def _forwarding_allowed(path: str, method: str) -> bool:
"""Gate which method/path pairs may reach an upstream at all.
The provider credential is attached during forwarding, so an unknown
endpoint must never be forwarded on the caller's say-so. The path is
reduced to its canonical form and looked up in the endpoint table; there is
no prefix match, so a known prefix no longer carries an unknown endpoint
(``v1/organization/api_keys`` is rejected even though ``v1/`` is familiar).
EHBP requests are gated by the same table. Their body is opaque to the
proxy, which is a reason to constrain the destination more tightly, not to
trust the caller's path: the encrypted contract covers the body, never the
endpoint the credential is spent against.
"""
if method not in _ALLOWED_METHODS:
return False
return method in _allowed_methods_for(_canonical_api_path(path))
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
async def proxy(
request: Request, path: str, session: AsyncSession = Depends(get_session)
@@ -255,14 +403,17 @@ async def proxy(
async def _proxy(
request: Request, path: str, session: AsyncSession
) -> Response | StreamingResponse:
# GET requests must hit a known API prefix; otherwise return a 404 (HTML
# for browsers, JSON for API clients). POST requests are always forwarded
# so that OpenAI-style endpoints work with or without the `v1/` prefix
# (e.g. `/chat/completions` as well as `/v1/chat/completions`).
if request.method == "GET" and not path.startswith(_API_PATH_PREFIXES):
# Screen the path before any routing decision: reject ambiguous spellings,
# then require a known API prefix so nothing unknown is forwarded with the
# provider credential attached.
if _is_ambiguously_spelled_path(path):
return build_not_found_response(request, path)
headers = dict(request.headers)
is_ehbp = "ehbp-encapsulated-key" in headers
if not _forwarding_allowed(path, request.method):
return build_not_found_response(request, path)
is_responses_api = path.startswith("v1/responses") or path.startswith("responses")
request_body = await request.body()
@@ -272,7 +423,6 @@ async def _proxy(
# extract the model id, so the SDK sends it in X-Routstr-Model. Forward the
# raw encrypted body to the upstream's /private/ endpoint and stream the
# encrypted response back untouched — the SDK's SecureClient decrypts it.
is_ehbp = "ehbp-encapsulated-key" in headers
if is_ehbp:
request_body_dict = {}
model_id = headers.get("x-routstr-model", "")
@@ -294,6 +444,13 @@ async def _proxy(
# without model/cost/auth lookups. Do not prefix-match here: paths such as
# /attestationjunk must continue through normal authentication.
if request.method == "GET" and _is_tinfoil_attestation_path(path):
if MODEL_PATH_HEADER in headers:
return create_error_response(
"unsupported_request",
"Model paths do not apply to attestation",
400,
request=request,
)
selected_upstreams = _select_unauthenticated_get_upstreams(path, _upstreams)
if not selected_upstreams:
return create_error_response(
@@ -334,6 +491,41 @@ async def _proxy(
"upstream_error", "All upstreams failed", 502, request=request
)
selector: ModelPathSelector | None = None
if MODEL_PATH_HEADER in headers:
selector = decode_model_path(headers[MODEL_PATH_HEADER])
if (
selector is None
or sum(
name.lower() == MODEL_PATH_HEADER for name, _ in request.headers.items()
)
!= 1
):
return create_error_response(
"invalid_request",
f"Malformed {MODEL_PATH_HEADER} header",
400,
request=request,
)
if not isinstance(model_id, str) or not _model_ids_match(
model_id, selector.model_id
):
return create_error_response(
"invalid_request",
f"{MODEL_PATH_HEADER} selects model '{selector.model_id}' but the "
f"request asks for '{model_id}'",
400,
request=request,
)
if "models" in request_body_dict:
return create_error_response(
"invalid_request",
"Model paths cannot be combined with model fallbacks",
400,
request=request,
)
model_id = selector.model_id
candidates = get_candidates(model_id)
if not candidates:
@@ -341,6 +533,51 @@ async def _proxy(
"invalid_model", f"Model '{model_id}' not found", 400, request=request
)
if selector is not None:
pinned = _candidate_for_selector(selector, candidates)
if pinned is None:
return create_error_response(
"invalid_model_path",
f"Model '{selector.model_id}' is not routable through provider "
f"{selector.provider_id}",
404,
request=request,
)
# Explicit routes must never enter cross-provider failover.
candidates = [pinned]
if selector.endpoint_tag:
if (
is_ehbp
or not request_body_dict
or not is_openrouter_base_url(pinned[1].base_url)
or _canonical_api_path(path)
not in {"chat/completions", "completions", "responses"}
):
return create_error_response(
"unsupported_request",
"Endpoint pinning requires an OpenRouter completion or Responses JSON request",
400,
request=request,
)
provider_options = request_body_dict.get("provider", {})
if not isinstance(provider_options, dict):
return create_error_response(
"invalid_request",
"provider must be an object",
400,
request=request,
)
request_body_dict = {
**request_body_dict,
"provider": {
**provider_options,
"order": [selector.endpoint_tag],
"allow_fallbacks": False,
},
}
request_body = json.dumps(request_body_dict).encode()
if is_ehbp:
candidates = [
(model, upstream)
@@ -391,11 +628,21 @@ async def _proxy(
)
elif is_responses_api:
return await upstream.handle_x_cashu_responses(
request, x_cashu, path, max_cost_for_model, model_obj
request,
x_cashu,
path,
max_cost_for_model,
model_obj,
request_body=request_body,
)
else:
return await upstream.handle_x_cashu(
request, x_cashu, path, max_cost_for_model, model_obj
request,
x_cashu,
path,
max_cost_for_model,
model_obj,
request_body=request_body,
)
except UpstreamError as e:
logger.warning(
@@ -593,17 +840,20 @@ async def _proxy(
)
raise
# Reactive recovery: some models reject one specific request
# param (e.g. newer Anthropic models deprecating `temperature`).
# When the upstream 400s naming such a param, strip it from the
# body and retry the SAME upstream. ``already_stripped`` bounds
# this to one retry per distinct param so it always terminates.
# Same-provider recovery must not relax an explicit route.
if response.status_code == 400 and not is_ehbp:
correction = correct_request(
request_body,
extract_error_message(response),
already_stripped,
)
if correction is not None and selector is not None:
corrected_body = json.loads(correction.body)
if any(
corrected_body.get(field) != request_body_dict.get(field)
for field in ("model", "provider")
):
correction = None
if correction is not None:
request_body, bad_param = correction.body, correction.label
already_stripped.add(bad_param)
+109
View File
@@ -0,0 +1,109 @@
"""In-memory negative cache for terminally failed Cashu token redemptions.
A dead token (already spent, malformed, zero value) presented as a bearer key
triggers a full redemption attempt against the issuing mint on *every* request,
because no ``api_keys`` row survives the failed attempt. Polling clients that
never back off turn one dead token into thousands of pointless mint calls per
day. This cache remembers terminal redemption failures by token hash so
repeated presentations are rejected locally with the same error the mint
attempt would have produced.
Only failures whose classification code is in :data:`TERMINAL_REDEMPTION_CODES`
are cached — transient failures (mint unreachable, rate-limited, cooldown) must
never be cached, or a brief mint outage would poison valid tokens.
The cache is deliberately in-memory (bounded LRU with TTL) rather than a
database row: persisting a row per failed token would let an attacker fill the
database with garbage tokens for free.
"""
import time
from collections import OrderedDict
from dataclasses import dataclass
from typing import Callable
# Redemption ``code`` values that can never succeed on retry. A token that was
# already spent, failed to decode, or redeemed to zero value stays that way
# forever; swap fees exceeding the token amount only changes if the mint
# lowers its fees, which the TTL covers.
TERMINAL_REDEMPTION_CODES: frozenset[str] = frozenset(
{
"cashu_token_already_spent",
"invalid_cashu_token",
"cashu_token_zero_value",
"cashu_token_swap_fees_exceed_amount",
}
)
DEFAULT_MAX_ENTRIES = 10_000
DEFAULT_TTL_SECONDS = 24 * 60 * 60
@dataclass(frozen=True)
class CachedRedemptionFailure:
"""Sanitized classification of a terminal redemption failure.
Mirrors the ``(type, status, message, code)`` tuple produced by
``classify_redemption_error`` so a cache hit yields a byte-identical error
envelope to the original mint-backed failure.
"""
status_code: int
error_type: str
message: str
code: str
class RedemptionNegativeCache:
"""Bounded TTL+LRU cache keyed by the SHA-256 hash of the bearer token.
Not thread-safe by design: all access happens on the asyncio event loop.
"""
def __init__(
self,
max_entries: int = DEFAULT_MAX_ENTRIES,
ttl_seconds: float = DEFAULT_TTL_SECONDS,
clock: Callable[[], float] = time.monotonic,
) -> None:
if max_entries <= 0:
raise ValueError("max_entries must be positive")
if ttl_seconds <= 0:
raise ValueError("ttl_seconds must be positive")
self._max_entries = max_entries
self._ttl_seconds = ttl_seconds
self._clock = clock
self._entries: OrderedDict[str, tuple[float, CachedRedemptionFailure]] = (
OrderedDict()
)
def get(self, hashed_key: str) -> CachedRedemptionFailure | None:
entry = self._entries.get(hashed_key)
if entry is None:
return None
expires_at, failure = entry
if self._clock() >= expires_at:
del self._entries[hashed_key]
return None
self._entries.move_to_end(hashed_key)
return failure
def put(self, hashed_key: str, failure: CachedRedemptionFailure) -> None:
if hashed_key in self._entries:
del self._entries[hashed_key]
elif len(self._entries) >= self._max_entries:
self._entries.popitem(last=False)
self._entries[hashed_key] = (self._clock() + self._ttl_seconds, failure)
def discard(self, hashed_key: str) -> None:
self._entries.pop(hashed_key, None)
def clear(self) -> None:
self._entries.clear()
def __len__(self) -> int:
return len(self._entries)
# Process-wide singleton used by the bearer-auth path.
redemption_negative_cache = RedemptionNegativeCache()
+613
View File
@@ -0,0 +1,613 @@
"""Refund claims: one open payout per API key, recorded before it is paid."""
import asyncio
import time
from typing import Any
import httpx
from fastapi import HTTPException
from sqlalchemy.exc import IntegrityError
from sqlmodel import col, func, select, update
from .core.db import (
REFUND_OPEN_STATUSES,
REFUND_UNRESOLVED_STATUSES,
ApiKey,
AsyncSession,
Refund,
create_session,
)
from .core.db import (
store_cashu_transaction_with_retry as store_cashu_transaction,
)
from .core.logging import get_logger
from .core.settings import settings
from .payment.lnurl import (
LNURLError,
MeltOutcomeAmbiguousError,
MeltUnpaidError,
get_lnurl_data,
)
from .wallet import (
check_bolt11_payment_status,
is_mint_connection_error,
send_to_lnurl,
send_token,
token_mint_url,
)
logger = get_logger(__name__)
RECONCILE_BATCH_LIMIT = 100
def refund_unit(key: ApiKey) -> str:
return key.refund_currency or "sat"
def amount_in_unit(amount_msats: int, unit: str) -> int:
return amount_msats // 1000 if unit == "sat" else amount_msats
def refund_mint(key: ApiKey) -> str:
if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints:
return key.refund_mint_url
return settings.primary_mint
async def validate_lightning_destination(destination: str) -> None:
try:
await get_lnurl_data(destination)
except LNURLError as e:
raise HTTPException(
status_code=400, detail=f"Invalid lightning destination: {e}"
)
except httpx.HTTPError as e:
raise HTTPException(
status_code=400, detail=f"Lightning destination unreachable: {e}"
)
async def open_claim(
session: AsyncSession,
key: ApiKey,
*,
method: str,
destination: str | None,
) -> Refund:
"""Zero the balance and insert the claim in one transaction."""
unit = refund_unit(key)
# created_at has second resolution; step past the previous claim so the
# newest claim for a key always sorts first.
latest = await session.exec(
select(func.max(col(Refund.created_at))).where(
Refund.api_key_hashed_key == key.hashed_key
)
)
created_at = max(int(time.time()), (latest.one() or 0) + 1)
refund = Refund(
api_key_hashed_key=key.hashed_key,
method=method,
destination=destination,
amount_msats=key.total_balance,
unit=unit,
mint_url=refund_mint(key),
claimed_at=int(time.time()),
created_at=created_at,
updated_at=created_at,
)
debit = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.balance) == key.balance)
.where(col(ApiKey.reserved_balance) == key.reserved_balance)
.values(balance=0, reserved_balance=0, reserved_at=None)
)
try:
debited = await session.exec(debit) # type: ignore[call-overload]
if debited.rowcount == 0:
await session.rollback()
raise HTTPException(
status_code=409,
detail="Balance changed concurrently. Please retry the refund.",
)
session.add(refund)
await session.commit()
except IntegrityError:
await session.rollback()
raise refund_in_progress_error()
return refund
async def _close(
session: AsyncSession,
refund: Refund,
*,
require_no_quote: bool = False,
require_no_token: bool = False,
from_statuses: tuple[str, ...] = REFUND_OPEN_STATUSES,
**values: object,
) -> bool:
stmt = (
update(Refund)
.where(col(Refund.id) == refund.id)
.where(col(Refund.status).in_(from_statuses))
)
if require_no_quote:
# A quote recorded since the row was read means a melt may be in flight.
stmt = stmt.where(col(Refund.quote_id).is_(None))
if require_no_token:
stmt = stmt.where(col(Refund.token).is_(None))
result = await session.exec( # type: ignore[call-overload]
stmt.values(claimed_at=None, updated_at=int(time.time()), **values)
)
return bool(result.rowcount)
async def renew_lease(session: AsyncSession, refund: Refund) -> None:
"""Push the reconciler lease forward before a slow mint step."""
await session.exec( # type: ignore[call-overload]
update(Refund)
.where(col(Refund.id) == refund.id)
.where(col(Refund.status).in_(REFUND_OPEN_STATUSES))
.values(claimed_at=int(time.time()))
)
await session.commit()
async def record_quote(refund: Refund, quote_id: str, mint_url: str) -> None:
"""Store the quote and its mint before the melt is sent; raises if the claim closed.
Mint fallback can issue the quote on a different mint than the claim's.
"""
async with create_session() as session:
result = await session.exec( # type: ignore[call-overload]
update(Refund)
.where(col(Refund.id) == refund.id)
.where(col(Refund.status).in_(REFUND_OPEN_STATUSES))
# Renew the lease so the reconciler leaves the payout alone.
.values(
quote_id=quote_id,
mint_url=mint_url,
claimed_at=int(time.time()),
updated_at=int(time.time()),
)
)
await session.commit()
if not result.rowcount:
raise LNURLError("Refund claim closed before the melt was dispatched")
refund.quote_id = quote_id
refund.mint_url = mint_url
async def settle(
session: AsyncSession,
refund: Refund,
*,
quote_id: str | None = None,
token: str | None = None,
mint_url: str | None = None,
) -> bool:
values: dict[str, Any] = {"status": "paid"}
if quote_id is not None:
values["quote_id"] = quote_id
if token is not None:
values["token"] = token
if mint_url is not None:
values["mint_url"] = mint_url
# The payout side knows the money moved, so a claim the reconciler gave up
# on (stuck) is closed as paid too.
settled = await _close(
session, refund, from_statuses=REFUND_UNRESOLVED_STATUSES, **values
)
await session.commit()
if not settled:
logger.warning(
"refund paid but its claim was already closed",
extra={"refund_id": refund.id, "prior_status": refund.status},
)
return settled
async def release(
session: AsyncSession, refund: Refund, *, require_no_quote: bool = False
) -> bool:
"""Mark the claim failed and restore the balance."""
if not await _close(
session, refund, require_no_quote=require_no_quote, status="failed"
):
# Commit, not rollback: rollback after an ORM UPDATE breaks later loads.
await session.commit()
return False
await session.exec( # type: ignore[call-overload]
update(ApiKey)
.where(col(ApiKey.hashed_key) == refund.api_key_hashed_key)
.values(balance=col(ApiKey.balance) + refund.amount_msats)
)
await session.commit()
logger.info(
"refund released; balance restored",
extra={
"refund_id": refund.id,
"key_hash": refund.api_key_hashed_key[:8],
"restored_msats": refund.amount_msats,
},
)
return True
async def hold(
session: AsyncSession,
refund: Refund,
quote_id: str | None,
*,
token: str | None = None,
) -> None:
"""Withhold the balance; the quote or token names what the mint may have paid."""
values: dict[str, Any] = {"status": "ambiguous", "quote_id": quote_id}
if token is not None:
values["token"] = token
values["mint_url"] = refund.mint_url
await _close(session, refund, **values)
await session.commit()
def refund_in_progress_error(refund: Refund | None = None) -> HTTPException:
"""The 409 raised when a key already has an unresolved refund claim."""
stuck = refund is not None and refund.status == "stuck"
error: dict[str, str] = {
"message": (
"A refund for this key is unresolved and requires operator reconciliation."
if stuck
else "A refund for this key is already in progress."
),
"type": "invalid_request_error",
"code": "refund_unresolved" if stuck else "refund_in_progress",
}
if refund is not None:
error["refund_id"] = refund.id
error["status"] = refund.status
return HTTPException(status_code=409, detail={"error": error})
async def _latest_with_status(
session: AsyncSession, key: ApiKey, statuses: tuple[str, ...]
) -> Refund | None:
result = await session.exec(
select(Refund)
.where(Refund.api_key_hashed_key == key.hashed_key)
.where(col(Refund.status).in_(statuses))
.order_by(col(Refund.created_at).desc(), col(Refund.updated_at).desc())
)
return result.first()
async def latest_open(session: AsyncSession, key: ApiKey) -> Refund | None:
"""Latest non-terminal (in-flight) claim for the key, if any.
An open claim means a prior refund already debited the balance and is still
settling, so the balance reads as zero even though a refund is under way.
"""
return await _latest_with_status(session, key, REFUND_OPEN_STATUSES)
async def latest_stuck(session: AsyncSession, key: ApiKey) -> Refund | None:
"""Latest claim the reconciler gave up on; needs operator recovery."""
return await _latest_with_status(session, key, ("stuck",))
async def latest_terminal(session: AsyncSession, key: ApiKey) -> Refund | None:
"""Latest paid refund of either method.
Cashu tokens are normally served from cashu_transactions, which tracks
collection and sweeping; the claim row is the fallback when that ledger
write failed after the token was already issued.
"""
return await _latest_with_status(session, key, ("paid",))
def describe(refund: Refund) -> dict[str, str]:
body: dict[str, str] = {"refund_id": refund.id, "status": refund.status}
if refund.token:
body["token"] = refund.token
if refund.destination:
body["recipient"] = refund.destination
if refund.unit == "sat":
body["sats"] = str(refund.amount_msats // 1000)
else:
body["msats"] = str(refund.amount_msats)
return body
async def _pay_lightning(session: AsyncSession, refund: Refund) -> bool:
async def capture_quote(quote: str, mint_url: str) -> None:
await record_quote(refund, quote, mint_url)
try:
await send_to_lnurl(
amount_in_unit(refund.amount_msats, refund.unit),
refund.unit,
refund.mint_url,
str(refund.destination),
on_melt_quote=capture_quote,
)
except MeltOutcomeAmbiguousError as e:
await hold(session, refund, refund.quote_id)
logger.error(
"refund outcome ambiguous; balance withheld pending reconciliation",
extra={
"refund_id": refund.id,
"error": str(e),
"key_hash": refund.api_key_hashed_key[:8],
"quote_id": refund.quote_id,
},
)
raise
return await settle(session, refund, quote_id=refund.quote_id)
async def _pay_cashu(session: AsyncSession, refund: Refund) -> bool:
amount = amount_in_unit(refund.amount_msats, refund.unit)
await renew_lease(session, refund)
token = await send_token(amount, refund.unit, refund.mint_url)
# From here the token is bearer money: keep it on the claim so a failed
# settle withholds the balance instead of restoring it.
refund.token = token
refund.mint_url = token_mint_url(token, refund.mint_url)
return await settle(session, refund, token=token, mint_url=refund.mint_url)
async def _record_cashu_payout(refund: Refund) -> None:
"""Ledger write for an issued token; the claim row already holds the token,
so a failure here must not fail the request or release the balance."""
try:
await store_cashu_transaction(
token=str(refund.token),
amount=amount_in_unit(refund.amount_msats, refund.unit),
unit=refund.unit,
mint_url=refund.mint_url,
typ="out",
collected=False,
source="apikey",
api_key_hashed_key=refund.api_key_hashed_key,
)
except Exception as e:
logger.error(
"refund token issued but cashu transaction was not recorded",
extra={
"refund_id": refund.id,
"error": str(e),
"error_type": type(e).__name__,
"key_hash": refund.api_key_hashed_key[:8],
},
)
def unresolved_refund_error() -> HTTPException:
return HTTPException(
status_code=502,
detail=(
"Refund was dispatched but its outcome is unconfirmed; the "
"balance is withheld until reconciliation completes"
),
)
async def _abort(session: AsyncSession, refund: Refund) -> None:
"""Fail the claim, or withhold it once a melt quote or token exists.
A recorded quote means the mint may already have paid; an issued token is
already bearer money. In both cases the balance must not be restored.
"""
if refund.quote_id is None and refund.token is None:
await release(session, refund)
return
await hold(session, refund, refund.quote_id, token=refund.token)
logger.error(
"refund failed after its payout was dispatched; balance withheld "
"pending reconciliation",
extra={
"refund_id": refund.id,
"key_hash": refund.api_key_hashed_key[:8],
"quote_id": refund.quote_id,
"has_token": refund.token is not None,
"mint_url": refund.mint_url,
},
)
raise unresolved_refund_error()
async def execute(session: AsyncSession, refund: Refund) -> dict[str, str]:
attached_refund = refund
# Keep payout evidence outside the identity map: a failed flush/commit can
# expire attached attributes, including the only copy of an issued token.
refund = Refund(**refund.model_dump())
try:
if refund.method == "lightning":
settled = await _pay_lightning(session, refund)
else:
settled = await _pay_cashu(session, refund)
except MeltOutcomeAmbiguousError:
# Already held by _pay_lightning; releasing here would pay out twice.
raise unresolved_refund_error()
except MeltUnpaidError as e:
# The mint answered the melt itself with unpaid: proof that nothing was
# sent, so the balance goes back now rather than after reconciliation.
await release(session, refund)
logger.warning(
"refund melt unpaid at the mint; balance restored",
extra={
"refund_id": refund.id,
"error": str(e),
"key_hash": refund.api_key_hashed_key[:8],
"quote_id": refund.quote_id,
},
)
raise HTTPException(
status_code=503,
detail="Lightning payment failed at the mint; balance restored. Retry later.",
)
except HTTPException:
await session.rollback()
await _abort(session, refund)
raise
except Exception as e:
await session.rollback()
await _abort(session, refund)
logger.error(
"refund payout failed",
extra={
"refund_id": refund.id,
"error": str(e),
"error_type": type(e).__name__,
"key_hash": refund.api_key_hashed_key[:8],
"method": refund.method,
"mint_url": refund.mint_url,
},
)
if is_mint_connection_error(e):
raise HTTPException(status_code=503, detail="Mint service unavailable")
raise HTTPException(status_code=500, detail="Refund failed")
if refund.method == "cashu":
await _record_cashu_payout(refund)
if settled:
refund.status = "paid"
refund.claimed_at = None
else:
# Report the row as it stands rather than a status that was not written.
await session.refresh(attached_refund)
refund = attached_refund
logger.info(
"refund paid",
extra={
"refund_id": refund.id,
"method": refund.method,
"amount_msats": refund.amount_msats,
"key_hash": refund.api_key_hashed_key[:8],
},
)
return describe(refund)
async def _lease(refund_id: str, now: int, lease_cutoff: int) -> bool:
async with create_session() as session:
result = await session.exec( # type: ignore[call-overload]
update(Refund)
.where(col(Refund.id) == refund_id)
.where(col(Refund.status).in_(REFUND_OPEN_STATUSES))
.where(
col(Refund.claimed_at).is_(None)
| (col(Refund.claimed_at) < lease_cutoff)
)
.values(claimed_at=now)
)
await session.commit()
return bool(result.rowcount)
async def _reconcile(refund: Refund, now: int) -> None:
if refund.method != "lightning":
if refund.token is not None:
# The token was issued and kept on the claim; the payout is done.
async with create_session() as session:
await settle(session, refund)
return
# No quote to query for cashu; withhold the balance and alert once.
async with create_session() as session:
if await _close(session, refund, require_no_token=True, status="stuck"):
await session.commit()
logger.critical(
"cashu refund stuck; balance withheld, manual reconciliation required",
extra={
"refund_id": refund.id,
"key_hash": refund.api_key_hashed_key[:8],
"amount_msats": refund.amount_msats,
},
)
else:
await session.commit()
current = await session.get(Refund, refund.id)
if current is not None and current.token is not None:
await settle(session, current)
return
if refund.quote_id is None:
# Never sent, unless a quote appeared since the row was read.
async with create_session() as session:
if not await release(session, refund, require_no_quote=True):
logger.info(
"refund gained a melt quote during reconciliation; left open",
extra={"refund_id": refund.id},
)
return
status = await check_bolt11_payment_status(
refund.mint_url, refund.unit, refund.quote_id
)
if status == "paid":
async with create_session() as session:
await settle(session, refund)
elif status == "unpaid":
# A fresh melt can report unpaid briefly; trust it only after a timeout.
if refund.updated_at > now - settings.refund_claim_timeout_seconds:
logger.info(
"refund unpaid at the mint but too recent to release; waiting",
extra={"refund_id": refund.id, "updated_at": refund.updated_at},
)
return
async with create_session() as session:
await release(session, refund)
else:
logger.warning(
"refund still unresolved at the mint",
extra={"refund_id": refund.id, "melt_status": status},
)
async def reconcile_once() -> None:
"""Resolve open claims whose lease has lapsed."""
now = int(time.time())
lease_cutoff = now - settings.refund_claim_timeout_seconds
async with create_session() as session:
result = await session.exec(
select(Refund)
.where(col(Refund.status).in_(REFUND_OPEN_STATUSES))
.where(
col(Refund.claimed_at).is_(None)
| (col(Refund.claimed_at) < lease_cutoff)
)
.order_by(col(Refund.created_at))
.limit(RECONCILE_BATCH_LIMIT)
)
stale = list(result.all())
for refund in stale:
if not await _lease(refund.id, now, lease_cutoff):
continue
try:
await _reconcile(refund, now)
except Exception as e:
logger.error(
"refund reconciliation failed",
extra={
"refund_id": refund.id,
"error": str(e),
"error_type": type(e).__name__,
},
)
async def periodic_refund_reconcile() -> None:
while True:
await asyncio.sleep(settings.refund_reconcile_interval_seconds)
try:
await reconcile_once()
except asyncio.CancelledError:
raise
except Exception as e:
logger.error(
"refund reconcile loop error",
extra={"error": str(e), "error_type": type(e).__name__},
)
+696 -71
View File
@@ -15,9 +15,6 @@ from ..core.db import (
UpstreamProviderRow,
create_session,
)
from ..core.db import (
store_cashu_transaction_with_retry as store_cashu_transaction,
)
from ..payment.price import sats_usd_price
from ..wallet import (
Bolt11PaymentAmbiguous,
@@ -27,7 +24,7 @@ from ..wallet import (
maximum_owner_cashu_balance_sats,
prepare_bolt11_payment,
release_token_reservation,
send_token,
send_token_from_owner_locked,
token_mint_url,
wallet_operation_guard,
)
@@ -49,6 +46,7 @@ PPQ_PHASES = frozenset({PPQ_PHASE_CLAIMED, PPQ_PHASE_IN_FLIGHT, PPQ_PHASE_RECONC
PPQ_SETTLEMENT_ATTEMPTS = 5
PPQ_SETTLEMENT_POLL_SECONDS = 2
PPQ_PENDING_TTL_SECONDS = 15 * 60
PPQ_SETTLED_COOLDOWN_SECONDS = 5 * 60
PPQ_MAX_INVOICE_PREMIUM = 1.10
PPQ_MIN_TOPUP_USD = 1
PPQ_MAX_TOPUP_USD = 500
@@ -58,6 +56,33 @@ PPQ_MAX_TOPUP_USD = 500
# the damage instead of letting the worker drain the owner's mint funds one
# per-transaction-capped payment at a time.
PPQ_MAX_DAILY_TOPUP_USD = 300
# Routstr-to-Routstr claim lifecycle. "claimed" holds the slot while the token
# is being minted; nothing has left the wallet yet. "sent" means a bearer token
# was handed to the peer and only the peer's balance can say whether it landed.
# "backoff" holds the failure count between attempts, and "halted" stops the
# provider entirely until an admin releases it.
ROUTSTR_PHASE_CLAIMED = "claimed"
ROUTSTR_PHASE_SENT = "sent"
ROUTSTR_PHASE_BACKOFF = "backoff"
ROUTSTR_PHASE_HALTED = "halted"
ROUTSTR_PHASES = frozenset(
{
ROUTSTR_PHASE_CLAIMED,
ROUTSTR_PHASE_SENT,
ROUTSTR_PHASE_BACKOFF,
ROUTSTR_PHASE_HALTED,
}
)
ROUTSTR_PENDING_TTL_SECONDS = 15 * 60
ROUTSTR_BACKOFF_BASE_SECONDS = 15 * 60
ROUTSTR_MAX_TOPUP_FAILURES = 3
ROUTSTR_MIN_TOPUP_SATS = 1
ROUTSTR_MAX_TOPUP_SATS = 1_000_000
# Rolling 24h ceiling on total Routstr auto top-up spend across all peers. The
# per-attempt claim already stops a peer from being paid twice for the same
# uncredited token; this bounds the total even when every attempt is credited
# and the peer simply keeps reporting a below-threshold balance.
ROUTSTR_MAX_DAILY_TOPUP_SATS = 2_000_000
async def periodic_auto_topup() -> None:
@@ -72,7 +97,7 @@ async def periodic_auto_topup() -> None:
except Exception as e:
logger.error(
"Auto top-up cycle failed",
extra={"error": str(e), "error_type": type(e).__name__},
extra={"error": repr(e), "error_type": type(e).__name__},
)
await asyncio.sleep(AUTO_TOPUP_INTERVAL_SECONDS)
@@ -103,7 +128,7 @@ async def _run_auto_topup_cycle() -> None:
extra={
"provider_id": row.id,
"base_url": row.base_url,
"error": str(e),
"error": repr(e),
"error_type": type(e).__name__,
},
)
@@ -134,12 +159,16 @@ async def _reconcile_all_ppq_claims() -> set[int]:
except Exception as e:
logger.error(
"PPQ claim reconciliation failed",
extra={"provider_id": row.id, "error": str(e)},
extra={
"provider_id": row.id,
"error": repr(e),
"error_type": type(e).__name__,
},
)
return active_provider_ids
def _invalid_ppq_number(value: object, *, integer: bool = False) -> bool:
def _invalid_topup_number(value: object, *, integer: bool = False) -> bool:
if isinstance(value, bool) or not isinstance(value, (int, float)):
return True
try:
@@ -160,9 +189,9 @@ def validate_ppq_auto_topup_settings(settings: dict | None) -> str | None:
threshold = settings.get("topup_threshold")
amount = settings.get("topup_amount_limit")
if _invalid_ppq_number(threshold):
if _invalid_topup_number(threshold):
return "PPQ auto top-up threshold must be a positive number"
if _invalid_ppq_number(amount, integer=True):
if _invalid_topup_number(amount, integer=True):
return "PPQ auto top-up amount must be a positive whole number"
amount_usd = int(typing.cast(int | float, amount))
if not PPQ_MIN_TOPUP_USD <= amount_usd <= PPQ_MAX_TOPUP_USD:
@@ -173,6 +202,62 @@ def validate_ppq_auto_topup_settings(settings: dict | None) -> str | None:
return None
_legacy_threshold_hinted: set[int] = set()
def _routstr_threshold_sats(row: UpstreamProviderRow, settings: dict) -> float:
"""Balance, in sats, below which a top-up fires.
``topup_threshold_sats`` is used as written. A legacy ``topup_threshold``
keeps the thousandfold it has always been compared with — reinterpreting it
as sats would drop an operator's trigger point by a factor of 1000 on
upgrade and leave the peer to run dry. The hint names the value to migrate
to, once per provider rather than once per scheduler tick.
"""
explicit = settings.get("topup_threshold_sats")
if explicit is not None:
return float(typing.cast(int | float, explicit))
threshold_sats = float(typing.cast(int | float, settings["topup_threshold"])) * 1000
if row.id is not None and row.id not in _legacy_threshold_hinted:
_legacy_threshold_hinted.add(row.id)
logger.warning(
"Routstr auto top-up uses the legacy unitless threshold; set "
"topup_threshold_sats to state the unit",
extra={"provider_id": row.id, "threshold_sats": threshold_sats},
)
return threshold_sats
def validate_routstr_auto_topup_settings(settings: dict | None) -> str | None:
"""Return why enabled Routstr auto top-up settings are invalid, if anything."""
if not settings or not settings.get("auto_topup"):
return None
threshold_sats = settings.get("topup_threshold_sats")
legacy_threshold = settings.get("topup_threshold")
amount = settings.get("topup_amount_limit")
mint_url = settings.get("topup_mint_url")
if threshold_sats is not None:
if _invalid_topup_number(threshold_sats):
return "Routstr auto top-up threshold must be a positive number of sats"
elif legacy_threshold is None:
return "Routstr auto top-up requires a threshold"
elif _invalid_topup_number(legacy_threshold):
return "Routstr auto top-up threshold must be a positive number"
if _invalid_topup_number(amount, integer=True):
return "Routstr auto top-up amount must be a positive whole number"
amount_sats = int(typing.cast(int | float, amount))
if not ROUTSTR_MIN_TOPUP_SATS <= amount_sats <= ROUTSTR_MAX_TOPUP_SATS:
return (
f"Routstr auto top-up amount must be between {ROUTSTR_MIN_TOPUP_SATS} "
f"and {ROUTSTR_MAX_TOPUP_SATS} sats"
)
if not isinstance(mint_url, str) or not mint_url.strip():
return "Routstr auto top-up requires a mint URL"
return None
async def _check_and_topup_ppq_from_row(row: UpstreamProviderRow) -> None:
settings: dict = {}
if row.provider_settings:
@@ -212,22 +297,18 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
if not settings.get("auto_topup"):
return
threshold = settings.get("topup_threshold")
amount = settings.get("topup_amount_limit")
mint_url = settings.get("topup_mint_url")
if not threshold or not amount or not mint_url:
problem = validate_routstr_auto_topup_settings(settings)
if problem is not None:
logger.warning(
"Auto top-up enabled but missing configuration",
extra={
"provider_id": row.id,
"has_threshold": bool(threshold),
"has_amount": bool(amount),
"has_mint": bool(mint_url),
},
"Auto top-up enabled but its configuration is invalid",
extra={"provider_id": row.id, "problem": problem},
)
return
threshold_sats = _routstr_threshold_sats(row, settings)
amount = int(settings["topup_amount_limit"])
mint_url = str(settings["topup_mint_url"])
if not row.api_key:
return
@@ -235,16 +316,40 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
provider = RoutstrUpstreamProvider.from_db_row(row)
if provider is None:
return
if await _reconcile_routstr_state(row, provider):
return
balance = await provider.get_balance()
if balance is None:
if balance is None or not math.isfinite(balance) or balance < 0:
logger.warning(
"Could not fetch balance for auto top-up",
extra={"provider_id": row.id, "base_url": row.base_url},
)
return
if balance >= threshold * 1000:
if balance >= threshold_sats:
return
spent_24h_sats = await _routstr_spent_last_24h_sats()
if spent_24h_sats + amount > ROUTSTR_MAX_DAILY_TOPUP_SATS:
logger.critical(
"Auto top-up skipped: rolling 24h spend cap reached",
extra={
"provider_id": row.id,
"spent_24h_sats": spent_24h_sats,
"topup_amount": amount,
"daily_cap_sats": ROUTSTR_MAX_DAILY_TOPUP_SATS,
},
)
return
# The balance the peer must report before another token may be sent. Any
# shortfall is treated as "not credited": the token is a bearer instrument
# and a peer that took one without crediting it must not be handed another.
expected_sats = math.floor(balance) + amount
operation_id = await _claim_routstr_topup(row, expected_sats=expected_sats)
if operation_id is None:
return
# Balance is below threshold - create token and top up
@@ -253,58 +358,64 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
extra={
"provider_id": row.id,
"balance": balance,
"threshold": threshold,
"threshold_sats": threshold_sats,
"topup_amount": amount,
"mint_url": mint_url,
},
)
try:
token = await send_token(amount, "sat", mint_url)
async with wallet_operation_guard():
# Keep the spend cap and audit mutation in one wallet lock.
spent_24h_sats = await _routstr_spent_last_24h_sats()
if spent_24h_sats + amount > ROUTSTR_MAX_DAILY_TOPUP_SATS:
raise ValueError("Routstr auto top-up daily spend cap reached")
token = await send_token_from_owner_locked(amount, "sat", mint_url)
actual_mint_url = token_mint_url(token, mint_url)
try:
await _persist_routstr_token_and_mark_sent(
row,
operation_id,
expected_sats=expected_sats,
token=token,
amount=amount,
mint_url=actual_mint_url,
)
except Exception:
logger.critical(
"Aborting auto top-up because its token and sent claim "
"could not be persisted atomically",
extra={"provider_id": row.id, "mint_url": actual_mint_url},
)
try:
await release_token_reservation(token)
except Exception as error:
logger.critical(
"Failed to release untracked auto-topup token",
extra={
"provider_id": row.id,
"mint_url": actual_mint_url,
"error": repr(error),
},
)
else:
logger.warning(
"Auto-topup token was released after persistence failed",
extra={"provider_id": row.id, "mint_url": actual_mint_url},
)
raise
except Exception as e:
logger.error(
"Failed to create cashu token for auto top-up",
logger.warning(
"Failed to create or persist cashu token for auto top-up",
extra={
"provider_id": row.id,
"amount": amount,
"mint_url": mint_url,
"error": str(e),
"error": repr(e),
"error_type": type(e).__name__,
},
)
return
actual_mint_url = token_mint_url(token, mint_url)
try:
await store_cashu_transaction(
token=token,
amount=amount,
unit="sat",
mint_url=actual_mint_url,
typ="out",
collected=False,
source="auto_topup",
)
except Exception:
logger.critical(
"Aborting auto top-up because its cashu token could not be persisted",
extra={"provider_id": row.id, "mint_url": actual_mint_url},
)
try:
await release_token_reservation(token)
except Exception as error:
logger.critical(
"Failed to release untracked auto-topup token",
extra={
"provider_id": row.id,
"mint_url": actual_mint_url,
"error": str(error),
},
)
else:
logger.warning(
"Auto-topup token was released after persistence failed",
extra={"provider_id": row.id, "mint_url": actual_mint_url},
)
await _release_routstr_claim(row, operation_id)
return
result = await provider.topup(token)
@@ -348,6 +459,481 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
)
def _routstr_state_id(row: UpstreamProviderRow) -> str:
if row.id is None:
raise ValueError("Routstr auto top-up requires a persisted provider row")
return _routstr_state_id_for_provider(row.id)
def _routstr_state_id_for_provider(provider_id: int | str) -> str:
return f"routstr-auto-topup-{provider_id}"
class RoutstrClaim(typing.NamedTuple):
operation_id: str
# Worker lease for "claimed"/"sent", retry-not-before for "backoff", and
# meaningless for "halted".
deadline: int
phase: str
# Peer balance in sats that proves this attempt was credited.
expected_sats: int
failures: int
def _routstr_request_id(
operation_id: str,
deadline: int,
phase: str,
expected_sats: int,
failures: int,
) -> str:
return f"routstr:{operation_id}:{deadline}:{phase}:{expected_sats}:{failures}"
def _parse_routstr_request_id(request_id: str | None) -> RoutstrClaim | None:
parts = (request_id or "").split(":", 5)
if len(parts) != 6 or parts[0] != "routstr" or parts[3] not in ROUTSTR_PHASES:
return None
try:
deadline = int(parts[2])
expected_sats = int(parts[4])
failures = int(parts[5])
except (TypeError, ValueError):
return None
return RoutstrClaim(parts[1], deadline, parts[3], expected_sats, failures)
async def _routstr_spent_last_24h_sats() -> int:
"""Total sats committed to Routstr auto top-ups in the last 24 hours.
Uncollected rows count too: a token whose delivery is unconfirmed is spent
for capping purposes. Rows marked ``collected=False, swept=True`` record
tokens that were provably returned to the wallet and are excluded.
"""
cutoff = int(time.time()) - 24 * 60 * 60
async with create_session() as session:
rows = (
await session.exec(
select(CashuTransaction.amount, CashuTransaction.unit).where(
col(CashuTransaction.source) == "auto_topup",
col(CashuTransaction.type) == "out",
col(CashuTransaction.created_at) >= cutoff,
or_(
col(CashuTransaction.collected) == True, # noqa: E712
col(CashuTransaction.swept) == False, # noqa: E712
),
)
)
).all()
return sum(
amount if unit == "sat" else math.ceil(amount / 1000) for amount, unit in rows
)
async def _routstr_provider_is_claimable(
session: AsyncSession, provider_id: int | str | None
) -> bool:
"""Re-read the provider inside the claim transaction.
Same reasoning as :func:`_ppq_provider_is_claimable`: a provider deleted or
retyped concurrently must either be visible here or lose the race against
the claim we are about to write.
"""
if provider_id is None:
return False
current = await session.get(UpstreamProviderRow, provider_id)
return current is not None and current.provider_type == "routstr"
async def _claim_routstr_topup(
row: UpstreamProviderRow, *, expected_sats: int
) -> str | None:
"""Acquire the provider's single durable auto top-up slot.
An expired backoff hands its failure count to the new attempt, so repeated
non-crediting peers still walk towards the halt instead of resetting the
counter every cycle.
"""
state_id = _routstr_state_id(row)
operation_id = uuid.uuid4().hex
deadline = int(time.time()) + ROUTSTR_PENDING_TTL_SECONDS
async with create_session() as session:
if not await _routstr_provider_is_claimable(session, row.id):
return None
existing = await session.get(CashuTransaction, state_id)
if existing is not None:
failures = 0
if not (existing.collected or existing.swept):
claim = _parse_routstr_request_id(existing.request_id)
if (
claim is None
or claim.phase != ROUTSTR_PHASE_BACKOFF
or time.time() < claim.deadline
):
return None
failures = claim.failures
result = await session.exec( # type: ignore[call-overload]
update(CashuTransaction)
.where(
col(CashuTransaction.id) == state_id,
# Fence on the exact row that was read: any concurrent
# writer that moved the claim must win instead of us.
col(CashuTransaction.request_id) == existing.request_id,
)
.values(
token="pending",
amount=0,
unit="sat",
mint_url=None,
request_id=_routstr_request_id(
operation_id,
deadline,
ROUTSTR_PHASE_CLAIMED,
expected_sats,
failures,
),
collected=False,
swept=False,
created_at=int(time.time()),
source="routstr_auto_topup_claim",
)
)
await session.commit()
if (getattr(result, "rowcount", 0) or 0) != 1:
return None
return operation_id
try:
async with create_session() as session:
if not await _routstr_provider_is_claimable(session, row.id):
return None
session.add(
CashuTransaction(
id=state_id,
token="pending",
amount=0,
unit="sat",
type="out",
request_id=_routstr_request_id(
operation_id,
deadline,
ROUTSTR_PHASE_CLAIMED,
expected_sats,
0,
),
collected=False,
source="routstr_auto_topup_claim",
)
)
await session.commit()
except IntegrityError:
return None
return operation_id
async def _advance_routstr_claim(
row: UpstreamProviderRow,
operation_id: str,
*,
deadline: int,
phase: str,
expected_sats: int,
failures: int,
token: str | None = None,
amount: int | None = None,
mint_url: str | None = None,
) -> bool:
"""Move this worker's claim to another phase, if it still owns it."""
values: dict[str, object] = {
"request_id": _routstr_request_id(
operation_id, deadline, phase, expected_sats, failures
)
}
if token is not None:
values.update(token=token, amount=amount, mint_url=mint_url)
async with create_session() as session:
result = await session.exec( # type: ignore[call-overload]
update(CashuTransaction)
.where(
col(CashuTransaction.id) == _routstr_state_id(row),
col(CashuTransaction.request_id).like(f"routstr:{operation_id}:%"),
col(CashuTransaction.collected) == False, # noqa: E712
col(CashuTransaction.swept) == False, # noqa: E712
)
.values(**values)
)
await session.commit()
return (getattr(result, "rowcount", 0) or 0) == 1
async def _set_routstr_state_terminal(
row: UpstreamProviderRow, operation_id: str, *, collected: bool, swept: bool
) -> bool:
"""Finish an attempt only if this worker still owns the claim."""
async with create_session() as session:
result = await session.exec( # type: ignore[call-overload]
update(CashuTransaction)
.where(
col(CashuTransaction.id) == _routstr_state_id(row),
col(CashuTransaction.request_id).like(f"routstr:{operation_id}:%"),
col(CashuTransaction.collected) == False, # noqa: E712
col(CashuTransaction.swept) == False, # noqa: E712
)
.values(collected=collected, swept=swept)
)
await session.commit()
return (getattr(result, "rowcount", 0) or 0) == 1
async def _release_routstr_claim(row: UpstreamProviderRow, operation_id: str) -> None:
"""Hand back a claim whose token never left the wallet."""
if not await _set_routstr_state_terminal(
row, operation_id, collected=False, swept=True
):
logger.warning(
"Could not release the auto top-up claim after a pre-send failure; "
"it is owned by another attempt",
extra={"provider_id": row.id},
)
async def _persist_routstr_token_and_mark_sent(
row: UpstreamProviderRow,
operation_id: str,
*,
expected_sats: int,
token: str,
amount: int,
mint_url: str,
) -> None:
state_id = _routstr_state_id(row)
async with create_session() as session:
state = await session.get(CashuTransaction, state_id)
claim = _parse_routstr_request_id(state.request_id if state else None)
if (
state is None
or state.collected
or state.swept
or claim is None
or claim.operation_id != operation_id
or claim.phase != ROUTSTR_PHASE_CLAIMED
):
raise RuntimeError("Routstr auto top-up claim ownership was lost")
result = await session.exec( # type: ignore[call-overload]
update(CashuTransaction)
.where(
col(CashuTransaction.id) == state_id,
col(CashuTransaction.request_id) == state.request_id,
col(CashuTransaction.collected) == False, # noqa: E712
col(CashuTransaction.swept) == False, # noqa: E712
)
.values(
request_id=_routstr_request_id(
operation_id,
int(time.time()) + ROUTSTR_PENDING_TTL_SECONDS,
ROUTSTR_PHASE_SENT,
expected_sats,
claim.failures,
),
token=token,
amount=amount,
unit="sat",
mint_url=mint_url,
)
)
if (getattr(result, "rowcount", 0) or 0) != 1:
await session.rollback()
raise RuntimeError("Routstr auto top-up claim ownership was lost")
session.add(
CashuTransaction(
id=uuid.uuid4().hex,
token=token,
amount=amount,
unit="sat",
mint_url=mint_url,
type="out",
collected=False,
source="auto_topup",
)
)
await session.commit()
async def _current_routstr_claim(row: UpstreamProviderRow) -> RoutstrClaim | None:
async with create_session() as session:
transaction = await session.get(CashuTransaction, _routstr_state_id(row))
if transaction is None:
return None
return _parse_routstr_request_id(transaction.request_id)
async def _reconcile_routstr_state(
row: UpstreamProviderRow, provider: RoutstrUpstreamProvider
) -> bool:
"""Return True while a prior attempt must suppress a new payment."""
async with create_session() as session:
transaction = await session.get(CashuTransaction, _routstr_state_id(row))
if transaction is None or transaction.collected or transaction.swept:
return False
claim = _parse_routstr_request_id(transaction.request_id)
if claim is None:
logger.critical(
"Malformed auto top-up state; suppressing duplicate payment",
extra={"provider_id": row.id},
)
return True
now = time.time()
if claim.phase == ROUTSTR_PHASE_HALTED:
return True
if claim.phase == ROUTSTR_PHASE_BACKOFF:
return now < claim.deadline
if claim.phase == ROUTSTR_PHASE_CLAIMED:
# Nothing left the wallet, so a dead worker's slot is free to reuse.
if now < claim.deadline:
return True
return not await _set_routstr_state_terminal(
row, claim.operation_id, collected=False, swept=True
)
balance = await provider.get_balance()
if (
balance is not None
and math.isfinite(balance)
and balance >= claim.expected_sats
):
if not await _set_routstr_state_terminal(
row, claim.operation_id, collected=True, swept=False
):
logger.critical(
"Auto top-up was credited but its claim was already released; "
"a duplicate top-up is possible on the next cycle",
extra={"provider_id": row.id},
)
return True
if now < claim.deadline:
return True
failures = claim.failures + 1
if failures >= ROUTSTR_MAX_TOPUP_FAILURES:
await _advance_routstr_claim(
row,
claim.operation_id,
deadline=claim.deadline,
phase=ROUTSTR_PHASE_HALTED,
expected_sats=claim.expected_sats,
failures=failures,
)
logger.critical(
"Auto top-up halted: the peer repeatedly failed to credit a token",
extra={
"provider_id": row.id,
"base_url": row.base_url,
"failures": failures,
"admin_action": (
f"POST /admin/api/upstream-providers/{row.id}"
"/routstr-auto-topup/release"
),
},
)
return True
await _advance_routstr_claim(
row,
claim.operation_id,
deadline=int(now) + ROUTSTR_BACKOFF_BASE_SECONDS * 2 ** (failures - 1),
phase=ROUTSTR_PHASE_BACKOFF,
expected_sats=claim.expected_sats,
failures=failures,
)
logger.warning(
"Auto top-up was not credited by the peer; backing off",
extra={
"provider_id": row.id,
"base_url": row.base_url,
"expected_sats": claim.expected_sats,
"failures": failures,
},
)
return True
async def get_routstr_auto_topup_state(provider_id: int) -> dict[str, object]:
"""Return admin-safe state for a provider's durable Routstr claim."""
async with create_session() as session:
transaction = await session.get(
CashuTransaction, _routstr_state_id_for_provider(provider_id)
)
if transaction is None or transaction.collected or transaction.swept:
return {"active": False}
claim = _parse_routstr_request_id(transaction.request_id)
return {
"active": True,
# Echoed back verbatim on release so a claim that moved on since the
# admin reviewed it fails the write instead of being swept unseen.
"state_token": transaction.request_id,
"operation_id": claim.operation_id if claim else None,
"phase": claim.phase if claim else None,
"expected_sats": claim.expected_sats if claim else None,
"failures": claim.failures if claim else None,
"deadline": claim.deadline if claim else None,
"created_at": transaction.created_at,
"amount": transaction.amount,
"unit": transaction.unit,
"mint_url": transaction.mint_url,
"malformed": claim is None,
}
class RoutstrReleaseOutcome(typing.NamedTuple):
released: bool
reason: str
async def release_routstr_auto_topup_state(
provider_id: int, *, state_token: str | None
) -> RoutstrReleaseOutcome:
"""Clear a halted or stuck claim after an admin reconciles the peer.
Unlike a Lightning melt there is no in-flight window to protect: the token
is already with the peer or still in the wallet either way. What the fence
does protect is the admin's decision — the row must be byte-identical to
the one they reviewed, so a claim that advanced in the meantime is not
swept on the strength of stale information.
"""
state_id = _routstr_state_id_for_provider(provider_id)
async with create_session() as session:
transaction = await session.get(CashuTransaction, state_id)
if transaction is None or transaction.collected or transaction.swept:
return RoutstrReleaseOutcome(False, "no_active_claim")
if transaction.request_id != state_token:
return RoutstrReleaseOutcome(False, "stale_state")
async with create_session() as session:
result = await session.exec( # type: ignore[call-overload]
update(CashuTransaction)
.where(
col(CashuTransaction.id) == state_id,
col(CashuTransaction.request_id) == state_token,
col(CashuTransaction.collected) == False, # noqa: E712
col(CashuTransaction.swept) == False, # noqa: E712
)
.values(swept=True)
)
if (getattr(result, "rowcount", 0) or 0) == 1:
await session.commit()
return RoutstrReleaseOutcome(True, "released")
await session.rollback()
return RoutstrReleaseOutcome(False, "claim_changed")
def _ppq_state_id(row: UpstreamProviderRow) -> str:
if row.id is None:
raise ValueError("PPQ auto top-up requires a persisted provider row")
@@ -528,7 +1114,11 @@ async def _set_ppq_state_terminal(
col(CashuTransaction.collected) == False, # noqa: E712
col(CashuTransaction.swept) == False, # noqa: E712
)
.values(collected=collected, swept=swept)
.values(
collected=collected,
swept=swept,
created_at=int(time.time()) if collected else CashuTransaction.created_at,
)
)
updated = (getattr(result, "rowcount", 0) or 0) == 1
if updated:
@@ -553,8 +1143,10 @@ async def _reconcile_ppq_state(
"""
async with create_session() as session:
transaction = await session.get(CashuTransaction, _ppq_state_id(row))
if transaction is None or transaction.collected or transaction.swept:
if transaction is None or transaction.swept:
return False
if transaction.collected:
return int(time.time()) - transaction.created_at < PPQ_SETTLED_COOLDOWN_SECONDS
claim = _parse_ppq_request_id(transaction.request_id)
if claim is None:
@@ -620,7 +1212,7 @@ async def _reconcile_ppq_state(
async def _ppq_provider_is_claimable(
session: AsyncSession, provider_id: int | None
session: AsyncSession, row: UpstreamProviderRow
) -> bool:
"""Re-read the provider inside the claim transaction.
@@ -631,10 +1223,16 @@ async def _ppq_provider_is_claimable(
this the worker could create a claim for a provider that no longer
exists, orphaning it forever.
"""
if provider_id is None:
if row.id is None:
return False
current = await session.get(UpstreamProviderRow, provider_id)
return current is not None and current.provider_type == "ppqai"
current = await session.get(UpstreamProviderRow, row.id)
return bool(
current is not None
and current.enabled
and current.provider_type == "ppqai"
and current.api_key == row.api_key
and current.provider_settings == row.provider_settings
)
async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
@@ -645,10 +1243,17 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
request_id = _ppq_request_id(operation_id, expires_at, PPQ_PHASE_CLAIMED, "pending")
async with create_session() as session:
if not await _ppq_provider_is_claimable(session, row.id):
if not await _ppq_provider_is_claimable(session, row):
return None
existing = await session.get(CashuTransaction, state_id)
if existing is not None:
if (
existing.collected
and not existing.swept
and int(time.time()) - existing.created_at
< PPQ_SETTLED_COOLDOWN_SECONDS
):
return None
result = await session.exec( # type: ignore[call-overload]
update(CashuTransaction)
.where(
@@ -679,7 +1284,7 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
async with create_session() as session:
# Same fencing as the update path: the provider must still exist
# inside the transaction that creates the claim.
if not await _ppq_provider_is_claimable(session, row.id):
if not await _ppq_provider_is_claimable(session, row):
return None
session.add(
CashuTransaction(
@@ -900,6 +1505,26 @@ async def _check_and_topup_ppq(row: UpstreamProviderRow, settings: dict) -> None
if balance >= threshold_usd:
return
# Require two low-balance reads before creating an invoice.
confirmed_balance = await provider.get_balance()
if (
confirmed_balance is None
or not math.isfinite(confirmed_balance)
or confirmed_balance < 0
or confirmed_balance >= threshold_usd
):
logger.info(
"PPQ auto top-up aborted by balance confirmation",
extra={
"provider_id": row.id,
"first_balance_usd": balance,
"confirmed_balance_usd": confirmed_balance,
"threshold_usd": threshold_usd,
},
)
return
balance = confirmed_balance
# Perform local pricing and owner-funds checks before asking PPQ to create
# an invoice. The exact mint quote still has to be checked afterward, but
# predictable local failures should not leave abandoned PPQ invoices.
+764 -468
View File
File diff suppressed because it is too large Load Diff
+230 -10
View File
@@ -23,7 +23,11 @@ import litellm
from fastapi.responses import Response
from ..core import get_logger
from ..payment.helpers import estimate_tokens
from ..payment.helpers import (
_count_prompt_token_ids,
estimate_prompt_tokens,
estimate_tokens,
)
from ..payment.models import Model
logger = get_logger(__name__)
@@ -39,11 +43,66 @@ def _parse_request_body(request_body: bytes | None) -> dict[str, Any]:
return parsed if isinstance(parsed, dict) else {}
def _count_with_litellm(model: str, body: dict[str, Any]) -> int:
def _model_name(model_obj: Model | None, body: dict[str, Any]) -> str:
if model_obj is not None:
return model_obj.forwarded_model_id or model_obj.id or ""
body_model = body.get("model")
return body_model if isinstance(body_model, str) else ""
def _count_with_litellm(
model: str, body: dict[str, Any], include_legacy_prompt: bool = False
) -> int:
messages = body.get("messages")
if not isinstance(messages, list):
messages = []
if "input" in body:
response_input = body["input"]
if isinstance(response_input, str):
messages = [{"role": "user", "content": response_input}]
elif isinstance(response_input, list):
messages = []
for item in response_input:
if not isinstance(item, dict) or "role" not in item:
raise ValueError(
"Responses input requires fallback token estimation"
)
content = item.get("content", "")
if isinstance(content, list):
parts = []
for part in content:
if not isinstance(part, dict) or part.get("type") not in (
"input_text",
"output_text",
"text",
):
raise ValueError(
"Non-text Responses input requires fallback token estimation"
)
parts.append({"type": "text", "text": part.get("text", "")})
content = parts
messages.append({"role": item["role"], "content": content})
else:
raise ValueError("Unsupported Responses input")
if body.get("instructions"):
messages.insert(0, {"role": "system", "content": body["instructions"]})
prompt_token_ids = 0
if include_legacy_prompt:
prompt = body.get("prompt")
if isinstance(prompt, str):
prompt_texts = [prompt]
elif isinstance(prompt, list):
prompt_texts = [item for item in prompt if isinstance(item, str)]
else:
prompt_texts = []
prompt_token_ids = _count_prompt_token_ids(prompt)
messages = [
*({"role": "user", "content": text} for text in prompt_texts if text),
*messages,
]
system = body.get("system")
if isinstance(system, str) and system:
messages = [{"role": "system", "content": system}, *messages]
@@ -58,7 +117,7 @@ def _count_with_litellm(model: str, body: dict[str, Any]) -> int:
tools = body.get("tools") if isinstance(body.get("tools"), list) else None
return int(
return prompt_token_ids + int(
litellm.token_counter(
model=model,
messages=messages,
@@ -67,6 +126,173 @@ def _count_with_litellm(model: str, body: dict[str, Any]) -> int:
)
def _count_text_with_litellm(model: str, text: str) -> int:
return int(
litellm.token_counter(
model=model,
text=text,
count_response_tokens=True,
)
)
def _generated_text(value: object) -> list[str]:
"""Extract generated text/tool arguments without counting response metadata."""
generated_keys = {
"arguments",
"content",
"delta",
"output_text",
"partial_json",
"reasoning",
"reasoning_content",
"text",
"thinking",
}
parts: list[str] = []
def walk(item: object, key: str | None = None) -> None:
if isinstance(item, str):
if key in generated_keys:
parts.append(item)
return
if isinstance(item, list):
for child in item:
walk(child, key)
return
if isinstance(item, dict):
for child_key, child in item.items():
walk(child, child_key)
walk(value)
return parts
class MissingUsageEstimator:
"""Estimate billable usage when an upstream omits its usage trailer.
The reservation is deliberately absent from this class: it is an
authorization ceiling, not an input to usage measurement.
"""
def __init__(self, request_body: bytes | None, model_obj: Model | None) -> None:
self.body = _parse_request_body(request_body)
self.model_name = _model_name(model_obj, self.body)
self._output_parts: list[str] = []
self._input_tokens: int | None = None
def _estimate_input_tokens(self) -> int:
if self._input_tokens is not None:
return self._input_tokens
try:
self._input_tokens = _count_with_litellm(
self.model_name, self.body, include_legacy_prompt=True
)
except Exception as exc:
self._input_tokens = estimate_prompt_tokens(self.body)
logger.debug(
"litellm request token count failed; using local estimator",
extra={
"model": self.model_name,
"error": str(exc),
"error_type": type(exc).__name__,
"estimated_tokens": self._input_tokens,
},
)
return self._input_tokens
@property
def output_text(self) -> str:
return "".join(self._output_parts)
def observe(self, response_data: object) -> None:
if isinstance(response_data, dict):
event_type = response_data.get("type")
if event_type in ("response.completed", "response.incomplete"):
response = response_data.get("response")
if isinstance(response, dict) and isinstance(
response.get("output"), list
):
# Terminal output is a snapshot, not another text delta.
self._output_parts = _generated_text(response["output"])
return
if isinstance(event_type, str) and event_type.endswith(".done"):
# Responses API ``*.done`` events repeat text already streamed
# via ``*.delta`` events; counting both would double-bill.
return
self._output_parts.extend(_generated_text(response_data))
def estimated_usage(self, model: str | None = None) -> dict[str, Any] | None:
"""Local usage estimate, or None when the upstream generated no text."""
if not self.output_text:
return None
return self.response_data(model)["usage"]
def billing_data(
self,
response_data: dict[str, Any] | None,
model: str | None = None,
) -> dict[str, Any]:
"""Use measured usage when present, otherwise return a local estimate."""
if isinstance(response_data, dict):
usage = response_data.get("usage")
if not isinstance(usage, dict):
nested = response_data.get("response")
usage = nested.get("usage") if isinstance(nested, dict) else None
if isinstance(usage, dict) and usage:
return {
"model": model or response_data.get("model") or self.model_name,
"usage": usage,
}
if not self._output_parts:
self.observe(response_data)
return self.response_data(model)
def response_data(self, model: str | None = None) -> dict[str, Any]:
text = self.output_text
try:
output_tokens = (
_count_text_with_litellm(self.model_name, text) if text else 0
)
except Exception as exc:
output_tokens = len(text) // 3
logger.debug(
"litellm response token count failed; using local estimator",
extra={
"model": self.model_name,
"error": str(exc),
"error_type": type(exc).__name__,
"estimated_tokens": output_tokens,
},
)
input_tokens = max(0, int(self._estimate_input_tokens()))
output_tokens = max(0, int(output_tokens))
return {
"model": model or self.model_name or "unknown",
"usage": {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": input_tokens + output_tokens,
"estimated": True,
},
}
def openai_response_data(self, model: str | None = None) -> dict[str, Any]:
"""Same estimate in the OpenAI chat-completions usage dialect."""
data = self.response_data(model)
usage = data["usage"]
return {
"model": data["model"],
"usage": {
"prompt_tokens": usage["input_tokens"],
"completion_tokens": usage["output_tokens"],
"total_tokens": usage["total_tokens"],
"estimated": True,
},
}
def count_tokens_locally(
request_body: bytes | None,
model_obj: Model | None,
@@ -75,13 +301,7 @@ def count_tokens_locally(
touching the upstream. Always returns 200; never raises."""
body = _parse_request_body(request_body)
model_name = ""
if model_obj is not None:
model_name = model_obj.forwarded_model_id or model_obj.id or ""
if not model_name:
body_model = body.get("model")
if isinstance(body_model, str):
model_name = body_model
model_name = _model_name(model_obj, body)
input_tokens: int
try:
+375 -240
View File
@@ -10,17 +10,17 @@ from urllib.parse import urlsplit, urlunsplit
from fastapi import Request
from fastapi.responses import Response, StreamingResponse
from sqlalchemy import case
from sqlmodel import col, update
from ..auth import (
ROUTSTR_FEE_PERCENT,
ReservationSnapshot,
_charge_reservation_rows,
_claim_reservation_for_charge,
_stop_reservation_heartbeat,
_validate_reservation_snapshot,
get_billing_key,
get_reservation_snapshot,
payments_logger,
release_reservation,
)
from ..core import get_logger
from ..core.db import (
@@ -31,7 +31,7 @@ from ..core.db import (
from ..core.db import (
store_cashu_transaction_with_retry as store_cashu_transaction,
)
from ..core.exceptions import UpstreamError
from ..core.exceptions import EhbpTimeoutError, UpstreamError
from ..core.settings import settings
from ..payment.cost_calculation import (
CostData,
@@ -40,8 +40,13 @@ from ..payment.cost_calculation import (
)
from ..payment.helpers import create_error_response
from ..payment.models import Model
from ..wallet import recieve_token, send_token
from .tinfoil_trailer import forward_with_trailer
from ..wallet import (
SPENT_TOKEN_CODES,
classify_redemption_error,
recieve_token,
send_token,
)
from .tinfoil_trailer import TrailerResponse, forward_with_trailer
logger = get_logger(__name__)
@@ -57,6 +62,55 @@ _TINFOIL_ALLOWED_ENCLAVE_HOST_SUFFIX = ".tinfoil.sh"
_TINFOIL_ALLOWED_ENCLAVE_HOSTS = frozenset({"tinfoil.sh"})
_KEY_CONFIG_PROBLEM_TYPE = "urn:ietf:params:ehbp:error:key-config"
def _is_ehbp_key_config_response(resp: TrailerResponse) -> bool:
"""Check whether an upstream EHBP response is a key-config mismatch.
The enclave returns ``422 application/problem+json`` with
``type=urn:ietf:params:ehbp:error:key-config`` when it cannot decrypt the
request body — meaning the client's HPKE key is stale (the enclave rotated
keys). The proxy must pass this response through with its original
content type so EHBP clients can detect it and trigger re-attestation.
"""
if resp.status_code != 422:
return False
ct = ""
for k, v in resp.headers:
if k.lower() == "content-type":
ct = v.lower()
break
media_type = ct.split(";", 1)[0].strip()
if media_type != "application/problem+json":
return False
try:
body = json.loads(resp.body)
except (json.JSONDecodeError, UnicodeDecodeError):
return False
return isinstance(body, dict) and body.get("type") == _KEY_CONFIG_PROBLEM_TYPE
def _passthrough_key_config_response(resp: TrailerResponse) -> Response:
"""Return the enclave's key-config 422 with its original body and content
type so the EHBP client's ``KeyConfigMismatchError`` detection fires.
Only the content type is forwarded. ``Ehbp-Response-Nonce`` must be
dropped: a nonce only carries meaning for an *encrypted* response body,
and the stock ``ehbp`` client (``shouldDecryptResponse``) checks for the
nonce *before* checking for a key-config mismatch — forwarding it would
send that client down the decrypt path on this plaintext error body, so
the re-attestation loop would never fire. Content-length is recomputed
from the body, and upstream-internal headers are filtered out.
"""
return Response(
content=resp.body,
status_code=422,
headers={"content-type": "application/problem+json"},
media_type="application/problem+json",
)
def _normalize_upstream_model_id(model_id: str | None) -> str:
"""Normalize casing and whitespace for upstream identity comparisons."""
if not model_id:
@@ -73,28 +127,43 @@ _PROXY_ONLY_HEADERS = frozenset(
}
)
# Namespace prefix the routstr catalog applies to Tinfoil models
# (e.g. ``tinfoil-deepseek-v4-1-flash``). The SDK strips this prefix for the
# encrypted body (``getTinfoilUpstreamModelId`` in client/TinfoilSecure.ts), so
# the enclave always reports the *bare* upstream model id in the usage-metrics
# header even though the routstr model id and ``forwarded_model_id`` carry it.
TINFOIL_MODEL_PREFIX = "tinfoil-"
def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None:
"""Parse ``X-Tinfoil-Usage-Metrics`` into an OpenAI-style usage dict.
The header format is::
prompt=<n>,completion=<n>,total=<n>[,model=<name>]
prompt=<n>,completion=<n>,total=<n>[,cached_prompt_tokens=<n>,
uncached_prompt_tokens=<n>][,model=<name>][,cost_usd=<usd>]
``prompt`` is the inclusive prompt total and ``cached_prompt_tokens`` is
the cache-read portion included within it. Routstr maps these to
``prompt_tokens`` and ``cache_read_input_tokens`` so ``normalize_usage``
can subtract the cached read from the prompt total (OpenAI-family
semantics). ``cost_usd`` is parsed as a float and kept for logging/
cross-checking only — billing uses the token path.
The ``model`` field (added in tinfoilsh/confidential-model-router PR #385)
is extracted as a string and included in the returned dict under the
``"model"`` key so callers can compare the served model against the
requested one and adjust pricing.
is extracted as a string so callers can compare the served model against
the requested one and adjust pricing.
Returns a dict like ``{"prompt_tokens": n, "completion_tokens": n,
"model": "<name>"}`` suitable for :func:`calculate_cost` (which ignores
the extra ``model`` key in the usage sub-dict), or ``None`` when the
Returns a dict suitable for :func:`calculate_cost`, or ``None`` when the
header is absent or malformed.
"""
if not header_value:
return None
parts: dict[str, int] = {}
int_parts: dict[str, int] = {}
model: str | None = None
cost_usd: float | None = None
for item in header_value.split(","):
key, sep, value = item.partition("=")
if not sep:
@@ -104,30 +173,44 @@ def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None:
if key == "model":
model = value
continue
if key == "cost_usd":
try:
cost_usd = float(value)
except (ValueError, TypeError):
cost_usd = None
continue
try:
parts[key] = int(value)
int_parts[key] = int(value)
except (ValueError, TypeError):
continue
prompt = parts.get("prompt")
completion = parts.get("completion")
if prompt is not None and completion is not None:
result: dict[str, int | str] = {
"prompt_tokens": prompt,
"completion_tokens": completion,
}
if "total" in parts:
result["total_tokens"] = parts["total"]
if model:
result["model"] = model
return result
logger.warning(
"Failed to parse X-Tinfoil-Usage-Metrics header",
extra={
"header_value": header_value,
"parsed_parts": parts,
},
)
return None
prompt = int_parts.get("prompt")
completion = int_parts.get("completion")
if prompt is None or completion is None:
logger.warning(
"Failed to parse X-Tinfoil-Usage-Metrics header",
extra={
"header_value": header_value,
"parsed_parts": int_parts,
},
)
return None
result: dict[str, int | float | str] = {
"prompt_tokens": prompt,
"completion_tokens": completion,
}
if "total" in int_parts:
result["total_tokens"] = int_parts["total"]
if "cached_prompt_tokens" in int_parts:
result["cache_read_input_tokens"] = int_parts["cached_prompt_tokens"]
if "uncached_prompt_tokens" in int_parts:
result["uncached_prompt_tokens"] = int_parts["uncached_prompt_tokens"]
if cost_usd is not None:
result["cost_usd"] = cost_usd
if model:
result["model"] = model
return result
def _get_header_case_insensitive(
@@ -191,7 +274,9 @@ def _resolve_ehbp_target_url(
otherwise the header is ignored so callers cannot redirect other providers
or leak upstream API keys.
"""
override_header = profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER
override_header = (
profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER
)
if not override_header:
return target_url
enclave_url = _get_header_case_insensitive(headers, override_header)
@@ -274,6 +359,11 @@ def _build_cost_info(
output_tokens: int = 0,
input_msats: int = 0,
output_msats: int = 0,
cache_read_input_tokens: int = 0,
cache_creation_input_tokens: int = 0,
cache_read_msats: int = 0,
cache_creation_msats: int = 0,
total_usd: float = 0.0,
actual_model: str | None = None,
) -> dict:
"""Build a cost-info dict with token counts and per-token-type costs.
@@ -282,22 +372,25 @@ def _build_cost_info(
one), it is included in the returned dict so callers can use it for billing
finalization and logging.
"""
result: dict[str, int | str | None] = {
result: dict[str, int | float | str | None] = {
"total_msats": total_msats,
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": input_tokens + output_tokens,
"input_msats": input_msats,
"output_msats": output_msats,
"cache_read_input_tokens": cache_read_input_tokens,
"cache_creation_input_tokens": cache_creation_input_tokens,
"cache_read_msats": cache_read_msats,
"cache_creation_msats": cache_creation_msats,
"total_usd": total_usd,
}
if actual_model:
result["actual_model"] = actual_model
return result
def _inject_cost_response_headers(
headers: dict[str, str], cost_info: dict
) -> None:
def _inject_cost_response_headers(headers: dict[str, str], cost_info: dict) -> None:
"""Add per-request cost headers to an EHBP response.
Since EHBP response bodies are opaque encrypted blobs, cost cannot be
@@ -305,8 +398,14 @@ def _inject_cost_response_headers(
the client/Tinfoil SDK can read without decrypting.
"""
headers["X-Routstr-Cost-Msats"] = str(cost_info["total_msats"])
if "computed_msats" in cost_info:
headers["X-Routstr-Computed-Cost-Msats"] = str(cost_info["computed_msats"])
headers["X-Routstr-Input-Cost-Msats"] = str(cost_info["input_msats"])
headers["X-Routstr-Output-Cost-Msats"] = str(cost_info["output_msats"])
headers["X-Routstr-Cache-Read-Msats"] = str(cost_info.get("cache_read_msats", 0))
headers["X-Routstr-Cache-Creation-Msats"] = str(
cost_info.get("cache_creation_msats", 0)
)
async def _compute_ehbp_actual_cost(
@@ -316,10 +415,10 @@ async def _compute_ehbp_actual_cost(
) -> dict:
"""Compute the actual cost in msats from Tinfoil usage metrics.
Falls back to ``max_cost_for_model`` when usage is absent (streaming) or
cannot be priced. The result is clamped to ``[min_request_msat,
max_cost_for_model]`` so the refund never exceeds the reservation and is
never zero.
When usage is present, the result is clamped to ``[min_request_msat,
max_cost_for_model]``. Missing or unpriceable usage returns zero: encrypted
EHBP bodies cannot be estimated locally, and the authorization ceiling is
not evidence of consumption.
When the usage-metrics header includes ``model=<name>`` and it differs
from ``model_obj.id``, the actual served model's pricing is used for the
@@ -332,7 +431,7 @@ async def _compute_ehbp_actual_cost(
"""
usage_dict = parse_tinfoil_usage_metrics(usage_header)
if usage_dict is None:
return _build_cost_info(max_cost_for_model)
return _build_cost_info(0)
# The enclave may serve a different model than the one requested (e.g.
# due to failover). The usage-metrics header's ``model=<name>`` carries
@@ -344,6 +443,14 @@ async def _compute_ehbp_actual_cost(
# look up the actual model's pricing.
actual_model: str | None = usage_dict.pop("model", None) # type: ignore[arg-type]
pricing_model_id = model_obj.id
# Bill the model we actually routed to. Passing only the model *string*
# to calculate_cost makes it re-derive pricing from the global alias map,
# which resolves the id to the best-ranked candidate — not the serving
# one. Tinfoil's catalog id (e.g. ``deepseek-v4-1-flash``) is also a
# cross-provider alias, and that cheaper candidate has no cache rate, so
# the cache discount silently disappeared (and the request was
# undercharged). Hand calculate_cost the identity it cannot reconstruct.
pricing_model_obj: Model = model_obj
expected_upstream_model = model_obj.forwarded_model_id or model_obj.id
expected_identity = _normalize_upstream_model_id(expected_upstream_model)
served_identity = _normalize_upstream_model_id(actual_model)
@@ -359,7 +466,24 @@ async def _compute_ehbp_actual_cost(
# the global model map. The resolved object can belong to a different
# provider and therefore have a different client-facing ``id`` while
# still representing the same upstream model.
actual_model_obj = get_model_instance(actual_model)
#
# The enclave reports the *bare* upstream id, but the routstr model is
# namespaced ``tinfoil-`` (and the SDK strips that prefix for the
# encrypted body). Resolve the served id within the same namespace
# first: a same-model report then maps back onto the requested Tinfoil
# model, and a genuine failover lands on the actually-served Tinfoil
# model — instead of the cheaper cross-provider model the bare id
# would resolve to in the global map.
namespaced_served = actual_model
if (
expected_upstream_model.startswith(TINFOIL_MODEL_PREFIX)
and not actual_model.startswith(TINFOIL_MODEL_PREFIX)
):
namespaced_served = TINFOIL_MODEL_PREFIX + actual_model
actual_model_obj = get_model_instance(namespaced_served)
if actual_model_obj is None and namespaced_served != actual_model:
actual_model_obj = get_model_instance(actual_model)
if actual_model_obj is None:
logger.warning(
"EHBP served model not found in registry, falling back "
@@ -375,9 +499,7 @@ async def _compute_ehbp_actual_cost(
resolved_upstream_model = (
actual_model_obj.forwarded_model_id or actual_model_obj.id
)
resolved_identity = _normalize_upstream_model_id(
resolved_upstream_model
)
resolved_identity = _normalize_upstream_model_id(resolved_upstream_model)
if resolved_identity != expected_identity:
logger.info(
"EHBP served model differs from requested, using actual "
@@ -390,6 +512,7 @@ async def _compute_ehbp_actual_cost(
},
)
pricing_model_id = actual_model_obj.id
pricing_model_obj = actual_model_obj
else:
# A different registry/client alias resolved to the same
# upstream model; retain the requested model's pricing.
@@ -402,22 +525,23 @@ async def _compute_ehbp_actual_cost(
cost = await calculate_cost(
{"model": pricing_model_id, "usage": usage_dict},
max_cost_for_model,
pricing_model_obj,
)
except Exception as e:
logger.warning(
"EHBP usage cost calculation failed, falling back to max cost",
"EHBP usage cost calculation failed; releasing instead of charging max cost",
extra={
"model": pricing_model_id,
"error": str(e),
"usage": usage_dict,
},
)
return _build_cost_info(max_cost_for_model, actual_model=actual_model)
return _build_cost_info(0, actual_model=actual_model)
if isinstance(cost, MaxCostData):
logger.warning(
"EHBP calculate_cost returned MaxCostData (no model pricing), "
"falling back to max cost",
"EHBP calculate_cost returned MaxCostData (no usable pricing); "
"releasing instead of charging max cost",
extra={
"model": pricing_model_id,
"max_cost_for_model": max_cost_for_model,
@@ -425,7 +549,7 @@ async def _compute_ehbp_actual_cost(
"cost_total_msats": cost.total_msats,
},
)
return _build_cost_info(max_cost_for_model, actual_model=actual_model)
return _build_cost_info(0, actual_model=actual_model)
if isinstance(cost, CostData):
actual = max(int(cost.total_msats), int(settings.min_request_msat))
clamped = min(actual, max_cost_for_model)
@@ -445,17 +569,22 @@ async def _compute_ehbp_actual_cost(
output_tokens=cost.output_tokens,
input_msats=cost.input_msats,
output_msats=cost.output_msats,
cache_read_input_tokens=cost.cache_read_input_tokens,
cache_creation_input_tokens=cost.cache_creation_input_tokens,
cache_read_msats=cost.cache_read_msats,
cache_creation_msats=cost.cache_creation_msats,
total_usd=cost.total_usd,
actual_model=actual_model,
)
# CostDataError
logger.warning(
"EHBP usage cost calculation error, falling back to max cost",
"EHBP usage cost calculation error; releasing instead of charging max cost",
extra={
"model": pricing_model_id,
"error": getattr(cost, "message", str(cost)),
},
)
return _build_cost_info(max_cost_for_model, actual_model=actual_model)
return _build_cost_info(0, actual_model=actual_model)
def _extract_usage_from_response(
@@ -500,6 +629,18 @@ class EHBPForwardingTarget:
profile: ConfidentialInferenceProfile | None = None
async def _release_failed_ehbp_charge(
reservation: ReservationSnapshot, session: AsyncSession
) -> None:
if await release_reservation(reservation, session, reservation.reserved_msats):
return
await _stop_reservation_heartbeat(reservation.release_id)
logger.critical(
"Failed to release EHBP reservation after rejected charge",
extra={"reservation_id": reservation.release_id},
)
async def finalize_ehbp_actual_cost_payment(
key: ApiKey,
session: AsyncSession,
@@ -507,61 +648,27 @@ async def finalize_ehbp_actual_cost_payment(
model_id: str,
cost_info: dict,
reservation_snapshot: ReservationSnapshot | None = None,
) -> None:
) -> int:
"""Finalize an EHBP bearer request using clamped provider usage metrics."""
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
await _validate_reservation_snapshot(key, reservation, session)
if not await _claim_reservation_for_charge(reservation, session):
return
return 0
reserved_cost_for_model = reservation.reserved_msats
billing_key = await get_billing_key(key, session)
key_hash = key.hashed_key
billing_key_hash = billing_key.hashed_key
total_cost_msats = max(0, int(cost_info.get("total_msats", reserved_cost_for_model)))
billing_key_hash = key_hash
total_cost_msats = max(
0, int(cost_info.get("total_msats", reserved_cost_for_model))
)
now = int(time.time())
safe_reserved = case(
(
col(ApiKey.reserved_balance) >= reserved_cost_for_model,
col(ApiKey.reserved_balance) - reserved_cost_for_model,
),
else_=0,
charged = await _charge_reservation_rows(
session,
billing_key_hash=billing_key_hash,
reserved_msats=reserved_cost_for_model,
charge_msats=total_cost_msats,
)
cleared_reserved_at = case(
(
col(ApiKey.reserved_balance) - reserved_cost_for_model > 0,
col(ApiKey.reserved_at),
),
else_=None,
)
stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=safe_reserved,
reserved_at=cleared_reserved_at,
balance=col(ApiKey.balance) - total_cost_msats,
total_spent=col(ApiKey.total_spent) + total_cost_msats,
)
)
result = await session.exec(stmt) # type: ignore[call-overload]
child_result = None
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(
reserved_balance=safe_reserved,
reserved_at=cleared_reserved_at,
total_spent=col(ApiKey.total_spent) + total_cost_msats,
)
)
child_result = await session.exec(child_stmt) # type: ignore[call-overload]
if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0):
await session.rollback()
if not charged:
logger.error(
"Failed to finalize EHBP usage-based payment",
extra={
@@ -570,16 +677,14 @@ async def finalize_ehbp_actual_cost_payment(
"model": model_id,
"reserved_cost_for_model": reserved_cost_for_model,
"total_cost_msats": total_cost_msats,
"parent_rowcount": result.rowcount,
"child_rowcount": getattr(child_result, "rowcount", None),
},
)
return
await _release_failed_ehbp_charge(reservation, session)
return 0
await session.commit()
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
await _stop_reservation_heartbeat(reservation.release_id)
await session.refresh(key)
if total_cost_msats > 0 and ROUTSTR_FEE_PERCENT > 0:
fee_msats = math.ceil(total_cost_msats * ROUTSTR_FEE_PERCENT / 100)
@@ -596,19 +701,30 @@ async def finalize_ehbp_actual_cost_payment(
extra={
"event": "finalize",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"billing_key_hash": key.hashed_key[:8] + "...",
"model": model_id,
"cost_reserved": reserved_cost_for_model,
"cost_charged": total_cost_msats,
"input_tokens": cost_info.get("input_tokens", 0),
"output_tokens": cost_info.get("output_tokens", 0),
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
# Cache splits are only knowable when the enclave reports
# ``cached_prompt_tokens``; absent that they are a measured zero on
# the token counts the provider did report (not an unknown), so the
# event key set stays stable for usage-analytics consumers.
"cache_read_input_tokens": cost_info.get("cache_read_input_tokens", 0),
"cache_creation_input_tokens": cost_info.get(
"cache_creation_input_tokens", 0
),
"cache_read_msats": cost_info.get("cache_read_msats", 0),
"cache_creation_msats": cost_info.get("cache_creation_msats", 0),
"balance": key.balance,
"reserved_balance": key.reserved_balance,
"total_spent": key.total_spent,
"finalize_type": "ehbp_usage",
"finalized_at": now,
},
)
return total_cost_msats
async def finalize_ehbp_max_cost_payment(
@@ -617,128 +733,26 @@ async def finalize_ehbp_max_cost_payment(
max_cost_for_model: int,
model_id: str,
reservation_snapshot: ReservationSnapshot | None = None,
) -> None:
"""Finalize an EHBP bearer request by charging the reserved max cost.
) -> int:
"""Release an unmeasured EHBP request without charging its reservation.
EHBP responses are encrypted, so Routstr cannot inspect token usage. Unlike
normal completion handlers, this intentionally charges the pre-reserved max
cost and releases the reservation.
The legacy name is retained for compatibility with internal callers. EHBP
responses are encrypted, so no local estimate is possible when the trusted
usage header/trailer is absent.
"""
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
await _validate_reservation_snapshot(key, reservation, session)
if not await _claim_reservation_for_charge(reservation, session):
return
max_cost_for_model = reservation.reserved_msats
billing_key = await get_billing_key(key, session)
key_hash = key.hashed_key
billing_key_hash = billing_key.hashed_key
total_cost_msats = max(0, int(max_cost_for_model))
now = int(time.time())
cleared_reserved_at = case(
(
col(ApiKey.reserved_balance) - max_cost_for_model > 0,
col(ApiKey.reserved_at),
),
else_=None,
)
safe_reserved = case(
(
col(ApiKey.reserved_balance) >= max_cost_for_model,
col(ApiKey.reserved_balance) - max_cost_for_model,
),
else_=0,
)
stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=safe_reserved,
reserved_at=cleared_reserved_at,
balance=col(ApiKey.balance) - total_cost_msats,
total_spent=col(ApiKey.total_spent) + total_cost_msats,
)
)
result = await session.exec(stmt) # type: ignore[call-overload]
if billing_key.hashed_key != key.hashed_key:
child_safe_reserved = case(
(
col(ApiKey.reserved_balance) >= max_cost_for_model,
col(ApiKey.reserved_balance) - max_cost_for_model,
),
else_=0,
)
child_cleared_reserved_at = case(
(
col(ApiKey.reserved_balance) - max_cost_for_model > 0,
col(ApiKey.reserved_at),
),
else_=None,
)
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(
reserved_balance=child_safe_reserved,
reserved_at=child_cleared_reserved_at,
total_spent=col(ApiKey.total_spent) + total_cost_msats,
)
)
child_result = await session.exec(child_stmt) # type: ignore[call-overload]
else:
child_result = None
if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0):
await session.rollback()
logger.error(
"Failed to finalize EHBP max-cost payment",
extra={
"key_hash": key_hash[:8] + "...",
"billing_key_hash": billing_key_hash[:8] + "...",
"model": model_id,
"max_cost_for_model": max_cost_for_model,
"parent_rowcount": result.rowcount,
"child_rowcount": getattr(child_result, "rowcount", None),
},
)
return
await session.commit()
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
if total_cost_msats > 0 and ROUTSTR_FEE_PERCENT > 0:
fee_msats = math.ceil(total_cost_msats * ROUTSTR_FEE_PERCENT / 100)
try:
await accumulate_routstr_fee(session, fee_msats)
except Exception as e:
logger.warning(
"Failed to accumulate Routstr fee for EHBP request",
extra={"error": str(e), "fee_msats": fee_msats},
)
payments_logger.info(
"FINALIZE",
key_log_hash = key.hashed_key[:8] + "..."
await release_reservation(reservation, session, reservation.reserved_msats)
logger.warning(
"Released unmeasured EHBP reservation without charging max cost",
extra={
"event": "finalize",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"key_hash": key_log_hash,
"model": model_id,
"cost_reserved": max_cost_for_model,
"cost_charged": total_cost_msats,
"input_tokens": 0,
"output_tokens": 0,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
"finalize_type": "ehbp_max_cost",
"finalized_at": now,
"max_cost_for_model": max_cost_for_model,
},
)
return 0
async def send_cashu_refund(
@@ -842,6 +856,25 @@ async def forward_ehbp_request(
"body_preview": body_preview,
},
)
# Key-config mismatch (stale client HPKE key): return the
# enclave's 422 problem+json directly so the SDK's
# KeyConfigMismatchError detection fires and triggers
# re-attestation. Wrapping it as application/json would
# destroy the signal and cause permanent failure.
if _is_ehbp_key_config_response(resp):
logger.warning(
"EHBP upstream %s returned key-config mismatch for model=%s, "
"passing through for client re-attestation",
provider_type,
model_obj.id,
extra={
"provider": provider_type,
"model": model_obj.id,
"path": path,
},
)
return _passthrough_key_config_response(resp)
raise UpstreamError(
f"EHBP upstream {provider_type} returned {resp.status_code} "
f"for model {model_obj.id}: {body_preview[:200] or '<empty>'}",
@@ -893,10 +926,9 @@ async def forward_ehbp_request(
cost_info = await _compute_ehbp_actual_cost(
usage_header, model_obj, max_cost_for_model
)
# Use the actual served model for billing when it differs from
# the requested model.
billing_model = cost_info.pop("actual_model", None) or model_obj.id
await finalize_ehbp_actual_cost_payment(
computed_msats = int(cost_info["total_msats"])
charged_msats = await finalize_ehbp_actual_cost_payment(
key,
session,
max_cost_for_model,
@@ -904,18 +936,25 @@ async def forward_ehbp_request(
cost_info,
reservation_snapshot,
)
cost_data = {**cost_info, "total_usd": 0.0}
cost_data = {
**cost_info,
"total_msats": charged_msats,
"charged_msats": charged_msats,
"total_usd": cost_info.get("total_usd", 0.0),
}
if computed_msats != charged_msats:
cost_data["computed_msats"] = computed_msats
else:
logger.warning(
"EHBP usage metrics not found in headers or trailers, "
"falling back to max-cost billing",
"EHBP usage metrics not found in headers or trailers; "
"releasing instead of charging the authorization ceiling",
extra={
"model": model_obj.id,
"provider": provider_type,
"key_hash": key.hashed_key[:8] + "...",
},
)
await finalize_ehbp_max_cost_payment(
charged_msats = await finalize_ehbp_max_cost_payment(
key,
session,
max_cost_for_model,
@@ -923,14 +962,15 @@ async def forward_ehbp_request(
reservation_snapshot,
)
cost_data = {
"total_msats": max_cost_for_model,
"total_msats": charged_msats,
"charged_msats": charged_msats,
"total_usd": 0.0,
"input_tokens": 0,
"output_tokens": 0,
}
# Build the cost_info dict from what adjust_payment_for_tokens returned
# or from the max-cost fallback. Fields match CostData/MaxCostData.dict().
# Build the cost_info dict from measured usage or the unmeasured-release
# fallback. Fields match CostData/MaxCostData.dict().
cost_info = {
"total_msats": cost_data.get("total_msats", max_cost_for_model),
"input_tokens": cost_data.get("input_tokens", 0),
@@ -939,7 +979,15 @@ async def forward_ehbp_request(
+ cost_data.get("output_tokens", 0),
"input_msats": cost_data.get("input_msats", 0),
"output_msats": cost_data.get("output_msats", 0),
"cache_read_input_tokens": cost_data.get("cache_read_input_tokens", 0),
"cache_creation_input_tokens": cost_data.get(
"cache_creation_input_tokens", 0
),
"cache_read_msats": cost_data.get("cache_read_msats", 0),
"cache_creation_msats": cost_data.get("cache_creation_msats", 0),
}
if "computed_msats" in cost_data:
cost_info["computed_msats"] = cost_data["computed_msats"]
cost_usd = cost_data.get("total_usd", 0.0)
# Build response headers, filtering out hop-by-hop headers
@@ -1028,7 +1076,9 @@ async def forward_ehbp_x_cashu_request(
target_url = _resolve_ehbp_target_url(
target.url, path, headers, provider_type, profile
)
upstream_headers = _prepare_ehbp_upstream_headers(headers, target.headers, profile)
upstream_headers = _prepare_ehbp_upstream_headers(
headers, target.headers, profile
)
request_body = await request.body()
# Merge query params into the target URL
@@ -1047,6 +1097,29 @@ async def forward_ehbp_x_cashu_request(
)
if resp.status_code != 200:
# Key-config mismatch (stale client HPKE key): refund the
# full token and pass the enclave's 422 problem+json through
# so the SDK's KeyConfigMismatchError detection fires.
if _is_ehbp_key_config_response(resp):
logger.warning(
"EHBP upstream %s returned key-config mismatch for "
"model=%s, refunding and passing through",
provider_type,
model_obj.id,
extra={
"provider": provider_type,
"model": model_obj.id,
"path": path,
"refunded_amount": amount,
},
)
refund_token = await send_cashu_refund(
amount, unit, mint, request_id
)
passthrough = _passthrough_key_config_response(resp)
passthrough.headers["X-Cashu"] = refund_token
return passthrough
refund_token = await send_cashu_refund(amount, unit, mint, request_id)
error_response = Response(
content=json.dumps(
@@ -1076,9 +1149,7 @@ async def forward_ehbp_x_cashu_request(
usage_source = (
"header"
if usage_header_name
and any(
k.lower() == usage_header_name.lower() for k, _ in resp.headers
)
and any(k.lower() == usage_header_name.lower() for k, _ in resp.headers)
else ("trailer" if usage_header else "none")
)
@@ -1153,6 +1224,46 @@ async def forward_ehbp_x_cashu_request(
except Exception:
raise
except EhbpTimeoutError as e:
logger.warning(
"EHBP X-Cashu upstream timed out",
extra={
"error": str(e),
"path": path,
"method": request.method,
"redeemed": redeemed,
},
)
if redeemed and amount > 0:
try:
refund_token = await send_cashu_refund(amount, unit, mint, request_id)
error_response = create_error_response(
"upstream_timeout",
str(e),
504,
request=request,
code="UPSTREAM_TIMEOUT",
)
error_response.headers["X-Cashu"] = refund_token
return error_response
except Exception as refund_error:
logger.error(
"Failed to refund EHBP X-Cashu token after timeout",
extra={
"error": str(refund_error),
"original_error": str(e),
},
)
return create_error_response(
"upstream_timeout",
str(e),
504,
request=request,
code="UPSTREAM_TIMEOUT",
)
except Exception as e:
error_message = str(e)
logger.error(
@@ -1186,6 +1297,30 @@ async def forward_ehbp_x_cashu_request(
},
)
if not redeemed:
classified = classify_redemption_error(e)
if classified is not None:
error_type, status_code, message, error_code = classified
# Never re-offer a spent/consumed token.
echo_token = None if error_code in SPENT_TOKEN_CODES else x_cashu_token
return create_error_response(
error_type,
message,
status_code,
request=request,
token=echo_token,
code=error_code,
)
# Raw exception text may contain the attacker-supplied mint URL.
return create_error_response(
"api_error",
"Internal error during token redemption",
500,
request=request,
token=x_cashu_token,
code="internal_error",
)
if "already spent" in error_message.lower():
return create_error_response(
"token_already_spent",
+38 -15
View File
@@ -32,6 +32,7 @@ from ..core.exceptions import UpstreamError
from ..core.redaction import redact_org_ids
from ..payment.models import Model
from .rate_limit import classify_rate_limit
from .reasoning_effort import adapt_messages_body_for_litellm
logger = get_logger(__name__)
@@ -40,6 +41,10 @@ logger = get_logger(__name__)
# unsupported params; these newer/extension fields get passed through
# verbatim and the upstream rejects them with a 400. Pop them here so the
# request reaches the upstream cleanly.
#
# Note: ``dispatch_anthropic_messages`` additionally enforces
# ``ALLOWED_MESSAGES_REQUEST_FIELDS``, which already excludes all of these.
# This tuple remains for ``gemini_messages``, which pops them explicitly.
ANTHROPIC_ONLY_FIELDS: tuple[str, ...] = (
"thinking",
"cache_control",
@@ -51,6 +56,27 @@ ANTHROPIC_ONLY_FIELDS: tuple[str, ...] = (
"anthropic_beta",
)
# Anthropic Messages API request fields forwarded from the client body
# into the upstream call. Only these are passed on; anything else is
# dropped so the forwarded request is deterministic and limited to the
# documented Messages surface.
ALLOWED_MESSAGES_REQUEST_FIELDS: frozenset[str] = frozenset(
{
"messages",
"max_tokens",
"system",
"temperature",
"top_p",
"top_k",
"stop_sequences",
"tools",
"tool_choice",
"metadata",
# OpenAI-shaped effort after thinking is lifted off Anthropic bodies.
"reasoning_effort",
}
)
def coerce_litellm_payload(payload: object) -> dict:
"""Convert a litellm event into a plain dict.
@@ -108,9 +134,7 @@ def parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]:
return events, buffer
def events_from_chunk(
chunk: object, sse_buffer: bytes
) -> tuple[list[dict], bytes]:
def events_from_chunk(chunk: object, sse_buffer: bytes) -> tuple[list[dict], bytes]:
"""Normalize a stream chunk into one or more event dicts.
``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE
@@ -201,9 +225,7 @@ async def aggregate_anthropic_events_to_message(
raw_json = partial_json.pop(idx, None)
if raw_json is not None and idx < len(blocks):
try:
blocks[idx]["input"] = (
json.loads(raw_json) if raw_json else {}
)
blocks[idx]["input"] = json.loads(raw_json) if raw_json else {}
except json.JSONDecodeError:
blocks[idx]["input"] = raw_json
elif etype == "message_delta":
@@ -445,9 +467,7 @@ async def dispatch_anthropic_messages(
on bad input or upstream failure.
"""
if not request_body:
raise UpstreamError(
"Missing request body for /v1/messages", status_code=400
)
raise UpstreamError("Missing request body for /v1/messages", status_code=400)
try:
body: dict = json.loads(request_body)
@@ -466,15 +486,18 @@ async def dispatch_anthropic_messages(
client_stream = bool(body.pop("stream", False))
upstream_stream = True
dropped: dict[str, Any] = {}
for field in ANTHROPIC_ONLY_FIELDS:
if field in body:
dropped[field] = body.pop(field)
adapt_messages_body_for_litellm(body, model_obj)
# Forward only allowlisted Anthropic Messages request fields. Any
# other client-supplied key is dropped so it cannot leak into the
# upstream request. See ALLOWED_MESSAGES_REQUEST_FIELDS.
dropped = sorted(set(body) - ALLOWED_MESSAGES_REQUEST_FIELDS)
if dropped:
logger.debug(
"Dropped anthropic-only fields before litellm dispatch",
extra={"dropped_keys": sorted(dropped.keys())},
"Dropped non-forwardable fields before litellm dispatch",
extra={"dropped_keys": dropped},
)
body = {k: v for k, v in body.items() if k in ALLOWED_MESSAGES_REQUEST_FIELDS}
# Convention: `model.id` is the canonical upstream model name;
# `forwarded_model_id` is the public alias the internal API exposes
+191 -24
View File
@@ -1,9 +1,6 @@
"""Model-path discovery service.
Exposes every selectable upstream route a Routstr model is reachable through.
This PR remains discovery-only: request-side routing will consume the opaque
selectors in a follow-up.
A path is a standard percent-encoded query string containing the configured
upstream URL, provider ID, client-visible model ID and, for an exact OpenRouter
endpoint, its machine-readable tag::
@@ -16,11 +13,12 @@ from __future__ import annotations
import asyncio
import ipaddress
import json
import random
import time
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Callable
from urllib.parse import urlencode, urlsplit
from urllib.parse import parse_qsl, urlencode, urlsplit
import httpx
from sqlalchemy.dialects.sqlite import insert
@@ -60,10 +58,11 @@ ModelKey = tuple[str, int]
@dataclass(frozen=True)
class EndpointIdentity:
"""Exact OpenRouter endpoint identity returned by ``/endpoints``."""
"""Exact OpenRouter endpoint and its provider-specific model metadata."""
tag: str
provider_name: str | None
model_metadata: dict[str, Any]
@dataclass(frozen=True)
@@ -83,6 +82,7 @@ class DiscoveredPath:
model_id: str
path: str
provider: ConfiguredProviderIdentity
model_metadata: dict[str, Any]
endpoint_tag: str | None = None
endpoint_name: str | None = None
@@ -132,21 +132,68 @@ def encode_model_path(
return urlencode(components)
@dataclass(frozen=True)
class ModelPathSelector:
"""Decoded client-supplied route selector."""
base_url: str
provider_id: int
model_id: str
endpoint_tag: str | None = None
def decode_model_path(path: str) -> ModelPathSelector | None:
"""Inverse of ``encode_model_path``; ``None`` when the selector is malformed."""
try:
pairs = parse_qsl(
path,
keep_blank_values=True,
strict_parsing=True,
max_num_fields=4,
errors="strict",
)
except ValueError:
return None
params = dict(pairs)
if len(params) != len(pairs) or params.keys() - {
"url",
"provider-id",
"model-id",
"endpoint",
}:
return None
if any(not value.strip() for value in params.values()):
return None
base_url = params.get("url", "")
model_id = params.get("model-id", "")
raw_provider_id = params.get("provider-id", "")
if not base_url or not model_id or not raw_provider_id:
return None
try:
provider_id = int(raw_provider_id)
except ValueError:
return None
if provider_id <= 0:
return None
return ModelPathSelector(
base_url=base_url,
provider_id=provider_id,
model_id=model_id,
endpoint_tag=params.get("endpoint") or None,
)
def _make_http_client() -> httpx.AsyncClient:
"""Client factory, separated so tests can substitute a mock transport."""
return httpx.AsyncClient()
def is_openrouter_base_url(base_url: str | None) -> bool:
"""True when ``base_url`` points at OpenRouter.
Deliberately separate from ``BaseUpstreamProvider._upstream_accepts_cache_control``:
that predicate also returns True for native Anthropic (correct for
cache-control, wrong for OpenRouter endpoint discovery). This one keys only
on the URL so a ``GenericUpstreamProvider`` aimed at OpenRouter is matched
while native Anthropic is not.
"""
return "openrouter.ai" in (base_url or "")
"""Match OpenRouter itself, not compatible providers or lookalike hosts."""
try:
return urlsplit(base_url or "").hostname == "openrouter.ai"
except ValueError:
return False
def exposed_model_id(model: object) -> str:
@@ -263,9 +310,14 @@ async def _fetch_openrouter_endpoint_subproviders(
try:
payload = resp.json()
data = payload.get("data") if isinstance(payload, dict) else None
endpoints = data.get("endpoints") if isinstance(data, dict) else None
if not isinstance(data, dict):
raise ValueError("data must be an object")
endpoints = data.get("endpoints")
if not isinstance(endpoints, list):
raise ValueError("endpoints must be a list")
common_metadata = {
key: value for key, value in data.items() if key != "endpoints"
}
identities: dict[str, EndpointIdentity] = {}
for endpoint in endpoints:
if not isinstance(endpoint, dict):
@@ -281,6 +333,7 @@ async def _fetch_openrouter_endpoint_subproviders(
provider_name=provider_name
if isinstance(provider_name, str) and provider_name
else None,
model_metadata={**common_metadata, **endpoint},
),
)
if endpoints and not identities:
@@ -341,6 +394,34 @@ async def _load_model_visibility() -> tuple[
return overrides_by_key, disabled_model_keys, provider_identities
def _serialize_model_metadata(model: object, model_id: str) -> dict[str, Any]:
"""Serialize provider-specific model details into the public API shape."""
model_dict = getattr(model, "dict", None)
if callable(model_dict):
metadata = dict(model_dict())
else:
metadata = {
key: value for key, value in vars(model).items() if not key.startswith("_")
}
for field in (
"architecture",
"pricing",
"sats_pricing",
"per_request_limits",
"top_provider",
"alias_ids",
):
value = metadata.get(field)
if isinstance(value, str):
try:
metadata[field] = json.loads(value)
except (TypeError, ValueError):
pass
metadata["id"] = model_id
return metadata
def _apply_model_visibility(
upstream: BaseUpstreamProvider,
overrides_by_key: dict[ModelKey, ModelRow] | None,
@@ -348,11 +429,10 @@ def _apply_model_visibility(
) -> list[object]:
"""Return provider models after DB disabled/override state is applied.
Only the identity fields (``id``, ``forwarded_model_id``,
``canonical_slug``) matter for path discovery, so DB override rows are used
directly rather than rebuilt into fully priced ``Model`` objects — the
pricing pipeline costs ~0.7ms of event-loop CPU per row for data this
module immediately discards.
DB override rows are used directly rather than rebuilt into priced
``Model`` objects. Their JSON metadata fields are decoded when each path is
collected, preserving the provider-specific stored values without running
the routing price-selection pipeline.
"""
overrides_by_key = overrides_by_key or {}
disabled_model_keys = disabled_model_keys or set()
@@ -414,6 +494,7 @@ async def _collect_provider_paths(
provider_identity.base_url, provider_identity.id, model_id
),
provider=provider_identity,
model_metadata=_serialize_model_metadata(model, model_id),
)
if not is_openrouter_base_url(upstream.base_url):
@@ -453,6 +534,7 @@ async def _collect_provider_paths(
endpoint.tag,
),
provider=provider_identity,
model_metadata={**endpoint.model_metadata, "id": model_id},
endpoint_tag=endpoint.tag,
endpoint_name=endpoint.provider_name,
)
@@ -512,6 +594,7 @@ async def _persist_provider_paths(
"provider_type": discovered.provider.provider_type,
"endpoint_tag": discovered.endpoint_tag,
"endpoint_name": discovered.endpoint_name,
"model_metadata": json.dumps(discovered.model_metadata),
"upstream_provider_id": upstream_provider_id,
"updated_at": now,
}
@@ -526,6 +609,7 @@ async def _persist_provider_paths(
"provider_type": insert_stmt.excluded.provider_type,
"endpoint_tag": insert_stmt.excluded.endpoint_tag,
"endpoint_name": insert_stmt.excluded.endpoint_name,
"model_metadata": insert_stmt.excluded.model_metadata,
"updated_at": insert_stmt.excluded.updated_at,
},
)
@@ -736,10 +820,83 @@ async def refresh_model_paths_periodically(
break
def _serialize_path(row: ModelPathRow) -> dict[str, Any]:
def _price_in_sats(model: dict[str, Any], provider_fee: float) -> None:
"""Run a path's USD rates through the ``/v1/models`` pricing pipeline.
Metadata copied from the provider model cache is already priced. OpenRouter
endpoint metadata is not: it carries that endpoint's own USD rates, which
still need the cache backfill, the provider fee and the sats conversion.
"""
pricing = model.get("pricing")
if model.get("sats_pricing") or not isinstance(pricing, dict):
return
from ..payment.models import (
Architecture,
Model,
Pricing,
TopProvider,
_calculate_usd_max_costs,
_update_model_sats_pricing,
backfill_cache_pricing,
)
from ..payment.price import sats_usd_price
try:
model_id = model.get("forwarded_model_id") or model["id"]
usd = backfill_cache_pricing(model_id, Pricing.parse_obj(pricing))
usd = Pricing.parse_obj({k: v * provider_fee for k, v in usd.dict().items()})
priced = Model(
id=model_id,
name=model.get("name") or model_id,
created=0,
description="",
context_length=model.get("context_length") or 0,
architecture=Architecture(
modality="text",
input_modalities=[],
output_modalities=[],
tokenizer="",
instruct_type=None,
),
pricing=usd,
top_provider=TopProvider(
context_length=model.get("context_length"),
max_completion_tokens=model.get("max_completion_tokens"),
),
)
(
usd.max_prompt_cost,
usd.max_completion_cost,
usd.max_cost,
) = _calculate_usd_max_costs(priced)
priced = _update_model_sats_pricing(priced, sats_usd_price())
except Exception as exc:
# An endpoint with rates we cannot price is still a usable route, so it
# is served with its raw upstream pricing rather than dropped.
logger.warning(
"Could not calculate sats pricing for model path",
extra={"model_id": model.get("id"), "error": str(exc)},
)
return
if priced.sats_pricing:
model["pricing"] = usd.dict()
model["sats_pricing"] = priced.sats_pricing.dict()
def _serialize_path(row: ModelPathRow, provider_fee: float) -> dict[str, Any]:
endpoint = None
if row.endpoint_tag or row.endpoint_name:
endpoint = {"tag": row.endpoint_tag, "name": row.endpoint_name}
try:
model = json.loads(row.model_metadata)
except (TypeError, ValueError):
model = {}
if not isinstance(model, dict):
model = {}
model.setdefault("id", row.model_id)
_price_in_sats(model, provider_fee)
return {
"path": row.path,
"provider": {
@@ -748,11 +905,17 @@ def _serialize_path(row: ModelPathRow) -> dict[str, Any]:
"type": row.provider_type,
},
"endpoint": endpoint,
"model": model,
}
async def _provider_fees(session: "AsyncSession") -> dict[int, float]:
rows = (await session.exec(select(UpstreamProviderRow))).all()
return {row.id: row.provider_fee for row in rows if row.id is not None}
async def get_all_model_paths() -> dict:
"""All models with their exact selectable routes."""
"""All models with exact routes and provider-specific model metadata."""
async with create_session() as session:
rows = (
await session.exec(
@@ -763,6 +926,7 @@ async def get_all_model_paths() -> dict:
)
)
).all()
fees = await _provider_fees(session)
grouped: dict[str, list[dict[str, Any]]] = {}
seen_paths: dict[str, set[str]] = {}
@@ -772,7 +936,9 @@ async def get_all_model_paths() -> dict:
if row.path in seen_paths.setdefault(row.model_id, set()):
continue
seen_paths[row.model_id].add(row.path)
grouped.setdefault(row.model_id, []).append(_serialize_path(row))
grouped.setdefault(row.model_id, []).append(
_serialize_path(row, fees.get(row.upstream_provider_id, 1.01))
)
data = [
{
"id": grouped_model_id,
@@ -806,6 +972,7 @@ async def get_paths_for_model(model_id: str) -> dict:
unprefixed_id = public_model_id(model_id)
if unprefixed_id != model_id:
rows = await load_rows(session, unprefixed_id)
fees = await _provider_fees(session)
seen: set[str] = set()
paths: list[dict] = []
@@ -815,5 +982,5 @@ async def get_paths_for_model(model_id: str) -> dict:
if row.path in seen:
continue
seen.add(row.path)
paths.append(_serialize_path(row))
paths.append(_serialize_path(row, fees.get(row.upstream_provider_id, 1.01)))
return {"data": paths, "updated_at": updated_at or None}
+8 -34
View File
@@ -66,9 +66,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
"""Strip 'ollama/' prefix for Ollama API compatibility."""
return model_id.removeprefix("ollama/")
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:
"""Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint."""
return f"{self.base_url.rstrip('/')}/v1"
@@ -185,7 +183,9 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
except Exception:
self._models_cache = models_with_fees
self._models_by_id = {m.forwarded_model_id or m.id: m for m in self._models_cache}
self._models_by_id = {
m.forwarded_model_id or m.id: m for m in self._models_cache
}
logger.info(
f"Refreshed models cache for {self.base_url}",
extra={"model_count": len(models)},
@@ -224,26 +224,14 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
Returns:
Model with provider fee applied to pricing and max costs calculated
"""
from ..payment.models import Model, Pricing, _calculate_usd_max_costs
from ..payment.models import Pricing, _calculate_usd_max_costs
adjusted_pricing = Pricing.parse_obj(
{k: v * self.provider_fee for k, v in model.pricing.dict().items()}
)
temp_model = Model(
id=model.id,
name=model.name,
created=model.created,
description=model.description,
context_length=model.context_length,
architecture=model.architecture,
pricing=adjusted_pricing,
sats_pricing=None,
per_request_limits=model.per_request_limits,
top_provider=model.top_provider,
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
temp_model = model.copy(
update={"pricing": adjusted_pricing, "sats_pricing": None}
)
(
@@ -252,18 +240,4 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
adjusted_pricing.max_cost,
) = _calculate_usd_max_costs(temp_model)
return Model(
id=model.id,
name=model.name,
created=model.created,
description=model.description,
context_length=model.context_length,
architecture=model.architecture,
pricing=adjusted_pricing,
sats_pricing=model.sats_pricing,
per_request_limits=model.per_request_limits,
top_provider=model.top_provider,
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
)
return model.copy(update={"pricing": adjusted_pricing})
+180 -108
View File
@@ -1,5 +1,9 @@
from __future__ import annotations
import asyncio
import random
import time
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Optional
import httpx
@@ -15,6 +19,87 @@ if TYPE_CHECKING:
logger = get_logger(__name__)
_PPQ_SAFE_READ_ATTEMPTS = 3
_PPQ_CIRCUIT_COOLDOWN_SECONDS = 30.0
class PPQCircuitOpenError(RuntimeError):
pass
@dataclass
class _PPQCircuitState:
consecutive_failures: int = 0
cooldown_until: float = 0.0
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
loop: asyncio.AbstractEventLoop | None = None
_ppq_circuits: dict[str, _PPQCircuitState] = {}
def _ppq_origin(url: str) -> str:
parsed = httpx.URL(url)
port = parsed.port or {"https": 443, "http": 80}.get(parsed.scheme, 0)
return f"{parsed.scheme}://{parsed.host}:{port}"
async def _safe_read_request(
client: httpx.AsyncClient,
method: str,
url: str,
*,
headers: dict[str, str],
json: dict[str, object] | None = None,
) -> httpx.Response:
state = _ppq_circuits.setdefault(_ppq_origin(url), _PPQCircuitState())
loop = asyncio.get_running_loop()
if state.loop is not loop:
# Locks cannot be reused across event loops.
state.lock = asyncio.Lock()
state.loop = loop
async with state.lock:
remaining = state.cooldown_until - time.monotonic()
if remaining > 0:
raise PPQCircuitOpenError(
f"PPQ.AI safe-read circuit is open; retry after {remaining:.2f}s"
)
for attempt in range(1, _PPQ_SAFE_READ_ATTEMPTS + 1):
try:
response = await client.request(method, url, headers=headers, json=json)
response.raise_for_status()
state.consecutive_failures = 0
state.cooldown_until = 0.0
return response
except (httpx.TransportError, httpx.HTTPStatusError) as error:
retryable_status = isinstance(error, httpx.HTTPStatusError) and (
error.response.status_code in {502, 503, 504}
)
if not isinstance(error, httpx.TransportError) and not retryable_status:
raise
state.consecutive_failures += 1
if attempt >= _PPQ_SAFE_READ_ATTEMPTS:
state.cooldown_until = (
time.monotonic() + _PPQ_CIRCUIT_COOLDOWN_SECONDS
)
raise
base_delay = 0.25 * (2 ** (attempt - 1))
delay = base_delay + random.uniform(0.0, base_delay)
logger.warning(
"PPQ.AI safe read failed; retrying",
extra={
"url": url,
"attempt": attempt,
"max_attempts": _PPQ_SAFE_READ_ATTEMPTS,
"backoff_seconds": round(delay, 3),
"error": repr(error),
"error_type": type(error).__name__,
},
)
await asyncio.sleep(delay)
raise RuntimeError("unreachable")
class PPQAIModelPricing(BaseModel):
ui: Optional[dict[str, float]] = None
@@ -123,122 +208,109 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
url = f"{self.base_url}/models"
headers = {"Authorization": f"Bearer {self.api_key}"}
try:
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.get(url, headers=headers)
response.raise_for_status()
data = response.json()
async with httpx.AsyncClient(timeout=30.0) as client:
response = await _safe_read_request(client, "GET", url, headers=headers)
data = response.json()
models_data = data.get("data", [])
models_data = data.get("data", [])
or_models = [
Model(**model) # type: ignore
for model in await async_fetch_openrouter_models()
]
or_models = [
Model(**model) # type: ignore
for model in await async_fetch_openrouter_models()
]
models = []
for model_data in models_data:
try:
ppqai_model = PPQAIModel.parse_obj(model_data)
if ppqai_model.id in self.IGNORED_MODEL_IDS:
continue
models = []
for model_data in models_data:
try:
ppqai_model = PPQAIModel.parse_obj(model_data)
if ppqai_model.id in self.IGNORED_MODEL_IDS:
continue
or_model = next(
(
model
for model in or_models
if (model.id == ppqai_model.id)
or (model.id.split("/")[-1] == ppqai_model.id)
or (model.id == ppqai_model.id.split("/")[-1])
),
None,
)
or_model = next(
(
model
for model in or_models
if (model.id == ppqai_model.id)
or (model.id.split("/")[-1] == ppqai_model.id)
or (model.id == ppqai_model.id.split("/")[-1])
),
None,
)
if or_model:
input_price = None
if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get(
"input_per_1M"
)
elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens
if or_model:
input_price = None
if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get("input_per_1M")
elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens
if input_price is not None:
or_model.pricing.prompt = input_price / 1_000_000
if input_price is not None:
or_model.pricing.prompt = input_price / 1_000_000
output_price = None
if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get(
"output_per_1M"
)
elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens
output_price = None
if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get("output_per_1M")
elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens
if output_price is not None:
or_model.pricing.completion = output_price / 1_000_000
if output_price is not None:
or_model.pricing.completion = output_price / 1_000_000
if cl := ppqai_model.context_length:
or_model.context_length = cl
models.append(or_model)
else:
input_price = 0.0
if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get(
"input_per_1M", 0.0
)
elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens
output_price = 0.0
if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get(
"output_per_1M", 0.0
)
elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens
models.append(
Model(
id=ppqai_model.id,
name=ppqai_model.name,
created=ppqai_model.created_at // 1000,
description=f"{ppqai_model.provider or 'PPQ.AI'} model",
context_length=ppqai_model.context_length,
architecture=Architecture(
modality="text->text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="Unknown",
instruct_type=None,
),
pricing=Pricing(
prompt=input_price / 1_000_000,
completion=output_price / 1_000_000,
request=0.0,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
),
)
if cl := ppqai_model.context_length:
or_model.context_length = cl
models.append(or_model)
else:
input_price = 0.0
if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get(
"input_per_1M", 0.0
)
elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens
output_price = 0.0
if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get(
"output_per_1M", 0.0
)
elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens
models.append(
Model(
id=ppqai_model.id,
name=ppqai_model.name,
created=ppqai_model.created_at // 1000,
description=f"{ppqai_model.provider or 'PPQ.AI'} model",
context_length=ppqai_model.context_length,
architecture=Architecture(
modality="text->text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="Unknown",
instruct_type=None,
),
pricing=Pricing(
prompt=input_price / 1_000_000,
completion=output_price / 1_000_000,
request=0.0,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
),
)
except Exception as e:
logger.warning(
"Failed to parse PPQ.AI model",
extra={
"model_id": model_data.get("id", "unknown"),
"error": str(e),
"error_type": type(e).__name__,
},
)
except Exception as e:
logger.warning(
"Failed to parse PPQ.AI model",
extra={
"model_id": model_data.get("id", "unknown"),
"error": str(e),
"error_type": type(e).__name__,
},
)
return models
except Exception as e:
logger.error(
"Error fetching models from PPQ.AI",
extra={"error": str(e), "error_type": type(e).__name__},
)
return []
return models
async def on_upstream_error_redirect(
self, status_code: int, error_message: str
@@ -360,8 +432,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
)
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.get(url, headers=headers)
response.raise_for_status()
response = await _safe_read_request(client, "GET", url, headers=headers)
status_data = response.json()
is_paid = status_data.get("status") == "Settled"
@@ -460,8 +531,9 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
logger.debug("Checking PPQ.AI account balance", extra={"url": url})
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.post(url, headers=headers, json={})
response.raise_for_status()
response = await _safe_read_request(
client, "POST", url, headers=headers, json={}
)
balance_data = response.json()
logger.debug(
+15 -15
View File
@@ -17,6 +17,8 @@ from __future__ import annotations
from dataclasses import dataclass, field
from ..payment.rates import coerce_rate
@dataclass
class ResolvedPricing:
@@ -65,11 +67,8 @@ def estimate_context_length(model_id: str) -> int:
def _as_float(value: object) -> float | None:
"""OpenRouter reports prices as strings; coerce, ``None`` if unparseable."""
try:
return float(value) # type: ignore[arg-type]
except (TypeError, ValueError):
return None
"""OpenRouter reports prices as strings; coerce, ``None`` if not a real rate."""
return coerce_rate(value)
def _as_int(value: object) -> int | None:
@@ -78,23 +77,24 @@ def _as_int(value: object) -> int | None:
def _from_litellm(model_id: str) -> ResolvedPricing | None:
# Lazy import so the resolver stays import-light and shares the exact
# lookup semantics used by cache-rate backfill.
# Lazy import so the resolver shares the exact lookup semantics used by
# cache-rate backfill without importing the models module at load time.
from ..payment.models import litellm_cost_entry
info = litellm_cost_entry(model_id)
if info is None:
return None
prompt = info.get("input_cost_per_token")
completion = info.get("output_cost_per_token")
if not isinstance(prompt, (int, float)) or not isinstance(completion, (int, float)):
prompt = coerce_rate(info.get("input_cost_per_token"))
completion = coerce_rate(info.get("output_cost_per_token"))
if prompt is None or completion is None:
return None
# A both-zero entry is litellm listing a model without a real price (free
# moderation/rerank tiers do this) — treating 0/0 as resolved would serve
# the model for free. Reject it (and any negative) so the caller falls
# through, mirroring async_fetch_openrouter_models' _has_valid_pricing.
if prompt < 0 or completion < 0 or (prompt == 0 and completion == 0):
# the model for free. Reject it so the caller falls through, mirroring
# async_fetch_openrouter_models' _has_valid_pricing. Coercion runs first:
# `NaN` would defeat this guard on its own, every comparison being False.
if prompt == 0 and completion == 0:
return None
input_modalities = ["text"]
@@ -102,8 +102,8 @@ def _from_litellm(model_id: str) -> ResolvedPricing | None:
input_modalities.append("image")
return ResolvedPricing(
prompt=float(prompt),
completion=float(completion),
prompt=prompt,
completion=completion,
# max_input_tokens is the context window; max_tokens is litellm's
# completion cap (it tracks max_output_tokens for ~94% of models), so
# it is never a context source. A missing window falls to the id-based
+197
View File
@@ -0,0 +1,197 @@
"""Map client reasoning/thinking effort onto a model's allowlist.
OpenRouter (and some other catalogs) publish per-model reasoning metadata:
which effort levels are legal, the default, and whether reasoning is
mandatory. Clients still send the generic OpenAI / Anthropic shapes
(``reasoning_effort``, ``reasoning.effort``, ``thinking``). This module
normalizes those into a supported effort and writes the fields the
upstream actually accepts, instead of dropping the parameter or
forwarding a value the model rejects.
"""
from __future__ import annotations
from typing import Any
from ..payment.models import Model, Reasoning
# Highest first. Unknown values are treated as unranked.
EFFORT_RANK: tuple[str, ...] = (
"max",
"xhigh",
"high",
"medium",
"low",
"minimal",
"none",
)
_RANK_INDEX: dict[str, int] = {name: i for i, name in enumerate(EFFORT_RANK)}
_REASONING_KEYS = ("reasoning", "reasoning_effort", "thinking")
def _normalize_effort(value: object) -> str | None:
if not isinstance(value, str):
return None
cleaned = value.strip().lower()
return cleaned or None
def closest_supported_effort(
requested: str | None,
supported: list[str],
*,
default_effort: str | None = None,
mandatory: bool = False,
) -> str | None:
"""Pick a legal effort for ``requested``.
Exact match wins. Otherwise the nearest rank in ``EFFORT_RANK`` is
used (preferring the higher neighbour on a tie). ``none`` is rejected
when ``mandatory`` is set. Missing / unmapped requests fall back to
``default_effort``, then the highest remaining supported level.
"""
allowed_efforts: list[str] = [
normalized
for item in supported
if (normalized := _normalize_effort(item)) is not None
]
if mandatory:
allowed_efforts = [item for item in allowed_efforts if item != "none"]
if not allowed_efforts:
if mandatory:
return _normalize_effort(default_effort)
return _normalize_effort(requested) or _normalize_effort(default_effort)
default = _normalize_effort(default_effort)
if default not in allowed_efforts:
default = allowed_efforts[0]
requested_norm = _normalize_effort(requested)
if requested_norm is None or (requested_norm == "none" and mandatory):
return default
if requested_norm in allowed_efforts:
return requested_norm
if requested_norm not in _RANK_INDEX:
return default
target = _RANK_INDEX[requested_norm]
return min(
allowed_efforts,
key=lambda effort: (
abs(_RANK_INDEX.get(effort, 10_000) - target),
_RANK_INDEX.get(effort, 10_000),
),
)
def resolve_effort(requested: str | None, reasoning: Reasoning | None) -> str | None:
"""Map ``requested`` through ``reasoning`` metadata when present."""
if reasoning is None:
return _normalize_effort(requested)
supported = reasoning.supported_efforts or []
return closest_supported_effort(
requested,
supported,
default_effort=reasoning.default_effort,
mandatory=bool(reasoning.mandatory),
)
def _effort_from_thinking(thinking: object) -> str | None:
if not isinstance(thinking, dict):
return None
effort = _normalize_effort(thinking.get("effort"))
if effort:
return effort
thinking_type = _normalize_effort(thinking.get("type"))
if thinking_type in {"disabled", "none"}:
return "none"
return None
def extract_requested_effort(data: dict[str, Any]) -> str | None:
"""Best-effort effort string from the OpenAI / Anthropic request shapes."""
if isinstance(data.get("reasoning"), dict):
nested = _normalize_effort(data["reasoning"].get("effort"))
if nested:
return nested
top_level = _normalize_effort(data.get("reasoning_effort"))
if top_level:
return top_level
return _effort_from_thinking(data.get("thinking"))
def _request_mentions_reasoning(data: dict[str, Any]) -> bool:
return any(key in data for key in _REASONING_KEYS)
def apply_reasoning_effort(
data: dict[str, Any],
model: Model,
*,
drop_thinking: bool = False,
) -> bool:
"""Rewrite ``data`` in place so effort matches the model allowlist.
Returns True when ``data`` changed. Leaves the body alone when the
caller did not send a reasoning field and the model does not require
one. Existing ``reasoning`` object keys (``max_tokens``, ``exclude``,
``enabled``) are preserved; only ``effort`` is mapped.
``drop_thinking`` is for OpenAI-compatible backends that reject the
Anthropic ``thinking`` object: the effort is lifted onto
``reasoning_effort`` / ``reasoning.effort`` and ``thinking`` is removed.
"""
if not isinstance(data, dict):
return False
reasoning_meta = getattr(model, "reasoning", None)
mentioned = _request_mentions_reasoning(data)
if not mentioned and not (reasoning_meta and reasoning_meta.mandatory):
return False
resolved = resolve_effort(extract_requested_effort(data), reasoning_meta)
changed = False
if drop_thinking and "thinking" in data:
data.pop("thinking", None)
changed = True
if resolved is None:
return changed
if "reasoning_effort" in data:
if data.get("reasoning_effort") != resolved:
data["reasoning_effort"] = resolved
changed = True
elif drop_thinking or (reasoning_meta and reasoning_meta.mandatory):
# Invent the OpenAI-shaped field when we stripped Anthropic
# ``thinking``, or when the model will reject a request with no
# effort at all.
if not isinstance(data.get("reasoning"), dict):
data["reasoning_effort"] = resolved
changed = True
existing = data.get("reasoning")
if isinstance(existing, dict):
if existing.get("effort") != resolved:
data["reasoning"] = {**existing, "effort": resolved}
changed = True
elif reasoning_meta and reasoning_meta.mandatory and "reasoning_effort" not in data:
data["reasoning"] = {"effort": resolved}
changed = True
return changed
def adapt_messages_body_for_litellm(data: dict[str, Any], model: Model) -> None:
"""Convert Anthropic ``thinking`` into OpenAI-shaped effort for litellm.
Litellm's Anthropic-messages adapter talking to an OpenAI-compatible
upstream will 400 on ``thinking``. Lift the effort onto
``reasoning_effort`` and drop the Anthropic-only object.
"""
apply_reasoning_effort(data, model, drop_thinking=True)
+30 -3
View File
@@ -84,13 +84,33 @@ def extract_error_message(response: Response) -> str:
return ""
def strip_unsupported_param(
body: dict, error_message: str
) -> tuple[dict, str] | None:
# Spend-shaping fields bound how much work — and therefore cost — the upstream
# may perform. The reservation was priced with these caps in place; dropping one
# and retrying would let the request run uncapped (or fan out) and bill above the
# caller's authorization. When the upstream names one of these, decline the strip
# and let the error propagate. Matched case-insensitively.
_SPEND_SHAPING_PARAMS = frozenset(
{
"max_tokens",
"max_completion_tokens",
"max_output_tokens",
"max_tokens_to_sample",
"n",
"best_of",
}
)
def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None:
"""Drop a top-level param the upstream named as unsupported/deprecated.
Returns ``(new_body, param)`` (a new dict, original untouched) when the
error names a top-level param present in the body, otherwise ``None``.
Spend-shaping fields (output caps, fan-out counts) are never stripped:
removing one after the reservation was priced would uncap the retry and
overcharge. When the upstream names such a field, decline so the original
error propagates rather than silently resizing the request's cost.
"""
match = _UNSUPPORTED_PARAM_RE.search(error_message)
if not match:
@@ -98,6 +118,13 @@ def strip_unsupported_param(
param = match.group("param")
if param not in body:
return None
if param.lower() in _SPEND_SHAPING_PARAMS:
logger.warning(
"Upstream rejected spend-shaping param '%s'; refusing to strip it "
"(retrying uncapped would overcharge) — surfacing the error",
param,
)
return None
new_body = {k: v for k, v in body.items() if k != param}
return new_body, param
+12 -1
View File
@@ -1,6 +1,6 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Optional
import httpx
from fastapi import Request
@@ -28,6 +28,7 @@ logger = get_logger(__name__)
class TinfoilModelPricing(BaseModel):
inputTokenPricePer1M: float = 0.0
outputTokenPricePer1M: float = 0.0
cachedInputTokenPricePer1M: Optional[float] = None
requestPrice: float = 0.0
@@ -186,6 +187,14 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider):
output_price = tf.pricing.outputTokenPricePer1M
request_price = tf.pricing.requestPrice
# Tinfoil bills cache reads at the cached rate when the
# model exposes one, otherwise at the full input rate.
# Cache writes are never priced separately — a miss is
# just regular input prefill.
cached_price = tf.pricing.cachedInputTokenPricePer1M
if cached_price is None or cached_price <= 0.0:
cached_price = input_price
modality = "text->text"
input_modalities = ["text"]
output_modalities = ["text"]
@@ -214,6 +223,8 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider):
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
input_cache_read=cached_price / 1_000_000,
input_cache_write=input_price / 1_000_000,
),
)
)
+28 -10
View File
@@ -20,11 +20,12 @@ from urllib.parse import urlsplit
import h11
from ..core import get_logger
from ..core.exceptions import EhbpTimeoutError
logger = get_logger(__name__)
_READ_BUFSIZE = 65536
_DEFAULT_TIMEOUT_SECONDS = 30.0
_DEFAULT_TIMEOUT_SECONDS = 60.0
_DEFAULT_CLOSE_TIMEOUT_SECONDS = 1.0
_DEFAULT_MAX_RESPONSE_BYTES = 25 * 1024 * 1024
_HOP_BY_HOP_HEADERS = {
@@ -100,10 +101,15 @@ async def forward_with_trailer(
headers = _strip_hop_by_hop_headers(headers)
ssl_ctx = ssl.create_default_context()
reader, writer = await asyncio.wait_for(
asyncio.open_connection(host, port, ssl=ssl_ctx),
timeout=timeout_seconds,
)
try:
reader, writer = await asyncio.wait_for(
asyncio.open_connection(host, port, ssl=ssl_ctx),
timeout=timeout_seconds,
)
except asyncio.TimeoutError as exc:
raise EhbpTimeoutError(
f"EHBP upstream {host} timed out after {timeout_seconds:g}s connecting"
) from exc
try:
# Build HTTP/1.1 request
@@ -126,7 +132,13 @@ async def forward_with_trailer(
request_data += body
writer.write(request_data)
await asyncio.wait_for(writer.drain(), timeout=timeout_seconds)
try:
await asyncio.wait_for(writer.drain(), timeout=timeout_seconds)
except asyncio.TimeoutError as exc:
raise EhbpTimeoutError(
f"EHBP upstream {host} timed out after "
f"{timeout_seconds:g}s sending request"
) from exc
# Parse response with h11
conn = h11.Connection(h11.CLIENT)
@@ -140,10 +152,16 @@ async def forward_with_trailer(
event = conn.next_event()
if event is h11.NEED_DATA:
data = await asyncio.wait_for(
reader.read(_READ_BUFSIZE),
timeout=timeout_seconds,
)
try:
data = await asyncio.wait_for(
reader.read(_READ_BUFSIZE),
timeout=timeout_seconds,
)
except asyncio.TimeoutError as exc:
raise EhbpTimeoutError(
f"EHBP upstream {host} timed out after "
f"{timeout_seconds:g}s waiting for response data"
) from exc
conn.receive_data(data if data else b"")
continue
+261 -856
View File
File diff suppressed because it is too large Load Diff
+89
View File
@@ -0,0 +1,89 @@
"""Redeem a cashu token into a balance and pay it out to a Lightning address.
Usage:
python scripts/refund_token_to_lightning.py <cashu-token> <lightning-address> [--url http://localhost:8000]
Steps:
1. POST /v1/balance/create redeems the token into a fresh API key
2. POST /v1/balance/refund pays the full balance to the Lightning address
A 502 from the refund means the melt was dispatched but unconfirmed; the
balance is withheld until the server reconciles it. Re-run with the printed
API key to check whether it settled.
"""
import argparse
import ipaddress
import sys
from urllib.parse import urlparse
import httpx
def _is_loopback(host: str) -> bool:
if host == "localhost":
return True
try:
return ipaddress.ip_address(host.strip("[]")).is_loopback
except ValueError:
return False
def check_url(url: str) -> str:
"""Reject a URL that would put the token and the API key on the wire."""
parsed = urlparse(url)
if parsed.scheme == "https":
return url
if parsed.scheme == "http" and _is_loopback(parsed.hostname or ""):
return url
raise SystemExit(
f"Refusing to send a cashu token and bearer key to {url!r}: "
"use https, or http only for a loopback host."
)
def create_balance(client: httpx.Client, token: str) -> str:
response = client.post("/v1/balance/create", json={"initial_balance_token": token})
response.raise_for_status()
data = response.json()
print(f"Redeemed token: balance {data['balance']} msats, key {data['api_key']}")
return str(data["api_key"])
def refund_to_lightning(client: httpx.Client, api_key: str, address: str) -> dict:
response = client.post(
"/v1/balance/refund",
headers={"Authorization": f"Bearer {api_key}"},
json={"lightning_address": address},
)
if response.status_code >= 400:
print(f"Refund failed ({response.status_code}): {response.text}")
sys.exit(1)
return dict(response.json())
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
parser.add_argument("token", help="cashu token, or sk-... key from a prior run")
parser.add_argument("lightning_address", help="Lightning address or LNURL")
parser.add_argument("--url", default="http://localhost:8000", help="routstr URL")
args = parser.parse_args()
with httpx.Client(base_url=check_url(args.url), timeout=120.0) as client:
api_key = (
args.token
if args.token.startswith("sk-")
else create_balance(client, args.token)
)
result = refund_to_lightning(client, api_key, args.lightning_address)
amount = result.get("sats") or result.get("msats")
unit = "sats" if "sats" in result else "msats"
print(
f"Refund {result['refund_id']} {result['status']}: "
f"{amount} {unit} -> {result.get('recipient', args.lightning_address)}"
)
if __name__ == "__main__":
main()
+18
View File
@@ -7,9 +7,27 @@ absent one) override this per-test via ``monkeypatch``.
"""
import os
from typing import Iterator
import pytest
# Valid Fernet keys; KEY_A is the suite default, KEY_B is for wrong-key tests.
TEST_SECRET_KEY = "l_Tkp-7xmjcQ-IFhr6qhILrU8HPRbEmYMrfSbo_5srU="
TEST_SECRET_KEY_ALT = "_Teyrky_iToeDK51Tj1FsI9MJ340_cqKGmeher-a7MQ="
os.environ.setdefault("ROUTSTR_SECRET_KEY", TEST_SECRET_KEY)
@pytest.fixture(autouse=True)
def _isolate_redemption_negative_cache() -> Iterator[None]:
"""Clear the process-wide negative cache between tests.
The cache deliberately persists terminal redemption failures across
requests; without this fixture a test that burns a token would poison
every later test reusing the same token string.
"""
from routstr.redemption_cache import redemption_negative_cache
redemption_negative_cache.clear()
yield
redemption_negative_cache.clear()
+15 -3
View File
@@ -1,7 +1,7 @@
import asyncio
import json
import os
from typing import Any, AsyncGenerator, Callable, Dict, List, Optional, Tuple
from typing import Any, AsyncGenerator, Callable, Dict, Iterator, List, Optional, Tuple
from unittest.mock import MagicMock, patch
import pytest
@@ -68,6 +68,14 @@ os.environ.pop("ADMIN_PASSWORD", None)
from routstr.core.db import ApiKey, get_session # noqa: E402
from routstr.core.main import app, lifespan # noqa: E402
from routstr.mint import MintRateGuard # noqa: E402
@pytest.fixture(autouse=True)
def isolate_mint_rate_guards() -> Iterator[None]:
MintRateGuard._guards.clear()
yield
MintRateGuard._guards.clear()
@pytest.fixture(scope="session")
@@ -508,6 +516,10 @@ async def integration_app(
# Copy all routes from the main app
test_app.router = app.router
# ...and its exception handlers, so a request that fails here fails the way
# it would in production rather than escaping as a bare exception.
test_app.exception_handlers.update(app.exception_handlers)
# Override the get_session dependency
async def override_get_session() -> AsyncGenerator[AsyncSession, None]:
yield integration_session
@@ -542,8 +554,8 @@ async def integration_app(
patch("routstr.wallet.send_to_lnurl", testmint_wallet.send_to_lnurl),
patch("routstr.wallet.recieve_token", testmint_wallet.redeem_token),
patch("routstr.wallet.get_balance", testmint_wallet.get_balance),
patch("routstr.balance.send_token", testmint_wallet.send_token),
patch("routstr.balance.send_to_lnurl", testmint_wallet.send_to_lnurl),
patch("routstr.refund.send_token", testmint_wallet.send_token),
patch("routstr.refund.send_to_lnurl", testmint_wallet.send_to_lnurl),
patch("websockets.connect") as mock_websockets,
patch("routstr.payment.price.btc_usd_price", return_value=50000.0),
patch("routstr.payment.price.sats_usd_price", return_value=0.0005),
@@ -0,0 +1,476 @@
"""Admin write edge: a rate that is not a number never becomes a stored price.
A billable rate is usable only when it is finite and non-negative. The admin
model endpoints are an entry point for rates the node will later bill on, and
they accept whatever a client sends: ``json`` parses the bare ``NaN``/
``Infinity`` literals into real floats and overflows ``1e999`` to ``inf``, a
non-numeric string coerced silently to ``$0``, and a negative rate is truthy so
it read back as a chargeable price that bills a negative amount.
These tests assert the edge answers a malformed rate with a 422 — a client bug
reported as a client bug — rather than persisting it or failing as a 500, and
that the operator can still open the listing that shows the row needing repair.
"""
from __future__ import annotations
import json
from datetime import datetime, timedelta, timezone
import pytest
from httpx import AsyncClient
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.admin import admin_sessions
from routstr.core.db import ModelRow, UpstreamProviderRow
from routstr.proxy import reinitialize_upstreams
def _admin_headers() -> dict[str, str]:
token = "test-admin-rate-validation-token"
admin_sessions[token] = int(
(datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp()
)
return {"Authorization": f"Bearer {token}"}
def _pricing(**overrides: object) -> dict[str, object]:
pricing: dict[str, object] = {
"prompt": 1.4e-7,
"completion": 2.8e-7,
"request": 0.0,
"image": 0.0,
"web_search": 0.0,
"internal_reasoning": 0.0,
"input_cache_read": 0.0,
"input_cache_write": 0.0,
}
pricing.update(overrides)
return pricing
def _payload(
provider_id: int,
*,
model_id: str = "rate-model",
pricing: dict[str, object] | None = None,
) -> dict[str, object]:
return {
"id": model_id,
"name": "Rate Model",
"description": "d",
"created": 0,
"context_length": 128000,
"architecture": {
"modality": "text",
"input_modalities": ["text"],
"output_modalities": ["text"],
"tokenizer": "unknown",
"instruct_type": None,
},
"pricing": pricing if pricing is not None else _pricing(),
"per_request_limits": None,
"top_provider": None,
"upstream_provider_id": provider_id,
"canonical_slug": None,
"alias_ids": [],
"enabled": True,
"forwarded_model_id": model_id,
}
async def _make_provider(session: AsyncSession) -> int:
provider = UpstreamProviderRow(
provider_type="generic",
base_url="https://rate-upstream.example/v1",
api_key="test-key",
provider_fee=1.0,
)
session.add(provider)
await session.commit()
await session.refresh(provider)
await reinitialize_upstreams()
assert provider.id is not None
return provider.id
def _raw_model_body(provider_id: int, model_id: str, prompt_literal: str) -> str:
"""A request body built as text, so it can carry a literal ``json`` accepts
but Python's own encoder would refuse to produce."""
return (
f'{{"id": "{model_id}", "name": "raw", "description": "d", "created": 0,'
' "context_length": 8192, "architecture": {"modality": "text"},'
f' "pricing": {{"prompt": {prompt_literal}, "completion": 2.8e-7}},'
f' "upstream_provider_id": {provider_id}}}'
)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_negative_price_is_rejected(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
"""A negative rate is not a valid price — accepting it would persist a row
that bills a negative amount, which settlement subtracts from the balance.
Being truthy, it also reads back as a chargeable price. Reject at the edge
rather than silently storing it."""
provider_id = await _make_provider(integration_session)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=_admin_headers(),
json=_payload(provider_id, model_id="neg-price", pricing=_pricing(prompt=-1.0)),
)
assert resp.status_code == 422
assert await integration_session.get(ModelRow, ("neg-price", provider_id)) is None
@pytest.mark.integration
@pytest.mark.asyncio
async def test_malformed_price_string_is_rejected(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
"""A present non-numeric rate is a client bug: it coerces to ``$0`` on the
read path, producing an unpriced-looking row indistinguishable from a
deliberate free price. Surface it as a 422 instead of accepting it."""
provider_id = await _make_provider(integration_session)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=_admin_headers(),
json=_payload(
provider_id, model_id="bad-price", pricing=_pricing(prompt="oops")
),
)
assert resp.status_code == 422
assert await integration_session.get(ModelRow, ("bad-price", provider_id)) is None
@pytest.mark.integration
@pytest.mark.asyncio
async def test_boolean_price_is_rejected(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
"""A JSON ``true`` coerces to a finite, positive ``1.0`` — a dollar per
token — so it passes every numeric guard. The write edge asks the same
coercion the catalog readers do, and answers a 422."""
provider_id = await _make_provider(integration_session)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=_admin_headers(),
json=_payload(
provider_id, model_id="bool-price", pricing=_pricing(prompt=True)
),
)
assert resp.status_code == 422
assert await integration_session.get(ModelRow, ("bool-price", provider_id)) is None
@pytest.mark.integration
@pytest.mark.asyncio
@pytest.mark.parametrize(
("model_id", "pricing"),
[
("null-prompt", _pricing(prompt=None)),
("null-aux-rate", _pricing(image=None)),
("no-prompt", {k: v for k, v in _pricing().items() if k != "prompt"}),
],
ids=["null-required", "null-auxiliary", "absent-required"],
)
async def test_a_rate_that_is_not_there_is_rejected(
model_id: str,
pricing: dict[str, object],
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""A rate given as ``null``, or a required rate left out, is not a price.
``dict.get`` cannot tell the two apart and skipped both, so a row
``Pricing`` cannot parse was committed and the response that reads it back
raised.
"""
provider_id = await _make_provider(integration_session)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=_admin_headers(),
json=_payload(provider_id, model_id=model_id, pricing=pricing),
)
assert resp.status_code == 422
assert await integration_session.get(ModelRow, (model_id, provider_id)) is None
@pytest.mark.integration
@pytest.mark.asyncio
async def test_an_absent_auxiliary_rate_is_still_accepted(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
"""Only ``prompt`` and ``completion`` are required; the rest carry defaults,
and a payload that omits them must still be accepted."""
provider_id = await _make_provider(integration_session)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=_admin_headers(),
json=_payload(
provider_id,
model_id="lean-price",
pricing={"prompt": 1.4e-7, "completion": 2.8e-7},
),
)
assert resp.status_code == 200
assert await integration_session.get(ModelRow, ("lean-price", provider_id))
@pytest.mark.integration
@pytest.mark.asyncio
async def test_numeric_string_price_is_still_accepted(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
"""The stored pricing JSON has always accepted numeric strings, and the UI
round-trips rates through text fields. Rejecting a *malformed* rate must not
also reject a well-formed one that arrives spelled as a string."""
provider_id = await _make_provider(integration_session)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=_admin_headers(),
json=_payload(
provider_id, model_id="string-price", pricing=_pricing(prompt="0.000005")
),
)
assert resp.status_code == 200
row = await integration_session.get(ModelRow, ("string-price", provider_id))
assert row is not None
assert json.loads(row.pricing)["prompt"] == "0.000005"
def test_non_finite_price_is_rejected_by_the_write_model() -> None:
"""``NaN``/``±inf`` are not billable rates: the carrier every write endpoint
shares must reject them before they can be persisted and read back as a
chargeable price."""
from pydantic import ValidationError
from routstr.core.admin import ModelCreate
for bad in (float("nan"), float("inf"), float("-inf")):
with pytest.raises(ValidationError):
ModelCreate.model_validate(
_payload(1, model_id="nonfinite", pricing=_pricing(prompt=bad))
)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_oversized_integer_price_is_rejected(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
"""A JSON integer too large for a float is a client bug, not a server fault.
``float()`` raises ``OverflowError`` for it, and pydantic converts only
``ValueError``/``AssertionError`` into validation errors, so it escaped the
edge as a 500. It must be answered with the same 422 as every other
unusable rate.
"""
provider_id = await _make_provider(integration_session)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers={**_admin_headers(), "Content-Type": "application/json"},
content=_raw_model_body(provider_id, "huge-price", "9" * 400),
)
assert resp.status_code == 422
assert await integration_session.get(ModelRow, ("huge-price", provider_id)) is None
@pytest.mark.integration
@pytest.mark.asyncio
async def test_non_finite_literal_price_is_rejected(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
"""A bare ``Infinity``/``NaN`` literal gets the same 422 as any other rate.
``json`` accepts both literals, so the edge sees a real float and rejects
it — but pydantic echoes the offending value back in the error's ``input``
field, and the response encoder runs with ``allow_nan=False``. Serializing
that reply raised "Out of range float values are not JSON compliant", so the
422 escaped as a 500 and reported a client's bad rate as a server fault.
"""
provider_id = await _make_provider(integration_session)
for literal in ("Infinity", "-Infinity", "NaN"):
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers={**_admin_headers(), "Content-Type": "application/json"},
content=_raw_model_body(provider_id, "odd-price", literal),
)
assert resp.status_code == 422, literal
assert (
await integration_session.get(ModelRow, ("odd-price", provider_id)) is None
)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_non_finite_literal_price_is_rejected_in_batch_override(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
"""The batch path shares the same carrier, so it must answer 422 too."""
provider_id = await _make_provider(integration_session)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/batch-override",
headers={**_admin_headers(), "Content-Type": "application/json"},
content=(
'{"models": ['
+ _raw_model_body(provider_id, "odd-batch", "Infinity")
+ "]}"
),
)
assert resp.status_code == 422
assert await integration_session.get(ModelRow, ("odd-batch", provider_id)) is None
@pytest.mark.integration
@pytest.mark.asyncio
async def test_admin_model_listing_shows_a_non_finite_stored_rate(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
"""The operator must be able to see the rate that needs fixing.
The admin listing is the one view that still carries a row the served
catalog holds back, and its encoder rendered a stored ``Infinity`` as
``null`` — indistinguishable from a rate the row never carried.
"""
provider_id = await _make_provider(integration_session)
integration_session.add(
ModelRow(
id="inf-rate",
name="inf-rate",
description="d",
created=0,
context_length=8192,
architecture=json.dumps(
{
"modality": "text",
"input_modalities": ["text"],
"output_modalities": ["text"],
"tokenizer": "unknown",
"instruct_type": None,
}
),
pricing=json.dumps({"prompt": float("inf"), "completion": 2e-06}),
upstream_provider_id=provider_id,
enabled=True,
forwarded_model_id="inf-rate",
)
)
await integration_session.commit()
resp = await integration_client.get(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=_admin_headers(),
)
assert resp.status_code == 200
listed = {m["id"]: m for m in resp.json()["db_models"]}
assert listed["inf-rate"]["pricing"]["prompt"] == "inf"
@pytest.mark.integration
@pytest.mark.asyncio
async def test_malformed_auxiliary_rate_is_rejected(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
"""Validation spans every billable rate, not just the token rates.
``prompt``/``completion`` are the rates most prices are built from, but the
request, image, search, reasoning and cache rates are billed too. A negative
or non-finite value in any of them is the same defect and must be answered
the same way.
"""
provider_id = await _make_provider(integration_session)
for field, bad in (
("request", -1.0),
("image", -0.5),
("web_search", float("inf")),
("internal_reasoning", float("nan")),
("input_cache_read", -1e-06),
("input_cache_write", float("-inf")),
("completion", -1.0),
):
# Send raw bytes rather than `json=`: httpx>=0.28 refuses to encode
# non-finite floats itself (allow_nan=False), but the point of this
# test is that the SERVER answers the bare NaN/Infinity literals
# with a 422, so the literals must still reach it.
body = json.dumps(
_payload(
provider_id, model_id="aux-rate", pricing=_pricing(**{field: bad})
)
).encode("utf-8")
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers={**_admin_headers(), "content-type": "application/json"},
content=body,
)
assert resp.status_code == 422, field
assert (
await integration_session.get(ModelRow, ("aux-rate", provider_id)) is None
), field
@pytest.mark.integration
@pytest.mark.asyncio
async def test_admin_single_model_shows_a_non_finite_stored_rate(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
"""The single-model view answers like the listing it is opened from.
It is the other view of a row the served-catalog backstop holds back, so it
has the same duty to name the rate that needs fixing rather than rendering
it as ``null``.
"""
provider_id = await _make_provider(integration_session)
integration_session.add(
ModelRow(
id="inf-one",
name="inf-one",
description="d",
created=0,
context_length=8192,
architecture=json.dumps(
{
"modality": "text",
"input_modalities": ["text"],
"output_modalities": ["text"],
"tokenizer": "unknown",
"instruct_type": None,
}
),
pricing=json.dumps({"prompt": float("inf"), "completion": 2e-06}),
upstream_provider_id=provider_id,
enabled=True,
forwarded_model_id="inf-one",
)
)
await integration_session.commit()
resp = await integration_client.get(
f"/admin/api/upstream-providers/{provider_id}/models/inf-one",
headers=_admin_headers(),
)
assert resp.status_code == 200
assert resp.json()["pricing"]["prompt"] == "inf"
-190
View File
@@ -1,190 +0,0 @@
import asyncio
import secrets
from typing import Any
import pytest
from fastapi import HTTPException
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.auth import adjust_payment_for_tokens, pay_for_request
from routstr.balance import ChildKeyRequest, create_child_key
from routstr.core.db import ApiKey, create_session
from routstr.core.settings import settings
@pytest.mark.asyncio
async def test_child_key_flow(integration_session: AsyncSession) -> None:
# 1. Create a parent key with balance
parent_raw = "parent_test_key_" + secrets.token_hex(4)
parent_key = ApiKey(
hashed_key=parent_raw,
balance=10000, # 10 sats
)
integration_session.add(parent_key)
await integration_session.commit()
await integration_session.refresh(parent_key)
# Mock settings
settings.child_key_cost = 1000 # 1 sat
# 2. Call create_child_key
result = await create_child_key(
ChildKeyRequest(count=1), parent_key, integration_session
)
assert "api_keys" in result
assert result["cost_msats"] == 1000
assert result["parent_balance"] == 9000
child_key_raw = result["api_keys"][0][3:] # remove sk-
# 3. Verify child key exists in DB
child_key_db = await integration_session.get(ApiKey, child_key_raw)
assert child_key_db is not None
assert child_key_db.parent_key_hash == parent_key.hashed_key
assert child_key_db.balance == 0
# 4. Test payment with child key
cost = 500
await pay_for_request(child_key_db, cost, integration_session)
# Refresh keys
await integration_session.refresh(parent_key)
await integration_session.refresh(child_key_db)
# Parent should be charged
assert parent_key.reserved_balance == 500
assert parent_key.total_requests == 1
# Child should have total_requests incremented
assert child_key_db.total_requests == 1
# 5. Test adjustment
response_data = {"model": "test-model", "usage": {"total_tokens": 10}}
# Mock calculate_cost
import routstr.auth
from routstr.payment.cost_calculation import CostData
async def mock_calculate_cost(*args: Any, **kwargs: Any) -> CostData:
return CostData(
base_msats=0, input_msats=200, output_msats=200, total_msats=400
)
# Patch calculate_cost
original_calculate_cost = routstr.auth.calculate_cost
routstr.auth.calculate_cost = mock_calculate_cost
try:
adjustment = await adjust_payment_for_tokens(
child_key_db, response_data, integration_session, 500, None, None
)
assert adjustment["total_msats"] == 400
# Refresh keys
await integration_session.refresh(parent_key)
await integration_session.refresh(child_key_db)
# Parent should have updated balance and total_spent
assert parent_key.reserved_balance == 0
assert parent_key.balance == 9000 - 400
assert (
parent_key.total_spent == 1400
) # 1000 for child key creation + 400 for request
# Child should also have total_spent updated
assert child_key_db.total_spent == 400
finally:
routstr.auth.calculate_cost = original_calculate_cost
@pytest.mark.asyncio
async def test_child_key_insufficient_balance(
integration_session: AsyncSession,
) -> None:
parent_key = ApiKey(
hashed_key="poor_parent_" + secrets.token_hex(4),
balance=500,
)
integration_session.add(parent_key)
await integration_session.commit()
await integration_session.refresh(parent_key)
settings.child_key_cost = 1000
with pytest.raises(HTTPException) as exc:
await create_child_key(
ChildKeyRequest(count=1), parent_key, integration_session
)
assert exc.value.status_code == 402
@pytest.mark.asyncio
async def test_concurrent_child_key_creation_is_atomic(
patched_db_engine: None,
) -> None:
"""Two concurrent create_child_key() calls with balance for exactly one must
result in exactly one success and one 402, with the parent balance deducted
only once."""
child_key_cost = 1000
settings.child_key_cost = child_key_cost
parent_hash = f"parent_concurrent_{secrets.token_hex(8)}"
async with create_session() as session:
parent = ApiKey(hashed_key=parent_hash, balance=child_key_cost)
session.add(parent)
await session.commit()
results: list[str] = []
async def attempt() -> None:
async with create_session() as session:
fresh_parent = await session.get(ApiKey, parent_hash)
assert fresh_parent is not None
try:
await create_child_key(ChildKeyRequest(count=1), fresh_parent, session)
results.append("success")
except HTTPException as exc:
assert exc.status_code == 402
results.append("blocked")
await asyncio.gather(attempt(), attempt())
assert sorted(results) == ["blocked", "success"], (
f"Expected exactly one success and one 402, got: {results}"
)
async with create_session() as session:
final = await session.get(ApiKey, parent_hash)
assert final is not None
assert final.balance == 0, (
f"Balance should be fully deducted once: expected 0, got {final.balance}"
)
assert final.total_spent == child_key_cost, (
f"total_spent should equal one deduction: expected {child_key_cost}, "
f"got {final.total_spent}"
)
@pytest.mark.asyncio
async def test_child_key_cannot_create_child(integration_session: AsyncSession) -> None:
parent_key = ApiKey(
hashed_key="parent_" + secrets.token_hex(4),
balance=10000,
)
child_key = ApiKey(
hashed_key="child_" + secrets.token_hex(4),
balance=0,
parent_key_hash=parent_key.hashed_key,
)
integration_session.add(parent_key)
integration_session.add(child_key)
await integration_session.commit()
await integration_session.refresh(child_key)
with pytest.raises(HTTPException) as exc:
await create_child_key(ChildKeyRequest(count=1), child_key, integration_session)
assert exc.value.status_code == 400
assert "Cannot create a child key for another child key" in str(exc.value.detail)
-102
View File
@@ -1,102 +0,0 @@
from typing import Any
import pytest
from httpx import AsyncClient
@pytest.mark.integration
@pytest.mark.asyncio
async def test_wallet_info_returns_child_keys(
integration_client: AsyncClient,
authenticated_client: AsyncClient,
integration_session: Any,
) -> None:
"""Test that GET /v1/wallet/info returns child keys for a parent key"""
# 1. Get parent info to find its hashed_key
response = await authenticated_client.get("/v1/wallet/info")
assert response.status_code == 200
parent_data = response.json()
parent_data["api_key"]
# 2. Create child keys for this parent
# We need to use the parent's authentication for this
child_payload = {"count": 2, "balance_limit": 1000, "balance_limit_reset": "daily"}
create_response = await authenticated_client.post(
"/v1/wallet/child-key", json=child_payload
)
assert create_response.status_code == 200
create_data = create_response.json()
child_keys = create_data["api_keys"]
assert len(child_keys) == 2
# 3. Call /info again and check for child_keys
info_response = await authenticated_client.get("/v1/wallet/info")
assert info_response.status_code == 200
info_data = info_response.json()
assert "child_keys" in info_data
assert len(info_data["child_keys"]) == 2
# Verify child key details
for ck in info_data["child_keys"]:
assert ck["api_key"] in child_keys
assert ck["balance_limit"] == 1000
assert ck["balance_limit_reset"] == "daily"
assert "total_spent" in ck
assert "total_requests" in ck
@pytest.mark.integration
@pytest.mark.asyncio
async def test_wallet_info_child_key_no_child_keys(
integration_client: AsyncClient,
authenticated_client: AsyncClient,
integration_session: Any,
) -> None:
"""Test that GET /v1/wallet/info for a child key does NOT return child_keys"""
# 1. Create a child key
child_payload = {"count": 1}
create_response = await authenticated_client.post(
"/v1/wallet/child-key", json=child_payload
)
assert create_response.status_code == 200
child_key = create_response.json()["api_keys"][0]
# 2. Use the child key to get its info
integration_client.headers["Authorization"] = f"Bearer {child_key}"
info_response = await integration_client.get("/v1/wallet/info")
assert info_response.status_code == 200
info_data = info_response.json()
parent_key = authenticated_client._test_api_key # type: ignore[attr-defined]
parent_key_hash = parent_key.removeprefix("sk-")
assert info_data["is_child"] is True
assert "child_keys" not in info_data
assert "parent_key" not in info_data
assert info_data["parent_key_preview"] == parent_key_hash[:8] + "..."
assert info_data["parent_key_preview"] not in {parent_key, parent_key_hash}
@pytest.mark.integration
@pytest.mark.asyncio
async def test_account_info_root_returns_child_keys(
authenticated_client: AsyncClient,
) -> None:
"""Test that GET / returns child keys for a parent key (root endpoint)"""
# 1. Create a child key
child_payload = {"count": 1}
await authenticated_client.post("/v1/wallet/child-key", json=child_payload)
# 2. Call root endpoint /v1/balance/
# Note: routstr/balance.py defines router = APIRouter()
# and it is included in balance_router with prefix /v1/balance
# The endpoint is @router.get("/")
response = await authenticated_client.get("/v1/balance/")
assert response.status_code == 200
data = response.json()
assert "child_keys" in data
assert len(data["child_keys"]) >= 1
@@ -31,7 +31,7 @@ class TestNetworkFailureScenarios:
AsyncMock(side_effect=ConnectError("Mint service unavailable")),
),
patch(
"routstr.balance.send_token",
"routstr.refund.send_token",
AsyncMock(side_effect=ConnectError("Mint service unavailable")),
),
):
+1 -98
View File
@@ -17,7 +17,7 @@ from httpx import AsyncClient
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ApiKey, ReservationRelease
from routstr.core.db import ReservationRelease
from routstr.payment.models import Architecture, Model, Pricing
from routstr.proxy import refresh_model_maps
from routstr.upstream.base import BaseUpstreamProvider
@@ -531,103 +531,6 @@ async def test_failover_beyond_balance_envelope_is_rejected(
# fallback must be rejected before its upstream is ever contacted.
assert response.status_code == 402
assert [r.url.host for r in sent_requests] == ["cheap.example.com"]
@pytest.fixture
async def three_candidate_child_maps(
patched_db_engine: None,
) -> AsyncGenerator[None, None]:
"""Second candidate cannot fit the child limit; third restores and serves."""
first = _StaticProvider(
CHEAP_BASE_URL,
"key-first",
1.0,
_make_model("dual-model", 0.001, 0.002, max_cost=50.0),
)
too_large = _StaticProvider(
EXPENSIVE_BASE_URL,
"key-too-large",
1.0,
_make_model("dual-model", 0.002, 0.003, max_cost=100.0),
)
third = _StaticProvider(
THIRD_BASE_URL,
"key-third",
1.0,
_make_model("dual-model", 0.003, 0.004, max_cost=50.0),
)
async for _ in _install_providers([first, too_large, third]):
yield
@pytest.mark.integration
@pytest.mark.asyncio
async def test_child_failover_rolls_back_failed_larger_reserve_before_restoring(
authenticated_client: AsyncClient,
three_candidate_child_maps: None,
integration_session: AsyncSession,
) -> None:
"""A failed child guard cannot leak its parent update into restoration."""
key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined]
child = await integration_session.get(ApiKey, key_hash)
assert child is not None
parent = ApiKey(hashed_key="failover-parent", balance=10_000_000)
child.parent_key_hash = parent.hashed_key
child.balance_limit = 75_000
integration_session.add(parent)
integration_session.add(child)
await integration_session.commit()
sent_requests: list[httpx.Request] = []
async def fake_transport(
request: httpx.Request, *args: Any, **kwargs: Any
) -> httpx.Response:
sent_requests.append(request)
return _upstream_response(request)
with (
patch(
"httpx.AsyncHTTPTransport.handle_async_request",
side_effect=fake_transport,
),
patch(
"routstr.payment.cost_calculation.sats_usd_price",
return_value=0.0005,
),
):
response = await authenticated_client.post(
"/v1/chat/completions",
json={
"model": "dual-model",
"messages": [{"role": "user", "content": "hello"}],
},
)
assert response.status_code == 200
# The 100-sat candidate is rejected before forwarding; the third serves.
assert [request.url.host for request in sent_requests] == [
"cheap.example.com",
"third.example.com",
]
await integration_session.refresh(parent)
await integration_session.refresh(child)
assert parent.reserved_balance == 0
assert child.reserved_balance == 0
assert parent.total_spent == response.json()["cost"]["total_msats"]
records = (
await integration_session.exec(
select(ReservationRelease).where(ReservationRelease.key_hash == key_hash)
)
).all()
assert len(records) == 2
assert sorted(record.status for record in records) == ["charged", "released"]
assert len({record.reserved_msats for record in records}) == 1
assert all(record.status != "active" for record in records)
@pytest.fixture
async def raised_envelope_provider_maps(
patched_db_engine: None,
@@ -7,7 +7,7 @@ from unittest.mock import patch
import pytest
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ApiKey
from routstr.core.db import ApiKey, ReservationRelease
from routstr.payment.cost_calculation import CostData
@@ -34,20 +34,30 @@ def _cost_data(total_msats: int) -> CostData:
@pytest.mark.asyncio
async def test_overrun_charges_after_reservation_swept(
async def test_overrun_with_corrupted_aggregate_releases_without_charging(
integration_session: AsyncSession,
) -> None:
"""Overrun finalize must charge even when the reservation was already released."""
from routstr.auth import adjust_payment_for_tokens, pay_for_request
"""An overrun whose aggregate reservation was externally zeroed must not
charge: subtracting the reservation would have to clamp, which can erase
sibling reservations. The reservation is released and the charge dropped.
"""
from routstr import auth
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
deducted_max_cost = 990 # discounted reservation
actual_token_cost = 1000 # actual cost overruns the reservation
# Sweeper has zeroed reserved_balance but left balance untouched.
# Something zeroed reserved_balance under an active durable reservation.
key = _make_key(balance=1000, reserved=0)
key_hash = key.hashed_key
integration_session.add(key)
await integration_session.commit()
await pay_for_request(key, deducted_max_cost, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
key.reserved_balance = 0
integration_session.add(key)
await integration_session.commit()
@@ -61,20 +71,67 @@ async def test_overrun_charges_after_reservation_swept(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
result = await adjust_payment_for_tokens(
key, response_data, integration_session, deducted_max_cost, None, None
)
await integration_session.refresh(key)
assert result["charged_msats"] == 0
integration_session.expunge_all()
key_row = await integration_session.get(ApiKey, key_hash)
assert key_row is not None
assert key.total_spent == actual_token_cost, (
f"Request was not billed (total_spent={key.total_spent}) — free response bug"
assert key_row.total_spent == 0, "corrupted aggregate must not be charged into"
assert key_row.balance == 1000
assert key_row.reserved_balance == 0
# The corrupt reservation must reach a terminal state — an active leftover
# would be renewed by its heartbeat forever and poison stale cleanup.
record = await integration_session.get(ReservationRelease, reservation.release_id)
assert record is not None and record.status == "released"
assert reservation.release_id not in auth._reservation_heartbeats
@pytest.mark.asyncio
async def test_missing_usage_never_turns_reservation_into_charge(
integration_session: AsyncSession,
) -> None:
"""A reservation is an authorization ceiling, not evidence of usage.
Upstream handlers should provide locally estimated usage when possible. If
no measurement or estimate reaches settlement, release the reservation
rather than charging its full value.
"""
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
assert key.balance == 1000 - actual_token_cost, (
f"Balance not charged: {key.balance}"
reserved = 4_000
key = _make_key(balance=10_000, reserved=0)
key_hash = key.hashed_key
integration_session.add(key)
await integration_session.commit()
await pay_for_request(key, reserved, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
# No `usage` key at all — the upstream stream dropped its final usage chunk.
response_data = {"model": "test-model"}
result = await adjust_payment_for_tokens(
key,
response_data,
integration_session,
reserved,
reservation_snapshot=reservation,
)
assert key.balance >= 0
assert key.reserved_balance == 0
assert result["charged_msats"] == 0
integration_session.expunge_all()
key_row = await integration_session.get(ApiKey, key_hash)
assert key_row is not None
assert key_row.total_spent == 0
assert key_row.balance == 10_000
assert key_row.reserved_balance == 0
@pytest.mark.asyncio
+1 -366
View File
@@ -1,14 +1,10 @@
import asyncio
import time
from datetime import datetime, timedelta
import pytest
from fastapi import HTTPException
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.auth import pay_for_request
from routstr.core.db import ApiKey, create_session
from routstr.core.db import ApiKey
@pytest.mark.asyncio
@@ -25,367 +21,6 @@ async def test_key_validity_date(integration_session: AsyncSession) -> None:
assert "expired" in str(excinfo.value).lower()
@pytest.mark.asyncio
async def test_key_balance_limit(integration_session: AsyncSession) -> None:
# 1. Create a key with a balance limit
key = ApiKey(
hashed_key="limited_key", balance=10000, balance_limit=500, total_spent=450
)
integration_session.add(key)
await integration_session.commit()
# 2. Try to pay for a request that exceeds the limit
with pytest.raises(Exception) as excinfo:
await pay_for_request(key, 100, integration_session)
assert "limit exceeded" in str(excinfo.value).lower()
# 3. Try to pay for a request that fits
await pay_for_request(key, 50, integration_session)
await integration_session.refresh(key)
# Note: total_spent is updated in adjust_payment_for_tokens,
# but pay_for_request checks it.
# In our current logic, pay_for_request checks (total_spent + cost) > balance_limit.
@pytest.mark.asyncio
async def test_key_daily_reset_policy(integration_session: AsyncSession) -> None:
# 1. Create a key with a daily reset policy and old reset date
yesterday = int((datetime.now() - timedelta(days=1)).timestamp())
key = ApiKey(
hashed_key="daily_reset_key",
balance=10000,
balance_limit=1000,
balance_limit_reset="daily",
balance_limit_reset_date=yesterday,
total_spent=900,
)
integration_session.add(key)
await integration_session.commit()
# 2. Pay for a request - should trigger reset first because it's a new day
# Request is 200, total_spent is 900. 900+200 > 1000,
# but reset should happen making total_spent 0, then 0+200 < 1000.
await pay_for_request(key, 200, integration_session)
await integration_session.refresh(key)
assert key.total_spent == 0 # Reset in pay_for_request happens before charging
# Wait, the charging logic in pay_for_request increments parent/billing_key's total_requests,
# but total_spent is updated in adjust_payment_for_tokens.
# However, the reset logic sets total_spent to 0.
assert key.balance_limit_reset_date is not None
assert key.balance_limit_reset_date > yesterday
@pytest.mark.asyncio
async def test_periodic_key_reset_job(integration_session: AsyncSession) -> None:
# 1. Create multiple keys needing reset
yesterday = int((datetime.now() - timedelta(days=1)).timestamp())
key1 = ApiKey(
hashed_key="job_reset_key_1",
balance=1000,
balance_limit=1000,
balance_limit_reset="daily",
balance_limit_reset_date=yesterday,
total_spent=500,
)
key2 = ApiKey(
hashed_key="job_reset_key_2",
balance=1000,
balance_limit=1000,
balance_limit_reset="daily",
balance_limit_reset_date=yesterday,
total_spent=800,
)
integration_session.add(key1)
integration_session.add(key2)
await integration_session.commit()
# 2. Run the periodic reset logic manually (mocking the background task loop)
# We can't easily run the actual loop because it has a sleep,
# but we can test the logic inside.
# Implementation of periodic_key_reset logic for testing:
stmt = select(ApiKey).where(ApiKey.balance_limit_reset != None) # noqa: E711
keys = (await integration_session.exec(stmt)).all()
now = int(time.time())
for k in keys:
if k.hashed_key in ["job_reset_key_1", "job_reset_key_2"]:
k.total_spent = 0
k.balance_limit_reset_date = now
integration_session.add(k)
await integration_session.commit()
# 3. Verify resets
await integration_session.refresh(key1)
await integration_session.refresh(key2)
assert key1.total_spent == 0
assert key2.total_spent == 0
@pytest.mark.asyncio
async def test_balance_limit_enforced_atomically_under_concurrency(
patched_db_engine: None,
) -> None:
parent_hash = "parent_limit_atomic"
child_hash = "child_limit_atomic"
cost = 300
async with create_session() as session:
parent = ApiKey(hashed_key=parent_hash, balance=10000)
child = ApiKey(
hashed_key=child_hash,
balance=0,
parent_key_hash=parent_hash,
balance_limit=cost,
total_spent=0,
)
session.add(parent)
session.add(child)
await session.commit()
results: list[str] = []
async def attempt() -> None:
async with create_session() as session:
fresh_child = await session.get(ApiKey, child_hash)
assert fresh_child is not None
try:
await pay_for_request(fresh_child, cost, session)
results.append("success")
except HTTPException as exc:
assert exc.status_code == 402
results.append("blocked")
await asyncio.gather(attempt(), attempt())
assert sorted(results) == ["blocked", "success"], (
f"Expected exactly one success and one 402, got: {results}"
)
async with create_session() as session:
final_child = await session.get(ApiKey, child_hash)
assert final_child is not None
assert final_child.reserved_balance == cost, (
f"Child reserved_balance should equal one reservation, "
f"got {final_child.reserved_balance}"
)
@pytest.mark.asyncio
async def test_parallel_payments_with_parent_and_child_key(
patched_db_engine: None,
) -> None:
parent_hash = "parent_parallel_mixed"
child_hash = "child_parallel_mixed"
cost = 300
async with create_session() as session:
parent = ApiKey(hashed_key=parent_hash, balance=10000)
child = ApiKey(
hashed_key=child_hash,
balance=0,
parent_key_hash=parent_hash,
balance_limit=2 * cost,
)
session.add(parent)
session.add(child)
await session.commit()
async def attempt(key_hash: str) -> str:
async with create_session() as session:
fresh_key = await session.get(ApiKey, key_hash)
assert fresh_key is not None
try:
await pay_for_request(fresh_key, cost, session)
return "success"
except HTTPException as exc:
assert exc.status_code == 402
return "blocked"
results = await asyncio.gather(attempt(parent_hash), attempt(child_hash))
assert results == ["success", "success"], (
f"Both parent and child payments should succeed, got: {results}"
)
async with create_session() as session:
final_parent = await session.get(ApiKey, parent_hash)
final_child = await session.get(ApiKey, child_hash)
assert final_parent is not None
assert final_child is not None
# Both requests bill the parent; only the child request reserves on the child.
assert final_parent.reserved_balance == 2 * cost
assert final_parent.total_requests == 2
assert final_child.reserved_balance == cost
assert final_child.total_requests == 1
@pytest.mark.asyncio
async def test_balance_limit_with_existing_total_spent_under_concurrency(
patched_db_engine: None,
) -> None:
parent_hash = "parent_total_spent"
child_hash = "child_total_spent"
cost = 300
async with create_session() as session:
parent = ApiKey(hashed_key=parent_hash, balance=10000)
# 700 already spent against a 1000 limit: only one more 300 request fits.
child = ApiKey(
hashed_key=child_hash,
balance=0,
parent_key_hash=parent_hash,
balance_limit=1000,
total_spent=700,
)
session.add(parent)
session.add(child)
await session.commit()
results: list[str] = []
async def attempt() -> None:
async with create_session() as session:
fresh_child = await session.get(ApiKey, child_hash)
assert fresh_child is not None
try:
await pay_for_request(fresh_child, cost, session)
results.append("success")
except HTTPException as exc:
assert exc.status_code == 402
results.append("blocked")
await asyncio.gather(attempt(), attempt())
assert sorted(results) == ["blocked", "success"], (
f"Expected exactly one success and one 402, got: {results}"
)
async with create_session() as session:
final_child = await session.get(ApiKey, child_hash)
assert final_child is not None
assert final_child.reserved_balance == cost
assert final_child.total_spent == 700
@pytest.mark.asyncio
async def test_balance_limit_with_existing_reserved_balance(
patched_db_engine: None,
) -> None:
parent_hash = "parent_reserved_set"
blocked_hash = "child_reserved_blocked"
allowed_hash = "child_reserved_allowed"
cost = 300
async with create_session() as session:
parent = ApiKey(hashed_key=parent_hash, balance=10000)
# 800 already reserved against a 1000 limit: another 300 must be rejected.
blocked_child = ApiKey(
hashed_key=blocked_hash,
balance=0,
parent_key_hash=parent_hash,
balance_limit=1000,
reserved_balance=800,
)
# 500 reserved against a 1000 limit: another 300 still fits.
allowed_child = ApiKey(
hashed_key=allowed_hash,
balance=0,
parent_key_hash=parent_hash,
balance_limit=1000,
reserved_balance=500,
)
session.add(parent)
session.add(blocked_child)
session.add(allowed_child)
await session.commit()
async with create_session() as session:
fresh_blocked = await session.get(ApiKey, blocked_hash)
assert fresh_blocked is not None
with pytest.raises(HTTPException) as exc_info:
await pay_for_request(fresh_blocked, cost, session)
assert exc_info.value.status_code == 402
async with create_session() as session:
fresh_allowed = await session.get(ApiKey, allowed_hash)
assert fresh_allowed is not None
await pay_for_request(fresh_allowed, cost, session)
async with create_session() as session:
final_blocked = await session.get(ApiKey, blocked_hash)
final_allowed = await session.get(ApiKey, allowed_hash)
final_parent = await session.get(ApiKey, parent_hash)
assert final_blocked is not None
assert final_allowed is not None
assert final_parent is not None
assert final_blocked.reserved_balance == 800, "Rejected request must not reserve"
assert final_blocked.total_requests == 0
assert final_allowed.reserved_balance == 500 + cost
assert final_allowed.total_requests == 1
# Only the allowed request should have billed the parent.
assert final_parent.reserved_balance == cost
assert final_parent.total_requests == 1
@pytest.mark.asyncio
async def test_child_reservation_discarded_when_parent_balance_depleted(
patched_db_engine: None,
) -> None:
parent_hash = "parent_depleted"
child_hash = "child_depleted"
cost = 300
async with create_session() as session:
# Parent can only afford one request; child has no balance_limit.
parent = ApiKey(hashed_key=parent_hash, balance=cost)
child = ApiKey(
hashed_key=child_hash,
balance=0,
parent_key_hash=parent_hash,
)
session.add(parent)
session.add(child)
await session.commit()
results: list[str] = []
async def attempt() -> None:
async with create_session() as session:
fresh_child = await session.get(ApiKey, child_hash)
assert fresh_child is not None
try:
await pay_for_request(fresh_child, cost, session)
results.append("success")
except HTTPException as exc:
assert exc.status_code == 402
results.append("blocked")
await asyncio.gather(attempt(), attempt())
assert sorted(results) == ["blocked", "success"], (
f"Expected exactly one success and one 402, got: {results}"
)
async with create_session() as session:
final_parent = await session.get(ApiKey, parent_hash)
final_child = await session.get(ApiKey, child_hash)
assert final_parent is not None
assert final_child is not None
assert final_parent.reserved_balance == cost
# The failed request must not leave a committed reservation on the child.
assert final_child.reserved_balance == cost, (
f"Child reserved_balance should reflect only the successful request, "
f"got {final_child.reserved_balance}"
)
assert final_child.total_requests == 1
@pytest.mark.asyncio
async def test_refund_does_not_delete_key(integration_session: AsyncSession) -> None:
# This requires mocking the router call or testing the logic in balance.py
@@ -1,10 +1,10 @@
"""Integration tests for Lightning invoice key constraint fields.
Covers two things:
- The three constraint fields (balance_limit, balance_limit_reset, validity_date)
are persisted on LightningInvoice and survive a DB round-trip.
- The production-path API-key record helper propagates those fields to the
created ApiKey, so the constraints are actually enforced when the key is used.
- The validity_date constraint field is persisted on LightningInvoice and
survives a DB round-trip.
- The production-path API-key record helper propagates it to the created
ApiKey, so the constraint is actually enforced when the key is used.
"""
from __future__ import annotations
@@ -67,32 +67,6 @@ def mock_wallet_mint() -> object:
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_invoice_persists_balance_limit(
integration_session: AsyncSession,
) -> None:
invoice = _make_invoice(balance_limit=5000)
integration_session.add(invoice)
await integration_session.commit()
stored = await integration_session.get(LightningInvoice, invoice.id)
assert stored is not None
assert stored.balance_limit == 5000
@pytest.mark.asyncio
async def test_invoice_persists_balance_limit_reset(
integration_session: AsyncSession,
) -> None:
invoice = _make_invoice(balance_limit=5000, balance_limit_reset="daily")
integration_session.add(invoice)
await integration_session.commit()
stored = await integration_session.get(LightningInvoice, invoice.id)
assert stored is not None
assert stored.balance_limit_reset == "daily"
@pytest.mark.asyncio
async def test_invoice_persists_validity_date(
integration_session: AsyncSession,
@@ -112,38 +86,6 @@ async def test_invoice_persists_validity_date(
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_created_key_receives_balance_limit(
integration_session: AsyncSession,
) -> None:
invoice = _make_invoice(balance_limit=8000)
integration_session.add(invoice)
await integration_session.flush()
api_key = await _create_api_key_record(invoice, integration_session)
await integration_session.commit()
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
assert stored_key is not None
assert stored_key.balance_limit == 8000
@pytest.mark.asyncio
async def test_created_key_receives_balance_limit_reset(
integration_session: AsyncSession,
) -> None:
invoice = _make_invoice(balance_limit=8000, balance_limit_reset="monthly")
integration_session.add(invoice)
await integration_session.flush()
api_key = await _create_api_key_record(invoice, integration_session)
await integration_session.commit()
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
assert stored_key is not None
assert stored_key.balance_limit_reset == "monthly"
@pytest.mark.asyncio
async def test_created_key_receives_validity_date(
integration_session: AsyncSession,
@@ -422,8 +364,6 @@ async def test_created_key_without_constraints_has_none_fields(
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
assert stored_key is not None
assert stored_key.balance_limit is None
assert stored_key.balance_limit_reset is None
assert stored_key.validity_date is None
@@ -0,0 +1,221 @@
"""Cover that an admin price write reaches the served catalogue (GET /v1/models).
A price edit changes the served price; the served price is fee-adjusted while the
admin read-back is raw; a disabled model leaves the catalogue but keeps its row.
Each test reaches the served model by id rather than iterating ``data["data"]``,
so an empty catalogue fails these tests instead of skipping them.
"""
from __future__ import annotations
from collections.abc import Iterator
from datetime import datetime, timedelta, timezone
from unittest.mock import patch
import pytest
from httpx import AsyncClient
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.admin import admin_sessions
from routstr.core.db import ModelRow, UpstreamProviderRow
from routstr.proxy import reinitialize_upstreams
# The conftest patches ``routstr.payment.price.sats_usd_price``, but
# ``models.py`` imports it as ``from .price import sats_usd_price`` — a
# local binding the conftest-level patch cannot reach. Pin it here so
# every test that goes through ``_row_to_model`` gets a real sats price.
@pytest.fixture(autouse=True)
def _pin_sats_usd() -> Iterator[None]:
with patch("routstr.payment.models.sats_usd_price", return_value=0.0005):
yield
def _admin_headers() -> dict[str, str]:
token = "test-propagation-token"
admin_sessions[token] = int(
(datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp()
)
return {"Authorization": f"Bearer {token}"}
def _model_payload(prompt: float, provider_id: int, enabled: bool = True) -> dict:
return {
"id": "propagation-test-model",
"name": "Propagation Test Model",
"description": "model used to verify price propagation",
"created": 0,
"context_length": 128000,
"architecture": {
"modality": "text",
"input_modalities": ["text"],
"output_modalities": ["text"],
"tokenizer": "unknown",
"instruct_type": None,
},
"pricing": {
"prompt": prompt,
"completion": prompt * 2,
"input_cache_read": 0.0,
"input_cache_write": 0.0,
"request": 0.0,
"image": 0.0,
"web_search": 0.0,
"internal_reasoning": 0.0,
},
"per_request_limits": None,
"top_provider": None,
"upstream_provider_id": provider_id,
"canonical_slug": None,
"alias_ids": [],
"enabled": enabled,
"forwarded_model_id": "propagation-test-model",
}
async def _seed_provider(session: AsyncSession, *, fee: float = 1.0) -> int:
"""Insert a provider, refresh the upstream map, and return its primary key."""
provider = UpstreamProviderRow(
provider_type="generic",
base_url="https://propagation-test.example/v1",
api_key="test-key",
provider_fee=fee,
)
session.add(provider)
await session.commit()
await session.refresh(provider)
await reinitialize_upstreams()
assert provider.id is not None
return provider.id
@pytest.mark.integration
@pytest.mark.asyncio
async def test_price_edit_propagates_to_served_catalogue(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""A price edit through the admin API must change the served /v1/models price."""
provider_id = await _seed_provider(integration_session)
headers = _admin_headers()
r = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=headers,
json=_model_payload(prompt=1.0e-7, provider_id=provider_id),
)
assert r.status_code == 200
# -- record the served price before edit -----------------------------------
public = await integration_client.get("/v1/models")
assert public.status_code == 200
public_data = public.json()
assert len(public_data["data"]) > 0, "catalogue must not be empty"
served_before = {
m["id"]: m.get("pricing", {}).get("prompt") for m in public_data["data"]
}
assert "propagation-test-model" in served_before
before = served_before["propagation-test-model"]
# -- edit the price and re-check -------------------------------------------
r = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=headers,
json=_model_payload(prompt=5.0e-7, provider_id=provider_id),
)
assert r.status_code == 200
public = await integration_client.get("/v1/models")
assert public.status_code == 200
served_after = {
m["id"]: m.get("pricing", {}).get("prompt") for m in public.json()["data"]
}
after = served_after["propagation-test-model"]
assert before != after, "served price did not change after admin edit"
# With provider_fee=1.0 the served price equals the stored raw price.
assert after == pytest.approx(5.0e-7)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_admin_readback_is_raw_served_is_fee_adjusted(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""Admin read-back returns the raw price; /v1/models returns the fee-adjusted one."""
provider_id = await _seed_provider(integration_session, fee=1.05)
headers = _admin_headers()
model_payload = _model_payload(prompt=1.0e-7, provider_id=provider_id)
r = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=headers,
json=model_payload,
)
assert r.status_code == 200
# Admin read-back: apply_provider_fee=False
admin_r = await integration_client.get(
f"/admin/api/upstream-providers/{provider_id}/models/propagation-test-model",
headers=headers,
)
assert admin_r.status_code == 200
admin_body = admin_r.json()
raw_prompt = admin_body["pricing"]["prompt"]
assert raw_prompt == pytest.approx(1.0e-7)
# Public /v1/models: fee-adjusted
public = await integration_client.get("/v1/models")
assert public.status_code == 200
served = {
m["id"]: m.get("pricing", {}).get("prompt") for m in public.json()["data"]
}
assert served["propagation-test-model"] == pytest.approx(1.0e-7 * 1.05)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_disabled_model_not_served(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""A disabled model must be absent from /v1/models but still present in the DB."""
provider_id = await _seed_provider(integration_session)
headers = _admin_headers()
r = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=headers,
json=_model_payload(prompt=1.0e-7, provider_id=provider_id, enabled=True),
)
assert r.status_code == 200
# Confirm it appears in the public catalogue.
public = await integration_client.get("/v1/models")
served_ids = {m["id"] for m in public.json()["data"]}
assert "propagation-test-model" in served_ids
# -- disable via upsert ----------------------------------------------------
r = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=headers,
json=_model_payload(prompt=1.0e-7, provider_id=provider_id, enabled=False),
)
assert r.status_code == 200
# Public catalogue must no longer list it.
public = await integration_client.get("/v1/models")
served_ids = {m["id"] for m in public.json()["data"]}
assert "propagation-test-model" not in served_ids
# DB row must still exist.
row = await integration_session.get(
ModelRow, ("propagation-test-model", provider_id)
)
assert row is not None
assert row.enabled is False
@@ -0,0 +1,435 @@
"""Characterization tests for the model-serialisation pipeline ``_row_to_model`` runs.
These pin what the three public surfaces that reach it serve today: ``GET
/v1/models``, the admin single-model read-back, and the admin provider model
listing. Covered are the plain model, litellm cache backfill, preservation of
explicit cache rates, the request-price floor, the provider-fee flag against
recomputed max costs, survival of a failed sats conversion, the full serialised
dict field for field, and agreement between the two admin views.
Nothing here asserts on the deterministic USD half in isolation, so the pins hold
whether or not it is split out from the live BTC-rate conversion.
"""
from __future__ import annotations
import json
from collections.abc import Iterator
from datetime import datetime, timedelta, timezone
from unittest.mock import patch
import pytest
from httpx import AsyncClient
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.admin import admin_sessions
from routstr.core.db import ModelRow, UpstreamProviderRow
from routstr.proxy import reinitialize_upstreams
# The conftest patches ``routstr.payment.price.sats_usd_price``, but
# ``models.py`` imports it as ``from .price import sats_usd_price`` — a
# local binding the conftest-level patch cannot reach. Pin it here so
# every test that goes through ``_row_to_model`` gets a real sats price.
@pytest.fixture(autouse=True)
def _pin_sats_usd() -> Iterator[None]:
with patch("routstr.payment.models.sats_usd_price", return_value=0.0005):
yield
def _admin_headers() -> dict[str, str]:
token = "test-serialisation-token"
admin_sessions[token] = int(
(datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp()
)
return {"Authorization": f"Bearer {token}"}
# -- helpers -------------------------------------------------------------------
_SEEDED_MODEL_ID = "ser-test-model"
async def _seed_provider(session: AsyncSession, fee: float = 1.0) -> int:
"""Insert a provider, refresh the upstream map, and return its primary key."""
provider = UpstreamProviderRow(
provider_type="generic",
base_url="https://serialisation-test.example/v1",
api_key="test-key",
provider_fee=fee,
)
session.add(provider)
await session.commit()
await session.refresh(provider)
await reinitialize_upstreams()
assert provider.id is not None
return provider.id
async def _seed_model(
session: AsyncSession,
provider_id: int,
*,
model_id: str = _SEEDED_MODEL_ID,
prompt: float = 1.0e-7,
completion: float = 2.0e-7,
cache_read: float = 0.0,
cache_write: float = 0.0,
request_price: float = 0.0,
enabled: bool = True,
) -> ModelRow:
row = ModelRow(
id=model_id,
name=f"SerTest {model_id}",
description="characterization model",
created=0,
context_length=128000,
architecture=json.dumps(
{
"modality": "text",
"input_modalities": ["text"],
"output_modalities": ["text"],
"tokenizer": "unknown",
"instruct_type": None,
}
),
pricing=json.dumps(
{
"prompt": prompt,
"completion": completion,
"input_cache_read": cache_read,
"input_cache_write": cache_write,
"request": request_price,
"image": 0.0,
"web_search": 0.0,
"internal_reasoning": 0.0,
}
),
upstream_provider_id=provider_id,
enabled=enabled,
forwarded_model_id=model_id,
)
session.add(row)
await session.commit()
return row
async def _raw_via_admin(client: AsyncClient, provider_id: int, model_id: str) -> dict:
"""Return the raw (``apply_provider_fee=False``) model dict from admin read-back."""
r = await client.get(
f"/admin/api/upstream-providers/{provider_id}/models/{model_id}",
headers=_admin_headers(),
)
assert r.status_code == 200
return r.json()
async def _served_via_public(client: AsyncClient, model_id: str) -> dict | None:
"""Return the served model dict from /v1/models, or None if absent."""
r = await client.get("/v1/models")
assert r.status_code == 200
return {m["id"]: m for m in r.json()["data"]}.get(model_id)
# -- test 1: plain model -------------------------------------------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_plain_model_serialisation(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""A model with prompt + completion prices has the expected serialised shape."""
provider_id = await _seed_provider(integration_session)
await _seed_model(integration_session, provider_id)
await reinitialize_upstreams()
# Admin read-back (raw, no fee)
body = await _raw_via_admin(
integration_client,
provider_id,
_SEEDED_MODEL_ID,
)
assert body["id"] == _SEEDED_MODEL_ID
assert body["pricing"]["prompt"] == pytest.approx(1.0e-7)
assert body["pricing"]["completion"] == pytest.approx(2.0e-7)
assert body["sats_pricing"] is not None
# With sats_usd_price = 0.0005 (the fixture)
assert body["sats_pricing"]["prompt"] == pytest.approx(1.0e-7 / 0.0005)
assert body["sats_pricing"]["completion"] == pytest.approx(2.0e-7 / 0.0005)
# Public /v1/models (fee applied)
s = await _served_via_public(integration_client, _SEEDED_MODEL_ID)
assert s is not None, f"{_SEEDED_MODEL_ID} not found in /v1/models"
# fee=1.0 so values match raw
assert s["pricing"]["prompt"] == pytest.approx(1.0e-7)
assert s["pricing"]["completion"] == pytest.approx(2.0e-7)
# -- test 2: no cache rates (litellm backfill) ---------------------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_model_without_cache_rates_gets_litellm_backfill(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""A well-known model without cache rates gets them from litellm's cost map."""
provider_id = await _seed_provider(integration_session)
# Use a real litellm-known id so backfill_cache_pricing can find it.
await _seed_model(
integration_session,
provider_id,
model_id="gpt-4o",
prompt=2.5e-6,
completion=1.0e-5,
cache_read=0.0,
cache_write=0.0,
)
await reinitialize_upstreams()
# Admin read-back (raw): cache_read should be present after backfill.
# (cache_write may not be in litellm's map for every model.)
body = await _raw_via_admin(
integration_client,
provider_id,
"gpt-4o",
)
assert body["pricing"]["input_cache_read"] > 0.0, (
"backfill_cache_pricing should have filled input_cache_read from litellm"
)
# -- test 3: cache rates already present (not overwritten) ---------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_existing_cache_rates_not_overwritten(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""A model with explicit cache rates must keep them; the backfill is a no-op."""
provider_id = await _seed_provider(integration_session)
await _seed_model(
integration_session,
provider_id,
prompt=1.0e-7,
completion=2.0e-7,
cache_read=9.99e-9,
cache_write=8.88e-9,
)
await reinitialize_upstreams()
body = await _raw_via_admin(
integration_client,
provider_id,
_SEEDED_MODEL_ID,
)
assert body["pricing"]["input_cache_read"] == pytest.approx(9.99e-9)
assert body["pricing"]["input_cache_write"] == pytest.approx(8.88e-9)
# -- test 4: request price floor -----------------------------------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_model_with_request_price(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""A model with a request price floor carries it through to the served model."""
provider_id = await _seed_provider(integration_session)
await _seed_model(integration_session, provider_id, request_price=0.01)
await reinitialize_upstreams()
body = await _raw_via_admin(
integration_client,
provider_id,
_SEEDED_MODEL_ID,
)
assert body["pricing"]["request"] == pytest.approx(0.01)
# -- test 5: provider_fee=True vs False, max costs recomputed ------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_fee_flag_changes_pricing_but_max_costs_are_recomputed(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""With provider_fee=1.5, fee-adjusted pricing is 1.5× raw, but max costs are NOT."""
provider_id = await _seed_provider(integration_session, fee=1.5)
await _seed_model(integration_session, provider_id)
await reinitialize_upstreams()
# Admin read-back: raw, no fee.
body = await _raw_via_admin(
integration_client,
provider_id,
_SEEDED_MODEL_ID,
)
assert body["pricing"]["prompt"] == pytest.approx(1.0e-7)
# Public /v1/models: fee applied.
s = await _served_via_public(integration_client, _SEEDED_MODEL_ID)
assert s is not None
assert s["pricing"]["prompt"] == pytest.approx(1.0e-7 * 1.5)
# max_prompt_cost must NOT just be multiplied by 1.5 — it is recomputed from
# the fee-inflated per-token rates and context_length.
cl = 128_000
expected_max_prompt = cl * 1.0e-7 * 1.5
assert s["pricing"]["max_prompt_cost"] == pytest.approx(expected_max_prompt)
# -- test 6: sats conversion failure keeps the model alive ---------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_model_survives_sats_conversion_failure(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""When the BTC feed fails, the model still returns — with no sats_pricing."""
provider_id = await _seed_provider(integration_session)
await _seed_model(integration_session, provider_id)
# Make sats_usd_price raise so _update_model_sats_pricing swallows it.
# The admin read-back must happen inside the patch block.
with patch(
"routstr.payment.models.sats_usd_price",
side_effect=RuntimeError("BTC feed down"),
):
await reinitialize_upstreams()
body = await _raw_via_admin(
integration_client,
provider_id,
_SEEDED_MODEL_ID,
)
assert body["id"] == _SEEDED_MODEL_ID
assert body["sats_pricing"] is None, (
"sats conversion failure must not crash — the model returns with no sats_pricing"
)
# -- test 7: the whole serialised dict -----------------------------------------
# The sats figures are the USD ones divided by the pinned 0.0005 rate; written
# as the division so the expectation carries the same float error the code does.
_SATS = 0.0005
def _expected_serialised_model(provider_id: int) -> dict:
"""Every field the raw admin read-back produces for the seeded model.
Note the max-cost asymmetry: with fees off the USD max costs stay at zero
while their sats counterparts are computed. That is what the code does
today, and pinning it is the point.
"""
return {
"alias_ids": None,
"architecture": {
"input_modalities": ["text"],
"instruct_type": None,
"modality": "text",
"output_modalities": ["text"],
"tokenizer": "unknown",
},
"canonical_slug": None,
"context_length": 128000,
"created": 0,
"description": "characterization model",
"enabled": True,
"forwarded_model_id": _SEEDED_MODEL_ID,
"id": _SEEDED_MODEL_ID,
"name": f"SerTest {_SEEDED_MODEL_ID}",
"per_request_limits": None,
"pricing": {
"completion": 2.0e-7,
"image": 0.0,
"input_cache_read": 0.0,
"input_cache_write": 0.0,
"internal_reasoning": 0.0,
"max_completion_cost": 0.0,
"max_cost": 0.0,
"max_prompt_cost": 0.0,
"prompt": 1.0e-7,
"request": 0.01,
"web_search": 0.0,
},
"sats_pricing": {
"completion": 2.0e-7 / _SATS,
"image": 0.0,
"input_cache_read": 0.0,
"input_cache_write": 0.0,
"internal_reasoning": 0.0,
"max_completion_cost": 0.0,
"max_cost": 0.001,
"max_prompt_cost": 0.0,
"prompt": 1.0e-7 / _SATS,
"request": 0.01 / _SATS,
"web_search": 0.0,
},
"top_provider": None,
"upstream_provider_id": provider_id,
}
@pytest.mark.integration
@pytest.mark.asyncio
async def test_full_serialised_model_is_unchanged(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""Pin every field of the serialised model, not just the interesting ones.
The tests above pin values. A refactor that dropped a field outright
would satisfy all of them and fail only here.
"""
provider_id = await _seed_provider(integration_session)
await _seed_model(integration_session, provider_id, request_price=0.01)
await reinitialize_upstreams()
body = await _raw_via_admin(
integration_client,
provider_id,
_SEEDED_MODEL_ID,
)
assert body == _expected_serialised_model(provider_id)
# -- test 8: the provider model listing ----------------------------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_provider_listing_matches_single_model_read_back(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""The listing is a third entry into the same builder — it must agree."""
provider_id = await _seed_provider(integration_session)
await _seed_model(integration_session, provider_id, request_price=0.01)
await reinitialize_upstreams()
single = await _raw_via_admin(
integration_client,
provider_id,
_SEEDED_MODEL_ID,
)
r = await integration_client.get(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=_admin_headers(),
)
assert r.status_code == 200
listed = {m["id"]: m for m in r.json()["db_models"]}
assert _SEEDED_MODEL_ID in listed, "seeded model missing from the provider listing"
assert listed[_SEEDED_MODEL_ID] == single
@@ -0,0 +1,747 @@
"""Regression tests for the production "negative available balance" 402.
Invariants protected:
* billing and the admin API agree on what "available" means
* an in-flight (heartbeaten) reservation is never swept; an abandoned one is
released and can never be charged afterwards
* a cost overrun spends only its own reservation plus unreserved balance
* corrupt reservations are repaired terminally instead of poisoning cleanup
* under concurrency, balance / reserved / available never go negative
"""
import asyncio
import random
import time
import uuid
from typing import Awaitable, Callable
from unittest.mock import patch
import pytest
from httpx import AsyncClient
from sqlmodel import col, func, select, update
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ApiKey, ReservationRelease
from routstr.payment.cost_calculation import CostData
pytestmark = pytest.mark.integration
# Realistic sweeper timeout: a renewed (heartbeaten) reservation stays alive,
# a reservation backdated past this is released.
STALE_TIMEOUT_SECONDS = 300
# created_at is whole seconds; -1 makes every reservation stale immediately.
# Only used in the fuzz test to stress the terminal-release path.
SWEEP_EVERYTHING = -1
def _cost_data(total_msats: int) -> CostData:
return CostData(
base_msats=0,
input_msats=total_msats // 2,
output_msats=total_msats - total_msats // 2,
total_msats=total_msats,
total_usd=0.0,
input_tokens=100,
output_tokens=100,
)
def _response(model: str = "test-model") -> dict:
return {
"model": model,
"usage": {"prompt_tokens": 100, "completion_tokens": 100},
}
async def _new_key(session: AsyncSession, balance: int) -> str:
key_hash = f"test_neg_{uuid.uuid4().hex}"
session.add(
ApiKey(
hashed_key=key_hash,
balance=balance,
reserved_balance=0,
total_spent=0,
total_requests=0,
)
)
await session.commit()
return key_hash
async def _backdate_reservation(
session: AsyncSession, release_id: str, seconds: int
) -> None:
"""Age a reservation's lease as if it had not been renewed for `seconds`."""
await session.exec( # type: ignore[call-overload]
update(ReservationRelease)
.where(col(ReservationRelease.id) == release_id)
.values(created_at=col(ReservationRelease.created_at) - seconds)
)
await session.commit()
async def _wait_for(
predicate: "Callable[[], Awaitable[bool]]",
timeout: float = 10.0,
interval: float = 0.1,
) -> bool:
"""Bounded polling instead of fixed sleeps for background-task effects."""
deadline = asyncio.get_event_loop().time() + timeout
while True:
if await predicate():
return True
if asyncio.get_event_loop().time() > deadline:
return False
await asyncio.sleep(interval)
@pytest.mark.asyncio
async def test_402_reports_negative_available_while_admin_shows_positive(
integration_session: AsyncSession,
integration_client: AsyncClient,
) -> None:
"""The production symptom: billing rejects on balance - reserved_balance,
so the admin endpoint must expose reserved/available, not just balance."""
from fastapi import HTTPException
from routstr.auth import _validate_bearer_key_locked
from routstr.core.db import set_admin_password
key_hash = await _new_key(integration_session, balance=263_000)
# Leaked reservations slightly exceeding the balance, as in production.
await integration_session.exec( # type: ignore[call-overload]
update(ApiKey)
.where(col(ApiKey.hashed_key) == key_hash)
.values(reserved_balance=267_215)
)
await integration_session.commit()
with pytest.raises(HTTPException) as exc:
await _validate_bearer_key_locked(
"sk-" + key_hash, integration_session, min_cost=1
)
assert exc.value.status_code == 402
message = exc.value.detail["error"]["message"] # type: ignore[index]
assert "-4.215 sats (-4215 msats) available" in message, message
await set_admin_password(integration_session, "test-admin-pw")
login = await integration_client.post(
"/admin/api/login", json={"password": "test-admin-pw"}
)
assert login.status_code == 200
token = login.json()["token"]
resp = await integration_client.get(
"/admin/api/temporary-balances",
params={"search": key_hash},
headers={"Authorization": f"Bearer {token}"},
)
assert resp.status_code == 200
rows = [row for row in resp.json()["balances"] if row["hashed_key"] == key_hash]
assert len(rows) == 1
row = rows[0]
assert row["balance"] == 263_000
assert row["reserved_balance"] == 267_215
assert row["available_balance"] == -4_215
assert resp.json()["totals"] == {
"total_balance": 263_000,
"total_reserved_balance": 267_215,
"total_available_balance": -4_215,
"total_spent": 0,
"total_requests": 0,
}
@pytest.mark.asyncio
async def test_abandoned_reservation_is_swept_and_cannot_finalize(
integration_session: AsyncSession,
) -> None:
"""Release is terminal: a swept reservation's late finalizer must not
charge, or it could spend funds since reserved by another request."""
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
from routstr.core.db import release_stale_reservations
cost = 5_000
key_hash = await _new_key(integration_session, balance=10_000)
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, cost, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
# Client vanished: the lease is never renewed and ages past the timeout.
await _backdate_reservation(
integration_session, reservation.release_id, STALE_TIMEOUT_SECONDS + 1
)
released = await release_stale_reservations(
integration_session, STALE_TIMEOUT_SECONDS
)
assert released == 1
# A zombie finalizer shows up afterwards; it must not charge.
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
await adjust_payment_for_tokens(
key,
_response(),
integration_session,
cost,
reservation_snapshot=reservation,
)
await integration_session.refresh(key)
record = await integration_session.get(ReservationRelease, reservation.release_id)
assert record is not None and record.status == "released"
assert key.total_spent == 0
assert key.balance == 10_000
assert key.reserved_balance == 0
@pytest.mark.asyncio
async def test_sweeper_cannot_release_a_reservation_renewed_after_selection(
integration_session: AsyncSession,
) -> None:
"""A heartbeat landing between the sweeper's select and its transition
must win — exercised through the public sweeper entry point."""
from routstr.auth import (
get_reservation_snapshot,
pay_for_request,
renew_reservation,
)
from routstr.core import db as core_db
from routstr.core.db import release_stale_reservations
key_hash = await _new_key(integration_session, balance=1_000)
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, 1_000, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
await _backdate_reservation(
integration_session, reservation.release_id, STALE_TIMEOUT_SECONDS + 100
)
real_transition = core_db._transition_stale_reservation
async def renew_then_transition(
session: AsyncSession, reservation_id: str, cutoff: int
) -> bool:
# The sweeper selected this reservation as stale; the heartbeat
# renews exactly between that select and the transition.
assert await renew_reservation(reservation, session)
return await real_transition(session, reservation_id, cutoff)
with patch.object(core_db, "_transition_stale_reservation", renew_then_transition):
released = await release_stale_reservations(
integration_session, STALE_TIMEOUT_SECONDS
)
assert released == 0
record = await integration_session.get(ReservationRelease, reservation.release_id)
assert record is not None and record.status == "active"
await integration_session.refresh(key)
assert key.reserved_balance == 1_000
@pytest.mark.asyncio
async def test_legacy_cleanup_cannot_erase_a_reservation_committed_after_its_read(
integration_session: AsyncSession,
) -> None:
"""A reservation committing between the legacy sweep's read and its
zeroing must survive — exercised through the public sweeper entry point."""
from routstr.core import db as core_db
from routstr.core.db import release_stale_reservations
key_hash = await _new_key(integration_session, balance=10_000)
stale_reserved_at = int(time.time()) - (STALE_TIMEOUT_SECONDS + 100)
await integration_session.exec( # type: ignore[call-overload]
update(ApiKey)
.where(col(ApiKey.hashed_key) == key_hash)
.values(reserved_balance=500, reserved_at=stale_reserved_at)
)
await integration_session.commit()
real_release = core_db._release_legacy_aggregate
async def commit_reservation_then_release(
session: AsyncSession,
target_key_hash: str,
observed_reserved: int,
observed_reserved_at: int | None,
) -> bool:
# The sweeper read the legacy aggregate and found no active durable
# owner; a new reservation commits exactly before the zeroing lands.
session.add(
ReservationRelease(
id=uuid.uuid4().hex,
key_hash=target_key_hash,
billing_key_hash=target_key_hash,
reserved_msats=700,
status="active",
)
)
await session.exec( # type: ignore[call-overload]
update(ApiKey)
.where(col(ApiKey.hashed_key) == target_key_hash)
.values(
reserved_balance=col(ApiKey.reserved_balance) + 700,
reserved_at=int(time.time()),
)
)
await session.commit()
return await real_release(
session, target_key_hash, observed_reserved, observed_reserved_at
)
with patch.object(
core_db, "_release_legacy_aggregate", commit_reservation_then_release
):
released = await release_stale_reservations(
integration_session, STALE_TIMEOUT_SECONDS
)
assert released == 0
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
assert key.reserved_balance == 1_200, "legacy cleanup erased a live reservation"
@pytest.mark.asyncio
async def test_sweeper_repairs_corrupt_reservation_and_continues_batch(
integration_session: AsyncSession,
) -> None:
"""One corrupt durable reservation (aggregate no longer holds its msats)
must be terminalized without aggregate subtraction, and must not stop the
rest of the batch from being released normally."""
from routstr.auth import get_reservation_snapshot, pay_for_request
from routstr.core.db import release_stale_reservations
cost = 1_000
corrupt_hash = await _new_key(integration_session, balance=cost)
healthy_hash = await _new_key(integration_session, balance=cost)
corrupt_key = await integration_session.get(ApiKey, corrupt_hash)
assert corrupt_key is not None
await pay_for_request(corrupt_key, cost, integration_session)
corrupt_reservation = await get_reservation_snapshot(
corrupt_key, integration_session
)
healthy_key = await integration_session.get(ApiKey, healthy_hash)
assert healthy_key is not None
await pay_for_request(healthy_key, cost, integration_session)
healthy_reservation = await get_reservation_snapshot(
healthy_key, integration_session
)
# Corrupt the first key: its aggregate no longer holds the reservation.
await integration_session.exec( # type: ignore[call-overload]
update(ApiKey)
.where(col(ApiKey.hashed_key) == corrupt_hash)
.values(reserved_balance=0)
)
await integration_session.commit()
for reservation in (corrupt_reservation, healthy_reservation):
await _backdate_reservation(
integration_session, reservation.release_id, STALE_TIMEOUT_SECONDS + 100
)
released = await release_stale_reservations(
integration_session, STALE_TIMEOUT_SECONDS
)
assert released == 2, "a corrupt record must not abort the sweep batch"
for reservation in (corrupt_reservation, healthy_reservation):
record = await integration_session.get(
ReservationRelease, reservation.release_id
)
assert record is not None and record.status == "released"
healthy_key = await integration_session.get(ApiKey, healthy_hash)
corrupt_key = await integration_session.get(ApiKey, corrupt_hash)
assert healthy_key is not None and corrupt_key is not None
assert healthy_key.reserved_balance == 0
assert corrupt_key.reserved_balance == 0
assert corrupt_key.balance == cost, "repair must not touch balances"
@pytest.mark.asyncio
async def test_heartbeat_survives_a_rolled_back_charge_attempt(
integration_session: AsyncSession,
patched_db_engine: None,
) -> None:
"""Claiming a reservation must not stop its heartbeat: a rollback restores
the active reservation, which then still needs lease renewal."""
from routstr import auth
from routstr.auth import (
_claim_reservation_for_charge,
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
from routstr.core.db import create_session
cost = 1_000
timeout = 3 # heartbeat interval = 1s
with patch.object(auth.settings, "stale_reservation_timeout_seconds", timeout):
async with create_session() as session:
key_hash = await _new_key(session, balance=2_000)
key = await session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, cost, session)
reservation = await get_reservation_snapshot(key, session)
try:
# A charge attempt claims the reservation, then its transaction
# fails and rolls back.
async with create_session() as session:
assert await _claim_reservation_for_charge(reservation, session)
await session.rollback()
async with create_session() as session:
record = await session.get(ReservationRelease, reservation.release_id)
assert record is not None and record.status == "active"
await _backdate_reservation(
session, reservation.release_id, timeout * 10
)
record = await session.get(ReservationRelease, reservation.release_id)
assert record is not None
backdated_lease = record.created_at
async def lease_renewed() -> bool:
async with create_session() as session:
record = await session.get(
ReservationRelease, reservation.release_id
)
return record is not None and record.created_at > backdated_lease
assert await _wait_for(lease_renewed), (
"heartbeat did not survive the rolled-back charge attempt"
)
# The restored reservation still finalizes normally.
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
async with create_session() as session:
key = await session.get(ApiKey, key_hash)
assert key is not None
result = await adjust_payment_for_tokens(
key,
_response(),
session,
cost,
reservation_snapshot=reservation,
)
finally:
await auth._stop_reservation_heartbeat(reservation.release_id)
assert result["charged_msats"] == cost
async with create_session() as session:
record = await session.get(ReservationRelease, reservation.release_id)
assert record is not None and record.status == "charged"
key = await session.get(ApiKey, key_hash)
assert key is not None
assert key.total_spent == cost
@pytest.mark.asyncio
async def test_heartbeat_dies_with_its_request_so_sweeper_can_recover(
integration_session: AsyncSession,
patched_db_engine: None,
) -> None:
"""A request that vanishes without finalizing must not renew forever —
its heartbeat stops with the owning task and the sweeper reclaims the
funds."""
from routstr import auth
from routstr.auth import get_reservation_snapshot, pay_for_request
from routstr.core.db import create_session, release_stale_reservations
cost = 1_000
timeout = 3 # heartbeat interval = 1s
async with create_session() as session:
key_hash = await _new_key(session, balance=cost)
holder: dict = {}
with patch.object(auth.settings, "stale_reservation_timeout_seconds", timeout):
async def doomed_request() -> None:
async with create_session() as session:
key = await session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, cost, session)
holder["reservation"] = await get_reservation_snapshot(key, session)
# ...request control dies here, no finalize and no release.
await asyncio.create_task(doomed_request())
release_id = holder["reservation"].release_id
try:
# While the heartbeat is still winding down it may renew once
# more; keep backdating until the sweeper wins, which it must as
# soon as the dead owner is noticed.
async def sweeper_recovered() -> bool:
async with create_session() as session:
await _backdate_reservation(session, release_id, timeout * 10)
return await release_stale_reservations(session, timeout) == 1
assert await _wait_for(sweeper_recovered), (
"sweeper never recovered the abandoned reservation"
)
finally:
await auth._stop_reservation_heartbeat(release_id)
async with create_session() as session:
record = await session.get(ReservationRelease, release_id)
assert record is not None and record.status == "released"
key = await session.get(ApiKey, key_hash)
assert key is not None
assert key.reserved_balance == 0
assert key.total_spent == 0
@pytest.mark.asyncio
async def test_reservation_heartbeat_covers_the_whole_request_lifecycle(
integration_session: AsyncSession,
patched_db_engine: None,
) -> None:
"""pay_for_request starts the heartbeat, finalization stops it, and a
backdated lease is renewed in the background without any manual call."""
from routstr import auth
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
from routstr.core.db import create_session, release_stale_reservations
cost = 1_000
timeout = 3 # heartbeat interval = 1s
async with create_session() as session:
key_hash = await _new_key(session, balance=2 * cost)
with patch.object(auth.settings, "stale_reservation_timeout_seconds", timeout):
async with create_session() as session:
key = await session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, cost, session)
reservation = await get_reservation_snapshot(key, session)
try:
# Only a background renewal can keep this alive now.
async with create_session() as session:
await _backdate_reservation(
session, reservation.release_id, timeout * 10
)
record = await session.get(ReservationRelease, reservation.release_id)
assert record is not None
backdated_lease = record.created_at
async def lease_renewed() -> bool:
async with create_session() as session:
record = await session.get(
ReservationRelease, reservation.release_id
)
return record is not None and record.created_at > backdated_lease
assert await _wait_for(lease_renewed), "heartbeat never renewed the lease"
async with create_session() as session:
assert await release_stale_reservations(session, timeout) == 0
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
async with create_session() as session:
key = await session.get(ApiKey, key_hash)
assert key is not None
result = await adjust_payment_for_tokens(
key,
_response(),
session,
cost,
reservation_snapshot=reservation,
)
finally:
await auth._stop_reservation_heartbeat(reservation.release_id)
assert result["charged_msats"] == cost
async with create_session() as session:
record = await session.get(ReservationRelease, reservation.release_id)
assert record is not None and record.status == "charged"
key = await session.get(ApiKey, key_hash)
assert key is not None
assert key.total_spent == cost
assert key.reserved_balance == 0
@pytest.mark.asyncio
async def test_overrun_cannot_spend_a_concurrent_reservation(
integration_session: AsyncSession,
) -> None:
"""An overrun is capped to its own reservation plus unreserved balance;
the sibling's reserved funds stay untouched and available stays >= 0."""
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
reserved_each = 100
overrun_cost = 150 # A's real token cost exceeds its reservation
# Balance covers exactly two reservations; nothing free on top.
key_hash = await _new_key(integration_session, balance=2 * reserved_each)
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, reserved_each, integration_session)
reservation_a = await get_reservation_snapshot(key, integration_session)
await pay_for_request(key, reserved_each, integration_session)
await get_reservation_snapshot(key, integration_session) # B stays in flight
await integration_session.refresh(key)
assert key.reserved_balance == 2 * reserved_each
# Only A finalizes; B is still streaming and its funds must stay reserved.
with patch("routstr.auth.calculate_cost", return_value=_cost_data(overrun_cost)):
result = await adjust_payment_for_tokens(
key,
_response(),
integration_session,
reserved_each,
reservation_snapshot=reservation_a,
)
assert result["charged_msats"] == reserved_each
assert result["total_msats"] == overrun_cost
await integration_session.refresh(key)
assert key.total_spent == reserved_each
assert key.reserved_balance == reserved_each
assert key.balance == reserved_each
assert key.total_balance >= 0, (
f"available balance went negative: balance={key.balance} "
f"reserved={key.reserved_balance} -> {key.total_balance} msats; "
"the overrun charge consumed the still-reserved funds of request B"
)
@pytest.mark.asyncio
async def test_concurrent_requests_with_sweeper_keep_balance_invariants(
integration_session: AsyncSession,
patched_db_engine: None,
) -> None:
"""Fuzz: concurrent requests against an everything-is-stale sweeper.
Some requests legitimately finish uncharged (release is terminal), but
balances must never go negative and no reservation may stay active."""
from fastapi import HTTPException
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
from routstr.core.db import create_session, release_stale_reservations
rng = random.Random(1337)
starting_balance = 200_000
n_requests = 24
async with create_session() as session:
key_hash = await _new_key(session, balance=starting_balance)
completed_costs: list[int] = []
rejected_requests = 0
async def one_request(index: int) -> None:
nonlocal rejected_requests
reserved = rng.randrange(1_000, 4_000)
actual = max(1, int(reserved * rng.uniform(0.5, 1.1)))
try:
async with create_session() as session:
key = await session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, reserved, session)
except HTTPException as exc:
# A depleted balance is the only legitimate rejection.
assert exc.status_code == 402, exc.detail
rejected_requests += 1
return
try:
async with create_session() as session:
key = await session.get(ApiKey, key_hash)
assert key is not None
reservation = await get_reservation_snapshot(key, session)
except RuntimeError:
# The everything-is-stale sweeper can release the reservation
# before the stream even starts; the request aborts uncharged.
return
await asyncio.sleep(rng.uniform(0, 0.02)) # the "stream"
async with create_session() as session:
key = await session.get(ApiKey, key_hash)
assert key is not None
await adjust_payment_for_tokens(
key,
_response(str(actual)),
session,
reserved,
reservation_snapshot=reservation,
)
completed_costs.append(actual)
sweeping = True
async def sweeper() -> None:
while sweeping:
async with create_session() as session:
await release_stale_reservations(session, SWEEP_EVERYTHING)
await asyncio.sleep(0.002)
sweep_task = asyncio.create_task(sweeper())
try:
with patch(
"routstr.auth.calculate_cost",
side_effect=lambda response_data, *a, **k: _cost_data(
int(response_data["model"])
),
):
await asyncio.gather(*(one_request(i) for i in range(n_requests)))
finally:
sweeping = False
await sweep_task
async with create_session() as session:
key = await session.get(ApiKey, key_hash)
assert key is not None
leftover_active = (
await session.exec( # type: ignore[call-overload]
select(func.count())
.select_from(ReservationRelease)
.where(col(ReservationRelease.status) == "active")
)
).one()
assert completed_costs or rejected_requests, "no request made any progress"
assert key.balance >= 0, f"balance went negative: {key.balance}"
assert key.reserved_balance >= 0, (
f"reserved_balance went negative: {key.reserved_balance}"
)
assert key.total_balance >= 0, (
f"available balance negative: balance={key.balance} "
f"reserved={key.reserved_balance} (this is the production symptom)"
)
assert key.total_spent <= starting_balance, (
f"spent {key.total_spent} of a {starting_balance} balance"
)
# Requests whose reservation was swept mid-flight finish uncharged, so the
# charged total can only be at most the sum of completed request costs.
max_expected_spend = sum(completed_costs)
assert key.total_spent <= max_expected_spend, (
f"charged {key.total_spent} msats but completed requests only cost "
f"{max_expected_spend}"
)
assert leftover_active == 0, (
f"{leftover_active} reservations still active after all requests finished"
)
@@ -0,0 +1,450 @@
"""Invariant coverage for the reserve → charge → release money path.
Every finalization branch of ``adjust_payment_for_tokens`` must respect the
same accounting rules: a completed request is charged exactly once, its
reported ``charged_msats`` matches the actual debit, it never spends more than
its own reservation leaves available, and concurrent requests cannot raid each
other's reservations.
"""
import asyncio
import uuid
from unittest.mock import patch
import pytest
from sqlmodel import col, select, update
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ApiKey, ReservationRelease
from routstr.payment.cost_calculation import (
CostData,
CostDataError,
MaxCostData,
)
pytestmark = pytest.mark.integration
def _cost_data(total_msats: int, cls: type[CostData] = CostData) -> CostData:
return cls(
base_msats=0,
input_msats=total_msats // 2,
output_msats=total_msats - total_msats // 2,
total_msats=total_msats,
total_usd=0.0,
input_tokens=100,
output_tokens=100,
)
def _response() -> dict:
return {
"model": "test-model",
"usage": {"prompt_tokens": 100, "completion_tokens": 100},
}
async def _new_key(session: AsyncSession, balance: int) -> str:
key_hash = f"test_inv_{uuid.uuid4().hex}"
session.add(
ApiKey(
hashed_key=key_hash,
balance=balance,
reserved_balance=0,
total_spent=0,
total_requests=0,
)
)
await session.commit()
return key_hash
async def _active_reservations(session: AsyncSession) -> int:
rows = await session.exec(
select(ReservationRelease).where(col(ReservationRelease.status) == "active")
)
return len(rows.all())
@pytest.mark.asyncio
async def test_corrupt_revert_still_decrements_request_count(
integration_session: AsyncSession,
) -> None:
from routstr.auth import (
get_reservation_snapshot,
pay_for_request,
revert_pay_for_request,
)
key_hash = await _new_key(integration_session, balance=10_000)
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, 3_000, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
await integration_session.exec( # type: ignore[call-overload]
update(ApiKey)
.where(col(ApiKey.hashed_key) == key_hash)
.values(reserved_balance=0)
)
await integration_session.commit()
assert await revert_pay_for_request(
key,
integration_session,
3_000,
reservation_snapshot=reservation,
)
updated = await integration_session.get(ApiKey, key_hash)
record = await integration_session.get(ReservationRelease, reservation.release_id)
assert updated is not None
assert updated.total_requests == 0
assert updated.reserved_balance == 0
assert record is not None and record.status == "released"
@pytest.mark.asyncio
async def test_exact_cost_branch_charges_the_reservation_once(
integration_session: AsyncSession,
) -> None:
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
cost = 3_000
key_hash = await _new_key(integration_session, balance=10_000)
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, cost, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
result = await adjust_payment_for_tokens(
key,
_response(),
integration_session,
cost,
reservation_snapshot=reservation,
)
assert result["charged_msats"] == cost
await integration_session.refresh(key)
assert key.balance == 10_000 - cost
assert key.total_spent == cost
assert key.reserved_balance == 0
assert key.total_requests == 1
assert await _active_reservations(integration_session) == 0
@pytest.mark.asyncio
async def test_max_cost_branch_charges_the_reservation_once(
integration_session: AsyncSession,
) -> None:
"""No token pricing configured -> flat MaxCostData charge of the reservation."""
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
cost = 2_500
key_hash = await _new_key(integration_session, balance=10_000)
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, cost, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(cost, cls=MaxCostData),
):
result = await adjust_payment_for_tokens(
key,
_response(),
integration_session,
cost,
reservation_snapshot=reservation,
)
assert result["charged_msats"] == cost
await integration_session.refresh(key)
assert key.balance == 10_000 - cost
assert key.total_spent == cost
assert key.reserved_balance == 0
@pytest.mark.asyncio
async def test_underrun_branch_refunds_the_unused_reservation(
integration_session: AsyncSession,
) -> None:
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
reserved = 5_000
actual = 1_200
key_hash = await _new_key(integration_session, balance=10_000)
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, reserved, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
with patch("routstr.auth.calculate_cost", return_value=_cost_data(actual)):
result = await adjust_payment_for_tokens(
key,
_response(),
integration_session,
reserved,
reservation_snapshot=reservation,
)
assert result["charged_msats"] == actual
await integration_session.refresh(key)
assert key.total_spent == actual, "user must pay the real cost, not the reservation"
assert key.balance == 10_000 - actual
assert key.reserved_balance == 0
@pytest.mark.asyncio
async def test_overrun_branch_charges_full_cost_when_balance_is_free(
integration_session: AsyncSession,
) -> None:
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
reserved = 1_000
actual = 1_400
key_hash = await _new_key(integration_session, balance=10_000)
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, reserved, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
with patch("routstr.auth.calculate_cost", return_value=_cost_data(actual)):
result = await adjust_payment_for_tokens(
key,
_response(),
integration_session,
reserved,
reservation_snapshot=reservation,
)
assert result["charged_msats"] == actual
await integration_session.refresh(key)
assert key.total_spent == actual
assert key.balance == 10_000 - actual
assert key.reserved_balance == 0
@pytest.mark.asyncio
async def test_zero_cost_response_is_free_and_releases_the_reservation(
integration_session: AsyncSession,
) -> None:
"""An empty/unusable upstream response must cost the user nothing."""
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
reserved = 4_000
key_hash = await _new_key(integration_session, balance=10_000)
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, reserved, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
with patch("routstr.auth.calculate_cost", return_value=_cost_data(0)):
await adjust_payment_for_tokens(
key,
{"model": "test-model"},
integration_session,
reserved,
reservation_snapshot=reservation,
)
await integration_session.refresh(key)
assert key.balance == 10_000
assert key.total_spent == 0
assert key.reserved_balance == 0
assert await _active_reservations(integration_session) == 0
@pytest.mark.asyncio
async def test_cost_error_releases_the_reservation_without_charging(
integration_session: AsyncSession,
) -> None:
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
reserved = 4_000
key_hash = await _new_key(integration_session, balance=10_000)
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, reserved, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
with patch(
"routstr.auth.calculate_cost",
return_value=CostDataError(message="no pricing", code="pricing_error"),
):
cost = await adjust_payment_for_tokens(
key,
_response(),
integration_session,
reserved,
reservation_snapshot=reservation,
)
assert cost["charged_msats"] == 0
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
assert key.balance == 10_000, "a pricing failure must not charge the user"
assert key.total_spent == 0
assert key.reserved_balance == 0, "funds must not stay locked"
assert await _active_reservations(integration_session) == 0
@pytest.mark.asyncio
async def test_repeated_finalization_charges_only_once(
integration_session: AsyncSession,
) -> None:
"""A retried finalizer (proxy retry, duplicate stream end) must not double-bill."""
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
cost = 3_000
key_hash = await _new_key(integration_session, balance=10_000)
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, cost, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
results = []
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
for _ in range(3):
# A declined re-charge rolls its session back, so re-load the key
# the way a fresh request would instead of reusing a stale instance.
integration_session.expunge_all()
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
results.append(
await adjust_payment_for_tokens(
key,
_response(),
integration_session,
cost,
reservation_snapshot=reservation,
)
)
# Only the first finalization debits; duplicates report a zero charge.
assert [r["charged_msats"] for r in results] == [cost, 0, 0]
integration_session.expunge_all()
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
assert key.total_spent == cost, f"charged {key.total_spent} for one request"
assert key.balance == 10_000 - cost
@pytest.mark.asyncio
async def test_concurrent_duplicate_finalization_charges_only_once(
integration_session: AsyncSession,
patched_db_engine: None,
) -> None:
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
from routstr.core.db import create_session
cost = 3_000
async with create_session() as session:
key_hash = await _new_key(session, balance=10_000)
key = await session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, cost, session)
reservation = await get_reservation_snapshot(key, session)
async def finalize() -> None:
async with create_session() as session:
fresh = await session.get(ApiKey, key_hash)
assert fresh is not None
await adjust_payment_for_tokens(
fresh,
_response(),
session,
cost,
reservation_snapshot=reservation,
)
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
await asyncio.gather(finalize(), finalize(), finalize())
async with create_session() as session:
key = await session.get(ApiKey, key_hash)
assert key is not None
assert key.total_spent == cost
assert key.balance == 10_000 - cost
assert key.reserved_balance == 0
@pytest.mark.asyncio
async def test_release_after_charge_does_not_credit_the_user_back(
integration_session: AsyncSession,
) -> None:
"""A late cleanup path must not turn a charged request into a free one."""
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
release_reservation,
)
cost = 3_000
key_hash = await _new_key(integration_session, balance=10_000)
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, cost, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
await adjust_payment_for_tokens(
key,
_response(),
integration_session,
cost,
reservation_snapshot=reservation,
)
released = await release_reservation(reservation, integration_session, cost)
assert released is False, "a charged reservation must not be releasable"
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
assert key.balance == 10_000 - cost
assert key.total_spent == cost
assert key.reserved_balance == 0
+34 -2
View File
@@ -23,6 +23,7 @@ from routstr.upstream.auto_topup import (
_ppq_request_id,
_ppq_spent_last_24h_usd,
_ppq_state_id_for_provider,
_reconcile_ppq_state,
_record_ppq_invoice,
_set_ppq_state_terminal,
get_ppq_auto_topup_state,
@@ -35,6 +36,8 @@ pytestmark = pytest.mark.asyncio
def _row(provider_id: int = 1) -> MagicMock:
row = MagicMock()
row.id = provider_id
row.api_key = "secret"
row.provider_settings = None
return row
@@ -111,12 +114,35 @@ async def test_claim_is_reusable_once_the_previous_attempt_finished(
await _seed_provider()
first = await _claim_ppq_topup(_row())
assert first is not None
assert await _set_ppq_state_terminal(_row(), first, collected=True, swept=False)
assert await _set_ppq_state_terminal(_row(), first, collected=False, swept=True)
second = await _claim_ppq_topup(_row())
assert second is not None and second != first
async def test_settled_claim_suppresses_immediate_duplicate(
patched_db_engine: Any,
) -> None:
await _seed_provider()
row = _row()
operation_id = await _claim_ppq_topup(row)
assert operation_id is not None
assert await _set_ppq_state_terminal(row, operation_id, collected=True, swept=False)
assert await _reconcile_ppq_state(row, provider=None) is True
assert await _claim_ppq_topup(row) is None
async def test_claim_rejects_stale_provider_configuration(
patched_db_engine: Any,
) -> None:
await _seed_provider()
stale = _row()
stale.provider_settings = '{"auto_topup":true}'
assert await _claim_ppq_topup(stale) is None
async def test_recording_the_invoice_moves_the_claim_in_flight(
patched_db_engine: Any,
) -> None:
@@ -287,7 +313,13 @@ async def test_ppq_payment_audit_row_is_visible_and_survives_next_claim(
assert audit["collected"] is True
assert "lnbc-secret-invoice" not in audit["token"]
# Reusing the deterministic claim lock must not overwrite history.
assert await _claim_ppq_topup(_row()) is None
async with create_session() as session:
state = await session.get(CashuTransaction, _ppq_state_id_for_provider(1))
assert state is not None
state.created_at = int(time.time()) - 301
session.add(state)
await session.commit()
assert await _claim_ppq_topup(_row()) is not None
async with create_session() as session:
assert await session.get(CashuTransaction, audit["id"]) is not None
+4 -17
View File
@@ -686,7 +686,7 @@ async def test_no_database_changes_during_provider_operations(
@pytest.mark.integration
@pytest.mark.asyncio
async def test_admin_routstr_topup_retries_transient_upstream_failure(
async def test_admin_routstr_topup_does_not_duplicate_invoice_on_upstream_failure(
integration_client: AsyncClient,
integration_session: Any,
) -> None:
@@ -739,16 +739,7 @@ async def test_admin_routstr_topup_retries_transient_upstream_failure(
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",
},
)
return MockResponse(500, {"detail": "ambiguous upstream failure"})
mock_client = MockAsyncClient()
@@ -759,11 +750,7 @@ async def test_admin_routstr_topup_retries_transient_upstream_failure(
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
assert response.status_code == 500
assert mock_client.calls == 1
finally:
admin_sessions.pop(admin_token, None)
@@ -289,8 +289,115 @@ async def test_proxy_post_unauthorized_access(integration_client: AsyncClient) -
assert response.status_code in [400, 401]
@pytest.mark.integration
@pytest.mark.asyncio
@pytest.mark.parametrize(
"bad_path",
[
"internal/admin", # unknown endpoint, no API prefix
"v1/../admin", # traversal onto a sibling path
"v1//models", # duplicate separator
"%2e%2e/secret", # encoded dot segment
],
)
async def test_authenticated_post_to_unknown_path_is_rejected(
authenticated_client: AsyncClient, bad_path: str
) -> None:
"""An authenticated POST to an unknown/traversal path must be rejected at
the edge (404) and never forwarded — the provider credential must not reach
an endpoint the caller merely spelled into the URL. If the guard let it
through, forwarding would raise and this would not be a clean 404."""
with patch(
"routstr.upstream.base.BaseUpstreamProvider.forward_request",
AsyncMock(side_effect=AssertionError("must not forward unknown path")),
):
response = await authenticated_client.post(
f"/{bad_path}",
json={"model": "gpt-3.5-turbo", "messages": []},
)
assert response.status_code == 404
@pytest.mark.integration
@pytest.mark.asyncio
@pytest.mark.parametrize(
"bad_path",
[
# Unambiguously spelled, under a prefix the proxy serves, but not an
# endpoint it offers. These are real upstream routes that manage keys,
# org membership, and billing on the same origin as inference.
"v1/organization/api_keys",
"v1/api_keys",
"v1/billing/usage",
"v1/admin/keys",
# An id segment is honoured one level deep, and only where an endpoint
# takes one at all.
"v1/chat/completions/abc",
"v1/models/gpt-4/secret",
],
)
async def test_known_prefix_does_not_carry_an_unknown_endpoint(
authenticated_client: AsyncClient, bad_path: str
) -> None:
"""A familiar prefix must not be a passport for the rest of the origin.
Nothing here is traversal-shaped, so the spelling screen lets it by; only
the endpoint allowlist stops it. Forwarding raises if the guard misses,
so a clean 404 also proves the credential never left."""
with patch(
"routstr.upstream.base.BaseUpstreamProvider.forward_request",
AsyncMock(side_effect=AssertionError("must not forward unknown endpoint")),
):
response = await authenticated_client.post(
f"/{bad_path}",
json={"model": "gpt-3.5-turbo", "messages": []},
)
assert response.status_code == 404
@pytest.mark.integration
@pytest.mark.asyncio
async def test_get_unknown_path_is_rejected(integration_client: AsyncClient) -> None:
"""A GET to an unknown (no API prefix) path is rejected at the edge."""
response = await integration_client.get("/internal/admin")
assert response.status_code == 404
@pytest.mark.integration
@pytest.mark.asyncio
@pytest.mark.parametrize(
"bad_path",
[
# Percent-encoded so the traversal survives the client to the server,
# which decodes it to "v1/../admin" before routing.
"v1/%2e%2e/admin",
# Well-spelled, so only the endpoint allowlist can stop it.
"v1/organization/api_keys",
"anything/encrypted",
],
)
async def test_ehbp_request_is_gated_by_the_endpoint_allowlist(
integration_client: AsyncClient, bad_path: str
) -> None:
"""EHBP hides the body from the proxy, not the destination.
The encrypted contract covers the request body; it says nothing about which
endpoint the provider credential gets spent against, so an EHBP request is
screened and allowlisted exactly like any other."""
with patch(
"routstr.proxy.forward_ehbp_request",
AsyncMock(side_effect=AssertionError("must not forward unknown endpoint")),
):
response = await integration_client.post(
f"/{bad_path}",
content=b"encrypted",
headers={
"ehbp-encapsulated-key": "x",
"x-routstr-model": "gpt-4",
},
)
assert response.status_code == 404
@pytest.mark.integration
@pytest.mark.asyncio
+31 -29
View File
@@ -45,7 +45,7 @@ def _dead_key(created_at: int | None) -> ApiKey:
@pytest.mark.asyncio
async def test_prunes_old_refunded_zero_key(patched_db_engine: None) -> None:
"""A funded-then-refunded key (0/0/0, NULL parent, old) is pruned."""
"""A funded-then-refunded zero-balance key with an old timestamp is pruned."""
key = _dead_key(LONG_AGO)
async with create_session() as session:
session.add(key)
@@ -98,34 +98,6 @@ async def test_used_key_never_pruned(patched_db_engine: None) -> None:
assert await _exists(k.hashed_key)
@pytest.mark.asyncio
async def test_parent_and_child_keys_are_not_pruned(
patched_db_engine: None,
) -> None:
"""Pruning must not orphan child keys or delete valid children."""
parent = _dead_key(LONG_AGO)
child = ApiKey(
hashed_key=f"child_{uuid.uuid4().hex}",
balance=0,
reserved_balance=0,
total_spent=0,
total_requests=0,
created_at=LONG_AGO,
parent_key_hash=parent.hashed_key,
)
async with create_session() as session:
session.add(parent)
session.add(child)
await session.commit()
async with create_session() as session:
pruned = await prune_dead_api_keys(session, OLD)
assert pruned == 0
assert await _exists(parent.hashed_key)
assert await _exists(child.hashed_key)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("status", "expires_at"),
@@ -321,3 +293,33 @@ async def test_periodic_prune_disabled_returns_immediately(
await auth.periodic_dead_key_prune()
sleep_mock.assert_not_called()
@pytest.mark.asyncio
async def test_refund_claim_protects_key(patched_db_engine: None) -> None:
"""A key with a refund row is the audit anchor for its payout and holds a
non-null FK, so the janitor must leave it alone."""
from routstr.core.db import Refund
key = _dead_key(LONG_AGO)
async with create_session() as session:
session.add(key)
await session.commit()
session.add(
Refund(
api_key_hashed_key=key.hashed_key,
method="lightning",
destination="user@ln.example.com",
amount_msats=1000,
unit="sat",
mint_url="https://mint.example.com",
status="paid",
)
)
await session.commit()
async with create_session() as session:
pruned = await prune_dead_api_keys(session, OLD)
assert pruned == 0
assert await _exists(key.hashed_key)
+891
View File
@@ -0,0 +1,891 @@
"""Refund claim lifecycle against a real SQLite database.
Covers the guarantees the ``refunds`` table exists to provide: one open claim
per key, a persisted melt quote before the melt is dispatched, and a
reconciler that never restores a balance whose payout may have settled.
"""
import time
from typing import Any, Awaitable, Callable
from unittest.mock import AsyncMock, patch
import httpx
import pytest
from fastapi import HTTPException
from sqlmodel import select
from routstr import refund
from routstr.balance import RefundRequest, refund_wallet_endpoint
from routstr.core.db import (
ApiKey,
AsyncSession,
CashuTransaction,
Refund,
store_cashu_transaction_with_retry,
total_user_liability,
)
from routstr.payment.lnurl import LNURLError, MeltOutcomeAmbiguousError
KEY_HASH = "refundclaimkey"
ADDRESS = "user@ln.example.com"
BALANCE_MSATS = 5_000_000
async def _seed_key(
session: AsyncSession, *, balance: int = BALANCE_MSATS, address: str | None = None
) -> ApiKey:
key = ApiKey(hashed_key=KEY_HASH)
key.balance = balance
key.reserved_balance = 0
key.refund_currency = "sat"
key.refund_address = address
key.total_spent = 0
key.total_requests = 0
session.add(key)
await session.commit()
await session.refresh(key)
return key
async def _load_key(session: AsyncSession) -> ApiKey:
key = await session.get(ApiKey, KEY_HASH)
assert key is not None
await session.refresh(key)
return key
async def _load_refund(session: AsyncSession, refund_id: str) -> Refund:
row = await session.get(Refund, refund_id)
assert row is not None
await session.refresh(row)
return row
async def _age_claim(session: AsyncSession, refund_id: str, seconds: int) -> None:
row = await _load_refund(session, refund_id)
row.claimed_at = int(time.time()) - seconds
row.created_at = int(time.time()) - seconds
row.updated_at = int(time.time()) - seconds
session.add(row)
await session.commit()
@pytest.fixture
def short_timeout() -> Any:
with patch.object(refund.settings, "refund_claim_timeout_seconds", 300):
yield
# --- exclusivity -----------------------------------------------------------
@pytest.mark.asyncio
async def test_second_claim_on_open_key_is_rejected(
integration_session: AsyncSession,
) -> None:
key = await _seed_key(integration_session)
first = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
key = await _load_key(integration_session)
assert key.balance == 0
with pytest.raises(HTTPException) as exc_info:
await refund.open_claim(
integration_session, key, method="cashu", destination=None
)
detail = exc_info.value.detail
assert exc_info.value.status_code == 409
assert isinstance(detail, dict)
assert detail["error"]["code"] == "refund_in_progress"
rows = (await integration_session.exec(select(Refund))).all()
assert [row.id for row in rows] == [first.id]
assert (await _load_key(integration_session)).balance == 0
@pytest.mark.asyncio
async def test_claim_from_second_session_hits_the_index(
integration_engine: Any, integration_session: AsyncSession
) -> None:
key = await _seed_key(integration_session)
await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
async with AsyncSession(integration_engine, expire_on_commit=False) as other:
other_key = await _load_key(other)
with pytest.raises(HTTPException) as exc_info:
await refund.open_claim(
other, other_key, method="lightning", destination=ADDRESS
)
assert exc_info.value.status_code == 409
@pytest.mark.asyncio
async def test_claim_rejects_stale_balance_snapshot(
integration_session: AsyncSession,
) -> None:
key = await _seed_key(integration_session)
integration_session.expunge(key) # a stale, detached snapshot
key.balance = BALANCE_MSATS + 1
with pytest.raises(HTTPException) as exc_info:
await refund.open_claim(
integration_session, key, method="cashu", destination=None
)
assert exc_info.value.status_code == 409
assert (await integration_session.exec(select(Refund))).all() == []
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
@pytest.mark.asyncio
async def test_retry_after_failed_claim_pays_once(
integration_session: AsyncSession,
) -> None:
key = await _seed_key(integration_session)
first = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
assert await refund.release(integration_session, first)
key = await _load_key(integration_session)
assert key.balance == BALANCE_MSATS
second = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
assert second.id != first.id
assert (await _load_key(integration_session)).balance == 0
@pytest.mark.asyncio
async def test_release_after_settle_is_a_noop(
integration_session: AsyncSession,
) -> None:
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
assert await refund.settle(integration_session, claim, quote_id="q1")
assert not await refund.release(integration_session, claim)
assert (await _load_key(integration_session)).balance == 0
row = await _load_refund(integration_session, claim.id)
assert row.status == "paid"
assert row.claimed_at is None
# --- execute ---------------------------------------------------------------
def _lnurl_stub(
outcome: BaseException | None = None,
*,
quoted: bool = True,
) -> Callable[..., Awaitable[int]]:
async def send(
amount: int,
unit: str,
mint: str,
address: str,
*,
on_melt_quote: Callable[[str, str], Awaitable[None]] | None = None,
) -> int:
if quoted and on_melt_quote is not None:
await on_melt_quote("quote-123", mint)
if outcome is not None:
raise outcome
return amount
return send
@pytest.mark.asyncio
async def test_execute_persists_quote_before_melt_and_settles(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
seen: list[str | None] = []
async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int:
await on_melt_quote("quote-123", claim.mint_url)
row = await _load_refund(integration_session, claim.id)
seen.append(row.quote_id)
return 5000
with patch("routstr.refund.send_to_lnurl", send):
body = await refund.execute(integration_session, claim)
assert seen == ["quote-123"], "quote must be on disk before the melt runs"
assert body["status"] == "paid"
assert body["recipient"] == ADDRESS
assert body["sats"] == "5000"
row = await _load_refund(integration_session, claim.id)
assert (row.status, row.quote_id, row.claimed_at) == ("paid", "quote-123", None)
@pytest.mark.asyncio
async def test_execute_ambiguous_holds_claim_with_quote(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
with patch(
"routstr.refund.send_to_lnurl", _lnurl_stub(MeltOutcomeAmbiguousError("?"))
):
with pytest.raises(HTTPException) as exc_info:
await refund.execute(integration_session, claim)
assert exc_info.value.status_code == 502
row = await _load_refund(integration_session, claim.id)
assert (row.status, row.quote_id, row.claimed_at) == (
"ambiguous",
"quote-123",
None,
)
assert (await _load_key(integration_session)).balance == 0
@pytest.mark.asyncio
async def test_execute_clean_failure_restores_balance(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
with patch(
"routstr.refund.send_to_lnurl",
_lnurl_stub(LNURLError("limits"), quoted=False),
):
with pytest.raises(HTTPException) as exc_info:
await refund.execute(integration_session, claim)
assert exc_info.value.status_code == 500
row = await _load_refund(integration_session, claim.id)
assert row.status == "failed"
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
@pytest.mark.asyncio
async def test_execute_failure_after_quote_withholds_balance(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""The mint may have paid the quote, so a later local failure must not restore."""
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
with patch("routstr.refund.send_to_lnurl", _lnurl_stub(RuntimeError("local"))):
with pytest.raises(HTTPException) as exc_info:
await refund.execute(integration_session, claim)
assert exc_info.value.status_code == 502
row = await _load_refund(integration_session, claim.id)
assert (row.status, row.quote_id) == ("ambiguous", "quote-123")
assert (await _load_key(integration_session)).balance == 0
@pytest.mark.asyncio
async def test_execute_records_mint_that_issued_the_quote(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""Mint fallback must leave reconciliation pointed at the issuing mint."""
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
fallback_mint = "https://fallback.mint.example"
async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int:
await on_melt_quote("quote-fallback", fallback_mint)
raise MeltOutcomeAmbiguousError("unknown")
with patch("routstr.refund.send_to_lnurl", send):
with pytest.raises(MeltOutcomeAmbiguousError):
await refund._pay_lightning(integration_session, claim)
row = await _load_refund(integration_session, claim.id)
assert (row.mint_url, row.quote_id, row.status) == (
fallback_mint,
"quote-fallback",
"ambiguous",
)
@pytest.mark.asyncio
async def test_execute_aborts_melt_when_claim_was_released(
integration_engine: Any, integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""A reconciler that released the claim first must stop the melt."""
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
melted = False
async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int:
nonlocal melted
async with AsyncSession(integration_engine, expire_on_commit=False) as other:
await refund.release(other, await _load_refund(other, claim.id))
await on_melt_quote("quote-123", claim.mint_url)
melted = True
return 5000
with patch("routstr.refund.send_to_lnurl", send):
with pytest.raises(HTTPException):
await refund.execute(integration_session, claim)
assert melted is False
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
# --- reconciler ------------------------------------------------------------
async def _open_ambiguous(session: AsyncSession, quote_id: str | None) -> Refund:
key = await _seed_key(session)
claim = await refund.open_claim(
session, key, method="lightning", destination=ADDRESS
)
await refund.hold(session, claim, quote_id)
return claim
@pytest.mark.asyncio
@pytest.mark.parametrize(
("mint_status", "expected_status", "expected_balance"),
[
("paid", "paid", 0),
("unpaid", "failed", BALANCE_MSATS),
("pending", "ambiguous", 0),
("unknown", "ambiguous", 0),
],
)
async def test_reconcile_ambiguous_claims(
integration_session: AsyncSession,
patched_db_engine: None,
short_timeout: None,
mint_status: str,
expected_status: str,
expected_balance: int,
) -> None:
claim = await _open_ambiguous(integration_session, "quote-123")
await _age_claim(integration_session, claim.id, 600)
with patch(
"routstr.refund.check_bolt11_payment_status",
AsyncMock(return_value=mint_status),
) as check:
await refund.reconcile_once()
check.assert_awaited_once_with(claim.mint_url, "sat", "quote-123")
row = await _load_refund(integration_session, claim.id)
assert row.status == expected_status
assert (await _load_key(integration_session)).balance == expected_balance
@pytest.mark.asyncio
async def test_reconcile_credits_balance_once_across_passes(
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
) -> None:
claim = await _open_ambiguous(integration_session, "quote-123")
await _age_claim(integration_session, claim.id, 600)
with patch(
"routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="unpaid")
):
await refund.reconcile_once()
await refund.reconcile_once()
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
@pytest.mark.asyncio
async def test_reconcile_waits_before_trusting_recent_unpaid(
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
) -> None:
"""A just-dispatched melt can report unpaid before it turns pending, so an
unpaid quote is only final once the claim has been quiet for a timeout."""
claim = await _open_ambiguous(integration_session, "quote-123")
with patch(
"routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="unpaid")
):
await refund.reconcile_once()
row = await _load_refund(integration_session, claim.id)
assert row.status == "ambiguous"
assert (await _load_key(integration_session)).balance == 0
await _age_claim(integration_session, claim.id, 600)
with patch(
"routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="unpaid")
):
await refund.reconcile_once()
row = await _load_refund(integration_session, claim.id)
assert row.status == "failed"
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
@pytest.mark.asyncio
async def test_reconcile_keeps_claim_that_gained_quote_mid_pass(
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
) -> None:
"""The payout records its quote between the reconciler's read and its
release: the stale ``quote_id is None`` must not restore the balance."""
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
await _age_claim(integration_session, claim.id, 600)
real_lease = refund._lease
async def lease_then_quote(refund_id: str, now: int, cutoff: int) -> bool:
leased = await real_lease(refund_id, now, cutoff)
await refund.record_quote(claim, "late-quote", claim.mint_url)
return leased
with patch("routstr.refund._lease", lease_then_quote):
await refund.reconcile_once()
row = await _load_refund(integration_session, claim.id)
assert row.status == "pending"
assert row.quote_id == "late-quote"
assert (await _load_key(integration_session)).balance == 0
@pytest.mark.asyncio
async def test_reconcile_leaves_fresh_pending_claim_alone(
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
) -> None:
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
with patch("routstr.refund.check_bolt11_payment_status", AsyncMock()) as check:
await refund.reconcile_once()
check.assert_not_awaited()
row = await _load_refund(integration_session, claim.id)
assert row.status == "pending"
assert row.claimed_at is not None
assert (await _load_key(integration_session)).balance == 0
@pytest.mark.asyncio
async def test_reconcile_releases_expired_claim_without_quote(
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
) -> None:
"""No quote on disk means the mint was never asked to pay."""
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
await _age_claim(integration_session, claim.id, 600)
with patch("routstr.refund.check_bolt11_payment_status", AsyncMock()) as check:
await refund.reconcile_once()
check.assert_not_awaited()
assert (await _load_refund(integration_session, claim.id)).status == "failed"
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
@pytest.mark.asyncio
async def test_reconcile_queries_mint_for_crashed_claim_with_quote(
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
) -> None:
"""Crash after the quote was persisted: the mint decides, not the timeout."""
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
await refund.record_quote(claim, "quote-crash", claim.mint_url)
await _age_claim(integration_session, claim.id, 600)
with patch(
"routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="paid")
) as check:
await refund.reconcile_once()
check.assert_awaited_once_with(claim.mint_url, "sat", "quote-crash")
assert (await _load_refund(integration_session, claim.id)).status == "paid"
assert (await _load_key(integration_session)).balance == 0
@pytest.mark.asyncio
async def test_reconcile_marks_expired_cashu_claim_stuck(
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
) -> None:
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="cashu", destination=None
)
await _age_claim(integration_session, claim.id, 600)
with patch("routstr.refund.logger") as log:
await refund.reconcile_once()
await refund.reconcile_once()
assert log.critical.call_count == 1
row = await _load_refund(integration_session, claim.id)
assert (row.status, row.claimed_at) == ("stuck", None)
assert (await _load_key(integration_session)).balance == 0
# A stuck claim is closed, so the key is not permanently locked out.
key = await _load_key(integration_session)
key.balance = 1000
integration_session.add(key)
await integration_session.commit()
await refund.open_claim(integration_session, key, method="cashu", destination=None)
@pytest.mark.asyncio
async def test_reconcile_survives_one_failing_row(
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
) -> None:
claim = await _open_ambiguous(integration_session, "quote-123")
with patch(
"routstr.refund.check_bolt11_payment_status",
AsyncMock(side_effect=RuntimeError("mint down")),
):
await refund.reconcile_once()
row = await _load_refund(integration_session, claim.id)
assert row.status == "ambiguous"
assert row.claimed_at is not None, "lease is kept until the next pass"
# --- endpoint --------------------------------------------------------------
@pytest.mark.asyncio
async def test_endpoint_uses_requested_address_over_persisted(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
await _seed_key(integration_session, address="stored@ln.example.com")
send = AsyncMock(side_effect=_lnurl_stub())
with (
patch("routstr.refund.get_lnurl_data", AsyncMock()) as resolve,
patch("routstr.refund.send_to_lnurl", send),
):
body = await refund_wallet_endpoint(
refund_request=RefundRequest(lightning_address=ADDRESS),
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
resolve.assert_awaited_once_with(ADDRESS)
assert isinstance(body, dict)
assert body["recipient"] == ADDRESS
assert send.await_args is not None
assert send.await_args.args[3] == ADDRESS
assert (await _load_key(integration_session)).balance == 0
@pytest.mark.asyncio
async def test_endpoint_rejects_bad_address_without_debit(
integration_session: AsyncSession,
) -> None:
await _seed_key(integration_session)
with patch(
"routstr.refund.get_lnurl_data", AsyncMock(side_effect=LNURLError("nope"))
):
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
refund_request=RefundRequest(lightning_address="bad@example"),
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert exc_info.value.status_code == 400
assert (await integration_session.exec(select(Refund))).all() == []
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
@pytest.mark.asyncio
async def test_endpoint_replays_paid_lightning_refund_on_empty_balance(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
await _seed_key(integration_session, address=ADDRESS)
with (
patch("routstr.refund.get_lnurl_data", AsyncMock()),
patch("routstr.refund.send_to_lnurl", _lnurl_stub()),
):
first = await refund_wallet_endpoint(
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
second = await refund_wallet_endpoint(
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert isinstance(first, dict) and isinstance(second, dict)
assert second["refund_id"] == first["refund_id"]
assert second["status"] == "paid"
@pytest.mark.asyncio
async def test_endpoint_refund_while_ambiguous_returns_409(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
await _open_ambiguous(integration_session, "quote-123")
key = await _load_key(integration_session)
key.balance = 2_000_000 # topped up while the melt is unresolved
integration_session.add(key)
await integration_session.commit()
send = AsyncMock()
with (
patch("routstr.refund.get_lnurl_data", AsyncMock()),
patch("routstr.refund.send_to_lnurl", send),
):
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
refund_request=RefundRequest(lightning_address=ADDRESS),
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert exc_info.value.status_code == 409
send.assert_not_awaited()
assert (await _load_key(integration_session)).balance == 2_000_000
@pytest.mark.asyncio
async def test_endpoint_zero_balance_with_open_claim_returns_409(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""A prior refund debited the balance and is still settling: the retry must
report refund_in_progress (409), not "no balance to refund" (400)."""
await _open_ambiguous(integration_session, "quote-123")
assert (await _load_key(integration_session)).balance == 0
send = AsyncMock()
with (
patch("routstr.refund.get_lnurl_data", AsyncMock()),
patch("routstr.refund.send_to_lnurl", send),
):
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert exc_info.value.status_code == 409
assert isinstance(exc_info.value.detail, dict)
assert exc_info.value.detail["error"]["code"] == "refund_in_progress"
send.assert_not_awaited()
@pytest.mark.asyncio
async def test_cashu_token_survives_failed_ledger_write(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""The token is issued once the mint signs it; a failed cashu_transactions
insert must neither fail the request nor release the balance, and a retry
must replay the token from the claim row."""
await _seed_key(integration_session)
with (
patch("routstr.refund.send_token", AsyncMock(return_value="cashuAtoken")),
patch("routstr.refund.token_mint_url", lambda token, mint: mint),
patch(
"routstr.refund.store_cashu_transaction",
AsyncMock(side_effect=RuntimeError("db down")),
),
):
first = await refund_wallet_endpoint(
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert isinstance(first, dict)
assert (first["token"], first["status"]) == ("cashuAtoken", "paid")
assert (await _load_key(integration_session)).balance == 0
second = await refund_wallet_endpoint(
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert isinstance(second, dict)
assert second["refund_id"] == first["refund_id"]
assert second["token"] == "cashuAtoken"
async def _topup(session: AsyncSession, amount: int = BALANCE_MSATS) -> None:
key = await _load_key(session)
key.balance = amount
session.add(key)
await session.commit()
async def _refund_cashu(
session: AsyncSession, token: str, *, ledger: bool
) -> dict[str, str]:
store = (
AsyncMock(side_effect=RuntimeError("db down"))
if not ledger
else store_cashu_transaction_with_retry
)
with (
patch("routstr.refund.send_token", AsyncMock(return_value=token)),
patch("routstr.refund.token_mint_url", lambda t, mint: mint),
patch("routstr.refund.store_cashu_transaction", store),
):
body = await refund_wallet_endpoint(
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=session,
)
assert isinstance(body, dict)
return body
@pytest.mark.asyncio
async def test_replay_prefers_the_token_of_the_newest_paid_claim(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""An older ledger row must not answer for a newer claim whose write failed."""
await _seed_key(integration_session)
await _refund_cashu(integration_session, "cashuAold", ledger=True)
await _topup(integration_session)
newest = await _refund_cashu(integration_session, "cashuAnew", ledger=False)
replay = await refund_wallet_endpoint(
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert isinstance(replay, dict)
assert replay["token"] == "cashuAnew"
assert replay["refund_id"] == newest["refund_id"]
@pytest.mark.asyncio
async def test_swept_older_ledger_row_does_not_reject_the_newest_claim(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
await _seed_key(integration_session)
await _refund_cashu(integration_session, "cashuAold", ledger=True)
await _topup(integration_session)
await _refund_cashu(integration_session, "cashuAnew", ledger=False)
result = await integration_session.exec(
select(CashuTransaction).where(CashuTransaction.token == "cashuAold")
)
old_tx = result.one()
old_tx.swept = True
integration_session.add(old_tx)
await integration_session.commit()
replay = await refund_wallet_endpoint(
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert isinstance(replay, dict)
assert replay["token"] == "cashuAnew"
@pytest.mark.asyncio
async def test_open_claim_is_reported_over_an_older_paid_claim(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""A paid claim from a previous cycle must not be replayed as the outcome
of the claim that is still settling."""
await _seed_key(integration_session, address=ADDRESS)
with (
patch("routstr.refund.get_lnurl_data", AsyncMock()),
patch("routstr.refund.send_to_lnurl", _lnurl_stub()),
):
paid = await refund_wallet_endpoint(
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert isinstance(paid, dict)
await _topup(integration_session)
key = await _load_key(integration_session)
open_claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
await refund.hold(integration_session, open_claim, "quote-open")
validate = AsyncMock()
with patch("routstr.refund.get_lnurl_data", validate):
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
refund_request=RefundRequest(lightning_address=ADDRESS),
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert exc_info.value.status_code == 409
assert isinstance(exc_info.value.detail, dict)
error = exc_info.value.detail["error"]
assert (error["refund_id"], error["status"]) == (open_claim.id, "ambiguous")
# A drained key must not be able to drive outbound destination lookups.
validate.assert_not_awaited()
@pytest.mark.asyncio
async def test_stuck_claim_is_reported_instead_of_no_balance(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="cashu", destination=None
)
await refund._close(integration_session, claim, status="stuck")
await integration_session.commit()
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert exc_info.value.status_code == 409
assert isinstance(exc_info.value.detail, dict)
error = exc_info.value.detail["error"]
assert (error["code"], error["refund_id"]) == ("refund_unresolved", claim.id)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("status", "counted"),
[
("pending", True),
("ambiguous", True),
("stuck", True),
("paid", False),
("failed", False),
],
)
async def test_liability_covers_claims_until_they_resolve(
integration_session: AsyncSession,
patched_db_engine: None,
status: str,
counted: bool,
) -> None:
"""Money in flight is still owed to the customer; an owner payout that read
only key balances could spend its backing."""
key = await _seed_key(integration_session)
before = await total_user_liability(integration_session)
assert before == BALANCE_MSATS
claim = await refund.open_claim(
integration_session, key, method="cashu", destination=None
)
assert await total_user_liability(integration_session) == BALANCE_MSATS
await refund._close(integration_session, claim, status=status)
await integration_session.commit()
expected = BALANCE_MSATS if counted else 0
assert await total_user_liability(integration_session) == expected
@pytest.mark.asyncio
async def test_unreachable_destination_is_a_client_error() -> None:
with patch(
"routstr.refund.get_lnurl_data",
AsyncMock(side_effect=httpx.ConnectError("All connection attempts failed")),
):
with pytest.raises(HTTPException) as exc_info:
await refund.validate_lightning_destination(ADDRESS)
assert exc_info.value.status_code == 400
@@ -0,0 +1,368 @@
"""Refund payouts that finish after the claim row stopped cooperating.
Each test pins one guarantee of the claim table that the happy path cannot
exercise: the payout side has authoritative knowledge of what the mint did,
and the claim row must end up agreeing with it.
"""
import time
from typing import Any
from unittest.mock import AsyncMock, patch
import pytest
from fastapi import HTTPException
from sqlalchemy import event
from sqlalchemy.exc import OperationalError
from routstr import refund
from routstr.balance import RefundRequest, refund_wallet_endpoint
from routstr.core.db import ApiKey, AsyncSession, Refund, total_user_liability
from routstr.payment.lnurl import MeltUnpaidError
KEY_HASH = "refundrecoverykey"
ADDRESS = "user@ln.example.com"
BALANCE_MSATS = 5_000_000
async def _seed_key(session: AsyncSession, *, address: str | None = None) -> ApiKey:
key = ApiKey(hashed_key=KEY_HASH)
key.balance = BALANCE_MSATS
key.reserved_balance = 0
key.refund_currency = "sat"
key.refund_address = address
key.total_spent = 0
key.total_requests = 0
session.add(key)
await session.commit()
await session.refresh(key)
return key
async def _load_key(session: AsyncSession) -> ApiKey:
key = await session.get(ApiKey, KEY_HASH)
assert key is not None
await session.refresh(key)
return key
async def _load_refund(session: AsyncSession, refund_id: str) -> Refund:
row = await session.get(Refund, refund_id)
assert row is not None
await session.refresh(row)
return row
def _cashu_payout(token: str, send: Any | None = None) -> Any:
return (
patch("routstr.refund.send_token", send or AsyncMock(return_value=token)),
patch("routstr.refund.token_mint_url", lambda t, mint: mint),
patch("routstr.refund.store_cashu_transaction", AsyncMock()),
)
# --- cashu token issued after the reconciler gave up -----------------------
@pytest.mark.asyncio
async def test_cashu_token_issued_after_reconciler_marked_claim_stuck_settles_it(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""A token in the customer's hands must leave its claim ``paid``: a row
left ``stuck`` keeps the amount in liability forever and tells the operator
to reconcile a payout that already happened."""
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="cashu", destination=None
)
async def slow_send_token(amount: int, unit: str, mint_url: str) -> str:
# The lease lapses while the mint is still working.
row = await _load_refund(integration_session, claim.id)
row.claimed_at = (row.claimed_at or 0) - 10_000
integration_session.add(row)
await integration_session.commit()
await refund.reconcile_once()
assert (await _load_refund(integration_session, claim.id)).status == "stuck"
return "cashuAlate"
send, mint, store = _cashu_payout("cashuAlate", slow_send_token)
with send, mint, store:
body = await refund.execute(integration_session, claim)
assert (body["token"], body["status"]) == ("cashuAlate", "paid")
row = await _load_refund(integration_session, claim.id)
assert (row.status, row.token, row.claimed_at) == ("paid", "cashuAlate", None)
assert await total_user_liability(integration_session) == 0
@pytest.mark.asyncio
async def test_cashu_payout_renews_lease_before_asking_the_mint(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""The reconciler leaves a claim alone while its lease is fresh, so the
payout renews it right before the slow step."""
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="cashu", destination=None
)
row = await _load_refund(integration_session, claim.id)
row.claimed_at = (row.claimed_at or 0) - 10_000
integration_session.add(row)
await integration_session.commit()
async def send_token(amount: int, unit: str, mint_url: str) -> str:
await refund.reconcile_once()
return "cashuAfresh"
send, mint, store = _cashu_payout("cashuAfresh", send_token)
with send, mint, store, patch("routstr.refund.logger") as log:
await refund.execute(integration_session, claim)
log.critical.assert_not_called()
assert (await _load_refund(integration_session, claim.id)).status == "paid"
# --- cashu token issued, claim write failed --------------------------------
@pytest.mark.asyncio
@pytest.mark.parametrize("failure_point", ["execute", "autoflush", "commit"])
async def test_cashu_claim_write_failure_after_token_creation_withholds_balance(
integration_session: AsyncSession,
patched_db_engine: None,
integration_engine: Any,
failure_point: str,
) -> None:
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="cashu", destination=None
)
claim_id = claim.id
fired = False
armed = False
def fail_once(*args: Any) -> None:
nonlocal fired
statement = args[2] if failure_point != "commit" else "COMMIT"
if (
armed
and not fired
and (statement.startswith("UPDATE refunds") or statement == "COMMIT")
):
fired = True
raise OperationalError(statement, {}, Exception("database is locked"))
async def send_token(amount: int, unit: str, mint_url: str) -> str:
nonlocal armed
if failure_point == "autoflush":
# A pending ORM write makes SQLAlchemy invalidate the transaction
# and expire attached objects when the actual SQL execution fails.
claim.updated_at -= 1
armed = True
return "cashuAstranded"
event_name = "commit" if failure_point == "commit" else "before_cursor_execute"
event.listen(integration_engine.sync_engine, event_name, fail_once)
send, _, store = _cashu_payout("cashuAstranded", send_token)
try:
with (
send,
store,
patch(
"routstr.refund.token_mint_url", return_value="https://fallback.mint"
),
):
with pytest.raises(HTTPException) as exc_info:
await refund.execute(integration_session, claim)
finally:
event.remove(integration_engine.sync_engine, event_name, fail_once)
assert fired
assert exc_info.value.status_code == 502
row = await _load_refund(integration_session, claim_id)
assert (row.status, row.token, row.mint_url) == (
"ambiguous",
"cashuAstranded",
"https://fallback.mint",
)
assert (await _load_key(integration_session)).balance == 0
assert await total_user_liability(integration_session) == BALANCE_MSATS
await refund.reconcile_once()
row = await _load_refund(integration_session, claim_id)
assert (row.status, row.token) == ("paid", "cashuAstranded")
assert await total_user_liability(integration_session) == 0
with patch("routstr.refund.send_token", AsyncMock()) as send_again:
replay = await refund_wallet_endpoint(
refund_request=RefundRequest(),
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert isinstance(replay, dict)
assert replay["token"] == "cashuAstranded"
send_again.assert_not_awaited()
assert (await _load_key(integration_session)).balance == 0
@pytest.mark.asyncio
async def test_reconciler_settles_held_cashu_claim_that_carries_a_token(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""A held cashu claim whose row carries the token is a completed payout;
the reconciler closes it as paid instead of escalating it to stuck."""
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="cashu", destination=None
)
await refund.hold(integration_session, claim, None, token="cashuAheld")
row = await _load_refund(integration_session, claim.id)
row.claimed_at = None
row.updated_at -= 10_000
integration_session.add(row)
await integration_session.commit()
with patch("routstr.refund.logger") as log:
await refund.reconcile_once()
log.critical.assert_not_called()
row = await _load_refund(integration_session, claim.id)
assert (row.status, row.token) == ("paid", "cashuAheld")
assert await total_user_liability(integration_session) == 0
@pytest.mark.asyncio
@pytest.mark.parametrize("hold_before_lease", [True, False])
async def test_stale_cashu_reconciliation_preserves_newly_held_token(
integration_session: AsyncSession,
patched_db_engine: None,
integration_engine: Any,
hold_before_lease: bool,
) -> None:
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="cashu", destination=None
)
claim_id = claim.id
now = int(time.time())
cutoff = now - refund.settings.refund_claim_timeout_seconds
claim.claimed_at = cutoff - 1
await integration_session.commit()
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
stale = await session.get(Refund, claim_id)
assert stale is not None and stale.token is None
if hold_before_lease:
await refund.hold(integration_session, claim, None, token="cashuAheld")
assert await refund._lease(claim_id, now, cutoff)
real_close = refund._close
async def hold_before_close(session: AsyncSession, row: Refund, **kw: Any) -> bool:
if not hold_before_lease and kw.get("status") == "stuck":
# Token arrives at the last moment, even after a potential reload.
await real_close(
integration_session, claim, status="ambiguous", token="cashuAheld"
)
await integration_session.commit()
return await real_close(session, row, **kw)
with patch.object(refund, "_close", hold_before_close):
await refund._reconcile(stale, now)
row = await _load_refund(integration_session, claim_id)
assert (row.status, row.token) == ("paid", "cashuAheld")
assert (await _load_key(integration_session)).balance == 0
assert await total_user_liability(integration_session) == 0
with patch("routstr.refund.send_token", AsyncMock()) as send_again:
replay = await refund_wallet_endpoint(
refund_request=RefundRequest(),
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert isinstance(replay, dict)
assert replay["token"] == "cashuAheld"
send_again.assert_not_awaited()
# --- mint proved the melt unpaid -------------------------------------------
@pytest.mark.asyncio
async def test_mint_confirmed_unpaid_melt_restores_balance_immediately(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""The mint answering ``unpaid`` to the melt itself is proof no payment
happened, so the customer gets the balance back now, not after the
reconciler timeout."""
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="lightning", destination=ADDRESS
)
async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int:
await on_melt_quote("quote-unpaid", claim.mint_url)
raise MeltUnpaidError("Cashu mint confirmed that the melt was unpaid")
with patch("routstr.refund.send_to_lnurl", send):
with pytest.raises(HTTPException) as exc_info:
await refund.execute(integration_session, claim)
assert exc_info.value.status_code == 503
row = await _load_refund(integration_session, claim.id)
assert (row.status, row.quote_id) == ("failed", "quote-unpaid")
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
# --- response reflects the persisted claim ---------------------------------
@pytest.mark.asyncio
async def test_response_status_reflects_persisted_claim(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""When ``settle`` closes nothing the response must not invent ``paid``."""
key = await _seed_key(integration_session)
claim = await refund.open_claim(
integration_session, key, method="cashu", destination=None
)
async def send_token(amount: int, unit: str, mint_url: str) -> str:
# Somebody closed the row as failed while the mint was working.
await refund._close(integration_session, claim, status="failed")
await integration_session.commit()
return "cashuAorphan"
send, mint, store = _cashu_payout("cashuAorphan", send_token)
with send, mint, store:
body = await refund.execute(integration_session, claim)
assert body["status"] == "failed"
assert (await _load_refund(integration_session, claim.id)).status == "failed"
@pytest.mark.asyncio
async def test_stored_refund_address_is_validated_before_debit(
integration_session: AsyncSession, patched_db_engine: None
) -> None:
"""A bad address stored on the key is a client error, not a payout failure."""
await _seed_key(integration_session, address="nobody@invalid.example")
send = AsyncMock()
with (
patch(
"routstr.refund.get_lnurl_data",
AsyncMock(side_effect=refund.LNURLError("no such user")),
),
patch("routstr.refund.send_to_lnurl", send),
):
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
refund_request=RefundRequest(),
authorization=f"Bearer sk-{KEY_HASH}",
x_cashu=None,
session=integration_session,
)
assert exc_info.value.status_code == 400
send.assert_not_awaited()
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
@@ -133,14 +133,11 @@ async def test_reserved_balance_with_successful_requests(
@pytest.mark.asyncio
async def test_revert_with_zero_reserved_balance_is_noop(
async def test_revert_with_zero_reserved_balance_repairs_terminally(
integration_session: AsyncSession,
) -> None:
"""Test that revert_pay_for_request is a no-op when reserved_balance is 0.
Previously this would drive reserved_balance negative. With the floor guard,
it should return False and leave reserved_balance at 0.
"""
"""Reverting after the aggregate was already zeroed must not drive it
negative: the corrupt durable reservation is released without subtraction."""
from routstr.auth import pay_for_request, revert_pay_for_request
unique_key = f"test_revert_key_{uuid.uuid4().hex[:8]}"
@@ -153,21 +150,24 @@ async def test_revert_with_zero_reserved_balance_is_noop(
await integration_session.commit()
await pay_for_request(test_key, 100, integration_session)
test_key.reserved_balance = 0
test_key.total_requests = 0
integration_session.add(test_key)
await integration_session.commit()
# A stale cleanup already released the aggregate reservation.
# A stale cleanup already released the aggregate reservation. The revert
# terminalizes the durable row (repair) without driving the aggregate
# negative.
result = await revert_pay_for_request(test_key, integration_session, 100)
await integration_session.refresh(test_key)
integration_session.expunge_all()
updated = await integration_session.get(ApiKey, unique_key)
assert updated is not None
assert result is False, "Revert should return False when reservation already released"
assert test_key.reserved_balance == 0, (
f"Reserved balance should remain 0, got: {test_key.reserved_balance}"
)
assert test_key.total_requests == 1, (
f"Total requests should remain 1, got: {test_key.total_requests}"
assert result is True, "Revert must terminalize the corrupt reservation"
assert updated.reserved_balance == 0, (
f"Reserved balance should remain 0, got: {updated.reserved_balance}"
)
assert updated.total_requests == 0
@pytest.mark.asyncio
@@ -203,10 +203,11 @@ async def test_revert_with_sufficient_reserved_balance_succeeds(
@pytest.mark.asyncio
async def test_revert_partial_reserved_balance_is_noop(
async def test_revert_partial_reserved_balance_repairs_terminally(
integration_session: AsyncSession,
) -> None:
"""Test that reverting more than the current reserved_balance is a no-op."""
"""Reverting more than the aggregate holds must not clamp or go negative:
the corrupt durable reservation is released without subtraction."""
from routstr.auth import pay_for_request, revert_pay_for_request
unique_key = f"test_revert_partial_{uuid.uuid4().hex[:8]}"
@@ -223,18 +224,19 @@ async def test_revert_partial_reserved_balance_is_noop(
integration_session.add(test_key)
await integration_session.commit()
# Try to revert 500 when only 50 is reserved — should be no-op
# Reverting 500 when only 50 is reserved cannot subtract; the corrupt
# reservation is terminalized and the aggregate left untouched.
result = await revert_pay_for_request(test_key, integration_session, 500)
await integration_session.refresh(test_key)
integration_session.expunge_all()
updated = await integration_session.get(ApiKey, unique_key)
assert updated is not None
assert result is False, "Revert should fail when cost > reserved_balance"
assert test_key.reserved_balance == 50, (
f"Reserved balance should stay at 50, got: {test_key.reserved_balance}"
)
assert test_key.total_requests == 1, (
f"Total requests should stay at 1, got: {test_key.total_requests}"
assert result is True, "Revert must terminalize the corrupt reservation"
assert updated.reserved_balance == 50, (
f"Reserved balance should stay at 50, got: {updated.reserved_balance}"
)
assert updated.total_requests == 0
@pytest.mark.asyncio
@@ -265,9 +267,7 @@ async def test_double_revert_prevented(
snapshot = await get_reservation_snapshot(test_key, integration_session)
# First revert — should succeed
result1 = await revert_pay_for_request(
test_key, integration_session, 500, snapshot
)
result1 = await revert_pay_for_request(test_key, integration_session, 500, snapshot)
await integration_session.refresh(test_key)
assert result1 is True
@@ -275,9 +275,7 @@ async def test_double_revert_prevented(
assert test_key.total_requests == 4
# Second revert of the same amount — should be no-op
result2 = await revert_pay_for_request(
test_key, integration_session, 500, snapshot
)
result2 = await revert_pay_for_request(test_key, integration_session, 500, snapshot)
await integration_session.refresh(test_key)
assert result2 is False, "Second revert should be a no-op"
@@ -319,9 +317,7 @@ async def test_sequential_reverts_never_go_negative(
# Run 5 sequential reverts for the same 500 reservation
results = []
for _ in range(5):
r = await revert_pay_for_request(
test_key, integration_session, 500, snapshot
)
r = await revert_pay_for_request(test_key, integration_session, 500, snapshot)
results.append(r)
await integration_session.refresh(test_key)
@@ -337,63 +333,3 @@ async def test_sequential_reverts_never_go_negative(
assert test_key.reserved_balance >= 0, (
f"Reserved balance went negative: {test_key.reserved_balance}"
)
@pytest.mark.asyncio
async def test_child_key_revert_floor_guard(
integration_session: AsyncSession,
) -> None:
"""Test that child key reserved_balance also has floor guard on revert."""
from routstr.auth import (
get_reservation_snapshot,
pay_for_request,
revert_pay_for_request,
)
parent_key_hash = f"test_parent_{uuid.uuid4().hex[:8]}"
child_key_hash = f"test_child_{uuid.uuid4().hex[:8]}"
parent_key = ApiKey(
hashed_key=parent_key_hash,
balance=10000,
reserved_balance=0,
total_requests=2,
)
child_key = ApiKey(
hashed_key=child_key_hash,
balance=0,
reserved_balance=0,
total_requests=2,
parent_key_hash=parent_key_hash,
)
integration_session.add(parent_key)
integration_session.add(child_key)
await integration_session.commit()
await pay_for_request(child_key, 500, integration_session)
snapshot = await get_reservation_snapshot(child_key, integration_session)
# First revert succeeds
result1 = await revert_pay_for_request(
child_key, integration_session, 500, snapshot
)
await integration_session.refresh(parent_key)
await integration_session.refresh(child_key)
assert result1 is True
assert parent_key.reserved_balance == 0
assert child_key.reserved_balance == 0
# Second revert is a no-op for both parent and child
result2 = await revert_pay_for_request(
child_key, integration_session, 500, snapshot
)
await integration_session.refresh(parent_key)
await integration_session.refresh(child_key)
assert result2 is False
assert parent_key.reserved_balance == 0, (
f"Parent reserved_balance should stay 0, got: {parent_key.reserved_balance}"
)
assert child_key.reserved_balance == 0, (
f"Child reserved_balance should stay 0, got: {child_key.reserved_balance}"
)
@@ -0,0 +1,380 @@
"""Real-database tests for the Routstr-to-Routstr auto top-up spend bound.
The bound has to survive a process restart and concurrent workers, so these
run against actual SQL instead of mocked sessions: an in-memory counter would
pass a mocked test and still let a non-crediting peer drain the wallet.
"""
import importlib
import time
from contextlib import ExitStack
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from sqlmodel import select
from routstr.core.db import CashuTransaction, UpstreamProviderRow, create_session
from routstr.upstream import auto_topup as auto_topup_module
from routstr.upstream.auto_topup import (
ROUTSTR_MAX_DAILY_TOPUP_SATS,
ROUTSTR_MAX_TOPUP_FAILURES,
ROUTSTR_PHASE_BACKOFF,
ROUTSTR_PHASE_HALTED,
ROUTSTR_PHASE_SENT,
_check_and_topup,
_claim_routstr_topup,
_parse_routstr_request_id,
_persist_routstr_token_and_mark_sent,
_routstr_spent_last_24h_sats,
_routstr_state_id_for_provider,
get_routstr_auto_topup_state,
release_routstr_auto_topup_state,
)
pytestmark = pytest.mark.asyncio
TOPUP_SATS = 50
async def _seed_provider(provider_id: int = 1) -> UpstreamProviderRow:
import json
row = UpstreamProviderRow(
id=provider_id,
slug=f"peer-{provider_id}",
provider_type="routstr",
base_url="https://peer.test",
api_key="secret",
enabled=True,
provider_settings=json.dumps(
{
"auto_topup": True,
"topup_threshold": 1,
"topup_amount_limit": TOPUP_SATS,
"topup_mint_url": "https://mint.test",
}
),
)
async with create_session() as session:
session.add(row)
await session.commit()
await session.refresh(row)
return row
def _peer(balance: float, *, topup: object = None) -> MagicMock:
provider = MagicMock()
provider.get_balance = AsyncMock(return_value=balance)
provider.topup = AsyncMock(return_value=topup or {"balance": balance})
return provider
def _patch_wallet(module: Any, peer: MagicMock, token: str) -> ExitStack:
stack = ExitStack()
stack.enter_context(
patch.object(module.RoutstrUpstreamProvider, "from_db_row", return_value=peer)
)
stack.enter_context(
patch.object(
module, "send_token_from_owner_locked", AsyncMock(return_value=token)
)
)
stack.enter_context(
patch.object(module, "token_mint_url", return_value="https://mint.test")
)
return stack
async def _claim_state(provider_id: int = 1) -> CashuTransaction | None:
async with create_session() as session:
return await session.get(
CashuTransaction, _routstr_state_id_for_provider(provider_id)
)
async def _sent_tokens() -> list[CashuTransaction]:
async with create_session() as session:
return list(
(
await session.exec(
select(CashuTransaction).where(
CashuTransaction.source == "auto_topup"
)
)
).all()
)
async def test_second_worker_cannot_claim_while_the_first_holds_one(
patched_db_engine: Any,
) -> None:
row = await _seed_provider()
assert await _claim_routstr_topup(row, expected_sats=TOPUP_SATS) is not None
assert await _claim_routstr_topup(row, expected_sats=TOPUP_SATS) is None
async def test_token_and_sent_claim_roll_back_together_on_commit_failure(
patched_db_engine: Any,
) -> None:
row = await _seed_provider()
operation_id = await _claim_routstr_topup(row, expected_sats=TOPUP_SATS)
assert operation_id is not None
with patch(
"sqlmodel.ext.asyncio.session.AsyncSession.commit",
new=AsyncMock(side_effect=RuntimeError("commit failed")),
):
with pytest.raises(RuntimeError, match="commit failed"):
await _persist_routstr_token_and_mark_sent(
row,
operation_id,
expected_sats=TOPUP_SATS,
token="cashu-token-atomic",
amount=TOPUP_SATS,
mint_url="https://mint.test",
)
claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr]
assert claim is not None and claim.phase != ROUTSTR_PHASE_SENT
assert await _sent_tokens() == []
async def test_token_is_persisted_before_it_reaches_the_peer(
patched_db_engine: Any,
) -> None:
row = await _seed_provider()
seen: list[CashuTransaction] = []
async def _record_then_accept(token: str) -> dict:
seen.extend(await _sent_tokens())
return {"balance": TOPUP_SATS}
peer = _peer(0.0)
peer.topup = AsyncMock(side_effect=_record_then_accept)
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
await _check_and_topup(row)
assert [tx.token for tx in seen] == ["cashu-token-1"]
assert seen[0].collected is False
assert (await _sent_tokens())[0].collected is True
async def test_untracked_token_is_returned_and_never_sent(
patched_db_engine: Any,
) -> None:
row = await _seed_provider()
peer = _peer(0.0)
with (
_patch_wallet(auto_topup_module, peer, "cashu-token-1"),
patch.object(
auto_topup_module,
"_persist_routstr_token_and_mark_sent",
AsyncMock(side_effect=RuntimeError("database unavailable")),
),
patch.object(
auto_topup_module, "release_token_reservation", AsyncMock()
) as reclaim,
):
await _check_and_topup(row)
reclaim.assert_awaited_once_with("cashu-token-1")
peer.topup.assert_not_awaited()
# Nothing left the wallet, so the slot must be free again immediately.
state = await _claim_state()
assert state is not None and state.swept is True
async def test_failed_topup_keeps_token_uncollected_and_suppresses_new_claims(
patched_db_engine: Any,
) -> None:
row = await _seed_provider()
peer = _peer(0.0, topup={"error": "rejected"})
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
await _check_and_topup(row)
tokens = await _sent_tokens()
assert len(tokens) == 1
assert tokens[0].collected is False
claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr]
assert claim is not None and claim.phase == ROUTSTR_PHASE_SENT
peer.topup.reset_mock()
with _patch_wallet(auto_topup_module, peer, "cashu-token-2"):
await _check_and_topup(row)
peer.topup.assert_not_awaited()
assert len(await _sent_tokens()) == 1
async def test_claim_blocks_a_restarted_process(patched_db_engine: Any) -> None:
row = await _seed_provider()
peer = _peer(0.0, topup={"error": "rejected"})
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
await _check_and_topup(row)
# A fresh module drops every process-local variable; only a durable row
# can still stop the next payment.
reloaded = importlib.reload(auto_topup_module)
try:
peer.topup.reset_mock()
with _patch_wallet(reloaded, peer, "cashu-token-2"):
await reloaded._check_and_topup(row)
peer.topup.assert_not_awaited()
finally:
importlib.reload(auto_topup_module)
assert len(await _sent_tokens()) == 1
async def _uncredited_attempt(
row: UpstreamProviderRow, peer: MagicMock, token: str
) -> None:
"""One full payment attempt against a peer that never credits it."""
with _patch_wallet(auto_topup_module, peer, token):
await _check_and_topup(row)
await _expire_claim()
with _patch_wallet(auto_topup_module, peer, f"{token}-retry"):
await _check_and_topup(row)
await _expire_claim()
async def test_non_crediting_peer_is_halted_after_repeated_failures(
patched_db_engine: Any,
) -> None:
row = await _seed_provider()
peer = _peer(0.0)
for attempt in range(ROUTSTR_MAX_TOPUP_FAILURES):
await _uncredited_attempt(row, peer, f"cashu-token-{attempt}")
claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr]
assert claim is not None and claim.phase == ROUTSTR_PHASE_HALTED
peer.topup.reset_mock()
with _patch_wallet(auto_topup_module, peer, "cashu-token-after-halt"):
await _check_and_topup(row)
peer.topup.assert_not_awaited()
assert len(await _sent_tokens()) == ROUTSTR_MAX_TOPUP_FAILURES
async def _expire_claim(provider_id: int = 1) -> None:
"""Age the claim's deadline so the reconciler treats it as timed out."""
async with create_session() as session:
state = await session.get(
CashuTransaction, _routstr_state_id_for_provider(provider_id)
)
assert state is not None and state.request_id is not None
claim = _parse_routstr_request_id(state.request_id)
assert claim is not None
state.request_id = auto_topup_module._routstr_request_id(
claim.operation_id,
int(time.time()) - 1,
claim.phase,
claim.expected_sats,
claim.failures,
)
session.add(state)
await session.commit()
async def test_crediting_peer_releases_the_claim_for_a_later_topup(
patched_db_engine: Any,
) -> None:
row = await _seed_provider()
peer = _peer(0.0)
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
await _check_and_topup(row)
tokens = await _sent_tokens()
assert len(tokens) == 1 and tokens[0].collected is True
peer.get_balance = AsyncMock(return_value=float(TOPUP_SATS))
with _patch_wallet(auto_topup_module, peer, "cashu-token-2"):
await _check_and_topup(row)
state = await _claim_state()
assert state is not None and state.collected is True
async def test_rolling_budget_refuses_a_topup_that_would_exceed_it(
patched_db_engine: Any,
) -> None:
row = await _seed_provider()
async with create_session() as session:
session.add(
CashuTransaction(
id="prior-spend",
token="cashu-prior",
amount=ROUTSTR_MAX_DAILY_TOPUP_SATS,
unit="sat",
type="out",
source="auto_topup",
)
)
await session.commit()
assert await _routstr_spent_last_24h_sats() == ROUTSTR_MAX_DAILY_TOPUP_SATS
peer = _peer(0.0)
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
await _check_and_topup(row)
peer.topup.assert_not_awaited()
assert await _claim_state() is None
async def test_admin_release_is_fenced_on_the_state_it_reviewed(
patched_db_engine: Any,
) -> None:
row = await _seed_provider()
peer = _peer(0.0, topup={"error": "rejected"})
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
await _check_and_topup(row)
state = await get_routstr_auto_topup_state(1)
assert state["active"] is True
assert state["phase"] == ROUTSTR_PHASE_SENT
stale = await release_routstr_auto_topup_state(
1, state_token="routstr:other:0:sent:0:0"
)
assert stale.released is False and stale.reason == "stale_state"
released = await release_routstr_auto_topup_state(
1, state_token=str(state["state_token"])
)
assert released.released is True
peer.topup.reset_mock()
with _patch_wallet(auto_topup_module, peer, "cashu-token-2"):
await _check_and_topup(row)
peer.topup.assert_awaited_once()
async def test_backoff_suppresses_retries_until_its_deadline(
patched_db_engine: Any,
) -> None:
row = await _seed_provider()
peer = _peer(0.0)
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
await _check_and_topup(row)
await _expire_claim()
# First reconciliation after the lease: the peer never credited, so the
# claim moves to backoff rather than paying again immediately.
peer.topup.reset_mock()
with _patch_wallet(auto_topup_module, peer, "cashu-token-2"):
await _check_and_topup(row)
peer.topup.assert_not_awaited()
claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr]
assert claim is not None and claim.phase == ROUTSTR_PHASE_BACKOFF
assert claim.deadline > int(time.time())
@@ -0,0 +1,164 @@
"""The Routstr top-up threshold has to state its unit.
``topup_amount_limit`` was already plain sats — it goes straight to
``send_token(amount, "sat", ...)`` and is added to the balance for logging — but
its sibling ``topup_threshold`` was compared as ``balance >= threshold * 1000``.
Two keys from one settings blob, two different units, and nothing naming either.
An operator asking for "top up below 1000 sats" got one below 1,000,000.
``topup_threshold_sats`` says what it means. Legacy values keep their effective
trigger point: reinterpreting them as sats would silently drop the trigger a
thousandfold and let a provider run dry.
"""
import json
from typing import Any
from unittest.mock import patch
import pytest
from routstr.core.db import UpstreamProviderRow, create_session
from routstr.upstream import auto_topup as auto_topup_module
from routstr.upstream.auto_topup import (
_check_and_topup,
validate_routstr_auto_topup_settings,
)
from .test_routstr_auto_topup_claim import _patch_wallet, _peer, _sent_tokens
# No module-level asyncio mark: the settings cases are sync, and the suite
# already runs asyncio in auto mode.
TOPUP_SATS = 50
@pytest.fixture(autouse=True)
def _forget_legacy_hints() -> Any:
# The hint fires once per provider for the life of the process. Reach it
# through the module: a sibling test reloads auto_topup, which rebinds the
# set, so a name imported here would go on clearing the old one.
auto_topup_module._legacy_threshold_hinted.clear()
yield
auto_topup_module._legacy_threshold_hinted.clear()
async def _seed(**topup_settings: Any) -> UpstreamProviderRow:
row = UpstreamProviderRow(
id=1,
slug="peer-1",
provider_type="routstr",
base_url="https://peer.test",
api_key="secret",
enabled=True,
provider_settings=json.dumps(
{
"auto_topup": True,
"topup_amount_limit": TOPUP_SATS,
"topup_mint_url": "https://mint.test",
**topup_settings,
}
),
)
async with create_session() as session:
session.add(row)
await session.commit()
await session.refresh(row)
return row
async def _topped_up(row: UpstreamProviderRow, balance: float) -> bool:
peer = _peer(balance)
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
await _check_and_topup(row)
return bool(await _sent_tokens())
async def test_explicit_sats_threshold_is_compared_against_a_sats_balance(
patched_db_engine: Any,
) -> None:
row = await _seed(topup_threshold_sats=1000)
assert await _topped_up(row, 1500.0) is False
async def test_explicit_sats_threshold_tops_up_below_the_stated_amount(
patched_db_engine: Any,
) -> None:
row = await _seed(topup_threshold_sats=1000)
assert await _topped_up(row, 500.0) is True
@pytest.mark.parametrize(
("balance", "expected"),
[(1500.0, False), (500.0, True)],
)
async def test_legacy_threshold_keeps_its_effective_trigger_point(
patched_db_engine: Any, balance: float, expected: bool
) -> None:
# 1 x 1000 == the 1000 sats the old comparison actually used. Reading it as
# 1 sat instead would leave the peer to run dry.
row = await _seed(topup_threshold=1)
assert await _topped_up(row, balance) is expected
async def test_explicit_sats_threshold_overrides_the_legacy_key(
patched_db_engine: Any,
) -> None:
row = await _seed(topup_threshold=1, topup_threshold_sats=100)
assert await _topped_up(row, 500.0) is False
async def test_legacy_threshold_reports_the_sats_value_to_migrate_to(
patched_db_engine: Any,
) -> None:
row = await _seed(topup_threshold=1)
# The app logger does not propagate to root, so caplog would see nothing.
with patch.object(auto_topup_module, "logger") as log:
await _topped_up(row, 5000.0)
hints = [
c for c in log.warning.call_args_list if "topup_threshold_sats" in c.args[0]
]
assert len(hints) == 1
assert hints[0].kwargs["extra"]["threshold_sats"] == 1000.0
# The scheduler runs every minute; the hint must not run with it.
await _topped_up(row, 5000.0)
assert (
sum("topup_threshold_sats" in c.args[0] for c in log.warning.call_args_list)
== 1
)
def test_settings_accept_the_sats_threshold_on_its_own() -> None:
assert (
validate_routstr_auto_topup_settings(
{
"auto_topup": True,
"topup_threshold_sats": 1000,
"topup_amount_limit": TOPUP_SATS,
"topup_mint_url": "https://mint.test",
}
)
is None
)
@pytest.mark.parametrize("value", [0, -1, True, float("inf"), "1000", None])
def test_settings_reject_an_unusable_sats_threshold(value: object) -> None:
assert validate_routstr_auto_topup_settings(
{
"auto_topup": True,
"topup_threshold_sats": value,
"topup_amount_limit": TOPUP_SATS,
"topup_mint_url": "https://mint.test",
}
)
def test_settings_require_one_of_the_threshold_keys() -> None:
assert validate_routstr_auto_topup_settings(
{
"auto_topup": True,
"topup_amount_limit": TOPUP_SATS,
"topup_mint_url": "https://mint.test",
}
)
@@ -0,0 +1,255 @@
"""The served catalog is the last guard between a stored row and a charge.
Stored pricing is JSON written by whatever produced the row — an upstream
import, an operator, a legacy migration, or a foreign writer that never passed
the admin edge. So the read path cannot assume a stored rate is a number: it
must decline to serve a row it cannot bill on, and it must survive a row it
cannot read at all rather than taking the whole catalog down with it.
The admin listing is deliberately exempt: it includes disabled models and is the
one view that still shows the operator the row that needs repair.
"""
from __future__ import annotations
import json
import pytest
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ModelRow, UpstreamProviderRow
from routstr.payment.models import list_models
_ARCHITECTURE = json.dumps(
{
"modality": "text",
"input_modalities": ["text"],
"output_modalities": ["text"],
"tokenizer": "unknown",
"instruct_type": None,
}
)
async def _make_provider(session: AsyncSession) -> int:
provider = UpstreamProviderRow(
provider_type="generic",
base_url="https://served-upstream.example/v1",
api_key="test-key",
provider_fee=1.0,
)
session.add(provider)
await session.commit()
await session.refresh(provider)
assert provider.id is not None
return provider.id
async def _insert_row(
session: AsyncSession,
provider_id: int,
*,
model_id: str,
pricing: dict[str, object],
) -> None:
session.add(
ModelRow(
id=model_id,
name=model_id,
description="d",
created=0,
context_length=8192,
architecture=_ARCHITECTURE,
pricing=json.dumps(pricing),
upstream_provider_id=provider_id,
enabled=True,
forwarded_model_id=model_id,
)
)
await session.commit()
@pytest.mark.integration
@pytest.mark.asyncio
@pytest.mark.parametrize(
"bad_rate",
[float("nan"), float("inf"), -1.0],
ids=["nan", "inf", "negative"],
)
async def test_served_catalog_excludes_a_malformed_stored_rate(
integration_session: AsyncSession, bad_rate: float
) -> None:
"""A stored rate that is not a number must not be advertised.
Zero is a real price and a free model is servable, but a negative or
non-finite rate is not a price at all: serving it advertises a rate the cost
calculation cannot bill on, so every request falls through to the flat
maximum reservation — or, for a negative rate, bills an amount settlement
credits back to the caller.
"""
provider_id = await _make_provider(integration_session)
await _insert_row(
integration_session,
provider_id,
model_id="good",
pricing={"prompt": 1e-06, "completion": 2e-06},
)
await _insert_row(
integration_session,
provider_id,
model_id="bad-rate",
pricing={"prompt": bad_rate, "completion": 2e-06},
)
served = {m.id for m in await list_models(integration_session, provider_id)}
assert served == {"good"}
@pytest.mark.integration
@pytest.mark.asyncio
async def test_free_stored_price_is_still_served(
integration_session: AsyncSession,
) -> None:
"""Zero is a real price. Rejecting malformed rates must not also drop a row
priced at zero, which is a free model and not a broken one."""
provider_id = await _make_provider(integration_session)
await _insert_row(
integration_session,
provider_id,
model_id="free",
pricing={"prompt": 0.0, "completion": 0.0},
)
served = {m.id for m in await list_models(integration_session, provider_id)}
assert served == {"free"}
@pytest.mark.integration
@pytest.mark.asyncio
async def test_admin_listing_still_shows_a_malformed_stored_rate(
integration_session: AsyncSession,
) -> None:
"""The operator has to be able to see the row that needs fixing.
The backstop keeps a malformed row out of the *served* catalog. The listing
that includes disabled models is the one view where the row must still
appear, or the operator loses the ability to repair it.
"""
provider_id = await _make_provider(integration_session)
await _insert_row(
integration_session,
provider_id,
model_id="bad-rate",
pricing={"prompt": -1.0, "completion": 2e-06},
)
listed = {
m.id
for m in await list_models(
integration_session, provider_id, include_disabled=True
)
}
assert listed == {"bad-rate"}
@pytest.mark.integration
@pytest.mark.asyncio
async def test_one_unreadable_stored_price_does_not_blank_the_catalog(
integration_session: AsyncSession,
) -> None:
"""A single unparseable row must cost that row, not every model on the node.
Stored pricing is JSON written by whatever produced the row, so a
non-numeric rate is reachable from a legacy import or a foreign writer.
Parsing it raises out of the row-to-model conversion, and because the
conversion ran inside the catalog loop the exception took the whole listing
with it — one bad row and the node advertised nothing at all.
"""
provider_id = await _make_provider(integration_session)
await _insert_row(
integration_session,
provider_id,
model_id="good",
pricing={"prompt": 1e-06, "completion": 2e-06},
)
await _insert_row(
integration_session,
provider_id,
model_id="unreadable",
pricing={"prompt": "not-a-number", "completion": 2e-06},
)
served = {m.id for m in await list_models(integration_session, provider_id)}
assert served == {"good"}
@pytest.mark.integration
@pytest.mark.asyncio
@pytest.mark.parametrize(
"field",
[
"image",
"web_search",
"internal_reasoning",
"input_cache_read",
"input_cache_write",
],
)
async def test_served_catalog_excludes_a_malformed_auxiliary_rate(
integration_session: AsyncSession, field: str
) -> None:
"""The backstop covers every billable rate, not only the token rates.
A price whose ``prompt``/``completion`` are sound can still carry a
malformed request, image, search, reasoning or cache rate — the catalog
import filter never inspects those — and the request that hits one is billed
against it just the same.
"""
provider_id = await _make_provider(integration_session)
await _insert_row(
integration_session,
provider_id,
model_id="good",
pricing={"prompt": 1e-06, "completion": 2e-06},
)
await _insert_row(
integration_session,
provider_id,
model_id="bad-aux",
pricing={"prompt": 1e-06, "completion": 2e-06, field: -1.0},
)
served = {m.id for m in await list_models(integration_session, provider_id)}
assert served == {"good"}
@pytest.mark.integration
@pytest.mark.asyncio
async def test_a_negative_request_rate_is_clamped_on_read_and_still_served(
integration_session: AsyncSession,
) -> None:
"""``request`` is the one billable rate the row-to-model conversion repairs.
It clamps a negative stored ``request`` to zero before the price is built,
so the backstop never sees one and the row is served at a zero request rate
— money-safe, and the reason ``request`` is absent from the list of rates
above. Pinned here so that if the clamp goes, this rate joins that list
rather than quietly becoming the one unguarded field.
"""
provider_id = await _make_provider(integration_session)
await _insert_row(
integration_session,
provider_id,
model_id="neg-request",
pricing={"prompt": 1e-06, "completion": 2e-06, "request": -1.0},
)
served = await list_models(integration_session, provider_id)
assert [m.id for m in served] == ["neg-request"]
assert served[0].pricing.request == 0.0
-210
View File
@@ -1,210 +0,0 @@
"""
Integration tests for reactive swap fee retries via the wallet topup endpoint.
Foreign-mint tokens are swapped to the primary mint using the foreign mint's
melt quote, whose fee_reserve is a non-binding estimate (NUT-05): the mint may
demand more when re-quoting or at melt execution. These tests cover the
endpoint behaviour in those cases:
1. The mint demands one sat more at melt time than every quote reported
(the mint.cubabitcoin.org incident): the swap retries with a smaller
invoice and the topup succeeds, crediting the recomputed amount.
2. The real melt quote reports a higher fee_reserve than the estimate: the
swap re-quotes from the observed fee and the topup succeeds.
3. The mint escalates its fee demands on every attempt: the retry budget is
exhausted and the endpoint returns 400 with a clear error (never 500),
without ever executing a melt.
"""
from collections.abc import Callable
from unittest.mock import AsyncMock, Mock, patch
import pytest
from cashu.core.base import MeltQuoteState
from httpx import AsyncClient, Response
from routstr.core.settings import settings
# Captured at collection time, before the integration_app fixture replaces it
# with the testmint stub that bypasses swapping (see conftest.py).
from routstr.wallet import recieve_token as _real_recieve_token
# Match the authenticated fixture's persisted refund mint: existing-key topups
# are intentionally constrained to that mint for collateral provenance.
PRIMARY_MINT = "http://localhost:3338"
def _make_swap_mocks(
token_amount: int,
fee_reserves: list[int],
input_fees: int = 0,
mint_url: str = "http://foreign-mint:3338",
) -> tuple[Mock, Mock, Mock]:
"""Return (token, token_wallet, primary_wallet) mocks that act like a mint.
Mint quotes pass the requested amount through their ``request`` field and
melt quotes echo that amount back, so the mocks stay consistent for
whatever amounts the implementation requests. ``fee_reserves`` supplies the
fee_reserve of each successive melt quote (the first serves the estimation
pass); requesting more quotes than provided fails the test.
"""
mock_token = Mock()
mock_token.mint = mint_url
mock_token.unit = "sat"
mock_token.amount = token_amount
mock_token.keysets = ["keyset1"]
mock_token.proofs = [Mock(amount=token_amount)]
mock_token_wallet = Mock()
mock_token_wallet.load_mint_keysets = AsyncMock()
mock_token_wallet.activate_keyset = AsyncMock()
mock_token_wallet._expand_short_keyset_ids = AsyncMock()
mock_token_wallet.load_proofs = AsyncMock()
mock_token_wallet.get_fees_for_proofs = Mock(return_value=input_fees)
mock_primary_wallet = Mock()
mock_primary_wallet.load_mint = AsyncMock()
mock_primary_wallet.load_proofs = AsyncMock()
mock_primary_wallet.available_balance = Mock(amount=0)
mock_primary_wallet.mint = AsyncMock(return_value=Mock())
fees = iter(fee_reserves)
def _next_fee() -> int:
try:
return next(fees)
except StopIteration:
raise AssertionError(
"more melt quotes requested than fee_reserves provided"
) from None
mock_primary_wallet.request_mint = AsyncMock(
side_effect=lambda amount: Mock(quote=f"mint_quote_{amount}", request=amount)
)
mock_token_wallet.melt_quote = AsyncMock(
side_effect=lambda invoice: Mock(
quote=f"melt_quote_{invoice}", amount=invoice, fee_reserve=_next_fee()
)
)
mock_token_wallet.melt = AsyncMock(
return_value=Mock(state=MeltQuoteState.paid)
)
return mock_token, mock_token_wallet, mock_primary_wallet
def _wallet_router(primary_wallet: Mock, token_wallet: Mock) -> Callable[..., Mock]:
"""Route get_wallet calls to the primary or foreign wallet mock by URL."""
def fake_get_wallet(
mint_url: str,
unit: str = "sat",
load: bool = True,
**kwargs: object,
) -> Mock:
return primary_wallet if mint_url == PRIMARY_MINT else token_wallet
return fake_get_wallet
async def _post_topup(
client: AsyncClient,
mock_token: Mock,
token_wallet: Mock,
primary_wallet: Mock,
) -> Response:
"""POST /v1/wallet/topup with the swap layer mocked at the mint boundary.
The conftest's testmint stub for recieve_token is swapped back for the
real implementation so the request exercises the actual swap path.
"""
with patch("routstr.wallet.recieve_token", _real_recieve_token):
with patch(
"routstr.wallet.deserialize_token_from_string", return_value=mock_token
):
with patch(
"routstr.wallet.get_wallet",
side_effect=_wallet_router(primary_wallet, token_wallet),
):
with patch.object(settings, "primary_mint", PRIMARY_MINT):
with patch.object(settings, "primary_mint_unit", "sat"):
with patch.object(settings, "cashu_mints", [PRIMARY_MINT]):
return await client.post(
"/v1/wallet/topup",
params={"cashu_token": "cashuAtest_foreign_token"},
)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_topup_retries_when_melt_demands_more_than_quoted(
authenticated_client: AsyncClient,
) -> None:
"""A 179-sat token where every quote reports fee_reserve=1 but the mint
rejects the first melt demanding 180. The retry shrinks the invoice to 177
and the topup credits 177 sats (177_000 msats)."""
mock_token, token_wallet, primary_wallet = _make_swap_mocks(
179, fee_reserves=[1, 1, 1], mint_url="http://mint.cubabitcoin.org"
)
token_wallet.melt.side_effect = [
Exception(
"Mint Error: not enough inputs provided for melt. "
"Provided: 179, needed: 180 (Code: 11000)"
),
Mock(state=MeltQuoteState.paid),
]
response = await _post_topup(
authenticated_client, mock_token, token_wallet, primary_wallet
)
assert response.status_code == 200
assert response.json()["msats"] == 177_000
assert token_wallet.melt.call_count == 2
@pytest.mark.integration
@pytest.mark.asyncio
async def test_topup_retries_when_quote_fee_exceeds_estimate(
authenticated_client: AsyncClient,
) -> None:
"""A 1000-sat token estimated at fee 20, but the real quote demands 23.
The retry recomputes 1000 - 23 = 977, which fits, and the topup credits
977 sats (977_000 msats) with a single melt."""
mock_token, token_wallet, primary_wallet = _make_swap_mocks(
1000, fee_reserves=[20, 23, 23]
)
response = await _post_topup(
authenticated_client, mock_token, token_wallet, primary_wallet
)
assert response.status_code == 200
assert response.json()["msats"] == 977_000
assert token_wallet.melt.call_count == 1
@pytest.mark.integration
@pytest.mark.asyncio
async def test_topup_returns_422_when_retries_exhausted(
authenticated_client: AsyncClient,
) -> None:
"""A mint that escalates fee_reserve on every re-quote (1 → 10 → 25 → 50)
exhausts the retry budget: clean 422 mint_error/too-small taxonomy, melt
never executed."""
mock_token, token_wallet, primary_wallet = _make_swap_mocks(
1000, fee_reserves=[1, 10, 25, 50]
)
response = await _post_topup(
authenticated_client, mock_token, token_wallet, primary_wallet
)
assert response.status_code == 422
raw_detail = response.json()["detail"]
message = (
raw_detail["error"]["message"] if isinstance(raw_detail, dict) else raw_detail
)
assert "too small to cover swap fees" in message
assert token_wallet.melt_quote.call_count == 4 # estimation + 3 attempts
token_wallet.melt.assert_not_called()
@@ -22,18 +22,18 @@ async def _add_key(
hashed_key: str,
*,
balance: int = 0,
reserved_balance: int = 0,
total_spent: int = 0,
total_requests: int = 0,
created_at: int | None = None,
parent_key_hash: str | None = None,
refund_address: str | None = None,
) -> ApiKey:
key = ApiKey(
hashed_key=hashed_key,
balance=balance,
reserved_balance=reserved_balance,
total_spent=total_spent,
total_requests=total_requests,
parent_key_hash=parent_key_hash,
refund_address=refund_address,
)
key.created_at = created_at
@@ -123,29 +123,18 @@ async def test_temporary_balances_pagination(
@pytest.mark.integration
@pytest.mark.asyncio
async def test_temporary_balances_totals_exclude_child_balance(
async def test_temporary_balances_totals(
integration_client: httpx.AsyncClient,
integration_session: AsyncSession,
) -> None:
await _add_key(
integration_session,
"parent",
"standalone_key",
balance=5000,
total_spent=100,
total_requests=3,
created_at=1000,
)
# Child draws from parent's balance, so its balance must NOT be summed,
# but its spent/requests still count.
await _add_key(
integration_session,
"child",
balance=0,
total_spent=200,
total_requests=7,
created_at=1001,
parent_key_hash="parent",
)
response = await integration_client.get(
"/admin/api/temporary-balances", headers=_admin_headers()
@@ -153,8 +142,8 @@ async def test_temporary_balances_totals_exclude_child_balance(
totals = response.json()["totals"]
assert totals["total_balance"] == 5000
assert totals["total_spent"] == 300
assert totals["total_requests"] == 10
assert totals["total_spent"] == 100
assert totals["total_requests"] == 3
@pytest.mark.integration

Some files were not shown because too many files have changed in this diff Show More