Compare commits

..
Author SHA1 Message Date
Paperclip Deployment Engineer 7ffebe6d91 sentrystr stuff 2026-06-27 20:55:38 +00:00
141 changed files with 1687 additions and 16502 deletions
+10 -17
View File
@@ -2,23 +2,7 @@
UPSTREAM_BASE_URL=https://api.openai.com/v1
UPSTREAM_API_KEY=your-upstream-api-key
# Tinfoil (confidential inference enclaves, EHBP)
# TINFOIL_API_KEY=your-tinfoil-api-key
# Secret key used to encrypt node secrets at rest (optional). If unset, the node
# generates one on first start, writes it to routstr_secret.key (override the path
# with ROUTSTR_SECRET_KEY_FILE), and prints it once — back that file up, because
# losing the key makes previously encrypted secrets unreadable. Set it explicitly
# to manage the key yourself (recommended in production). Generate one with:
# uv run python -c "from cryptography.fernet import Fernet; print(Fernet.generate_key().decode())"
ROUTSTR_SECRET_KEY=
# The admin password and the Nostr identity (nsec) are NOT set here. The admin
# password is generated and logged once on first start (read it from the logs to
# sign in); both are managed afterwards from the admin UI and stored encrypted in
# the database. ADMIN_PASSWORD / NSEC are still read once as a legacy seed for
# existing deployments, but new nodes should set them in the UI — a value left in
# .env is ignored once the node has been configured.
# ADMIN_PASSWORD=secure-admin-password
# Database
# DATABASE_URL=sqlite+aiosqlite:///keys.db
@@ -26,6 +10,7 @@ ROUTSTR_SECRET_KEY=
# Node Information
# NAME=My Routstr Node
# DESCRIPTION=Fast AI API access with Bitcoin payments
# NSEC=nsec1...
# HTTP_URL=https://api.mynode.com
# ONION_URL=http://mynode.onion (auto fetched from compose)
# RELAYS="wss://relay.damus.io,wss://relay.nostr.band,wss://eden.nostr.land,wss://relay.routstr.com"
@@ -49,6 +34,14 @@ ROUTSTR_SECRET_KEY=
# LOG_LEVEL=INFO
# ENABLE_CONSOLE_LOGGING=true
# SentryStr: mirror logs to Nostr relays (requires `pip install sentrystr`).
# Announcements run on a background worker thread so they never block requests.
# ENABLE_SENTRYSTR=false
# SENTRYSTR_LOG_LEVEL=ERROR # defaults to ERROR; keep >= LOG_LEVEL
# SENTRYSTR_RELAYS="wss://relay.damus.io,wss://nos.lol" # defaults to RELAYS
# SENTRYSTR_NSEC=nsec1... # defaults to NSEC; ephemeral key if unset
# SENTRYSTR_RECIPIENT_NPUB=npub1... # if set, log events are NIP-44 encrypted
# Custom Model Management
# BASE_URL=https://openrouter.ai/api/v1
# MODELS_PATH=models.json
-5
View File
@@ -1,7 +1,6 @@
__pycache__
.env
keys.db
routstr_secret.key
wallet.sqlite3
# Python build artifacts
@@ -17,9 +16,6 @@ dist/
*.db-shm
*.db-wal
.*wallet.sqlite3
.wallet/
AGENTS.md
TEST_SUITE_OVERVIEW.md
*models.json
.cashu
.relay
@@ -42,4 +38,3 @@ proof_backups
*.todo
ui_out
.worktrees
+1 -1
View File
@@ -7,7 +7,7 @@ WORKDIR /app/ui
RUN corepack enable pnpm && corepack prepare pnpm@10.15.0 --activate
# Copy UI source
COPY ui/package.json ui/pnpm-lock.yaml* ui/pnpm-workspace.yaml* ./
COPY ui/package.json ui/pnpm-lock.yaml* ./
RUN pnpm install --frozen-lockfile
COPY ui/ ./
+3 -25
View File
@@ -55,41 +55,19 @@ If you are a node runner, start a Routstr Core instance using Docker Compose:
1. **Prepare your `.env`**:
```bash
# Optional: encrypts node secrets at rest. If unset, the node generates a key
# on first start, writes it to routstr_secret.key, and prints it once — back
# up that file. Set it explicitly to manage the key yourself (recommended in
# production).
ROUTSTR_SECRET_KEY=<generated-key>
ADMIN_PASSWORD=mysecretpassword
NAME="My AI Node"
DESCRIPTION="Fast access to models"
NSEC=yournsec
RECEIVE_LN_ADDRESS=yourname@wallet.com
```
Your Nostr identity (`nsec`) is not set in `.env` — configure it from the admin
UI after first start, where it's stored encrypted in the database. (`NSEC` in
`.env` is still read once as a legacy seed for existing deployments.)
If you don't set one, a key is generated and printed on first start — save it
somewhere safe (losing it makes previously encrypted secrets unreadable). To
supply your own, generate it once and keep it stable:
```bash
uv run python -c "from cryptography.fernet import Fernet; print(Fernet.generate_key().decode())"
```
2. **Start the services**:
```bash
docker compose up -d
```
3. **Get your admin password**:
On first start the node generates an admin password and logs it once with the
`/admin` URL. Read it from the logs:
```bash
docker compose logs routstr | grep -i admin
```
(Lost it? Reset with `docker compose exec routstr /.venv/bin/python scripts/reset_admin_password.py --regenerate`.)
4. **Configure**:
3. **Configure**:
Open [http://localhost:8000/admin/](http://localhost:8000/admin/) to connect your AI providers and set pricing.
For full instructions, see the **[Provider Quick Start Guide](https://docs.routstr.com/provider/quickstart/)**.
-3
View File
@@ -13,7 +13,6 @@ services:
- ./ui_out:/output:z
command:
["sh", "-c", "mkdir -p /output && cp -r /app/built/. /output/ && echo 'UI build copied to mounted volume' && ls -la /output/ && echo 'UI built and ready' && tail -f /dev/null"]
restart: unless-stopped
routstr:
build: .
@@ -32,7 +31,6 @@ services:
- 8000:8000
extra_hosts: # Needed to access locally running models
- "host.docker.internal:host-gateway"
restart: unless-stopped
tor:
image: ghcr.io/hundehausen/tor-hidden-service:latest
@@ -43,7 +41,6 @@ services:
- HS_ROUTER=routstr:8000:80
depends_on:
- routstr
restart: unless-stopped
volumes:
tor-data:
+20 -87
View File
@@ -113,109 +113,42 @@ All errors follow a consistent JSON structure:
**Status:** 402
**Resolution:** Top up API key balance
### Cashu Token Redemption Errors
These errors are returned when a Cashu token you pay with cannot be redeemed.
They apply to every endpoint that accepts a token:
- **Per-request payment** via the `X-Cashu` header (chat completions + Responses API).
- **API key top-up** via `POST /v1/wallet/topup`.
- **Minting an API key** from a token sent in `Authorization: Bearer <cashu-token>`.
All three share one classifier, so the same failure yields the same HTTP status
and sanitized message everywhere. Structured error envelopes (`X-Cashu` and
`Authorization: Bearer <cashu-token>`) also expose the same `type` and `code`
branch on `type` (or `code` for finer granularity). `POST /v1/wallet/topup`
keeps its existing plain-string `detail` envelope, so branch on status there.
| `type` | Status | `code` | Retryable | Meaning |
|--------|--------|--------|-----------|---------|
| `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_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. |
| `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 `mint_unreachable`"
Only `mint_unreachable` (503) means the same token will work again later —
everything else is a permanent property of the token and must not be
blindly retried. Use exponential backoff for the 503. In particular, a
`token_consumed` 500 means the mint already spent the token, so a retry
would fail as `token_already_spent`.
#### Mint Unreachable (retryable)
#### Invalid Token
```json
{
"error": {
"type": "mint_unreachable",
"message": "Cashu mint is unreachable",
"code": "cashu_mint_unreachable"
"type": "payment_error",
"message": "Invalid Cashu token",
"code": "invalid_token",
"details": {
"reason": "Token already spent"
}
}
}
```
**Status:** 503
**Status:** 400
**Resolution:** Use a valid, unspent token
**Resolution:** The token is valid — the mint is temporarily down. Retry with
backoff, or pay with a token from a different mint.
#### Token Already Spent
#### Mint Unavailable
```json
{
"error": {
"type": "token_already_spent",
"message": "Cashu token already spent",
"code": "cashu_token_already_spent"
"type": "payment_error",
"message": "Cannot connect to Cashu mint",
"code": "mint_unavailable",
"details": {
"mint_url": "https://mint.example.com",
"retry_after": 60
}
}
}
```
**Status:** 400
**Resolution:** Use a fresh, unspent token. Do not retry with the same token.
#### Response envelope differs by endpoint
The `error` object above is identical everywhere, but the surrounding envelope
depends on how you paid:
- **`X-Cashu` header payments** (chat + Responses API) return the object at the
top level, alongside a `request_id`:
```json
{
"error": { "type": "mint_unreachable", "message": "Cashu mint is unreachable", "code": "cashu_mint_unreachable" },
"request_id": "req-abc123"
}
```
The original token is echoed back in the `X-Cashu` **response header only when
it is still spendable** (e.g. `mint_unreachable`, `invalid_cashu_token`, fee
errors) so you can recover/retry it. It is **not** echoed for spent/consumed
tokens (`cashu_token_already_spent`, `cashu_token_consumed`,
`cashu_token_zero_value`, `internal_error`) — retrying those can never succeed.
- **`Authorization: Bearer <cashu-token>`** (API key minting) wraps it in
FastAPI's `detail` field:
```json
{ "detail": { "error": { "type": "mint_unreachable", "message": "Cashu mint is unreachable", "code": "cashu_mint_unreachable" } } }
```
- **`POST /v1/wallet/topup`** returns a plain string message under `detail` —
it carries the shared HTTP **status** and **message** (e.g. `503` for an
unreachable mint) but not the structured `type`/`code`, so branch on the
status code here:
```json
{ "detail": "Cashu mint is unreachable" }
```
**Status:** 503
**Resolution:** Try again later or use different mint
### Validation Errors
@@ -414,7 +347,7 @@ class ErrorHandler:
'rate_limit',
'upstream_timeout',
'model_overloaded',
'cashu_mint_unreachable'
'mint_unavailable'
}
# Errors requiring user action
-150
View File
@@ -1,150 +0,0 @@
# EHBP Proxy Support for Tinfoil Models
## Problem
The SDK's `SecureClient.fetch` encrypts request bodies with HPKE (EHBP protocol)
and sends them to the Routstr provider. The Routstr proxy had no EHBP handling:
1. It tried to `json.loads()` the binary HPKE-sealed body → failed with a 400
2. The upstream PPQ.AI public endpoint (`/v1/chat/completions`) doesn't speak
EHBP, so the response had no `Ehbp-Response-Nonce` header
3. `SecureClient` threw `Missing Ehbp-Response-Nonce header` because it expects
every response from an EHBP-configured `baseURL` to carry that header
## Root cause
PPQ.AI exposes EHBP-aware inference at `/private/`, separate from the public
`/v1/` endpoint. The Routstr proxy was forwarding to `/v1/` (the public
endpoint) instead of `/private/` (the enclave endpoint). The public endpoint
can't decrypt the body, returns a normal HTTP response, and the SDK can't
decrypt it because there's no nonce header.
The PPQ private-mode proxy (`ppq-private-mode-proxy/lib/proxy.ts`) shows the
correct pattern: `SecureClient` talks to `api.ppq.ai/private/v1/chat/completions`,
which decrypts inside the attested enclave and returns an EHBP-encrypted
response with the `Ehbp-Response-Nonce` header.
## What was changed
### `routstr/proxy.py`
Detects EHBP requests by checking for the `Ehbp-Encapsulated-Key` header (set
by the EHBP transport on every encrypted request). For EHBP requests:
- Skips JSON body parsing (the body is binary ciphertext, not JSON)
- Reads the model ID from the `X-Routstr-Model` header (set by the SDK) instead
of from `body.model`
- Routes through new `forward_ehbp_request` (bearer auth) and
`forward_ehbp_x_cashu_request` (x-cashu auth) methods
- Skips reactive 400 param correction (can't parse encrypted response body)
- Still charges the user via `pay_for_request` (uses `max_cost_for_model` from
the model registry, not the body)
### `routstr/upstream/base.py`
Keeps EHBP as an explicit opt-in provider capability instead of making every
upstream provider appear EHBP-capable:
- `supports_ehbp = False` by default
- `get_ehbp_forwarding_target(path, model_obj)` raises `NotImplementedError`
unless a provider opts in and returns a provider-specific EHBP target
The actual EHBP forwarding logic does **not** live in `base.py`.
### `routstr/upstream/ehbp.py`
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)
- `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)
### Provider support
EHBP is currently enabled only for `TinfoilUpstreamProvider`. It forwards to
Tinfoil's attested enclave and requests `X-Tinfoil-Usage-Metrics` for billing.
PPQ.AI retains its private-target implementation, but `supports_ehbp = False`
until it has a provider-specific trusted usage/model-binding strategy.
## Why it's done this way
The proxy is a **blind relay** for EHBP requests. It cannot decrypt the body
(only the attested enclave can), so it must:
1. Get the model ID from a header, not the body
2. Forward the raw bytes without parsing or transformation
3. Stream the response back without SSE/cost parsing
4. Pass through EHBP protocol headers (`Ehbp-Encapsulated-Key` on request,
`Ehbp-Response-Nonce` on response)
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.
## End-to-end flow
```
SDK Routstr Proxy PPQ.AI /private/
│ │ │
│── X-Routstr-Model: tinfoil-kimi-k2-6 ─│ │
│── Ehbp-Encapsulated-Key: <hex> ───────│ │
│── Authorization: Bearer <cashu> ──────│ │
│── body = HPKE-encrypted(kimi-k2-6) ───│ │
│ │ │
│ detects Ehbp-Encapsulated-Key │
│ reads model from X-Routstr-Model │
│ does billing/routing │
│ │ │
│ adds X-Private-Model: private/kimi-k2-6
│ forwards raw body to /private/v1/... │
│ │──────────────────────────────▶│
│ │ enclave decrypts
│ │ runs inference
│ │◀── Ehbp-Response-Nonce ──────│
│ │◀── encrypted response ────────│
│ │ │
│ streams response back untouched │
│◀── encrypted response ────────────────│ │
│ │ │
SecureClient reads nonce, decrypts │ │
SDK SSE processing sees plaintext │ │
```
## Model ID mapping
Three parties see three different model IDs:
| Party | Header/Body | Value | Source |
|---|---|---|---|
| Routstr proxy | `X-Routstr-Model` header | `tinfoil-kimi-k2-6` | SDK sends full caller-facing id |
| Tinfoil usage metrics | `model` field | `kimi-k2-6` | Enclave reports the model actually served |
| Tinfoil enclave | `body.model` (encrypted) | `kimi-k2-6` | SDK strips `tinfoil-` prefix before encryption |
## Implementation status
A dedicated `TinfoilUpstreamProvider` (`routstr/upstream/tinfoil.py`) now
implements the direct blind-upstream pattern described above. The shared EHBP
helpers in `routstr/upstream/ehbp.py` were extended to:
- Request usage metrics via `X-Tinfoil-Request-Usage-Metrics: true`.
- Parse `X-Tinfoil-Usage-Metrics` from the response header (non-streaming) or
HTTP trailer (streaming).
- Override the forwarding URL with a validated `X-Tinfoil-Enclave-Url` when the
SDK sends it.
- Finalize bearer billing with the dedicated EHBP actual-cost finalizer.
- Compute X-Cashu refunds from actual cost instead of max cost.
See `docs/tinfoil-direct-integration.md` for the full implementation notes.
## Verification status
Unit coverage includes usage parsing, target validation, HTTP trailer capture,
response-size limits, and bearer payment finalization. End-to-end requests have
verified both non-streaming usage headers and streaming usage trailers against
Tinfoil. SDK behavior on proxy-generated non-2xx responses without an
`Ehbp-Response-Nonce` still merits explicit end-to-end coverage.
+6 -25
View File
@@ -13,10 +13,7 @@ Before running your node, you should create a `.env` file in the project root. T
### Example .env
```bash
# Encrypts node secrets at rest. Optional — if unset, the node generates a key on
# first start and prints it once (back it up). Set it to manage the key yourself
# (recommended in production). See "Secrets at Rest" below.
ROUTSTR_SECRET_KEY=
ADMIN_PASSWORD=your-secure-password
# Node Identity
NAME="My AI Node"
@@ -28,10 +25,10 @@ RECEIVE_LN_ADDRESS=yourname@wallet.com
### Setting the UI Password
On first start the node generates an admin password and logs it once — read it from the container logs to sign in. You can then change it two ways:
There are two ways to set or change your Admin Dashboard password:
1. **Via Dashboard**: Once logged in, go to **Settings****Security** to update your password.
2. **Via Environment Variable (legacy seed)**: Setting `ADMIN_PASSWORD` in `.env` before the first start seeds the initial password instead of generating one. It's read only once, for existing deployments; a value left in `.env` is ignored after the node has been configured.
1. **Via Environment Variable**: Set `ADMIN_PASSWORD` in your `.env` file before starting the container. This will be the password used for the first login.
2. **Via Dashboard**: Once logged in, go to **Settings****Security** to update your password. Dashboard settings override the `.env` file once saved.
---
@@ -126,14 +123,12 @@ Use environment variables for:
| -------------------- | --------------------------------- | ------------------------------------ |
| `UPSTREAM_BASE_URL` | Upstream API endpoint | — |
| `UPSTREAM_API_KEY` | Upstream API key | — |
| `ADMIN_PASSWORD` | Legacy seed for the dashboard password (otherwise generated + logged on first start) | (auto-generated) |
| `ROUTSTR_SECRET_KEY` | Master key encrypting node secrets at rest. Auto-generated to a key file if unset | (auto-generated) |
| `ROUTSTR_SECRET_KEY_FILE` | Path to the generated key file (used when `ROUTSTR_SECRET_KEY` is unset) | `routstr_secret.key` beside the database |
| `ADMIN_PASSWORD` | Dashboard password | (none) |
| `DATABASE_URL` | Database connection string | `sqlite+aiosqlite:///keys.db` |
| `NAME` | Node display name | `ARoutstrNode` |
| `DESCRIPTION` | Node description | `A Routstr Node` |
| `NPUB` | Nostr public key (bech32) | — |
| `NSEC` | Legacy seed for the Nostr private key (otherwise set from the admin UI) | — |
| `NSEC` | Nostr private key | — |
| `ENABLE_ANALYTICS_SHARING` | Enable usage analytics sharing to Nostr | `true` |
| `CASHU_MINTS` | Comma-separated mint URLs | `https://mint.minibits.cash/Bitcoin` |
| `RECEIVE_LN_ADDRESS` | Lightning address for withdrawals | — |
@@ -147,20 +142,6 @@ Use environment variables for:
Environment variables are read on startup. Dashboard settings override them and persist in the database. Once you change a setting in the dashboard, the env var is ignored for that setting.
### Secrets at Rest
The node's Nostr private key (`nsec`) is encrypted in the database using
`ROUTSTR_SECRET_KEY`. You don't have to set it: if it's unset, the node generates a
key on first start, writes it **beside the database** (the file named by
`ROUTSTR_SECRET_KEY_FILE`, default `routstr_secret.key`) so it persists on the same
volume as your data, and prints it once.
**Back up that key** — it lives on the same volume as your database, so include it
in your backups. If it is lost or changed, previously encrypted secrets can't be
decrypted and must be re-entered — there is no rotation. To keep the key off the
data volume, set `ROUTSTR_SECRET_KEY` explicitly (an env value always takes
precedence over the file). See also [Deployment](deployment.md).
---
## Models
+6 -24
View File
@@ -38,6 +38,7 @@ services:
- routstr-data:/app/data
environment:
DATABASE_URL: "sqlite:////app/data/routstr.db"
ADMIN_KEY: "your-secure-admin-key"
LOG_LEVEL: "info"
volumes:
@@ -86,8 +87,6 @@ services:
- ./logs:/app/logs
environment:
- TOR_PROXY_URL=socks5://tor:9050
# Keep the database (and the key file generated beside it) on the volume.
- DATABASE_URL=sqlite:////app/data/routstr.db
depends_on:
- tor
@@ -126,8 +125,8 @@ services:
- UPSTREAM_BASE_URL=https://api.openai.com/v1
- UPSTREAM_API_KEY=sk-proj-...
# The admin password is generated and logged once on first start; set
# ADMIN_PASSWORD here only as a legacy seed for an existing deployment.
# Secure the dashboard (recommended)
- ADMIN_PASSWORD=your-secure-password
# Node identity
- NAME=My Provider Node
@@ -135,9 +134,6 @@ services:
# Lightning withdrawals
- RECEIVE_LN_ADDRESS=me@walletofsatoshi.com
# Keep the database (and the key file generated beside it) on the volume.
- DATABASE_URL=sqlite:////app/data/routstr.db
volumes:
- ./data:/app/data
```
@@ -159,36 +155,22 @@ Example `.env`:
```bash
UPSTREAM_BASE_URL=https://api.openai.com/v1
UPSTREAM_API_KEY=sk-proj-...
# Keep the database (and the key file generated beside it) on the mounted volume.
DATABASE_URL=sqlite:////app/data/routstr.db
# Encrypts node secrets at rest. Optional — if unset, a key is generated next to
# your database (on the same volume) and its file is named once for backup. Set
# it explicitly to manage the key yourself.
ROUTSTR_SECRET_KEY=
ADMIN_PASSWORD=change-me
NAME=My Provider Node
RECEIVE_LN_ADDRESS=me@walletofsatoshi.com
```
!!! note "Secret key persistence"
If you leave `ROUTSTR_SECRET_KEY` unset, the node generates one and stores it
as `routstr_secret.key` **next to your database**, so it persists on the same
volume as your data — just include that volume in your backups. For stronger
isolation (keeping the key off the data volume), set `ROUTSTR_SECRET_KEY` from
a secrets manager instead.
See [Configuration](configuration.md) for all available options.
---
## Persistence
Point `DATABASE_URL` inside `/app/data` (as the examples above do) so everything
Routstr persists lands on the mounted volume:
Routstr stores all data in `/app/data`:
| Path | Contents |
|------|----------|
| `routstr.db` | SQLite database (settings, API keys, sessions) |
| `routstr_secret.key` | Auto-generated master key, written beside the database when `ROUTSTR_SECRET_KEY` is unset |
| `keys.db` | SQLite database (settings, API keys, sessions) |
| `.wallet/` | Cashu wallet data (your Bitcoin!) |
!!! warning "Back Up Your Data"
+4 -10
View File
@@ -29,25 +29,19 @@ In future versions, you'll be able to run a node that connects to other Routstr
Create a `.env` file in the root of the project to store your secrets:
```bash
# Encrypts node secrets at rest. Optional — if unset, the node generates a key on
# first start and prints it once (back it up).
ROUTSTR_SECRET_KEY=
# Initial Admin Password
ADMIN_PASSWORD=mysecretpassword
# Node Identity
NAME="My AI Node"
DESCRIPTION="Fast access to models"
NSEC=yournsec
# Lightning Payouts
RECEIVE_LN_ADDRESS=yourname@wallet.com
```
The admin password is generated and logged once on first start (read it from the
logs to sign in), and your Nostr identity (`nsec`) is configured afterwards from
the admin UI — both are stored encrypted in the database, not in `.env`.
(`ADMIN_PASSWORD` / `NSEC` are still read once as a legacy seed for existing
deployments.)
## 2. Start the Node
The recommended way to run Routstr is using Docker Compose, which handles the node, the UI, and optional services like Tor.
@@ -78,7 +72,7 @@ docker compose up -d
Open the **Admin Dashboard** at [http://localhost:8000/admin/](http://localhost:8000/admin/).
!!! note "Login"
On first start the node generates an admin password and logs it once — read it from the container logs to sign in. You can change it afterwards from **Settings****Security**.
Use the `ADMIN_PASSWORD` you defined in your `.env` file to log in. If you didn't set one, the dashboard will prompt you to set one on first visit.
### Connect Your AI Providers
-493
View File
@@ -1,493 +0,0 @@
# Tinfoil / PPQ Private-Mode Integration Notes
This document summarizes the current options for integrating Tinfoil/PPQ private models with Routstr, based on the EHBP work in this branch and local testing against `ppq-private-mode-proxy`.
## Background
PPQ private models run behind a Tinfoil/EHBP flow:
- Request bodies are HPKE-encrypted by a Tinfoil client.
- The PPQ `/private/` endpoint routes ciphertext to the attested enclave.
- The enclave decrypts, runs inference, and returns an encrypted response.
- The caller's Tinfoil client decrypts the response locally.
The PPQ private-mode proxy (`~/projects/ppq-private-mode-proxy`) uses this pattern with the JavaScript `tinfoil` SDK:
```ts
import { SecureClient } from "tinfoil";
const apiBase = "https://api.ppq.ai";
const client = new SecureClient({
baseURL: `${apiBase}/private/`,
attestationBundleURL: `${apiBase}/private`,
transport: "ehbp",
});
await client.ready();
const response = await client.fetch(`${apiBase}/private/v1/chat/completions`, {
method: "POST",
headers: {
"Content-Type": "application/json",
Authorization: `Bearer ${process.env.PPQ_API_KEY}`,
"X-Private-Model": "private/gpt-oss-120b",
"x-query-source": "api",
},
body: JSON.stringify({
model: "gpt-oss-120b", // enclave-internal model id
messages: [{ role: "user", content: "Hello" }],
}),
});
const json = await response.json();
console.log(json.usage);
```
`X-Private-Model` carries the PPQ-facing private model id, while the encrypted JSON body uses the enclave-internal model id without the `private/` prefix.
## PPQ private model pricing
PPQ exposes private model pricing through:
```text
GET https://api.ppq.ai/v1/models?type=all
```
Filter models whose IDs start with `private/`.
Example private pricing observed:
| Model | Input USD / 1M tokens | Output USD / 1M tokens |
|---|---:|---:|
| `private/gpt-oss-120b` | `0.79125` | `1.31875` |
| `private/llama3-3-70b` | `1.84625` | `2.90125` |
| `private/qwen3-vl-30b` | `1.31875` | `4.22` |
| `private/glm-5-2` | `1.5825` | `5.53875` |
| `private/gemma4-31b` | `0.47475` | `1.055` |
| `private/kimi-k2-6` | `1.5825` | `5.53875` |
The `ppq-private-mode-proxy` commit `ba984214793d3bca0f7d046b6955d42abc1c6843` changed OpenClaw display metadata to align with a 5% API margin. The live PPQ model endpoint returns the more precise rates above, which already include that margin.
Actual PPQ private billing is:
```text
price_usd =
input_tokens * input_per_1M_tokens / 1_000_000
+ output_tokens * output_per_1M_tokens / 1_000_000
```
A test request to `private/gpt-oss-120b` produced:
```json
{
"input_count": 74,
"output_count": 3,
"price_in_usd": 0.00006250875
}
```
which exactly matches:
```text
74 * 0.79125 / 1_000_000 + 3 * 1.31875 / 1_000_000
= 0.00006250875 USD
```
## Integration architectures
There are three materially different ways Routstr could integrate Tinfoil/PPQ private inference.
## Option A: User integrates Tinfoil directly
```text
User app / SDK
-> Tinfoil SecureClient / TinfoilAI
-> EHBP-encrypted request
-> Tinfoil/PPQ private enclave
```
Properties:
- Best privacy for the user.
- Routstr is not in the request path.
- User's client encrypts requests and decrypts responses.
- Usage is visible to the user's app after decryption.
- Billing is handled directly by Tinfoil/PPQ.
Example with Tinfoil's OpenAI-compatible client:
```ts
import { TinfoilAI } from "tinfoil";
const client = new TinfoilAI({
apiKey: process.env.TINFOIL_API_KEY,
transport: "ehbp",
});
const res = await client.chat.completions.create({
model: "llama3-3-70b",
messages: [{ role: "user", content: "Hello" }],
});
console.log(res.choices[0].message.content);
console.log(res.usage);
```
This is not a Routstr marketplace flow unless Routstr only acts as discovery/UI around direct Tinfoil/PPQ usage.
## Option B: Routstr integrates Tinfoil as an upstream client
```text
User -> Routstr plaintext request
-> Routstr Tinfoil SecureClient encrypts to PPQ/Tinfoil
-> PPQ private enclave
-> Routstr receives decrypted response
-> Routstr bills from decrypted usage
-> Routstr returns plaintext response to user
```
Properties:
- Easier exact billing.
- Routstr can read the decrypted OpenAI response and `usage` object.
- Routstr can charge exact PPQ token pricing.
- Privacy is different: the user sends plaintext to Routstr, and Routstr sees prompts/responses.
- End-to-end encryption is only Routstr-to-enclave, not user-to-enclave.
This should be considered a separate product/provider mode, not the same as an end-to-end private relay.
Practical implementation options:
1. Run a Node sidecar that uses the `tinfoil` npm package and expose it as a local HTTP upstream to Routstr.
2. Port EHBP client behavior to Python.
3. Reuse or adapt `ppq-private-mode-proxy` as a local upstream.
A Node sidecar is probably the quickest implementation path because `ppq-private-mode-proxy` already demonstrates the full flow.
## Option C: Routstr uses Tinfoil as a direct blind upstream
This is the current branch's design intent and is likely the best fit if Tinfoil/PPQ exposes usage metadata headers:
```text
User Tinfoil SecureClient
-> encrypted request body
-> Routstr proxy
-> Tinfoil/PPQ private API
with X-Tinfoil-Request-Usage-Metrics: true
-> PPQ private enclave
<- encrypted response
plus X-Tinfoil-Usage-Metrics / cost headers
<- Routstr proxy
-> user decrypts response
```
In this mode Routstr is still a normal upstream proxy from the user's point of view, but the upstream is Tinfoil/PPQ private inference and the body remains opaque to Routstr.
Properties:
- Strongest privacy with Routstr in the path.
- Routstr never sees plaintext prompt or plaintext response.
- Routstr can authenticate and route based on plaintext headers.
- Routstr should request Tinfoil usage metadata by adding `X-Tinfoil-Request-Usage-Metrics: true` to the upstream request.
- Tinfoil documents `X-Tinfoil-Usage-Metrics` as an upstream response header for non-streaming requests.
- For streaming requests, Tinfoil documents `X-Tinfoil-Usage-Metrics` as an HTTP trailer available only after the response body completes.
- If Tinfoil/PPQ returns `X-Tinfoil-Usage-Metrics` or a cost header on the encrypted response, Routstr can bill exactly without decrypting the body.
- If usage is only present inside the encrypted response body, Routstr still cannot read it and exact billing is not possible without a separate metadata path.
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
The current EHBP implementation charges successful EHBP requests at `max_cost_for_model` because Routstr cannot decrypt the response body:
```text
successful EHBP request -> charge full reserved max cost
```
That is incorrect for PPQ private models. Max cost should be only a reservation/solvency ceiling. Final charge should use actual PPQ private token pricing.
Desired behavior:
```text
reserve max cost
forward encrypted request
obtain actual usage/cost metadata
finalize actual cost
refund/release the difference
```
## Usage/cost metadata requirement
For blind-relay exact billing, PPQ should return one of the following outside the encrypted response body:
```http
X-PPQ-Cost-USD: 0.00006250875
```
or:
```http
X-Private-Usage-Metrics: input=74,output=3
```
or:
```http
X-Tinfoil-Usage-Metrics: prompt=74,completion=3,total=77
```
Tinfoil proxy documentation references a usage-metrics flow where a proxy can request usage via:
```http
X-Tinfoil-Request-Usage-Metrics: true
```
and read usage from:
```http
X-Tinfoil-Usage-Metrics
```
Docs/example references:
- https://docs.tinfoil.sh/guides/proxy-server
- https://github.com/tinfoilsh/encrypted-request-proxy-example
During local PPQ testing, PPQ responses included this CORS exposure header:
```http
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.
The decrypted body did include normal OpenAI usage, but only the decrypting Tinfoil client can see that body.
## Query-history fallback
PPQ's query history endpoint exposes actual usage and cost:
```text
GET https://api.ppq.ai/queries/history?page=1&page_count=...
```
A record includes:
```json
{
"timestamp": "...",
"model": "private/gpt-oss-120b",
"input_count": 74,
"output_count": 3,
"price_in_usd": 0.00006250875,
"query_type": "chat_completion",
"query_source": "api"
}
```
This could be used as a fallback, but it is less robust than response headers/trailers because matching a request to a history row can be race-prone under concurrency. It would need a reliable request identifier or metadata field that PPQ stores in history.
## Recommended Routstr direction
Use Tinfoil/PPQ as a direct blind upstream and have Routstr explicitly request usage metadata:
```text
User encrypts body
Routstr reserves max cost
Routstr forwards encrypted body to Tinfoil/PPQ /private/
Routstr includes X-Tinfoil-Request-Usage-Metrics: true
Tinfoil/PPQ returns X-Tinfoil-Usage-Metrics as a response header for non-streaming,
or as an HTTP trailer after the body completes for streaming
Routstr finalizes exact charge
User decrypts encrypted response
```
This preserves both:
- privacy: Routstr cannot read prompts/responses;
- exact billing: Routstr can charge actual PPQ private model cost, assuming Tinfoil/PPQ returns usage/cost metadata outside the encrypted body.
Implementation steps:
1. Update PPQ model fetching to include private models:
```text
GET https://api.ppq.ai/v1/models?type=all
```
2. Register `private/*` models and any Routstr-facing aliases with correct `forwarded_model_id`.
3. Keep max-cost reservation for bearer keys and X-Cashu solvency checks.
4. Add parsing support for possible usage/cost headers:
```http
X-PPQ-Cost-USD
X-Private-Usage-Metrics
X-Tinfoil-Usage-Metrics
X-Encrypted-Usage-Metrics
```
5. Finalize by actual cost instead of max cost:
```text
actual_msats = ceil((actual_usd / sats_usd_price()) * 1000)
actual_msats = max(actual_msats, settings.min_request_msat)
actual_msats = min(actual_msats, reserved_msats)
```
or, if only token counts are available:
```text
actual_usd =
input_tokens * input_per_1M_tokens / 1_000_000
+ output_tokens * output_per_1M_tokens / 1_000_000
```
6. If usage/cost metadata is missing, choose an explicit policy:
- fail closed and refund/revert;
- query PPQ history as a fallback;
- fallback to max-cost billing only if explicitly configured and clearly disclosed.
Silent max-cost billing should not be the default for PPQ private requests.
## X-Cashu consideration
For bearer-auth requests, finalization can happen after the response stream completes if usage is delivered as a trailer.
For `X-Cashu`, Routstr needs to return the refund token in the response headers. If usage/cost is only available after consuming the encrypted response stream, Routstr may need to buffer EHBP responses before sending them to the client so it can compute the refund amount first.
Possible approaches:
1. Prefer a non-trailer response header with actual cost, available before streaming body starts.
2. Buffer EHBP X-Cashu responses and then return `X-Cashu` refund.
3. Introduce a later/refund-claim mechanism, which would be a larger protocol change.
## Summary
- 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.
- 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.
## Implementation status
Direct Tinfoil upstream integration is implemented in `routstr/upstream/tinfoil.py`
and `routstr/upstream/ehbp.py`.
### What was built
- `TinfoilUpstreamProvider` (`provider_type = "tinfoil"`):
- 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.
- `supports_ehbp = True` — acts as a blind EHBP relay.
- `get_ehbp_forwarding_target()` returns a target that includes
`X-Tinfoil-Request-Usage-Metrics: true`.
- `forward_get_request()` proxies `/attestation` to `https://atc.tinfoil.sh/attestation`
so the SDK can fetch attestation bundles through Routstr.
- Registered in `routstr/upstream/__init__.py` and seeded from
`TINFOIL_API_KEY` env var.
- `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.
- `_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`,
`X-Tinfoil-Enclave-Url`, and `X-Tinfoil-Request-Usage-Metrics` before
forwarding to the enclave.
- `_compute_ehbp_actual_cost()` converts the usage header into msats via
`calculate_cost()`, clamped to `[min_request_msat, max_cost_for_model]`.
When the header's `model=<name>` differs from the requested model, the
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.
- `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.
- `routstr/proxy.py`: `/attestation` and `/tee/attestation` paths are forwarded
to Tinfoil upstreams without model/cost/auth lookups.
### Billing behavior
| Request shape | Usage source | Billing |
|---|---|---|
| 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 |
| 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` |
### Cost response headers
Since EHBP response bodies are opaque encrypted blobs, per-request cost cannot
be injected into the JSON body (as done in the normal proxy flow). Instead,
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-Input-Cost-Msats` | Bearer, X-Cashu | msats attributed to input tokens |
| `X-Routstr-Output-Cost-Msats` | Bearer, X-Cashu | msats attributed to output tokens |
The client/Tinfoil SDK can read these headers from the HTTP response without
needing to decrypt the body.
### Setup
```bash
TINFOIL_API_KEY=your-tinfoil-api-key
```
The provider is auto-seeded on first startup.
### Usage metrics header format
Tinfoil returns usage metrics in the `X-Tinfoil-Usage-Metrics` response header
(non-streaming) or HTTP trailer (streaming) when `X-Tinfoil-Request-Usage-Metrics:
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>
```
The `model` field carries the actual model name served by the enclave.
Routstr uses this to:
- Verify the served model matches the expected upstream model. The comparison
uses ``model_obj.forwarded_model_id`` (the actual upstream ID, e.g.
``glm-5-2``) rather than ``model_obj.id`` (the client-facing alias, e.g.
``tinfoil-glm-5-2``), so aliased models don't trigger a spurious mismatch.
- When they genuinely differ (Tinfoil served a different upstream model than
expected), look up the actual served model's pricing and use it for billing.
The reverse lookup uses ``get_model_instance``, which resolves
``forwarded_model_id`` values registered as routable aliases.
- Log the discrepancy for observability.
If the actual model is not found in Routstr's model registry, billing falls
back to the requested model's pricing.
### What still needs verification
- ~~End-to-end test with a real Tinfoil SDK client against a Routstr node with
`TINFOIL_API_KEY` set.~~ Verified: both non-streaming (header) and streaming
(trailer) responses include `model=<name>`.
- Streaming trailer capture is implemented by buffering the encrypted response
in `forward_with_trailer()` and then using the dedicated EHBP payment
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.
@@ -1,62 +0,0 @@
"""add reservation release idempotency records
Revision ID: 7f2843d3f4e4
Revises: fc4fa29630d2
Create Date: 2026-07-24 02:06:06.066726
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "7f2843d3f4e4"
down_revision = "fc4fa29630d2"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"reservation_releases",
sa.Column("id", sa.String(), nullable=False),
sa.Column("key_hash", sa.String(), nullable=False),
sa.Column("billing_key_hash", sa.String(), nullable=False),
sa.Column("reserved_msats", sa.Integer(), nullable=False),
sa.Column(
"status", sa.String(), nullable=False, server_default="active"
),
sa.Column("created_at", sa.Integer(), nullable=False),
sa.PrimaryKeyConstraint("id"),
)
op.create_index(
"ix_reservation_releases_key_hash",
"reservation_releases",
["key_hash"],
)
op.create_index(
"ix_reservation_releases_billing_key_hash",
"reservation_releases",
["billing_key_hash"],
)
op.create_index(
"ix_reservation_releases_status_created_at",
"reservation_releases",
["status", "created_at"],
)
def downgrade() -> None:
op.drop_index(
"ix_reservation_releases_status_created_at",
table_name="reservation_releases",
)
op.drop_index(
"ix_reservation_releases_billing_key_hash",
table_name="reservation_releases",
)
op.drop_index(
"ix_reservation_releases_key_hash",
table_name="reservation_releases",
)
op.drop_table("reservation_releases")
@@ -1,88 +0,0 @@
"""add slug to upstream_providers
Revision ID: c6d7e8f9a0b1
Revises: b5e7c9d1f3a2
Create Date: 2026-06-29 00:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
from routstr.core.provider_slugs import provider_slug_base, provider_slug_candidate
revision = "c6d7e8f9a0b1"
down_revision = "b5e7c9d1f3a2"
branch_labels = None
depends_on = None
def _allocate_backfill_slug(provider_type: str, reserved_slugs: set[str]) -> str:
base = provider_slug_base(provider_type)
suffix_number = 1
while True:
candidate = provider_slug_candidate(base, suffix_number)
if candidate not in reserved_slugs:
reserved_slugs.add(candidate)
return candidate
suffix_number += 1
def _backfill_provider_slugs(conn: sa.Connection) -> None:
existing_rows = conn.execute(
sa.text(
"SELECT slug FROM upstream_providers "
"WHERE slug IS NOT NULL AND slug != ''"
)
)
reserved_slugs = {str(row.slug).lower() for row in existing_rows}
rows_to_backfill = conn.execute(
sa.text(
"SELECT id, provider_type FROM upstream_providers "
"WHERE slug IS NULL OR slug = '' "
"ORDER BY id"
)
)
for row in rows_to_backfill:
slug = _allocate_backfill_slug(str(row.provider_type), reserved_slugs)
conn.execute(
sa.text("UPDATE upstream_providers SET slug = :slug WHERE id = :id"),
{"slug": slug, "id": row.id},
)
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = {c["name"] for c in inspector.get_columns("upstream_providers")}
if "slug" not in columns:
op.add_column(
"upstream_providers",
sa.Column("slug", sa.String(), nullable=True),
)
_backfill_provider_slugs(conn)
existing_indexes = {idx["name"] for idx in inspector.get_indexes("upstream_providers")}
if "ix_upstream_providers_slug" not in existing_indexes:
op.create_index(
"ix_upstream_providers_slug",
"upstream_providers",
["slug"],
unique=True,
)
def downgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
existing_indexes = {idx["name"] for idx in inspector.get_indexes("upstream_providers")}
if "ix_upstream_providers_slug" in existing_indexes:
op.drop_index("ix_upstream_providers_slug", table_name="upstream_providers")
columns = {c["name"] for c in inspector.get_columns("upstream_providers")}
if "slug" in columns:
op.drop_column("upstream_providers", "slug")
@@ -1,37 +0,0 @@
"""add fee payout checkpoint
Revision ID: d7e8f9a0b1c2
Revises: c6d7e8f9a0b1
Create Date: 2026-07-18 00:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "d7e8f9a0b1c2"
down_revision = "c6d7e8f9a0b1"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"routstr_fees",
sa.Column(
"payout_in_progress_msats",
sa.Integer(),
nullable=False,
server_default="0",
),
)
op.add_column(
"routstr_fees",
sa.Column("payout_started_at", sa.Integer(), nullable=True),
)
def downgrade() -> None:
op.drop_column("routstr_fees", "payout_started_at")
op.drop_column("routstr_fees", "payout_in_progress_msats")
@@ -1,51 +0,0 @@
"""add secrets table
Revision ID: fc4fa29630d2
Revises: d7e8f9a0b1c2
Create Date: 2026-07-23 00:00:00.000000
Creates the node-level singleton secret store (issue #553). Schema only; moving
any legacy plaintext into the encrypted/hashed columns happens at bootstrap,
where the live ROUTSTR_SECRET_KEY is available. ``nsec_state`` records the vault's
ownership of the nsec (legacy | encrypted | cleared), so a cleared identity is
never resurrected from a stale legacy ``NSEC`` env var / settings blob on the next
boot.
"""
import sqlalchemy as sa
import sqlmodel
from alembic import op
revision = "fc4fa29630d2"
down_revision = "d7e8f9a0b1c2"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"secrets",
sa.Column("id", sa.Integer(), nullable=False),
sa.Column(
"admin_password_hash",
sqlmodel.sql.sqltypes.AutoString(),
nullable=True,
),
sa.Column(
"encrypted_nsec",
sqlmodel.sql.sqltypes.AutoString(),
nullable=True,
),
sa.Column(
"nsec_state",
sqlmodel.sql.sqltypes.AutoString(),
nullable=False,
server_default="legacy",
),
sa.Column("updated_at", sa.Integer(), nullable=True),
sa.PrimaryKeyConstraint("id"),
)
def downgrade() -> None:
op.drop_table("secrets")
+6 -2
View File
@@ -1,6 +1,6 @@
[project]
name = "routstr"
version = "0.4.4"
version = "0.4.3"
description = "Payment proxy for your LLM endpoint using cashu and nostr."
readme = "README.md"
requires-python = ">=3.11"
@@ -10,7 +10,6 @@ dependencies = [
"aiosqlite>=0.20",
"sqlmodel>=0.0.24",
"httpx[socks]>=0.25.2",
"h11>=0.14",
"greenlet>=3.2.1",
"alembic>=1.13",
"python-json-logger>=2.0.0",
@@ -24,6 +23,11 @@ dependencies = [
"litellm>=1.55.0",
]
[project.optional-dependencies]
# Decentralized log mirroring to Nostr relays. Enable at runtime with
# ENABLE_SENTRYSTR=true (see routstr/core/logging.py).
sentrystr = ["sentrystr>=0.2.0"]
[dependency-groups]
dev = [
"mypy>=1.15.0",
+20 -30
View File
@@ -86,12 +86,10 @@ def get_provider_penalty(provider: "BaseUpstreamProvider") -> float:
def create_model_mappings(
upstreams: list["BaseUpstreamProvider"],
overrides_by_key: dict[tuple[str, int], tuple],
disabled_model_keys: set[tuple[str, int]],
overrides_by_id: dict[str, tuple],
disabled_model_ids: set[str],
) -> tuple[
dict[str, "Model"],
dict[str, list[tuple["Model", "BaseUpstreamProvider"]]],
dict[str, "Model"],
dict[str, "Model"], dict[str, list["BaseUpstreamProvider"]], dict[str, "Model"]
]:
"""Create optimal model mappings based on cost and provider preferences.
@@ -99,9 +97,7 @@ def create_model_mappings(
and creates three mappings based on cost optimization:
1. model_instances: alias -> Model (all model aliases mapped to their Model objects)
2. provider_map: alias -> List[(Model, UpstreamProvider)] (sorted candidate
list for each alias; each provider is paired with ITS OWN model so
failover can forward and bill the candidate that actually serves)
2. provider_map: alias -> List[UpstreamProvider] (sorted list of providers for each alias)
3. unique_models: base_id -> Model (unique models without provider prefixes)
The algorithm:
@@ -111,9 +107,8 @@ def create_model_mappings(
Args:
upstreams: List of all upstream provider instances
overrides_by_key: Dict of model overrides from database
{(model_id_lower, upstream_provider_id): (ModelRow, fee)}
disabled_model_keys: Set of provider-scoped model keys that should be excluded
overrides_by_id: Dict of model overrides from database {model_id: (ModelRow, fee)}
disabled_model_ids: Set of model IDs that should be excluded
Returns:
Tuple of (model_instances, provider_map, unique_models)
@@ -184,22 +179,14 @@ def create_model_mappings(
"""Process all models from a given provider."""
upstream_prefix = getattr(upstream, "upstream_name", None)
provider_key = get_provider_identity(upstream)
upstream_db_id = getattr(upstream, "db_id", None)
for model in upstream.get_cached_models():
model_key = (
(model.id.lower(), upstream_db_id)
if isinstance(upstream_db_id, int)
else None
)
if not model.enabled or (
model_key is not None and model_key in disabled_model_keys
):
if not model.enabled or model.id in disabled_model_ids:
continue
# 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]
# Apply overrides if present
if model.id in overrides_by_id:
override_row, provider_fee = overrides_by_id[model.id]
model_to_use = _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee
)
@@ -250,10 +237,13 @@ def create_model_mappings(
# Include enabled DB overrides even when provider discovery misses models.
# This is important for deployment-based providers like Azure.
for (model_id, upstream_provider_id), override_data in overrides_by_key.items():
if (model_id, upstream_provider_id) in disabled_model_keys:
for model_id, override_data in overrides_by_id.items():
if model_id in disabled_model_ids:
continue
override_row, provider_fee = override_data
upstream_provider_id = getattr(override_row, "upstream_provider_id", None)
if not isinstance(upstream_provider_id, int):
continue
upstream_for_override = providers_by_db_id.get(upstream_provider_id)
if upstream_for_override is None:
@@ -331,7 +321,7 @@ def create_model_mappings(
# Sort candidates and build final maps
model_instances: dict[str, "Model"] = {}
provider_map: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {}
provider_map: dict[str, list["BaseUpstreamProvider"]] = {}
def alias_priority(model: "Model", alias: str) -> int:
"""Rank how strong the mapping of alias->model is.
@@ -378,13 +368,13 @@ def create_model_mappings(
best_model, best_provider = items[0]
model_instances[alias] = best_model
provider_map[alias] = list(items)
provider_map[alias] = [p for _, p in items]
# Log provider distribution (using top provider for stats)
provider_counts: dict[str, int] = {}
for candidate_list in provider_map.values():
if candidate_list:
provider = candidate_list[0][1]
for providers in provider_map.values():
if providers:
provider = providers[0]
provider_name = getattr(provider, "upstream_name", "unknown")
provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1
+158 -443
View File
@@ -3,25 +3,16 @@ import hashlib
import math
import random
import time
import uuid
from contextvars import ContextVar
from dataclasses import dataclass
from datetime import datetime
from typing import TYPE_CHECKING, Optional
from typing import Optional
from fastapi import HTTPException
from sqlalchemy import case, inspect
from sqlalchemy import case
from sqlalchemy.exc import IntegrityError
from sqlmodel import col, select, update
from .core import get_logger
from .core.db import (
ApiKey,
AsyncSession,
ReservationRelease,
accumulate_routstr_fee,
create_session,
)
from .core.db import ApiKey, AsyncSession, accumulate_routstr_fee
from .core.settings import settings
from .payment.cost_calculation import (
CostData,
@@ -29,46 +20,17 @@ from .payment.cost_calculation import (
MaxCostData,
calculate_cost,
)
from .wallet import (
classify_redemption_error,
credit_balance,
deserialize_token_from_string,
)
if TYPE_CHECKING:
from .payment.models import Model
from .wallet import credit_balance, deserialize_token_from_string
logger = get_logger(__name__)
payments_logger = get_logger("routstr.payments")
# Routstr platform fee constants
ROUTSTR_FEE_PERCENT: float = 2.1
ROUTSTR_LN_ADDRESS: str = (
"npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash"
)
ROUTSTR_LN_ADDRESS: str = "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash"
ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS: int = 900
ROUTSTR_FEE_DEFAULT_PAYOUT: int = 200
@dataclass(frozen=True)
class ReservationSnapshot:
release_id: str
key_hash: str
billing_key_hash: str
reserved_msats: int
_current_reservation: ContextVar[ReservationSnapshot | None] = ContextVar(
"current_billing_reservation", default=None
)
def _clear_current_reservation(snapshot: ReservationSnapshot) -> None:
current = _current_reservation.get()
if current is not None and current.release_id == snapshot.release_id:
_current_reservation.set(None)
# TODO: implement prepaid api key (not like it was before)
# PREPAID_API_KEY = os.environ.get("PREPAID_API_KEY", None)
# PREPAID_BALANCE = int(os.environ.get("PREPAID_BALANCE", "0")) * 1000 # Convert to msats
@@ -116,37 +78,6 @@ async def check_and_reset_limit(key: ApiKey, session: AsyncSession) -> bool:
return False
def redemption_error_to_http_exception(error: Exception) -> HTTPException:
"""Map a Cashu token redemption failure to a sanitized client-facing error.
Thin wrapper over the shared :func:`classify_redemption_error` so the bearer
path stays identical to the X-Cashu and top-up paths.
"""
classified = classify_redemption_error(error)
if classified is None:
return HTTPException(
status_code=500,
detail={
"error": {
"message": "Internal error during token redemption",
"type": "api_error",
"code": "internal_error",
}
},
)
error_type, status_code, message, error_code = classified
return HTTPException(
status_code=status_code,
detail={
"error": {
"message": message,
"type": error_type,
"code": error_code,
}
},
)
async def validate_bearer_key(
bearer_key: str,
session: AsyncSession,
@@ -285,17 +216,7 @@ async def validate_bearer_key(
try:
hashed_key = hashlib.sha256(bearer_key.encode()).hexdigest()
try:
token_obj = deserialize_token_from_string(bearer_key)
except Exception as decode_error:
# A malformed token is a bad token (400 invalid_cashu_token via
# the shared taxonomy), not an auth failure (401) — otherwise it
# would fall through to the generic "Invalid API key" handler.
raise redemption_error_to_http_exception(
ValueError(
f"Invalid Cashu token: could not decode token ({decode_error})"
)
) from decode_error
token_obj = deserialize_token_from_string(bearer_key)
logger.debug(
"Generated token hash", extra={"hash_preview": hashed_key[:16] + "..."}
)
@@ -413,32 +334,19 @@ async def validate_bearer_key(
"error_type": type(credit_error).__name__,
},
)
await session.rollback()
raise redemption_error_to_http_exception(credit_error) from credit_error
raise credit_error
if msats <= 0:
logger.error(
"Token redemption returned zero or negative amount",
extra={"msats": msats, "key_hash": hashed_key[:8] + "..."},
)
# Defense-in-depth: credit_balance already raises
# ValueError("Redeemed token amount must be positive…") before
# returning (wallet.py), so this branch is only reachable if a
# zero/negative row was somehow persisted; drop it so we never
# leave an orphan zero-balance key. Reuse the shared taxonomy
# (cashu_error) so the envelope matches the mapper above.
# Defense-in-depth: credit_balance now refuses to commit on a
# zero/negative redemption, but if a row was nonetheless
# persisted, drop it so we never leave an orphan zero-balance key.
await session.delete(new_key)
await session.commit()
raise HTTPException(
status_code=400,
detail={
"error": {
"message": "Failed to redeem Cashu token: token yielded no value",
"type": "cashu_error",
"code": "cashu_token_zero_value",
}
},
)
raise Exception("Token redemption failed")
await session.refresh(new_key)
await session.commit()
@@ -471,7 +379,7 @@ async def validate_bearer_key(
status_code=401,
detail={
"error": {
"message": "Invalid or expired Cashu key",
"message": f"Invalid or expired Cashu key: {str(e)}",
"type": "invalid_request_error",
"code": "invalid_api_key",
}
@@ -626,16 +534,6 @@ async def pay_for_request(
},
)
# Create the durable reservation identity before changing aggregate balances.
# The row and balance updates commit together, so every reserved amount has one
# owner that can reach exactly one terminal state.
reservation = ReservationSnapshot(
release_id=uuid.uuid4().hex,
key_hash=key.hashed_key,
billing_key_hash=billing_key.hashed_key,
reserved_msats=cost_per_request,
)
# Charge the base cost for the request atomically to avoid race conditions
reserved_at_now = int(time.time())
stmt = (
@@ -697,83 +595,22 @@ async def pay_for_request(
child_result = await session.exec(child_stmt) # type: ignore[call-overload]
if child_result.rowcount == 0:
# Build the error before rollback expires ORM attributes.
limit_message = (
f"Balance limit exceeded: {key.balance_limit} mSats limit. "
f"{key.total_spent} already spent ({key.reserved_balance} reserved), "
f"{cost_per_request} required for this request."
)
# The parent reservation update already ran in this transaction.
# Roll it back before failover code attempts to restore the previous
# reservation; otherwise that later commit can persist both updates.
await session.rollback()
raise HTTPException(
status_code=402,
detail={
"error": {
"message": limit_message,
"message": f"Balance limit exceeded: {key.balance_limit} mSats limit. {key.total_spent} already spent ({key.reserved_balance} reserved), {cost_per_request} required for this request.",
"type": "insufficient_quota",
"code": "balance_limit_exceeded",
}
},
)
session.add(
ReservationRelease(
id=reservation.release_id,
key_hash=reservation.key_hash,
billing_key_hash=reservation.billing_key_hash,
reserved_msats=reservation.reserved_msats,
status="active",
)
)
# Publish the identity before commit. If the commit succeeds but its
# acknowledgement is interrupted, exact cleanup can still recover the
# durable row. A definitely failed commit is harmless because every
# terminal transition validates that row before touching balances.
_current_reservation.set(reservation)
try:
await session.commit()
except BaseException:
# The database may have committed even if acknowledgement was cancelled
# or the connection failed. Reconcile using a fresh transaction and the
# exact durable identity; no upstream request has started yet.
try:
await session.rollback()
except Exception:
pass
try:
async with create_session() as cleanup_session:
record = await cleanup_session.get(
ReservationRelease, reservation.release_id
)
if record is not None and record.status == "active":
await _transition_reservation_to_released(
reservation,
cleanup_session,
decrement_requests=True,
idempotent_success=True,
)
except Exception:
logger.exception(
"Failed to reconcile ambiguous reservation commit",
extra={"reservation_id": reservation.release_id},
)
finally:
_clear_current_reservation(reservation)
raise
await session.commit()
try:
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
except Exception:
# The reservation transaction is already committed and durable. Logging
# refresh failures must not make the caller treat it as unreserved.
logger.exception(
"Reservation committed but post-commit refresh failed",
extra={"reservation_id": reservation.release_id},
)
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
logger.info(
"Payment processed successfully",
@@ -803,175 +640,72 @@ async def pay_for_request(
async def revert_pay_for_request(
key: ApiKey,
session: AsyncSession,
cost_per_request: int,
reservation_snapshot: ReservationSnapshot | None = None,
key: ApiKey, session: AsyncSession, cost_per_request: int
) -> bool:
"""Revert the current request's durable reservation exactly once."""
snapshot = reservation_snapshot or await get_reservation_snapshot(key, session)
await _validate_reservation_snapshot(key, snapshot, session, require_active=False)
if cost_per_request != snapshot.reserved_msats:
return False
return await _transition_reservation_to_released(
snapshot,
session,
decrement_requests=True,
idempotent_success=False,
"""Revert a previously reserved payment. Returns True if revert succeeded,
False if the reservation was already released (prevents negative reserved_balance)."""
billing_key = await get_billing_key(key, session)
# Keep reserved_at while other reservations remain
cleared_reserved_at = case(
(col(ApiKey.reserved_balance) - cost_per_request > 0, col(ApiKey.reserved_at)),
else_=None,
)
async def _validate_reservation_snapshot(
key: ApiKey,
snapshot: ReservationSnapshot,
session: AsyncSession,
*,
require_active: bool = True,
) -> None:
"""Reject cross-request or forged reservation handles before any mutation."""
state = inspect(key)
identity = state.identity if state is not None else None
key_hash = str(identity[0]) if identity else key.__dict__.get("hashed_key")
if snapshot.key_hash != key_hash:
raise RuntimeError("Billing reservation does not belong to this key")
persisted_key = await session.get(ApiKey, snapshot.key_hash)
if persisted_key is None:
raise RuntimeError("Billing reservation key no longer exists")
expected_billing_hash = persisted_key.parent_key_hash or persisted_key.hashed_key
if snapshot.billing_key_hash != expected_billing_hash:
raise RuntimeError("Billing reservation does not belong to this billing key")
record = await session.get(ReservationRelease, snapshot.release_id)
if (
record is None
or (require_active and record.status != "active")
or record.key_hash != snapshot.key_hash
or record.billing_key_hash != snapshot.billing_key_hash
or record.reserved_msats != snapshot.reserved_msats
):
raise RuntimeError("Billing reservation record does not match the request")
async def get_reservation_snapshot(
key: ApiKey, session: AsyncSession
) -> ReservationSnapshot:
"""Return the durable reservation created for the current request."""
snapshot = _current_reservation.get()
if snapshot is None:
raise RuntimeError("No billing reservation is associated with this request")
await _validate_reservation_snapshot(key, snapshot, session)
return snapshot
async def _transition_reservation_to_released(
snapshot: ReservationSnapshot,
session: AsyncSession,
*,
decrement_requests: bool,
idempotent_success: bool,
) -> bool:
transition = (
update(ReservationRelease)
.where(col(ReservationRelease.id) == snapshot.release_id)
.where(col(ReservationRelease.status) == "active")
.where(col(ReservationRelease.key_hash) == snapshot.key_hash)
.where(col(ReservationRelease.billing_key_hash) == snapshot.billing_key_hash)
.where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats)
.values(status="released")
)
transition_result = await session.exec(transition) # type: ignore[call-overload]
if transition_result.rowcount != 1:
await session.rollback()
existing = await session.get(ReservationRelease, snapshot.release_id)
return bool(
idempotent_success
and existing is not None
and existing.status == "released"
and existing.key_hash == snapshot.key_hash
and existing.billing_key_hash == snapshot.billing_key_hash
and existing.reserved_msats == snapshot.reserved_msats
)
values: dict[str, object] = {
"reserved_balance": col(ApiKey.reserved_balance) - snapshot.reserved_msats,
"reserved_at": case(
(
col(ApiKey.reserved_balance) - snapshot.reserved_msats > 0,
col(ApiKey.reserved_at),
),
else_=None,
),
}
if decrement_requests:
values["total_requests"] = col(ApiKey.total_requests) - 1
release_stmt = (
stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == snapshot.billing_key_hash)
.where(col(ApiKey.reserved_balance) >= snapshot.reserved_msats)
.values(**values)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.where(col(ApiKey.reserved_balance) >= cost_per_request)
.values(
reserved_balance=col(ApiKey.reserved_balance) - cost_per_request,
reserved_at=cleared_reserved_at,
total_requests=col(ApiKey.total_requests) - 1,
)
)
result = await session.exec(release_stmt) # type: ignore[call-overload]
if result.rowcount != 1:
await session.rollback()
return False
if snapshot.billing_key_hash != snapshot.key_hash:
child_release_stmt = (
result = await session.exec(stmt) # type: ignore[call-overload]
# Also decrement total_requests and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == snapshot.key_hash)
.where(col(ApiKey.reserved_balance) >= snapshot.reserved_msats)
.values(**values)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.reserved_balance) >= cost_per_request)
.values(
total_requests=col(ApiKey.total_requests) - 1,
reserved_balance=col(ApiKey.reserved_balance) - cost_per_request,
reserved_at=cleared_reserved_at,
)
)
child_result = await session.exec( # type: ignore[call-overload]
child_release_stmt
)
if child_result.rowcount != 1:
await session.rollback()
return False
await session.exec(child_stmt) # type: ignore[call-overload]
await session.commit()
_clear_current_reservation(snapshot)
return True
async def release_reservation(
snapshot: ReservationSnapshot,
session: AsyncSession,
reserved_msats: int,
) -> bool:
"""Release one durable reservation exactly once without charging."""
if reserved_msats <= 0 or reserved_msats != snapshot.reserved_msats:
if result.rowcount == 0:
logger.warning(
"Revert skipped - reservation already released (no-op to prevent negative reserved_balance)",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"cost_to_revert": cost_per_request,
"current_reserved_balance": billing_key.reserved_balance,
},
)
return False
return await _transition_reservation_to_released(
snapshot,
session,
decrement_requests=False,
idempotent_success=True,
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
payments_logger.info(
"REVERT",
extra={
"event": "revert",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"cost_reverted": cost_per_request,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
},
)
async def _claim_reservation_for_charge(
snapshot: ReservationSnapshot, session: AsyncSession
) -> bool:
"""Claim an active reservation in the caller's charge transaction."""
statement = (
update(ReservationRelease)
.where(col(ReservationRelease.id) == snapshot.release_id)
.where(col(ReservationRelease.status) == "active")
.where(col(ReservationRelease.key_hash) == snapshot.key_hash)
.where(col(ReservationRelease.billing_key_hash) == snapshot.billing_key_hash)
.where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats)
.values(status="charged")
)
result = await session.exec(statement) # type: ignore[call-overload]
if result.rowcount == 1:
_clear_current_reservation(snapshot)
return True
await session.rollback()
return False
return True
async def adjust_payment_for_tokens(
@@ -979,30 +713,16 @@ async def adjust_payment_for_tokens(
response_data: dict,
session: AsyncSession,
deducted_max_cost: int,
model_obj: "Model | None" = None,
provider_fee: float | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
) -> dict:
"""
Adjusts the payment based on token usage in the response.
This is called after the initial payment and the upstream request is complete.
Returns cost data to be included in the response.
``model_obj`` is the model that actually served the request; it is passed
through to ``calculate_cost`` so billing uses the serving candidate's
pricing instead of re-deriving it from the response's model string.
The response's usage object is normalized with the default union parser in
``calculate_cost``.
"""
billing_key = await get_billing_key(key, session)
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
await _validate_reservation_snapshot(
key, reservation, session, require_active=False
)
# The persisted amount is authoritative if request-level minimum pricing
# changed the caller's original estimate.
deducted_max_cost = reservation.reserved_msats
model = response_data.get("model", "unknown")
logger.debug(
@@ -1018,21 +738,50 @@ async def adjust_payment_for_tokens(
)
async def release_reservation_only() -> None:
"""Fallback to release this request's reservation without charging."""
"""Fallback to release reservation without charging when main update fails."""
try:
released = await release_reservation(
reservation, session, reservation.reserved_msats
)
logger.warning(
"Released reservation without charging (fallback)"
if released
else "Reservation was already finalized; fallback skipped",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"deducted_max_cost": deducted_max_cost,
},
release_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.where(col(ApiKey.reserved_balance) >= deducted_max_cost)
.values(
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost
)
)
result = await session.exec(release_stmt) # type: ignore[call-overload]
# Also release on child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_release_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.reserved_balance) >= deducted_max_cost)
.values(
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost
)
)
await session.exec(child_release_stmt) # type: ignore[call-overload]
await session.commit()
if result.rowcount == 0: # type: ignore[union-attr]
logger.warning(
"Release reservation skipped - already released (no-op to prevent negative reserved_balance)",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"deducted_max_cost": deducted_max_cost,
},
)
else:
logger.warning(
"Released reservation without charging (fallback)",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"deducted_max_cost": deducted_max_cost,
},
)
except Exception as e:
logger.error(
"Failed to release reservation in fallback",
@@ -1054,17 +803,7 @@ async def adjust_payment_for_tokens(
extra={"error": str(e), "fee_msats": fee_msats},
)
calculated_cost = await calculate_cost(
response_data, deducted_max_cost, model_obj, provider_fee
)
if not isinstance(calculated_cost, CostDataError):
if not await _claim_reservation_for_charge(reservation, session):
# A prior charge or release already owns this reservation. Returning
# the calculated metadata is safe; the aggregate balances must not
# be modified a second time.
return calculated_cost.dict()
match calculated_cost:
match await calculate_cost(response_data, deducted_max_cost):
case MaxCostData() as cost:
logger.debug(
"Using max cost data (no token adjustment)",
@@ -1092,10 +831,8 @@ async def adjust_payment_for_tokens(
)
safe_reserved = case(
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
@@ -1113,10 +850,8 @@ async def adjust_payment_for_tokens(
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_safe_reserved = case(
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
child_stmt = (
@@ -1226,10 +961,8 @@ async def adjust_payment_for_tokens(
)
exact_safe_reserved = case(
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
@@ -1247,10 +980,8 @@ async def adjust_payment_for_tokens(
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_exact_safe_reserved = case(
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
child_stmt = (
@@ -1289,45 +1020,31 @@ async def adjust_payment_for_tokens(
# actual cost exceeded discounted reservation (due to tolerance_percentage)
if cost_difference > 0:
# Lock the billing row so the parent and child record the same
# database-determined charge under concurrent finalizations.
actual_charge_msats = 0
for attempt in range(5):
locked_billing_key = (
await session.exec(
select(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.with_for_update()
.execution_options(populate_existing=True)
)
).one()
observed_balance = locked_billing_key.balance
actual_charge_msats = min(observed_balance, total_cost_msats)
overrun_safe_reserved = case(
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
else_=0,
# Always release the reservation and charge min(actual_cost, balance).
# CASE expressions keep this atomic and safe even when the
# stale-reservation sweeper has already released the reservation.
chargeable = case(
(col(ApiKey.balance) >= total_cost_msats, total_cost_msats),
else_=col(ApiKey.balance),
)
overrun_safe_reserved = case(
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
else_=0,
)
finalize_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=overrun_safe_reserved,
balance=col(ApiKey.balance) - chargeable,
total_spent=col(ApiKey.total_spent) + chargeable,
)
finalize_result = await session.exec( # type: ignore[call-overload]
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.where(col(ApiKey.balance) == observed_balance)
.values(
reserved_balance=overrun_safe_reserved,
balance=col(ApiKey.balance) - actual_charge_msats,
total_spent=col(ApiKey.total_spent) + actual_charge_msats,
)
)
if finalize_result.rowcount == 1:
break
await session.rollback()
if not await _claim_reservation_for_charge(reservation, session):
return cost.dict()
else:
await session.rollback()
raise RuntimeError("Could not atomically finalize cost overrun")
)
await session.exec(finalize_stmt) # type: ignore[call-overload]
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
@@ -1335,7 +1052,7 @@ async def adjust_payment_for_tokens(
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(
reserved_balance=overrun_safe_reserved,
total_spent=col(ApiKey.total_spent) + actual_charge_msats,
total_spent=col(ApiKey.total_spent) + min(billing_key.balance, total_cost_msats),
)
)
await session.exec(child_stmt) # type: ignore[call-overload]
@@ -1345,18 +1062,18 @@ async def adjust_payment_for_tokens(
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
cost.total_msats = actual_charge_msats
cost.total_msats = total_cost_msats
logger.info(
"Finalized payment with additional charge",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"charged_amount": actual_charge_msats,
"charged_amount": total_cost_msats,
"new_balance": billing_key.balance,
"model": model,
},
)
await _accumulate_fee(actual_charge_msats)
await _accumulate_fee(total_cost_msats)
payments_logger.info(
"FINALIZE",
extra={
@@ -1365,7 +1082,7 @@ async def adjust_payment_for_tokens(
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model,
"cost_reserved": deducted_max_cost,
"cost_charged": actual_charge_msats,
"cost_charged": total_cost_msats,
"input_tokens": cost.input_tokens,
"output_tokens": cost.output_tokens,
"balance": billing_key.balance,
@@ -1405,10 +1122,8 @@ async def adjust_payment_for_tokens(
)
refund_safe_reserved = case(
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
@@ -1426,10 +1141,8 @@ async def adjust_payment_for_tokens(
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_refund_safe_reserved = case(
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
child_stmt = (
@@ -1604,7 +1317,9 @@ async def periodic_dead_key_prune() -> None:
try:
async with create_session() as session:
await prune_dead_api_keys(session, settings.dead_key_min_age_seconds)
await prune_dead_api_keys(
session, settings.dead_key_min_age_seconds
)
except asyncio.CancelledError:
break
except Exception as e:
+53 -49
View File
@@ -7,7 +7,7 @@ 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, or_, select, update
from .auth import get_billing_key, validate_bearer_key
from .core.db import (
@@ -15,22 +15,12 @@ from .core.db import (
AsyncSession,
CashuTransaction,
get_session,
release_stale_reservations,
)
from .core.db import (
store_cashu_transaction_with_retry as store_cashu_transaction,
store_cashu_transaction,
)
from .core.logging import get_logger
from .core.settings import settings
from .lightning import lightning_router
from .wallet import (
classify_redemption_error,
credit_balance,
is_mint_connection_error,
recieve_token,
send_to_lnurl,
send_token,
)
from .wallet import credit_balance, recieve_token, send_to_lnurl, send_token
router = APIRouter()
balance_router = APIRouter(prefix="/v1/balance")
@@ -166,18 +156,30 @@ async def topup_wallet_endpoint(
raise HTTPException(status_code=400, detail="Invalid token format")
try:
amount_msats = await credit_balance(cashu_token, billing_key, session)
except Exception as e:
# Shared taxonomy so top-up matches the bearer/X-Cashu paths (503 for an
# unreachable mint, 422 for fee/swap failures, 400 for token faults).
classified = classify_redemption_error(e)
if classified is None:
logger.error(
"topup_wallet_endpoint: unhandled error",
extra={"error": str(e), "error_type": type(e).__name__},
except ValueError as e:
error_msg = str(e)
if "already spent" in error_msg.lower():
raise HTTPException(status_code=400, detail="Token already spent")
elif "invalid" in error_msg.lower() or "decode" in error_msg.lower():
raise HTTPException(status_code=400, detail="Invalid token format")
elif "insufficient" in error_msg.lower() or "melt fee" in error_msg.lower():
raise HTTPException(
status_code=400,
detail=f"Token value is too small to cover swap fees. {error_msg}",
)
raise HTTPException(status_code=500, detail="Internal server error")
_type, status_code, message, _code = classified
raise HTTPException(status_code=status_code, detail=message)
elif "failed to melt" in error_msg.lower():
raise HTTPException(
status_code=400,
detail=f"Failed to swap foreign mint token. {error_msg}",
)
else:
raise HTTPException(status_code=400, detail=f"Failed to redeem token: {error_msg}")
except Exception as e:
logger.error(
"topup_wallet_endpoint: unhandled error",
extra={"error": str(e), "error_type": type(e).__name__},
)
raise HTTPException(status_code=500, detail="Internal server error")
return {"msats": amount_msats}
@@ -272,20 +274,7 @@ async def refund_wallet_endpoint(
)
out_tx = out_tx_result.first()
if out_tx is None:
# The "in" row exists with a request_id, but the "out" (refund)
# row hasn't been written yet — the upstream request is still in
# flight and the refund will be minted once it completes. Tell the
# client to retry instead of 404ing permanently (race condition
# where /v1/wallet/refund is polled before the refund exists).
logger.debug(
"refund_wallet_endpoint: refund pending (in row exists, out row not yet created)",
extra={"request_id": in_tx.request_id},
)
raise HTTPException(
status_code=425,
detail="Refund is pending; retry shortly.",
headers={"Retry-After": "2"},
)
raise HTTPException(status_code=404, detail="Refund not found")
if out_tx.swept:
raise HTTPException(status_code=410, detail="Refund has been swept")
@@ -324,19 +313,30 @@ async def refund_wallet_endpoint(
)
if key.reserved_balance > 0:
# Release only durable reservations old enough to be stale. A newer
# request on the same aggregate balance must remain reserved.
await release_stale_reservations(
session,
settings.stale_reservation_timeout_seconds,
key_hash=key.hashed_key,
# Release the reservation if it is stale
cutoff = int(time.time()) - settings.stale_reservation_timeout_seconds
stale_release_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.reserved_balance) > 0)
.where(
or_(
col(ApiKey.reserved_at).is_(None),
col(ApiKey.reserved_at) < cutoff,
)
)
.values(reserved_balance=0, reserved_at=None)
)
await session.refresh(key)
if key.reserved_balance > 0:
stale_result = await session.exec(stale_release_stmt) # type: ignore[call-overload]
await session.commit()
if stale_result.rowcount == 0:
raise HTTPException(
status_code=400,
detail="Cannot refund key. There are ongoing requests for this api key.",
)
await session.refresh(key)
logger.warning(
"refund_wallet_endpoint: released stale reservation before refund",
extra={
@@ -441,10 +441,14 @@ async def refund_wallet_endpoint(
"has_refund_address": bool(key.refund_address),
},
)
if is_mint_connection_error(e):
raise HTTPException(status_code=503, detail="Mint service unavailable")
if (
"mint" in error_msg.lower()
or "connection" in error_msg.lower()
or "ConnectError" in str(type(e))
):
raise HTTPException(status_code=503, detail=f"Mint service unavailable: {error_msg}")
else:
raise HTTPException(status_code=500, detail="Refund failed")
raise HTTPException(status_code=500, detail=f"Refund failed: {error_msg}")
await _refund_cache_set(bearer_value, result)
+171 -307
View File
@@ -1,6 +1,5 @@
import asyncio
import json
import re
import secrets
from datetime import datetime, timezone
from pathlib import Path
@@ -9,7 +8,6 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Request
from pydantic import BaseModel, RootModel
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 ..proxy import refresh_model_maps, reinitialize_upstreams
@@ -20,7 +18,6 @@ from ..wallet import (
send_token,
slow_filter_spend_proofs,
)
from . import vault
from .db import (
ApiKey,
CashuTransaction,
@@ -29,17 +26,10 @@ from .db import (
ModelRow,
UpstreamProviderRow,
create_session,
get_secret,
set_admin_password,
set_nsec,
)
from .db import (
store_cashu_transaction_with_retry as store_cashu_transaction,
)
from .log_manager import log_manager
from .logging import get_logger
from .provider_slugs import allocate_unique_provider_slug
from .settings import SettingsService, derive_npub_from_nsec, settings
from .settings import SettingsService, settings
logger = get_logger(__name__)
@@ -210,6 +200,8 @@ async def get_settings(request: Request) -> dict:
data = settings.dict()
if "upstream_api_key" in data:
data["upstream_api_key"] = "[REDACTED]" if data["upstream_api_key"] else ""
if "admin_password" in data:
data["admin_password"] = "[REDACTED]" if data["admin_password"] else ""
if "nsec" in data:
data["nsec"] = "[REDACTED]" if data["nsec"] else ""
return data
@@ -226,10 +218,9 @@ class PasswordUpdate(BaseModel):
@admin_router.patch("/api/settings", dependencies=[Depends(require_admin_api)])
async def update_settings(request: Request, update: SettingsUpdate) -> dict:
# Secrets are not editable through the general settings endpoint; they have
# dedicated rotation paths and never reach the settings blob.
# Remove sensitive fields from general settings update
settings_data = update.root.copy()
sensitive_fields = ["upstream_api_key", "nsec"]
sensitive_fields = ["admin_password", "upstream_api_key", "nsec"]
for field in sensitive_fields:
if field in settings_data:
del settings_data[field]
@@ -244,6 +235,8 @@ async def update_settings(request: Request, update: SettingsUpdate) -> dict:
data = new_settings.dict()
if "upstream_api_key" in data:
data["upstream_api_key"] = "[REDACTED]" if data["upstream_api_key"] else ""
if "admin_password" in data:
data["admin_password"] = "[REDACTED]" if data["admin_password"] else ""
if "nsec" in data:
data["nsec"] = "[REDACTED]" if data["nsec"] else ""
return data
@@ -251,63 +244,44 @@ async def update_settings(request: Request, update: SettingsUpdate) -> dict:
@admin_router.patch("/api/password", dependencies=[Depends(require_admin_api)])
async def update_password(request: Request, password_update: PasswordUpdate) -> dict:
current_password = settings.admin_password
if not current_password:
raise HTTPException(status_code=500, detail="Admin password not configured")
if password_update.current_password != current_password:
raise HTTPException(status_code=401, detail="Current password is incorrect")
# Validate new password
new_password = password_update.new_password.strip()
if len(new_password) < 6:
raise HTTPException(
status_code=400, detail="New password must be at least 6 characters"
)
# Update password
async with create_session() as session:
secret = await get_secret(session)
if not secret.admin_password_hash:
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"
)
# Validate new password
new_password = password_update.new_password.strip()
if len(new_password) < vault.MIN_PASSWORD_LENGTH:
raise HTTPException(
status_code=400,
detail=(
"New password must be at least "
f"{vault.MIN_PASSWORD_LENGTH} characters"
),
)
await set_admin_password(session, new_password)
await SettingsService.update({"admin_password": new_password}, session)
return {"ok": True, "message": "Password updated successfully"}
class NsecUpdate(BaseModel):
nsec: str
class SetupRequest(BaseModel):
password: str
@admin_router.patch("/api/nsec", dependencies=[Depends(require_admin_api)])
async def update_nsec(request: Request, payload: NsecUpdate) -> dict[str, object]:
# The node's Nostr identity is a secret: it is stored encrypted in the
# Secret store, never in the settings blob, so it gets its own endpoint
# rather than riding the general settings PATCH (which strips it). An empty
# nsec clears the identity.
nsec = payload.nsec.strip()
npub = ""
if nsec:
derived = derive_npub_from_nsec(nsec)
if not derived:
raise HTTPException(status_code=400, detail="Invalid nsec")
npub = derived
@admin_router.post("/api/setup")
async def initial_setup(request: Request, payload: SetupRequest) -> dict[str, object]:
if settings.admin_password:
raise HTTPException(status_code=409, detail="Admin password already set")
pw = (payload.password or "").strip()
if len(pw) < 8:
raise HTTPException(
status_code=400, detail="Password must be at least 8 characters"
)
async with create_session() as session:
await set_nsec(session, nsec)
# Reflect the change in the live runtime so Nostr signing/announcements pick
# it up without a restart (mirrors what bootstrap_secrets sets at boot).
settings.nsec = nsec
settings.npub = npub
return {"ok": True, "npub": npub}
await SettingsService.update({"admin_password": pw}, session)
return {"ok": True}
class AdminLoginRequest(BaseModel):
@@ -318,16 +292,12 @@ class AdminLoginRequest(BaseModel):
async def admin_login(
request: Request, payload: AdminLoginRequest
) -> dict[str, object]:
async with create_session() as session:
secret = await get_secret(session)
# Read the hash while the session is open; the ORM object is detached
# once the context exits and its attributes can no longer be loaded.
password_hash = secret.admin_password_hash
admin_pw = settings.admin_password
if not password_hash:
if not admin_pw:
raise HTTPException(status_code=500, detail="Admin password not configured")
if not vault.verify_password(payload.password, password_hash):
if payload.password != admin_pw:
raise HTTPException(status_code=401, detail="Invalid password")
token = secrets.token_urlsafe(32)
@@ -438,11 +408,12 @@ async def withdraw(
# Get wallet and check balance
from .settings import settings as global_settings
effective_mint = withdraw_request.mint_url or global_settings.primary_mint
wallet = await get_wallet(effective_mint, withdraw_request.unit)
wallet = await get_wallet(
withdraw_request.mint_url or global_settings.primary_mint, withdraw_request.unit
)
proofs = get_proofs_per_mint_and_unit(
wallet,
effective_mint,
withdraw_request.mint_url or global_settings.primary_mint,
withdraw_request.unit,
not_reserved=True,
)
@@ -458,27 +429,8 @@ async def withdraw(
raise HTTPException(status_code=400, detail="Insufficient wallet balance")
token = await send_token(
withdraw_request.amount, withdraw_request.unit, effective_mint
withdraw_request.amount, withdraw_request.unit, withdraw_request.mint_url
)
try:
await store_cashu_transaction(
token=token,
amount=withdraw_request.amount,
unit=withdraw_request.unit,
mint_url=effective_mint,
typ="out",
collected=False,
source="admin",
)
except Exception:
logger.critical(
"Admin withdrawal token issued without a persisted audit record",
extra={
"amount": withdraw_request.amount,
"unit": withdraw_request.unit,
"mint_url": effective_mint,
},
)
return {"token": token}
@@ -504,17 +456,19 @@ class ModelCreate(BaseModel):
dependencies=[Depends(require_admin_api)],
)
async def upsert_provider_model(
provider_id: str, payload: ModelCreate
provider_id: int, payload: ModelCreate
) -> dict[str, object]:
print(payload)
logger.info(
f"UPSERT_PROVIDER_MODEL called: provider_id={provider_id}, model_id={payload.id}"
)
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
provider_pk = _provider_pk(provider)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
# Try to get existing model
existing_row = await session.get(ModelRow, (payload.id, provider_pk))
existing_row = await session.get(ModelRow, (payload.id, provider_id))
if existing_row:
# Update existing model
@@ -570,7 +524,7 @@ async def upsert_provider_model(
alias_ids=(
json.dumps(payload.alias_ids) if payload.alias_ids else None
),
upstream_provider_id=provider_pk,
upstream_provider_id=provider_id,
enabled=payload.enabled,
forwarded_model_id=payload.forwarded_model_id or payload.id,
)
@@ -589,7 +543,7 @@ async def upsert_provider_model(
dependencies=[Depends(require_admin_api)],
)
async def update_provider_model_legacy(
provider_id: str, model_id: str, payload: ModelCreate
provider_id: int, model_id: str, payload: ModelCreate
) -> dict[str, object]:
"""Legacy PATCH endpoint - redirects to upsert POST endpoint for backward compatibility."""
logger.info(
@@ -602,12 +556,13 @@ async def update_provider_model_legacy(
"/api/upstream-providers/{provider_id}/models/{model_id:path}",
dependencies=[Depends(require_admin_api)],
)
async def get_provider_model(provider_id: str, model_id: str) -> dict[str, object]:
async def get_provider_model(provider_id: int, model_id: str) -> dict[str, object]:
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
provider_pk = _provider_pk(provider)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
row = await session.get(ModelRow, (model_id, provider_pk))
row = await session.get(ModelRow, (model_id, provider_id))
if not row:
raise HTTPException(
status_code=404, detail="Model not found for this provider"
@@ -621,11 +576,9 @@ async def get_provider_model(provider_id: str, model_id: str) -> dict[str, objec
"/api/upstream-providers/{provider_id}/models/{model_id:path}",
dependencies=[Depends(require_admin_api)],
)
async def delete_provider_model(provider_id: str, model_id: str) -> dict[str, object]:
async def delete_provider_model(provider_id: int, model_id: str) -> dict[str, object]:
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
provider_pk = _provider_pk(provider)
row = await session.get(ModelRow, (model_id, provider_pk))
row = await session.get(ModelRow, (model_id, provider_id))
if not row:
raise HTTPException(
status_code=404, detail="Model not found for this provider"
@@ -640,12 +593,10 @@ async def delete_provider_model(provider_id: str, model_id: str) -> dict[str, ob
"/api/upstream-providers/{provider_id}/models",
dependencies=[Depends(require_admin_api)],
)
async def delete_all_provider_models(provider_id: str) -> dict[str, object]:
async def delete_all_provider_models(provider_id: int) -> dict[str, object]:
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
provider_pk = _provider_pk(provider)
result = await session.exec(
select(ModelRow).where(ModelRow.upstream_provider_id == provider_pk)
select(ModelRow).where(ModelRow.upstream_provider_id == provider_id)
) # type: ignore
rows = result.all()
for row in rows:
@@ -664,7 +615,7 @@ class BatchOverrideRequest(BaseModel):
dependencies=[Depends(require_admin_api)],
)
async def batch_override_provider_models(
provider_id: str, payload: BatchOverrideRequest
provider_id: int, payload: BatchOverrideRequest
) -> dict[str, object]:
"""Batch override models for a specific provider."""
logger.info(
@@ -672,14 +623,15 @@ async def batch_override_provider_models(
)
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
provider_pk = _provider_pk(provider)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
overridden_count = 0
for model_data in payload.models:
# Try to get existing model regardless of whether it's enabled or not
existing_row = await session.get(ModelRow, (model_data.id, provider_pk))
existing_row = await session.get(ModelRow, (model_data.id, provider_id))
if existing_row:
# Update existing
@@ -733,7 +685,7 @@ async def batch_override_provider_models(
if model_data.alias_ids
else None
),
upstream_provider_id=provider_pk,
upstream_provider_id=provider_id,
enabled=model_data.enabled,
)
session.add(row)
@@ -750,85 +702,6 @@ async def batch_override_provider_models(
}
_SLUG_PATTERN = re.compile(r"^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$")
def _validate_slug(value: str) -> str:
candidate = value.strip().lower()
if not _SLUG_PATTERN.fullmatch(candidate):
raise HTTPException(
status_code=400,
detail=(
"slug must be 3-64 chars, lowercase letters/digits/hyphens, "
"and may not start or end with a hyphen"
),
)
if candidate.isdigit():
raise HTTPException(
status_code=400,
detail="slug must not be all digits",
)
return candidate
async def _ensure_unique_slug(
session: AsyncSession, slug: str, exclude_id: int | None = None
) -> None:
stmt = select(UpstreamProviderRow).where(UpstreamProviderRow.slug == slug)
result = await session.exec(stmt)
existing = result.first()
if existing and existing.id != exclude_id:
raise HTTPException(
status_code=409,
detail="Provider with this slug already exists",
)
async def _get_upstream_provider_by_ref(
session: AsyncSession, provider_ref: str
) -> UpstreamProviderRow:
if provider_ref.isdigit():
provider = await session.get(UpstreamProviderRow, int(provider_ref))
else:
slug = _validate_slug(provider_ref)
result = await session.exec(
select(UpstreamProviderRow).where(UpstreamProviderRow.slug == slug)
)
provider = result.first()
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
return provider
def _provider_pk(provider: UpstreamProviderRow) -> int:
if provider.id is None:
raise HTTPException(status_code=500, detail="Provider has no database id")
return provider.id
def _serialize_provider(
provider: UpstreamProviderRow, redact_api_key: bool = True
) -> dict[str, object]:
return {
"id": provider.id,
"slug": provider.slug,
"provider_type": provider.provider_type,
"base_url": provider.base_url,
"api_key": "[REDACTED]"
if (redact_api_key and provider.api_key)
else provider.api_key
if not redact_api_key
else "",
"api_version": provider.api_version,
"enabled": provider.enabled,
"provider_fee": provider.provider_fee,
"provider_settings": json.loads(provider.provider_settings)
if provider.provider_settings
else None,
}
class UpstreamProviderCreate(BaseModel):
provider_type: str
base_url: str
@@ -837,7 +710,6 @@ class UpstreamProviderCreate(BaseModel):
enabled: bool = True
provider_fee: float = 1.01
provider_settings: dict | None = None
slug: str | None = None
class UpstreamProviderUpdate(BaseModel):
@@ -848,50 +720,6 @@ class UpstreamProviderUpdate(BaseModel):
enabled: bool | None = None
provider_fee: float | None = None
provider_settings: dict | None = None
slug: str | None = None
class UpstreamProviderUpdateBySlug(BaseModel):
slug: str
new_slug: str | None = None
provider_type: str | None = None
base_url: str | None = None
api_key: str | None = None
api_version: str | None = None
enabled: bool | None = None
provider_fee: float | None = None
provider_settings: dict | None = None
async def _apply_provider_update(
session: AsyncSession,
provider: UpstreamProviderRow,
payload: UpstreamProviderUpdate,
new_slug: str | None = None,
) -> None:
if new_slug is not None:
validated = _validate_slug(new_slug)
await _ensure_unique_slug(session, validated, exclude_id=provider.id)
provider.slug = validated
if payload.provider_type is not None:
provider.provider_type = payload.provider_type
if payload.base_url is not None:
provider.base_url = payload.base_url
if payload.api_key is not None:
provider.api_key = payload.api_key
if payload.api_version is not None:
provider.api_version = payload.api_version
if payload.enabled is not None:
provider.enabled = payload.enabled
if payload.provider_fee is not None:
provider.provider_fee = payload.provider_fee
if payload.provider_settings is not None:
provider.provider_settings = json.dumps(payload.provider_settings)
session.add(provider)
await session.commit()
await session.refresh(provider)
@admin_router.get("/api/upstream-providers", dependencies=[Depends(require_admin_api)])
@@ -899,7 +727,21 @@ async def get_upstream_providers() -> list[dict[str, object]]:
async with create_session() as session:
result = await session.exec(select(UpstreamProviderRow))
providers = result.all()
return [_serialize_provider(p) for p in providers]
return [
{
"id": p.id,
"provider_type": p.provider_type,
"base_url": p.base_url,
"api_key": "[REDACTED]" if p.api_key else "",
"api_version": p.api_version,
"enabled": p.enabled,
"provider_fee": p.provider_fee,
"provider_settings": json.loads(p.provider_settings)
if p.provider_settings
else None,
}
for p in providers
]
@admin_router.post("/api/upstream-providers", dependencies=[Depends(require_admin_api)])
@@ -919,14 +761,7 @@ async def create_upstream_provider(
detail="Provider with this base URL and API key already exists",
)
if payload.slug:
slug = _validate_slug(payload.slug)
await _ensure_unique_slug(session, slug)
else:
slug = await allocate_unique_provider_slug(session, payload.provider_type)
provider = UpstreamProviderRow(
slug=slug,
provider_type=payload.provider_type,
base_url=payload.base_url,
api_key=payload.api_key,
@@ -943,81 +778,99 @@ async def create_upstream_provider(
await reinitialize_upstreams()
await refresh_model_maps()
return _serialize_provider(provider)
return {
"id": provider.id,
"provider_type": provider.provider_type,
"base_url": provider.base_url,
"api_key": "[REDACTED]",
"api_version": provider.api_version,
"enabled": provider.enabled,
"provider_fee": provider.provider_fee,
"provider_settings": payload.provider_settings,
}
@admin_router.get(
"/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)]
)
async def get_upstream_provider(provider_id: str) -> dict[str, object]:
async def get_upstream_provider(provider_id: int) -> dict[str, object]:
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
return _serialize_provider(provider)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
return {
"id": provider.id,
"provider_type": provider.provider_type,
"base_url": provider.base_url,
"api_key": "[REDACTED]" if provider.api_key else "",
"api_version": provider.api_version,
"enabled": provider.enabled,
"provider_fee": provider.provider_fee,
"provider_settings": json.loads(provider.provider_settings)
if provider.provider_settings
else None,
}
@admin_router.patch(
"/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)]
)
async def update_upstream_provider(
provider_id: str, payload: UpstreamProviderUpdate
provider_id: int, payload: UpstreamProviderUpdate
) -> dict[str, object]:
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
await _apply_provider_update(session, provider, payload, new_slug=payload.slug)
await reinitialize_upstreams()
await refresh_model_maps()
return _serialize_provider(provider)
@admin_router.patch(
"/api/upstream-providers", dependencies=[Depends(require_admin_api)]
)
async def update_upstream_provider_by_slug(
payload: UpstreamProviderUpdateBySlug,
) -> dict[str, object]:
lookup = _validate_slug(payload.slug)
async with create_session() as session:
result = await session.exec(
select(UpstreamProviderRow).where(
UpstreamProviderRow.slug == lookup
)
)
provider = result.first()
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
update_payload = UpstreamProviderUpdate(
provider_type=payload.provider_type,
base_url=payload.base_url,
api_key=payload.api_key,
api_version=payload.api_version,
enabled=payload.enabled,
provider_fee=payload.provider_fee,
provider_settings=payload.provider_settings,
)
await _apply_provider_update(
session, provider, update_payload, new_slug=payload.new_slug
)
if payload.provider_type is not None:
provider.provider_type = payload.provider_type
if payload.base_url is not None:
provider.base_url = payload.base_url
if payload.api_key is not None:
provider.api_key = payload.api_key
if payload.api_version is not None:
provider.api_version = payload.api_version
if payload.enabled is not None:
provider.enabled = payload.enabled
if payload.provider_fee is not None:
provider.provider_fee = payload.provider_fee
if payload.provider_settings is not None:
provider.provider_settings = json.dumps(payload.provider_settings)
session.add(provider)
await session.commit()
await session.refresh(provider)
await reinitialize_upstreams()
await refresh_model_maps()
return _serialize_provider(provider)
return {
"id": provider.id,
"provider_type": provider.provider_type,
"base_url": provider.base_url,
"api_key": "[REDACTED]",
"api_version": provider.api_version,
"enabled": provider.enabled,
"provider_fee": provider.provider_fee,
"provider_settings": json.loads(provider.provider_settings)
if provider.provider_settings
else None,
}
@admin_router.delete(
"/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)]
)
async def delete_upstream_provider(provider_id: str) -> dict[str, object]:
async def delete_upstream_provider(provider_id: int) -> dict[str, object]:
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
deleted_id = _provider_pk(provider)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
await session.delete(provider)
await session.commit()
await reinitialize_upstreams()
await refresh_model_maps()
return {"ok": True, "deleted_id": deleted_id}
return {"ok": True, "deleted_id": provider_id}
@admin_router.get("/api/provider-types", dependencies=[Depends(require_admin_api)])
@@ -1032,16 +885,17 @@ async def get_provider_types() -> list[dict[str, object]]:
"/api/upstream-providers/{provider_id}/models",
dependencies=[Depends(require_admin_api)],
)
async def get_provider_models(provider_id: str) -> dict[str, object]:
async def get_provider_models(provider_id: int) -> dict[str, object]:
from ..upstream.helpers import _instantiate_provider
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
provider_pk = _provider_pk(provider)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
db_models = await list_models(
session=session,
upstream_id=provider_pk,
upstream_id=provider_id,
include_disabled=True,
apply_fees=False,
)
@@ -1131,11 +985,13 @@ class TopupTokenRequest(BaseModel):
dependencies=[Depends(require_admin_api)],
)
async def topup_provider_with_token(
provider_id: str, payload: TopupTokenRequest
provider_id: int, payload: TopupTokenRequest
) -> dict:
"""Redeem a Cashu token for an upstream provider."""
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
import httpx
@@ -1166,13 +1022,15 @@ async def topup_provider_with_token(
dependencies=[Depends(require_admin_api)],
)
async def initiate_provider_topup(
provider_id: str, payload: TopupRequest
provider_id: int, payload: TopupRequest
) -> dict[str, object]:
"""Initiate a Lightning Network top-up for the upstream provider account."""
from ..upstream.helpers import _instantiate_provider
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
try:
logger.info(
@@ -1292,13 +1150,15 @@ async def initiate_provider_topup(
"/api/upstream-providers/{provider_id}/topup/{invoice_id}/status",
dependencies=[Depends(require_admin_api)],
)
async def check_topup_status(provider_id: str, invoice_id: str) -> dict[str, object]:
async def check_topup_status(provider_id: int, invoice_id: str) -> dict[str, object]:
"""Check the status of a Lightning Network top-up invoice."""
from ..upstream.helpers import _instantiate_provider
from ..upstream.ppqai import PPQAIUpstreamProvider
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
# For Routstr providers, proxy the status check
if provider.provider_type == "routstr":
@@ -1345,12 +1205,14 @@ async def check_topup_status(provider_id: str, invoice_id: str) -> dict[str, obj
"/api/upstream-providers/{provider_id}/balance",
dependencies=[Depends(require_admin_api)],
)
async def get_provider_balance(provider_id: str) -> dict[str, object]:
async def get_provider_balance(provider_id: int) -> dict[str, object]:
"""Get the current balance for an upstream provider account."""
from ..upstream.helpers import _instantiate_provider
async with create_session() as session:
provider = await _get_upstream_provider_by_ref(session, provider_id)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
# For Routstr providers, proxy the balance check
if provider.provider_type == "routstr":
@@ -1729,13 +1591,15 @@ async def get_lightning_invoices_api(
"/api/upstream-providers/{provider_id}/routstr/refund",
dependencies=[Depends(require_admin_api)],
)
async def refund_routstr_provider_balance(provider_id: str) -> dict[str, object]:
async def refund_routstr_provider_balance(provider_id: int) -> dict[str, object]:
"""Refund balance from an upstream Routstr provider back to the local wallet."""
from ..upstream.helpers import _instantiate_provider
from ..upstream.routstr import RoutstrUpstreamProvider
async with create_session() as session:
provider_row = await _get_upstream_provider_by_ref(session, provider_id)
provider_row = await session.get(UpstreamProviderRow, provider_id)
if not provider_row:
raise HTTPException(status_code=404, detail="Provider not found")
if provider_row.provider_type != "routstr":
raise HTTPException(
+25 -342
View File
@@ -1,19 +1,16 @@
import asyncio
import hashlib
import os
import pathlib
import sqlite3
import time
import uuid
from contextlib import asynccontextmanager
from enum import Enum
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, or_
from sqlalchemy.exc import IntegrityError, OperationalError
from sqlalchemy import UniqueConstraint, delete
from sqlalchemy.exc import OperationalError
from sqlalchemy.ext.asyncio.engine import create_async_engine
from sqlalchemy.orm import aliased
from sqlmodel import Field, Relationship, SQLModel, col, func, select, update
@@ -99,130 +96,32 @@ class ApiKey(SQLModel, table=True): # type: ignore
async def reset_all_reserved_balances(session: AsyncSession) -> None:
"""Release every active durable reservation during explicit startup reset."""
await session.exec( # type: ignore[call-overload]
update(ReservationRelease)
.where(col(ReservationRelease.status) == "active")
.values(status="released")
)
await session.exec( # type: ignore[call-overload]
update(ApiKey).values(reserved_balance=0, reserved_at=None)
)
stmt = update(ApiKey).values(reserved_balance=0, reserved_at=None)
await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
logger.info("Reset reserved balances on startup")
async def release_stale_reservations(
session: AsyncSession,
max_age_seconds: int,
*,
key_hash: str | None = None,
session: AsyncSession, max_age_seconds: int
) -> int:
"""Release stale durable reservations without touching newer reservations."""
"""Release reservations whose last reserve is older than max_age_seconds.
"""
cutoff = int(time.time()) - max_age_seconds
query = (
select(ReservationRelease)
.where(col(ReservationRelease.status) == "active")
.where(col(ReservationRelease.created_at) < cutoff)
stmt = (
update(ApiKey)
.where(col(ApiKey.reserved_balance) > 0)
.where(col(ApiKey.reserved_at).is_not(None))
.where(col(ApiKey.reserved_at) < cutoff)
.values(reserved_balance=0, reserved_at=None)
)
if key_hash is not None:
query = query.where(
or_(
col(ReservationRelease.key_hash) == key_hash,
col(ReservationRelease.billing_key_hash) == key_hash,
)
)
reservations = (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:
continue
values = {
"reserved_balance": col(ApiKey.reserved_balance)
- reservation.reserved_msats,
"reserved_at": case(
(
col(ApiKey.reserved_balance) - reservation.reserved_msats > 0,
col(ApiKey.reserved_at),
),
else_=None,
),
}
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)
.values(**values)
)
if parent_result.rowcount != 1:
await session.rollback()
return 0
if reservation.billing_key_hash != reservation.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)
.values(**values)
)
if child_result.rowcount != 1:
await session.rollback()
return 0
released += 1
# Rolling upgrades can leave aggregate reservations created before durable
# reservation rows existed. Release only stale aggregates that have no active
# durable owner; targeted refund cleanup also heals legacy NULL timestamps.
legacy_query = select(ApiKey).where(col(ApiKey.reserved_balance) > 0)
if key_hash is None:
legacy_query = legacy_query.where(col(ApiKey.reserved_at).is_not(None)).where(
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(
or_(col(ApiKey.reserved_at).is_(None), col(ApiKey.reserved_at) < cutoff)
)
for legacy_key in (await session.exec(legacy_query)).all():
active_owner = (
await session.exec(
select(ReservationRelease.id)
.where(col(ReservationRelease.status) == "active")
.where(
or_(
col(ReservationRelease.key_hash) == legacy_key.hashed_key,
col(ReservationRelease.billing_key_hash)
== legacy_key.hashed_key,
)
)
.limit(1)
)
).first()
if active_owner is not None:
continue
legacy_key.reserved_balance = 0
legacy_key.reserved_at = None
session.add(legacy_key)
released += 1
result = await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
released = int(result.rowcount or 0)
if released:
logger.warning(
"Released stale reservations",
extra={"released_reservations": released, "max_age_seconds": max_age_seconds},
"Released stale balance reservations",
extra={"released_keys": released, "max_age_seconds": max_age_seconds},
)
return released
@@ -388,13 +287,10 @@ async def store_cashu_transaction(
created_at: int | None = None,
source: str = "x-cashu",
api_key_hashed_key: str | None = None,
transaction_id: str | None = None,
log_failure: bool = True,
) -> bool:
) -> None:
try:
async with create_session() as session:
tx = CashuTransaction(
id=transaction_id or uuid.uuid4().hex,
token=token,
amount=amount,
unit=unit,
@@ -408,93 +304,11 @@ async def store_cashu_transaction(
)
session.add(tx)
await session.commit()
except Exception:
if log_failure:
logger.critical(
"Failed to store Cashu transaction",
extra={"type": typ, "request_id": request_id, "source": source},
exc_info=True,
)
raise
return True
async def _cashu_transaction_exists(transaction_id: str) -> bool:
async with create_session() as session:
return await session.get(CashuTransaction, transaction_id) is not None
async def store_cashu_transaction_with_retry(
token: str,
amount: int,
unit: str,
mint_url: str | None = None,
typ: str = "out",
request_id: str | None = None,
collected: bool = False,
created_at: int | None = None,
source: str = "x-cashu",
api_key_hashed_key: str | None = None,
max_attempts: int = 3,
) -> bool:
"""Retry a critical Cashu transaction write with bounded backoff."""
transaction_id = hashlib.sha256(f"{typ}\0{token}".encode()).hexdigest()
last_error: Exception | None = None
for attempt in range(1, max_attempts + 1):
try:
return await store_cashu_transaction(
token=token,
amount=amount,
unit=unit,
mint_url=mint_url,
typ=typ,
request_id=request_id,
collected=collected,
created_at=created_at,
source=source,
api_key_hashed_key=api_key_hashed_key,
transaction_id=transaction_id,
log_failure=False,
)
except IntegrityError as error:
try:
if await _cashu_transaction_exists(transaction_id):
return True
except Exception as lookup_error:
last_error = lookup_error
else:
last_error = error
except Exception as error:
last_error = error
if last_error is not None:
if attempt == max_attempts:
break
delay = 0.25 * (2 ** (attempt - 1))
logger.warning(
"Cashu transaction storage failed; retrying",
extra={
"type": typ,
"request_id": request_id,
"attempt": attempt,
"max_attempts": max_attempts,
"retry_delay_seconds": delay,
},
)
await asyncio.sleep(delay)
logger.critical(
"Cashu transaction storage failed after bounded retries",
extra={
"type": typ,
"request_id": request_id,
"attempts": max_attempts,
"error": str(last_error),
},
)
if last_error is None:
raise RuntimeError("Cashu transaction storage failed without an exception")
raise last_error
except Exception as e:
logger.warning(
f"Failed to store cashu transaction: {e} (type={typ})",
extra={"error": str(e), "type": typ},
)
class UpstreamProviderRow(SQLModel, table=True): # type: ignore
@@ -505,12 +319,6 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore
),
)
id: int | None = Field(default=None, primary_key=True)
slug: str | None = Field(
default=None,
unique=True,
index=True,
description="Stable external slug used for updates via API key.",
)
provider_type: str = Field(
description="Provider type: custom, openai, anthropic, azure, openrouter, etc."
)
@@ -532,64 +340,12 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore
)
class ReservationRelease(SQLModel, table=True): # type: ignore
__tablename__ = "reservation_releases"
__table_args__ = (
Index("ix_reservation_releases_status_created_at", "status", "created_at"),
)
id: str = Field(primary_key=True)
key_hash: str = Field(index=True)
billing_key_hash: str = Field(index=True)
reserved_msats: int
status: str = Field(default="active")
created_at: int = Field(default_factory=lambda: int(time.time()))
class RoutstrFee(SQLModel, table=True): # type: ignore
__tablename__ = "routstr_fees"
id: int = Field(default=1, primary_key=True)
accumulated_msats: int = Field(default=0)
total_paid_msats: int = Field(default=0)
last_paid_at: int | None = Field(default=None)
payout_in_progress_msats: int = Field(default=0)
payout_started_at: int | None = Field(default=None)
class NsecState(str, Enum):
"""Ownership state of the node's nsec — an explicit 3-state machine.
The single ``encrypted_nsec`` column cannot distinguish "never migrated" from
"intentionally cleared" (both leave it empty), which let a cleared identity be
resurrected from a stale legacy ``NSEC``. This names the three states so the
bootstrap branches on ownership rather than inferring it:
* ``legacy`` — the vault has not taken ownership; a plaintext ``NSEC`` (env or
old settings blob) may still exist and should be migrated in once.
* ``encrypted`` — the vault owns a ciphertext; decrypt it, never re-read env.
* ``cleared`` — the vault owns it but the operator emptied it; stay empty,
never re-import from a stale legacy copy.
"""
legacy = "legacy"
encrypted = "encrypted"
cleared = "cleared"
class Secret(SQLModel, table=True): # type: ignore
"""Node-level secrets, stored encrypted/hashed at rest (singleton, id=1).
The asymmetric column names document the encoding: ``_hash`` is one-way
(scrypt, verify only) while ``encrypted_`` is reversible (Fernet). Per-provider
upstream keys live on ``upstream_providers``, not here. See ``routstr.core.vault``.
"""
__tablename__ = "secrets"
id: int = Field(default=1, primary_key=True)
admin_password_hash: str | None = Field(default=None)
encrypted_nsec: str | None = Field(default=None)
nsec_state: NsecState = Field(default=NsecState.legacy)
updated_at: int | None = Field(default=None)
class CliToken(SQLModel, table=True): # type: ignore
@@ -630,91 +386,18 @@ async def get_routstr_fee(session: AsyncSession) -> RoutstrFee:
return fee
async def get_secret(session: AsyncSession) -> Secret:
secret = await session.get(Secret, 1)
if secret is None:
secret = Secret(id=1)
session.add(secret)
try:
await session.commit()
except IntegrityError:
# Another worker created the singleton row between our read and
# insert (multiple workers booting against one shared DB). Roll back
# and read the row they committed instead of failing startup.
await session.rollback()
secret = await session.get(Secret, 1)
if secret is None:
raise
return secret
await session.refresh(secret)
return secret
async def set_admin_password(session: AsyncSession, password: str) -> None:
"""Store the admin password as a one-way hash on the Secret singleton."""
from .vault import hash_password
secret = await get_secret(session)
secret.admin_password_hash = hash_password(password)
secret.updated_at = int(time.time())
session.add(secret)
await session.commit()
async def set_nsec(session: AsyncSession, nsec: str) -> None:
"""Store the node's nsec, Fernet-encrypted, on the Secret singleton.
An empty string clears it (the node then holds no Nostr identity and signs
no events). Either way the vault now owns the nsec, so the state moves off
``legacy``: a cleared identity (``cleared``) must not be resurrected from a
stale legacy ``NSEC`` on the next boot.
"""
from .vault import encrypt
secret = await get_secret(session)
secret.encrypted_nsec = encrypt(nsec) if nsec else None
secret.nsec_state = NsecState.encrypted if nsec else NsecState.cleared
secret.updated_at = int(time.time())
session.add(secret)
await session.commit()
async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> bool:
"""Checkpoint a fee payout before making the external payment."""
async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> None:
stmt = (
update(RoutstrFee)
.where(col(RoutstrFee.id) == 1)
.where(col(RoutstrFee.payout_in_progress_msats) == 0)
.where(col(RoutstrFee.accumulated_msats) >= paid_msats)
.values(
accumulated_msats=RoutstrFee.accumulated_msats - paid_msats,
payout_in_progress_msats=paid_msats,
payout_started_at=int(time.time()),
)
)
result = await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
return result.rowcount == 1
async def complete_routstr_fee_payout(
session: AsyncSession, paid_msats: int
) -> bool:
"""Mark a checkpointed payout complete after the external payment succeeds."""
stmt = (
update(RoutstrFee)
.where(col(RoutstrFee.id) == 1)
.where(col(RoutstrFee.payout_in_progress_msats) == paid_msats)
.values(
payout_in_progress_msats=0,
payout_started_at=None,
total_paid_msats=RoutstrFee.total_paid_msats + paid_msats,
last_paid_at=int(time.time()),
)
)
result = await session.exec(stmt) # type: ignore[call-overload]
await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
return result.rowcount == 1
async def balances_for_mint_and_unit(
+150 -21
View File
@@ -232,32 +232,47 @@ class SecurityFilter(logging.Filter):
"refund_address",
}
def _scrub(self, message: str) -> str:
"""Redact secrets from a single string."""
standalone_patterns = [
r"Bearer\s+([a-zA-Z0-9_\-\.]{10,})", # Bearer token (must be 10 characters or more to reduce false-positives)
r"cashu[A-Z]+([a-zA-Z0-9_\-\.=/+]+)", # Cashu tokens
r"nsec[a-z0-9]+", # Nostr Public / Private Key
]
for pattern in standalone_patterns:
message = re.sub(pattern, "[REDACTED]", message, flags=re.IGNORECASE)
for key in self.SENSITIVE_KEYS:
if key in message.lower():
key_patterns = [
rf"{key}\s*[:=]\s*([a-zA-Z0-9_\-\.=/+]+)", # key:value or key=value (including any variant with spaces)
rf'{key}\s*[:=]\s*["\']([^"\']+)["\']', # key:"value" or key='value' (including any variant with spaces)
]
for pattern in key_patterns:
message = re.sub(
pattern, f"{key}: [REDACTED]", message, flags=re.IGNORECASE
)
return message
def filter(self, record: logging.LogRecord) -> bool:
"""Filter out sensitive information from log records."""
try:
message = record.getMessage()
message = redact_org_ids(message)
standalone_patterns = [
r"Bearer\s+([a-zA-Z0-9_\-\.]{10,})", # Bearer token (must be 10 characters or more to reduce false-positives)
r"cashu[A-Z]+([a-zA-Z0-9_\-\.=/+]+)", # Cashu tokens
r"nsec[a-z0-9]+", # Nostr Public / Private Key
]
for pattern in standalone_patterns:
message = re.sub(pattern, "[REDACTED]", message, flags=re.IGNORECASE)
for key in self.SENSITIVE_KEYS:
if key in message.lower():
key_patterns = [
rf"{key}\s*[:=]\s*([a-zA-Z0-9_\-\.=/+]+)", # key:value or key=value (including any variant with spaces)
rf'{key}\s*[:=]\s*["\']([^"\']+)["\']', # key:"value" or key='value' (including any variant with spaces)
]
for pattern in key_patterns:
message = re.sub(
pattern, f"{key}: [REDACTED]", message, flags=re.IGNORECASE
)
record.msg = message
message = redact_org_ids(record.getMessage())
record.msg = self._scrub(message)
record.args = ()
# The exception traceback is appended by the *formatter* (e.g. when a
# QueueHandler flattens the record before it is announced to the
# public, permanent SentryStr relays), so the message scrub above
# never sees it. Pre-render and redact it here and cache the result
# in `exc_text` so every formatter reuses the scrubbed version.
if record.exc_info and record.exc_info[0] is not None and not record.exc_text:
record.exc_text = logging.Formatter().formatException(record.exc_info)
if record.exc_text:
record.exc_text = self._scrub(redact_org_ids(record.exc_text))
if record.stack_info:
record.stack_info = self._scrub(redact_org_ids(record.stack_info))
# 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
@@ -452,6 +467,120 @@ def setup_logging() -> None:
logging.config.dictConfig(LOGGING_CONFIG)
_install_sentrystr_logging(LOGGING_CONFIG)
# Handle for the background SentryStr logging worker (None when disabled).
_sentrystr_handle: Any = None
# Bound on the in-memory queue feeding the background SentryStr worker. Records
# beyond this are dropped rather than blocking the caller or growing memory
# without limit when a relay is slow or unreachable.
_SENTRYSTR_QUEUE_SIZE = 10_000
# Max seconds to wait for the queued backlog to flush on shutdown before
# abandoning it, so shutdown cannot hang on a slow or unreachable relay.
_SENTRYSTR_SHUTDOWN_TIMEOUT = 10.0
def _install_sentrystr_logging(logging_config: dict[str, Any]) -> None:
"""Mirror routstr application logs to Nostr relays via SentryStr.
Announcing to relays blocks, so SentryStr's installer runs the publish on a
dedicated background worker thread fed by an in-memory queue; the request /
event-loop threads only pay the cost of enqueueing a record. Disabled unless
`ENABLE_SENTRYSTR` is set. Any failure here is swallowed so optional logging
can never take down startup.
"""
global _sentrystr_handle
if _sentrystr_handle is not None:
return
try:
from .settings import settings
except Exception:
return
if not getattr(settings, "enable_sentrystr", False):
return
routstr_logger = logging.getLogger("routstr")
relays = list(settings.sentrystr_relays or settings.relays or [])
if not relays:
routstr_logger.warning(
"SentryStr logging is enabled but no relays are configured; set "
"SENTRYSTR_RELAYS or RELAYS. Skipping SentryStr logging."
)
return
try:
from sentrystr import install_sentrystr_logging
except Exception:
routstr_logger.warning(
"SentryStr logging is enabled but the `sentrystr` package is not "
"installed. Install it with `pip install sentrystr`. Skipping "
"SentryStr logging."
)
return
# Default to ERROR: these events are announced to public, permanent relays,
# so only genuine problems should be mirrored unless the operator opts into
# something more verbose via SENTRYSTR_LOG_LEVEL. (Keep this at or above
# LOG_LEVEL; a lower value is gated out by the routstr loggers' own level.)
level = (settings.sentrystr_log_level or "ERROR").upper()
# routstr application loggers set propagate=False, so the handler must be
# attached to each of them rather than relying on propagation to the root.
routstr_loggers = [
name
for name in logging_config.get("loggers", {})
if name == "routstr" or name.startswith("routstr.")
] or ["routstr"]
try:
_sentrystr_handle = install_sentrystr_logging(
relays=relays,
private_key=(settings.sentrystr_nsec or settings.nsec or None),
level=level,
loggers=routstr_loggers,
recipient_pubkey=(settings.sentrystr_recipient_npub or None),
platform="routstr",
formatter=logging.Formatter("%(message)s"),
# Run the same scrubbing/enrichment filters used by the other
# handlers, on the calling thread, before records are handed to the
# background worker. RequestIdFilter must run here because it reads a
# ContextVar that is only set on the request thread; SecurityFilter
# must run here so secrets are redacted before anything is announced
# to the (public, permanent) relays. Passing them to the installer
# attaches them before the handler goes live, so no record can slip
# through unfiltered during startup.
filters=[RequestIdFilter(), SecurityFilter()],
# Bound the queue so a slow/unreachable relay can't grow memory.
queue_size=_SENTRYSTR_QUEUE_SIZE,
)
except Exception as exc:
routstr_logger.warning("Failed to initialize SentryStr logging: %s", exc)
return
routstr_logger.info(
"SentryStr logging enabled",
extra={"sentrystr_level": level, "sentrystr_relay_count": len(relays)},
)
def shutdown_sentrystr_logging() -> None:
"""Flush and stop the background SentryStr logging worker, if running."""
global _sentrystr_handle
handle = _sentrystr_handle
_sentrystr_handle = None
if handle is not None:
try:
handle.close(timeout=_SENTRYSTR_SHUTDOWN_TIMEOUT)
except Exception:
pass
def get_logger(name: str) -> logging.Logger:
"""Get a logger instance for the given module name."""
+14 -23
View File
@@ -37,7 +37,7 @@ from .exceptions import general_exception_handler, http_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
from .settings import SettingsService, bootstrap_secrets
from .settings import SettingsService
from .settings import settings as global_settings
from .version import __version__
@@ -85,20 +85,17 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
# Initialize application settings (env -> computed -> DB precedence)
async with create_session() as session:
# Move secrets into the encrypted/hashed store and decrypt the nsec
# into the in-memory settings BEFORE initializing settings: the
# initialize step strips secrets from the persisted blob, so legacy
# plaintext (env or old blob) must be migrated into the Secret store
# first or the only copy of a blob-only secret would be lost.
# Generates and logs an admin password on a fresh node; fails fast if
# a stored secret can't be decrypted.
await bootstrap_secrets(session)
s = await SettingsService.initialize(session)
if s.reset_reserved_balance_on_startup:
from .db import reset_all_reserved_balances
await reset_all_reserved_balances(session)
if not s.admin_password:
logger.warning(
f"Admin password is not set. Visit {s.http_url or 'http://localhost:8000'}/admin to set the password."
)
# Apply app metadata from settings
try:
app.title = s.name
@@ -233,6 +230,13 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
extra={"error": str(e), "error_type": type(e).__name__},
)
# Flush and stop the background SentryStr logging worker (no-op when
# SentryStr logging is disabled). Run off the event loop: the drain is
# blocking (though bounded) and must not stall the loop during shutdown.
from .logging import shutdown_sentrystr_logging
await asyncio.to_thread(shutdown_sentrystr_logging)
class _ImmutableStaticFiles(StaticFiles):
"""Static files with long Cache-Control for content-hashed Next.js assets.
@@ -260,20 +264,7 @@ app.add_middleware(
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
expose_headers=[
"x-routstr-request-id",
"x-cashu",
"x-routstr-cost-msats",
"x-routstr-cost-usd",
"x-routstr-input-cost-msats",
"x-routstr-output-cost-msats",
# EHBP (Tinfoil) protocol headers must be exposed so browser clients
# can detect and decrypt encrypted responses. Without these, the
# browser hides them via CORS and the SDK treats the response as a
# plaintext proxy error, returning raw ciphertext.
"Ehbp-Response-Nonce",
"Ehbp-Encapsulated-Key",
],
expose_headers=["x-routstr-request-id", "x-cashu"],
)
# Add logging middleware
-65
View File
@@ -1,65 +0,0 @@
from __future__ import annotations
import re
from itertools import count
from typing import Collection
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from .db import UpstreamProviderRow
_SLUG_BASE_PATTERN = re.compile(r"[^a-z0-9]+")
_MAX_SLUG_LENGTH = 64
def provider_slug_base(provider_type: str) -> str:
"""Return a deterministic slug base for a provider type."""
base = _SLUG_BASE_PATTERN.sub("-", provider_type.lower()).strip("-")
if not base:
base = "provider"
elif base.isdigit():
base = f"provider-{base}"
elif len(base) < 3:
base = f"{base}-provider"
if len(base) > _MAX_SLUG_LENGTH:
base = base[:_MAX_SLUG_LENGTH].rstrip("-") or "provider"
return base
def provider_slug_candidate(base: str, suffix_number: int) -> str:
if suffix_number == 1:
return base
suffix = f"-{suffix_number}"
max_base_length = _MAX_SLUG_LENGTH - len(suffix)
return f"{base[:max_base_length].rstrip('-')}{suffix}"
async def allocate_unique_provider_slug(
session: AsyncSession,
provider_type: str,
reserved_slugs: Collection[str] = (),
) -> str:
"""Allocate a stable, deterministic provider slug.
The first provider of a type gets ``openai``; later collisions get
``openai-2``, ``openai-3``, etc. ``reserved_slugs`` covers rows staged in
memory but not flushed yet, such as settings/env seeding.
"""
base = provider_slug_base(provider_type)
reserved = {slug.lower() for slug in reserved_slugs}
for suffix_number in count(1):
candidate = provider_slug_candidate(base, suffix_number)
if candidate in reserved:
continue
result = await session.exec(
select(UpstreamProviderRow).where(UpstreamProviderRow.slug == candidate)
)
if result.first() is None:
return candidate
raise RuntimeError("unreachable")
+47 -239
View File
@@ -3,8 +3,6 @@ from __future__ import annotations
import asyncio
import json
import os
import secrets
import time
from datetime import datetime, timezone
from typing import Any
@@ -18,7 +16,12 @@ class Settings(BaseSettings):
@classmethod
def parse_env_var(cls, field_name: str, raw_value: str) -> Any: # type: ignore[override]
if field_name in {"cashu_mints", "cors_origins", "relays"}:
if field_name in {
"cashu_mints",
"cors_origins",
"relays",
"sentrystr_relays",
}:
v = str(raw_value).strip()
if v == "":
return []
@@ -28,6 +31,7 @@ class Settings(BaseSettings):
# Core
upstream_base_url: str = Field(default="", env="UPSTREAM_BASE_URL")
upstream_api_key: str = Field(default="", env="UPSTREAM_API_KEY")
admin_password: str = Field(default="", env="ADMIN_PASSWORD")
# Node info
name: str = Field(default="ARoutstrNode", env="NAME")
@@ -104,6 +108,21 @@ class Settings(BaseSettings):
log_level: str = Field(default="INFO", env="LOG_LEVEL")
enable_console_logging: bool = Field(default=True, env="ENABLE_CONSOLE_LOGGING")
# SentryStr (decentralized logging to Nostr relays). When enabled, routstr
# application logs at `sentrystr_log_level` and above are mirrored to Nostr
# relays on a background worker thread (see routstr/core/logging.py).
enable_sentrystr: bool = Field(default=False, env="ENABLE_SENTRYSTR")
# Minimum level forwarded to relays. Empty -> defaults to ERROR. Should be
# at or above `log_level`, since the routstr loggers gate lower records.
sentrystr_log_level: str = Field(default="", env="SENTRYSTR_LOG_LEVEL")
# Relays to announce to. Empty -> fall back to the discovery `relays`.
sentrystr_relays: list[str] = Field(default_factory=list, env="SENTRYSTR_RELAYS")
# Secret key (nsec or hex) used to sign log events. Empty -> fall back to
# `nsec`; if that is also empty an ephemeral key is generated per process.
sentrystr_nsec: str = Field(default="", env="SENTRYSTR_NSEC")
# Optional recipient npub. When set, log events are NIP-44 encrypted to it.
sentrystr_recipient_npub: str = Field(default="", env="SENTRYSTR_RECIPIENT_NPUB")
# Other
chat_completions_api_version: str = Field(
default="", env="CHAT_COMPLETIONS_API_VERSION"
@@ -133,70 +152,10 @@ def _normalize_settings_data(data: dict[str, Any]) -> dict[str, Any]:
return normalized
# Secrets are credentials, not config: they live in the encrypted/hashed Secret
# store (and decrypted in-memory for runtime use), never in the persisted
# settings blob. ``admin_password`` is gone from the model entirely; ``nsec``
# remains a live field but is stripped from every blob write so it is never
# written back to plaintext. ``upstream_api_key`` is intentionally *not* here:
# it has no encrypted home yet (it is node-scoped today but really belongs on a
# provider), so stripping it would lose it on the next restart. It stays in the
# blob as before; encrypting it is follow-up work. See ``bootstrap_secrets`` and
# ``routstr.core.vault``.
SECRET_FIELDS = frozenset({"admin_password", "nsec"})
def _strip_secret_fields(data: dict[str, Any]) -> dict[str, Any]:
"""Return a copy of ``data`` without any secret fields (for persistence)."""
return {k: v for k, v in data.items() if k not in SECRET_FIELDS}
def _apply_to_live_settings(data: dict[str, Any]) -> None:
"""Apply ``data`` onto the live ``settings`` for all in-process importers.
Secrets are owned exclusively by ``bootstrap_secrets``, which runs first and
has already decrypted the authoritative nsec into memory (importing any
legacy plaintext on the way). Never re-apply secret fields from env/blob
here: a non-empty but stale ``NSEC`` env var would otherwise override an nsec
the vault has taken ownership of (e.g. after the operator rotates it in the
UI), and an empty one would wipe the live value. Skip them entirely.
"""
for k, v in data.items():
if k in SECRET_FIELDS:
continue
setattr(settings, k, v)
def _compute_primary_mint(cashu_mints: list[str]) -> str:
return cashu_mints[0] if cashu_mints else "https://mint.minibits.cash/Bitcoin"
def derive_npub_from_nsec(nsec: str) -> str | None:
"""Derive the npub (bech32) from an nsec or 64-char hex private key, or None.
Parsing is delegated to :func:`routstr.nostr.listing.nsec_to_keypair`, the
single place that knows the nsec/hex formats (and already returns ``None`` on
any unusable input); this only bech32-encodes the resulting public key. The
contract stays "return None on unusable input", so a bad key never crashes
boot.
"""
try:
from nostr.key import PublicKey # type: ignore
from ..nostr.listing import nsec_to_keypair
except ImportError:
return None
keypair = nsec_to_keypair(nsec)
if keypair is None:
return None
_privkey_hex, pubkey_hex = keypair
try:
return PublicKey(bytes.fromhex(pubkey_hex)).bech32()
except (ValueError, AttributeError):
return None
def resolve_bootstrap() -> Settings:
base = Settings() # Reads env with custom parse_env_var
# Back-compat env mapping
@@ -251,9 +210,23 @@ def resolve_bootstrap() -> Settings:
pass
# Derive NPUB from NSEC if not provided
if not base.npub and base.nsec:
npub = derive_npub_from_nsec(base.nsec)
if npub:
base.npub = npub
try:
from nostr.key import PrivateKey # type: ignore
if base.nsec.startswith("nsec"):
pk = PrivateKey.from_nsec(base.nsec)
elif len(base.nsec) == 64:
pk = PrivateKey(bytes.fromhex(base.nsec))
else:
pk = None
if pk is not None:
try:
base.npub = pk.public_key.bech32()
except Exception:
# Fallback to hex if bech32 not available
base.npub = pk.public_key.hex()
except Exception:
pass
if not base.cors_origins:
base.cors_origins = ["*"]
if not base.primary_mint:
@@ -303,14 +276,15 @@ class SettingsService:
text(
"INSERT INTO settings (id, data, updated_at) VALUES (1, :data, :updated_at)"
).bindparams(
data=json.dumps(_strip_secret_fields(env_resolved.dict())),
data=json.dumps(env_resolved.dict()),
updated_at=datetime.now(timezone.utc),
)
)
await db_session.commit()
cls._current = settings
# Update the existing instance in-place for all live importers
_apply_to_live_settings(env_resolved.dict())
for k, v in env_resolved.dict().items():
setattr(settings, k, v)
return cls._current
db_id, db_data, _updated_at = row
@@ -337,37 +311,20 @@ class SettingsService:
merged_dict.get("cashu_mints", [])
)
# Keep npub consistent with the live nsec. bootstrap_secrets has
# already run and holds the single authoritative nsec (decrypted from
# the encrypted store, or freshly imported). merged_dict starts from
# the env/blob, which may carry a STALE nsec — and therefore a stale
# derived npub — after the vault took ownership. Derive from the live
# value and OVERRIDE, not just fill: otherwise the node keeps the
# vault's private key but announces the old env key's npub (npub is a
# pure derivation of nsec, never configured independently of it).
if settings.nsec:
derived_npub = derive_npub_from_nsec(settings.nsec)
if derived_npub:
merged_dict["npub"] = derived_npub
# Persist without secrets; compare against the stripped target so a
# legacy blob that still carries plaintext secrets gets rewritten
# (and thereby sunset) even when its non-secret values are unchanged.
persisted = _strip_secret_fields(merged_dict)
if db_json_raw != persisted:
if db_json_raw != merged_dict:
await db_session.exec( # type: ignore
text(
"UPDATE settings SET data = :data, updated_at = :updated_at WHERE id = 1"
).bindparams(
data=json.dumps(persisted),
data=json.dumps(merged_dict),
updated_at=datetime.now(timezone.utc),
)
)
await db_session.commit()
# Update the existing instance in-place for all live importers
# (keeps the decrypted nsec live in memory).
_apply_to_live_settings(merged_dict)
for k, v in merged_dict.items():
setattr(settings, k, v)
cls._current = settings
return cls._current
@@ -389,7 +346,7 @@ class SettingsService:
text(
"UPDATE settings SET data = :data, updated_at = :updated_at WHERE id = 1"
).bindparams(
data=json.dumps(_strip_secret_fields(candidate.dict())),
data=json.dumps(candidate.dict()),
updated_at=datetime.now(timezone.utc),
)
)
@@ -418,152 +375,3 @@ class SettingsService:
setattr(settings, k, v)
cls._current = settings
return settings
async def _read_raw_settings_blob(db_session: AsyncSession) -> dict[str, Any]:
"""Best-effort read of the raw persisted settings JSON (may not exist yet)."""
from sqlmodel import text
try:
result = await db_session.exec( # type: ignore
text("SELECT data FROM settings WHERE id = 1")
)
row = result.first()
except Exception:
return {}
if row is None:
return {}
(data_str,) = row
try:
data = json.loads(data_str) if isinstance(data_str, str) else dict(data_str)
except Exception:
return {}
return data if isinstance(data, dict) else {}
def _legacy_plaintext(
raw_blob: dict[str, Any], env_name: str, blob_key: str
) -> str | None:
"""Legacy plaintext for a secret: env first, then the old settings blob."""
env_value = os.environ.get(env_name)
if env_value:
return env_value
blob_value = raw_blob.get(blob_key)
if isinstance(blob_value, str) and blob_value:
return blob_value
return None
async def bootstrap_secrets(db_session: AsyncSession) -> None:
"""Move node secrets into the encrypted/hashed Secret store at startup.
Per secret:
* column already set -> use it (the nsec is decrypted into the in-memory
``settings``; a wrong ROUTSTR_SECRET_KEY surfaces as a clear fail-fast).
* column empty but legacy plaintext exists (env, or the old settings
blob) -> transform it (hash the password / encrypt the nsec) into the
column.
* nothing (admin password only) -> generate a strong random password,
hash it, and log it once with the /admin URL.
"""
from cryptography.fernet import InvalidToken
from sqlmodel import col, update
from . import vault
from .db import NsecState, Secret, get_secret
raw_blob = await _read_raw_settings_blob(db_session)
secret = await get_secret(db_session)
changed = False
# Admin password — one-way scrypt hash.
if secret.admin_password_hash is None:
legacy_password = _legacy_plaintext(
raw_blob, "ADMIN_PASSWORD", "admin_password"
)
if legacy_password:
secret.admin_password_hash = vault.hash_password(legacy_password)
changed = True
else:
generated = secrets.token_urlsafe(24)
# Claim the empty slot atomically: only the worker whose UPDATE flips
# NULL -> hash owns the generated password and announces it. On a
# shared DB a racing worker gets rowcount 0, so it neither clobbers
# the winner's hash (which the operator may already be using) nor
# prints a second password that would never work.
claim_stmt = (
update(Secret)
.where(col(Secret.id) == 1)
.where(col(Secret.admin_password_hash).is_(None))
.values(
admin_password_hash=vault.hash_password(generated),
updated_at=int(time.time()),
)
)
result = await db_session.exec(claim_stmt) # type: ignore[call-overload]
await db_session.commit()
await db_session.refresh(secret)
if result.rowcount == 1:
admin_url = (settings.http_url or "http://localhost:8000").rstrip("/")
# Print to stdout rather than the logger: the operator must see
# this once (e.g. `docker compose logs`), but it must not be
# persisted into the on-disk log files the logger also writes to.
print(
"No admin password set; generated a temporary one (shown "
f"only now): {generated}\nLog in at {admin_url}/admin and "
"change it from the dashboard settings.",
flush=True,
)
# Nostr nsec — reversible Fernet encryption. ``nsec_state`` is the single
# source of truth for ownership, so "intentionally cleared" is never
# conflated with "never migrated" (the bug the old bool could not encode).
if secret.nsec_state == NsecState.encrypted:
# The vault owns the identity: decrypt the ciphertext, never re-read
# env/blob. A missing ciphertext here means the row is inconsistent (a
# failed write or manual edit); fail fast rather than silently dropping
# the identity and falling back to a stale legacy copy.
if secret.encrypted_nsec is None:
raise RuntimeError(
"nsec_state is 'encrypted' but no ciphertext is stored; the "
"secrets row is inconsistent. Refusing to boot rather than "
"silently resurrecting a stale legacy NSEC."
)
try:
settings.nsec = vault.decrypt(secret.encrypted_nsec)
except InvalidToken as exc:
raise RuntimeError(
"Stored nsec cannot be decrypted with the current "
"ROUTSTR_SECRET_KEY. The key changed, or this database came from "
"another node. Restore the original ROUTSTR_SECRET_KEY to recover."
) from exc
elif secret.nsec_state == NsecState.cleared:
# The operator emptied the identity via the admin API. A fresh process
# has already reloaded a stale ``NSEC`` from env/blob into the live
# settings (and may have derived its npub); actively clear both so the
# cleared store wins rather than silently resurrecting the old identity.
settings.nsec = ""
settings.npub = ""
else: # NsecState.legacy — the vault has not taken ownership yet.
# Import any legacy plaintext (env, or the old settings blob) exactly
# once. Encryption at rest is mandatory, but a missing key is
# provisioned, not fatal: vault.encrypt generates and persists a master
# key (with a loud one-time operator notice) when none was supplied, so
# an upgrading node keeps running. The nsec is never stored in plaintext.
legacy_nsec = _legacy_plaintext(raw_blob, "NSEC", "nsec")
if legacy_nsec:
secret.encrypted_nsec = vault.encrypt(legacy_nsec)
secret.nsec_state = NsecState.encrypted
settings.nsec = legacy_nsec
changed = True
# Derive npub from whatever nsec we now hold, if not already known.
if settings.nsec and not settings.npub:
npub = derive_npub_from_nsec(settings.nsec)
if npub:
settings.npub = npub
if changed:
secret.updated_at = int(time.time())
db_session.add(secret)
await db_session.commit()
-320
View File
@@ -1,320 +0,0 @@
"""Encrypt/hash/fingerprint helpers for secrets at rest (issue #553).
Thin wrapper over ``cryptography`` so nothing else in the codebase touches
Fernet/scrypt/HMAC directly:
- :func:`encrypt`/:func:`decrypt` — Fernet symmetric encryption, keyed by the
mandatory master key. Ciphertext is self-describing (``fernet:v1:`` prefix) so
a value can be told apart from legacy plaintext and so reading it under the
wrong key surfaces as a hard error rather than silent corruption.
- :func:`hash_password`/:func:`verify_password` — salted scrypt hashing. This is
*key-independent*: it never reads the master key, so password login and the
recovery script keep working even when the key is missing.
Key custody is flexible but encryption is not optional. The key comes from the
``ROUTSTR_SECRET_KEY`` env var, else a persisted key file
(``ROUTSTR_SECRET_KEY_FILE``, defaulting beside the SQLite database so it persists
on the same volume as the data); when neither is set, :func:`encrypt` generates
one to the key file and prints a one-time notice, so an existing node upgrades
without breaking instead of refusing to boot. Reading is strict — :func:`decrypt`
never generates a key (a new key could not match existing ciphertext) and fails
fast with the generation command when none is configured. A malformed
``ROUTSTR_SECRET_KEY`` is an operator error and always fails fast.
"""
import base64
import hashlib
import hmac
import os
import secrets
import tempfile
from pathlib import Path
from cryptography.fernet import Fernet
from sqlalchemy.engine import make_url
from sqlalchemy.exc import ArgumentError
_PREFIX = "fernet:v1:"
_GEN_COMMAND = (
'python -c "from cryptography.fernet import Fernet; '
'print(Fernet.generate_key().decode())"'
)
# Where an auto-generated master key is persisted when the operator supplies no
# ``ROUTSTR_SECRET_KEY``. Defaults beside the SQLite database so it rides whatever
# volume already persists the data (a container recreate would otherwise generate
# a fresh key and be unable to decrypt existing secrets); falls back to the
# working directory when the DB location is unknown. Override the exact path with
# ``ROUTSTR_SECRET_KEY_FILE``.
_KEY_FILE_ENV = "ROUTSTR_SECRET_KEY_FILE"
_DEFAULT_KEY_FILE = "routstr_secret.key"
# Minimum admin-password length, enforced wherever a password is set/changed
# (admin endpoints + the recovery script) so the policy lives in one place.
MIN_PASSWORD_LENGTH = 8
# scrypt parameters; packed into each hash so verification is parameter-free.
_SCRYPT_N = 2**14
_SCRYPT_R = 8
_SCRYPT_P = 1
_SCRYPT_DKLEN = 32
_SCRYPT_SALT_BYTES = 16
def _database_dir() -> Path | None:
"""Directory of the SQLite database file, or ``None`` when it has no on-disk
location (a non-SQLite URL or ``:memory:``).
Read from ``DATABASE_URL`` at call time and parsed here rather than importing
``routstr.core.db`` — that module builds the engine at import, which the
crypto layer must not drag in. Mirrors db.py's ``DATABASE_URL`` default.
"""
url_str = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db")
try:
url = make_url(url_str)
except ArgumentError:
return None
if url.get_backend_name() != "sqlite" or not url.database:
return None
if url.database == ":memory:":
return None
return Path(url.database).parent
def _key_file_path() -> Path:
"""Where the auto-generated master key is read from / written to.
``ROUTSTR_SECRET_KEY_FILE`` wins; otherwise the key sits beside the SQLite
database so it persists on the same volume as the data, falling back to the
working directory when the DB location is unknown.
"""
override = os.environ.get(_KEY_FILE_ENV)
if override:
return Path(override)
directory = _database_dir()
return (directory or Path()) / _DEFAULT_KEY_FILE
def _read_key_file(path: Path) -> str | None:
try:
stored = path.read_text().strip()
except OSError:
return None
if not stored:
return None
_repair_key_file_perms(path)
return stored
def _repair_key_file_perms(path: Path) -> None:
# A master key must never be group/other-readable. On POSIX, tighten loose
# permissions to owner-only (0600) rather than trust — or hard-fail on — a
# world-readable key; a friendlier repair keeps an upgrading node booting.
if os.name != "posix":
return
try:
mode = path.stat().st_mode
except OSError:
return
if mode & 0o077:
try:
os.chmod(path, 0o600)
except OSError:
pass
def _load_secret_key() -> str | None:
"""The configured key without provisioning: env var, then the key file."""
return os.environ.get("ROUTSTR_SECRET_KEY") or _read_key_file(_key_file_path())
def _warn_generated_key(path: Path) -> None:
# stdout, not the logger: the operator must see this once (e.g. in
# ``docker compose logs``), but it must never be persisted into the on-disk
# log files the logger also writes. Mirrors the generated-admin-password
# notice so an upgrade cannot silently create an unbacked key.
print(
"No ROUTSTR_SECRET_KEY was set; generated one to encrypt node secrets at "
f"rest and saved it to {path}.\n"
"!! BACK UP THIS FILE. If it is lost, the encrypted secrets cannot be "
"recovered and will have to be re-entered.\n"
"To manage the key yourself (e.g. from a secrets manager) set "
"ROUTSTR_SECRET_KEY in the environment instead; the value is in the file "
"above.",
flush=True,
)
def _generate_and_persist_key(path: Path) -> str:
"""Generate a Fernet key, persist it owner-only and atomically, warn once.
The key is written to a temp file in the same directory, flushed durably,
then ``os.link``-ed into place. ``os.link`` publishes the complete file in a
single atomic step — a crash mid-write leaves only the temp file (which is
removed), never a half-written or empty key at the final path that a later
boot would read as corrupt. It also refuses to overwrite an existing key, so
a racing worker that generated first keeps ownership (secrets may already be
encrypted under its key); the loser adopts that key instead of clobbering it.
"""
key = Fernet.generate_key().decode()
path.parent.mkdir(parents=True, exist_ok=True)
fd, tmp_name = tempfile.mkstemp(
dir=path.parent, prefix=".routstr_secret.", suffix=".tmp"
)
tmp = Path(tmp_name)
try:
with os.fdopen(fd, "w") as handle:
handle.write(key) # mkstemp already created it 0600
handle.flush()
os.fsync(handle.fileno())
try:
os.link(tmp, path)
except FileExistsError:
# A concurrent worker linked its key in first; adopt theirs rather
# than clobber a key that secrets may already be encrypted under.
existing = _read_key_file(path)
if existing:
return existing
raise
_fsync_dir(path.parent)
finally:
tmp.unlink(missing_ok=True)
_warn_generated_key(path)
return key
def _fsync_dir(directory: Path) -> None:
# Persist the new directory entry so the linked key survives a crash right
# after publish. Best-effort: not every platform lets you fsync a directory.
try:
dir_fd = os.open(directory, os.O_RDONLY)
except OSError:
return
try:
os.fsync(dir_fd)
except OSError:
pass
finally:
os.close(dir_fd)
def ensure_secret_key() -> str:
"""Return the master key, provisioning one if the operator supplied none.
Precedence: the ``ROUTSTR_SECRET_KEY`` env var, then the persisted key file,
otherwise a freshly generated key written to the key file (with a one-time
operator notice). This keeps encryption at rest mandatory while letting an
existing node upgrade without setting a key first. A malformed env key is
left to fail at :func:`get_fernet` — it is an operator error, not an unset
key, so it must not trigger silent self-provisioning.
"""
env_key = os.environ.get("ROUTSTR_SECRET_KEY")
if env_key:
return env_key
path = _key_file_path()
return _read_key_file(path) or _generate_and_persist_key(path)
def _fernet_from_key(key: str) -> Fernet:
try:
return Fernet(key.encode())
except (ValueError, TypeError) as exc:
raise RuntimeError(
"ROUTSTR_SECRET_KEY is malformed; it must be a url-safe base64 "
"32-byte Fernet key. Generate one with:\n " + _GEN_COMMAND
) from exc
def get_fernet() -> Fernet:
"""Build a :class:`Fernet` from the configured key (env var or key file).
Strict: this never generates a key, so already-encrypted ciphertext is never
shadowed by a fresh key. A read with no key configured fails fast with the
generation command.
"""
key = _load_secret_key()
if not key:
raise RuntimeError(
"ROUTSTR_SECRET_KEY is not set. It is required to encrypt secrets at "
"rest. Generate one with:\n " + _GEN_COMMAND
)
return _fernet_from_key(key)
def encrypt(plaintext: str) -> str:
"""Encrypt ``plaintext`` into a self-describing ``fernet:v1:`` token.
Provisions a master key (env var, key file, or a freshly generated one) so an
upgrading node never has to set one before its first secret is stored; the
value is always encrypted, never persisted in plaintext.
"""
fernet = _fernet_from_key(ensure_secret_key())
return _PREFIX + fernet.encrypt(plaintext.encode()).decode()
def is_encrypted(value: str) -> bool:
"""True if ``value`` carries the ``fernet:v1:`` prefix this module emits."""
return value.startswith(_PREFIX)
def decrypt(ciphertext: str) -> str:
"""Decrypt a ``fernet:v1:`` token.
Raises ``ValueError`` for an unprefixed value (so legacy plaintext is never
mistaken for ciphertext) and ``InvalidToken`` when the value was written
under a different ``ROUTSTR_SECRET_KEY``.
"""
if not is_encrypted(ciphertext):
raise ValueError("value is not fernet:v1: ciphertext")
token = ciphertext[len(_PREFIX) :]
return get_fernet().decrypt(token.encode()).decode()
def hash_password(password: str) -> str:
"""Salted scrypt hash, self-describing as ``scrypt:n:r:p:salt:hash``."""
salt = secrets.token_bytes(_SCRYPT_SALT_BYTES)
derived = hashlib.scrypt(
password.encode(),
salt=salt,
n=_SCRYPT_N,
r=_SCRYPT_R,
p=_SCRYPT_P,
dklen=_SCRYPT_DKLEN,
)
return ":".join(
[
"scrypt",
str(_SCRYPT_N),
str(_SCRYPT_R),
str(_SCRYPT_P),
base64.b64encode(salt).decode(),
base64.b64encode(derived).decode(),
]
)
def verify_password(password: str, stored: str) -> bool:
"""Constant-time check of ``password`` against a :func:`hash_password` value."""
try:
scheme, n, r, p, salt_b64, hash_b64 = stored.split(":")
if scheme != "scrypt":
return False
n_int, r_int, p_int = int(n), int(r), int(p)
# Cap the work factor at the parameters this module emits. scrypt's
# memory cost grows with N*r, so an oversized N/r in a tampered or
# corrupt stored hash could turn a single login into an OOM/DoS.
if n_int > _SCRYPT_N or r_int > _SCRYPT_R or p_int > _SCRYPT_P:
return False
salt = base64.b64decode(salt_b64)
expected = base64.b64decode(hash_b64)
derived = hashlib.scrypt(
password.encode(),
salt=salt,
n=n_int,
r=r_int,
p=p_int,
dklen=len(expected),
)
except (ValueError, TypeError):
return False
return hmac.compare_digest(derived, expected)
+1 -1
View File
@@ -20,7 +20,7 @@ import subprocess
from functools import lru_cache
from pathlib import Path
BASE_VERSION = "0.4.4"
BASE_VERSION = "0.4.3"
_REPO_ROOT = Path(__file__).resolve().parents[2]
_GIT_TIMEOUT_SECONDS = 2.0
+56 -187
View File
@@ -1,5 +1,4 @@
import math
from typing import TYPE_CHECKING
from pydantic.v1 import BaseModel
@@ -8,9 +7,6 @@ from ..core.settings import settings
from .price import sats_usd_price
from .usage import normalize_usage, parse_token_count
if TYPE_CHECKING:
from .models import Model
__all__ = [
"CostData",
"CostDataError",
@@ -45,48 +41,15 @@ class CostDataError(BaseModel):
code: str
def _empty_cost(cls: type[CostData] = CostData) -> CostData:
"""Build an all-zero cost object — a full refund for an empty response.
Shared by the two paths that must not bill: an upstream response with no
usage data at all, and one that reports a USD cost but carries zero tokens
in every bucket.
"""
return cls(
base_msats=0,
input_msats=0,
output_msats=0,
total_msats=0,
total_usd=0.0,
input_tokens=0,
output_tokens=0,
cache_read_input_tokens=0,
cache_creation_input_tokens=0,
cache_read_msats=0,
cache_creation_msats=0,
)
async def calculate_cost(
response_data: dict,
max_cost: int,
model_obj: "Model | None" = None,
provider_fee: float | None = None,
) -> CostData | MaxCostData | CostDataError:
"""Calculate the cost of an API request based on token usage.
Args:
response_data: Response data containing usage information
max_cost: Maximum cost in millisats
model_obj: The model that actually served the request. When given,
its pricing is billed directly; without it, pricing is re-derived
from the response's model string via the alias map, which resolves
to the best-ranked candidate — not necessarily the serving one.
provider_fee: The serving provider's fee multiplier, applied on the
USD-cost path and the litellm pricing fallback (configured model
pricing already carries the fee baked in). Without it, the fee is
re-derived from the response's model string, which yields the
best-ranked provider's fee.
Returns:
Cost data or error information
@@ -120,7 +83,19 @@ async def calculate_cost(
else None,
},
)
return _empty_cost(MaxCostData)
return MaxCostData(
base_msats=0,
input_msats=0,
output_msats=0,
total_msats=0,
total_usd=0.0,
input_tokens=0,
output_tokens=0,
cache_read_input_tokens=0,
cache_creation_input_tokens=0,
cache_read_msats=0,
cache_creation_msats=0,
)
usage_data = response_data.get("usage") or {}
if not isinstance(usage_data, dict):
@@ -134,27 +109,6 @@ async def calculate_cost(
# Try USD cost first
usd_cost = _resolve_usd_cost(usage_data, response_data)
if usd_cost > 0:
truly_empty = (
input_tokens == 0
and output_tokens == 0
and cache_read_tokens == 0
and cache_creation_tokens == 0
)
if truly_empty:
logger.warning(
"Upstream reported a USD cost but the response carries no "
"tokens at all (input, output, cache-read and cache-creation "
"are all zero) — refunding in full rather than billing the "
"USD-derived cost for an empty response.",
extra={
"model": response_data.get("model", "unknown"),
"usd_cost": usd_cost,
"usage_keys": sorted(usage_data.keys())
if isinstance(usage_data, dict)
else None,
},
)
return _empty_cost()
if input_tokens == 0 and output_tokens == 0:
logger.warning(
"Upstream reported a USD cost but no token counts — "
@@ -172,16 +126,11 @@ async def calculate_cost(
},
)
try:
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")
usage_data.get("cost_details", {}).get("input_cost", 0)
)
output_usd = _coerce_usd(
cost_details.get("output_cost")
or cost_details.get("upstream_inference_completions_cost")
usage_data.get("cost_details", {}).get("output_cost", 0)
)
return _calculate_from_usd_cost(
usd_cost,
@@ -192,7 +141,6 @@ async def calculate_cost(
cache_creation_tokens,
output_tokens,
response_data,
provider_fee,
)
except Exception as e:
logger.warning(
@@ -206,7 +154,7 @@ async def calculate_cost(
# Fall back to token-based pricing
try:
pricing_rates = _get_pricing_rates(response_data, model_obj, provider_fee)
pricing_rates = _get_pricing_rates(response_data)
except ValueError as e:
return CostDataError(message=str(e), code="pricing_error")
@@ -280,17 +228,7 @@ def _coerce_usd(value: object) -> float:
def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
"""Resolve USD cost with clear priority order.
Priority:
1. ``cost_details.total_cost``
2. ``cost_details.upstream_inference_cost`` (BYOK — see below)
3. ``total_cost`` → ``cost`` (in both usage and response)
**BYOK path (PPQ.AI):** when ``is_byok`` is true the ``usage.cost`` field
is only a small (~5 %) routing fee, not the inference cost. The real cost
lives in ``cost_details.upstream_inference_cost`` and the provider's
balance is debited by ``upstream_inference_cost + byok_fee``. Billing just
the fee under-charges by ~20×.
Priority: cost_details.total_cost → total_cost → cost (in both usage and response).
"""
cost_details = usage_data.get("cost_details")
if isinstance(cost_details, dict):
@@ -298,18 +236,6 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
if cost > 0:
return cost
# PPQ.AI BYOK: upstream_inference_cost is the real inference cost;
# usage.cost is only a ~5 % BYOK routing fee. Bill the sum — what PPQ
# 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")
)
if upstream_cost > 0 and usage_data.get("is_byok"):
byok_fee = _coerce_usd(usage_data.get("cost"))
return upstream_cost + byok_fee
for source in [usage_data, response_data]:
if not isinstance(source, dict):
continue
@@ -323,106 +249,55 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
def _get_pricing_rates(
response_data: dict,
model_obj: "Model | None",
provider_fee: float | None,
) -> tuple[float, float, float, float] | None:
"""Get configured rates, falling back to LiteLLM's model cost map.
"""Get model-based pricing rates or None if using fixed pricing.
The served ``model_obj`` (when the caller has it) is billed directly;
otherwise the response's model string is resolved through the alias map,
which yields the best-ranked candidate rather than the serving one.
Returns: (input_rate, output_rate, cache_read_rate, cache_write_rate).
``None`` means configured fixed pricing should be used by the caller.
Returns: (input_rate, output_rate, cache_read_rate, cache_write_rate)
"""
if settings.fixed_pricing and (
settings.fixed_per_1k_input_tokens
or settings.fixed_per_1k_output_tokens
):
if settings.fixed_pricing:
return None
from ..proxy import get_model_instance
from .models import litellm_cost_entry
response_model = response_data.get("model", "")
if model_obj is None:
logger.warning(
"Settling without routed model identity — re-deriving pricing "
"from the response's model string via the alias map",
extra={"response_model": response_model},
)
model_obj = get_model_instance(response_model)
model_obj = get_model_instance(response_model)
if model_obj and model_obj.sats_pricing:
try:
mspp = float(model_obj.sats_pricing.prompt)
mspc = float(model_obj.sats_pricing.completion)
mscr = float(model_obj.sats_pricing.input_cache_read or 0)
mscw = float(model_obj.sats_pricing.input_cache_write or 0)
if not model_obj:
logger.error("Invalid model in response", extra={"response_model": response_model})
raise ValueError(f"Invalid model: {response_model}")
mspp_1k = mspp * 1_000_000.0
mspc_1k = mspc * 1_000_000.0
mscr_1k = mscr * 1_000_000.0 if mscr > 0 else mspp_1k
mscw_1k = mscw * 1_000_000.0 if mscw > 0 else mspp_1k
source = "configured"
except Exception as e:
logger.error("Invalid pricing data", extra={"error": str(e)})
raise ValueError("Invalid pricing data") from e
else:
pricing_model = (
model_obj.forwarded_model_id if model_obj else None
) or response_model
pricing = litellm_cost_entry(pricing_model)
if pricing is None:
logger.error(
"Model pricing not found in configured models or LiteLLM",
extra={
"response_model": response_model,
"pricing_model": pricing_model,
},
)
raise ValueError(f"Pricing not found for model: {response_model}")
if not model_obj.sats_pricing:
logger.error(
"Model pricing not defined",
extra={"model": response_model, "model_id": response_model},
)
raise ValueError("Model pricing not defined")
input_usd = _coerce_usd(pricing.get("input_cost_per_token"))
output_usd = _coerce_usd(pricing.get("output_cost_per_token"))
if input_usd <= 0 or output_usd <= 0:
raise ValueError(f"Incomplete LiteLLM pricing for model: {pricing_model}")
try:
mspp = float(model_obj.sats_pricing.prompt)
mspc = float(model_obj.sats_pricing.completion)
mscr = float(model_obj.sats_pricing.input_cache_read or 0)
mscw = float(model_obj.sats_pricing.input_cache_write or 0)
if provider_fee is None:
provider_fee = _resolve_provider_fee(response_model)
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")
)
mscr_1k = (
cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat
if cache_read_usd > 0
else mspp_1k
)
mscw_1k = (
cache_write_usd * provider_fee * 1_000_000.0 / usd_per_sat
if cache_write_usd > 0
else mspp_1k
)
source = "litellm"
mspp_1k = mspp * 1_000_000.0
mspc_1k = mspc * 1_000_000.0
mscr_1k = mscr * 1_000_000.0 if mscr > 0 else mspp_1k
mscw_1k = mscw * 1_000_000.0 if mscw > 0 else mspp_1k
logger.info(
"Applied model-specific pricing",
extra={
"model": response_model,
"pricing_source": source,
"input_price_msats_per_1k": mspp_1k,
"output_price_msats_per_1k": mspc_1k,
"cache_read_price_msats_per_1k": mscr_1k,
"cache_write_price_msats_per_1k": mscw_1k,
},
)
return mspp_1k, mspc_1k, mscr_1k, mscw_1k
logger.info(
"Applied model-specific pricing",
extra={
"model": response_model,
"input_price_msats_per_1k": mspp_1k,
"output_price_msats_per_1k": mspc_1k,
"cache_read_price_msats_per_1k": mscr_1k,
"cache_write_price_msats_per_1k": mscw_1k,
},
)
return mspp_1k, mspc_1k, mscr_1k, mscw_1k
except Exception as e:
logger.error("Invalid pricing data", extra={"error": str(e)})
raise ValueError("Invalid pricing data") from e
def _resolve_provider_fee(model_id: str) -> float:
@@ -450,11 +325,9 @@ def _calculate_from_usd_cost(
cache_creation_tokens: int,
output_tokens: int,
response_data: dict,
provider_fee: float | None,
) -> CostData:
"""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", ""))
provider_fee = _resolve_provider_fee(response_data.get("model", ""))
usd_cost = usd_cost * provider_fee
input_usd = input_usd * provider_fee
output_usd = output_usd * provider_fee
@@ -463,12 +336,8 @@ def _calculate_from_usd_cost(
cost_in_msats = math.ceil(cost_in_sats * 1000)
if input_usd > 0 or output_usd > 0:
# The total is the authoritative billed amount. Allocating that integer
# total proportionally avoids losing sub-millisatoshi remainders when
# input and output components are each truncated independently.
component_usd = input_usd + output_usd
input_msats = math.floor(cost_in_msats * input_usd / component_usd)
output_msats = cost_in_msats - input_msats
input_msats = int((input_usd * sats_per_usd) * 1000)
output_msats = int((output_usd * sats_per_usd) * 1000)
else:
effective_input_tokens = (
input_tokens + cache_read_tokens + cache_creation_tokens
+6 -20
View File
@@ -1,17 +1,11 @@
from __future__ import annotations
import asyncio
import math
from typing import TypedDict
import httpx
from cashu.wallet.wallet import Proof, Wallet
# The Cashu library issues POST /v1/melt/bolt11 with timeout=None, so a hung or
# very slow mint can block a melt (and any caller, e.g. the payout loop)
# indefinitely. Bound it here so callers fail instead of hanging forever.
MELT_TIMEOUT_SECONDS = 60
try:
from bech32 import bech32_decode, convertbits # type: ignore
except ModuleNotFoundError: # pragma: no cover allow runtime miss
@@ -226,18 +220,10 @@ async def raw_send_to_lnurl(
if amount:
proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True)
try:
_ = await asyncio.wait_for(
wallet.melt(
proofs=proofs,
invoice=bolt11_invoice,
fee_reserve_sat=melt_quote_resp.fee_reserve,
quote_id=melt_quote_resp.quote,
),
timeout=MELT_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError as e:
raise LNURLError(
f"Melt timed out after {MELT_TIMEOUT_SECONDS}s (mint unresponsive)"
) from e
_ = await wallet.melt(
proofs=proofs,
invoice=bolt11_invoice,
fee_reserve_sat=melt_quote_resp.fee_reserve,
quote_id=melt_quote_resp.quote,
)
return final_amount
+14 -46
View File
@@ -85,30 +85,6 @@ class Model(BaseModel):
return hash(self.id)
def litellm_cost_entry(model_id: str) -> dict | None:
"""Look up ``model_id`` in litellm's bundled cost map.
litellm ships per-model USD rates keyed by the exact OpenRouter id
(``deepseek/deepseek-chat``) or the bare model name (``gpt-4o``,
``claude-sonnet-4-5``), so both spellings are tried. Keys are lowercase, so
a mixed-case upstream id (``deepseek-ai/DeepSeek-V4-Flash``) is retried via
a case-insensitive scan. Returns the matched cost dict, or ``None``.
"""
import litellm
candidates = (model_id, model_id.split("/", 1)[-1])
for key in candidates:
info = litellm.model_cost.get(key)
if isinstance(info, dict):
return info
lowered = {c.lower() for c in candidates}
for key, info in litellm.model_cost.items():
if isinstance(key, str) and key.lower() in lowered and isinstance(info, dict):
return info
return None
def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing:
"""Fill missing cache rates from litellm's bundled cost map.
@@ -116,8 +92,9 @@ def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing:
for many models (most DeepSeek entries, openai/gpt-4o, ...). Without a
cache rate, billing falls back to the full input rate, which overcharges
cache reads (DeepSeek hits are 10x cheaper) and undercharges Anthropic
cache writes (1.25x). The lookup (see ``litellm_cost_entry``) tries both
id spellings and a case-insensitive fallback.
cache writes (1.25x). litellm ships per-model USD rates keyed by the exact
OpenRouter id (deepseek/deepseek-chat) or by the bare model name
(gpt-4o, claude-sonnet-4-5), so both spellings are tried.
Rates already present (e.g. provided by OpenRouter) are authoritative and
never overwritten. Unknown models are returned unchanged.
@@ -127,7 +104,14 @@ def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing:
if not (needs_read or needs_write):
return pricing
info = litellm_cost_entry(model_id)
import litellm
info: dict | None = None
for key in (model_id, model_id.split("/", 1)[-1]):
candidate = litellm.model_cost.get(key)
if isinstance(candidate, dict):
info = candidate
break
if info is None:
return pricing
@@ -231,29 +215,13 @@ def _row_to_model(
)
top_provider_dict = json.loads(row.top_provider) if row.top_provider else None
if apply_provider_fee and isinstance(pricing, dict):
pricing = {k: float(v) * provider_fee for k, v in pricing.items()}
if isinstance(pricing, dict) and float(pricing.get("request", 0.0)) <= 0.0:
pricing["request"] = max(pricing.get("request", 0.0), 0.0)
parsed_pricing = Pricing.parse_obj(pricing)
# Fill missing cache-read/write rates from litellm's cost map BEFORE applying
# the provider fee, so they carry the same markup as every other component.
# DB-stored override pricing (e.g. generic providers) omits cache rates;
# without this, ``_row_to_model`` bills cache reads at the full input rate —
# the ``_apply_provider_fee_to_model`` path backfills, but the override path
# used for admin-configured providers did not.
#
# Key on ``forwarded_model_id`` (the actual upstream model name litellm
# prices) when set: an alias row (id="local-alias",
# forwarded_model_id="deepseek-v4-flash") would otherwise look up the alias
# and miss the cache rate.
pricing_model_id = getattr(row, "forwarded_model_id", None) or row.id
parsed_pricing = backfill_cache_pricing(pricing_model_id, parsed_pricing)
if apply_provider_fee:
parsed_pricing = Pricing.parse_obj(
{k: float(v) * provider_fee for k, v in parsed_pricing.dict().items()}
)
model = Model(
id=row.id,
name=row.name,
+70 -232
View File
@@ -7,13 +7,7 @@ from fastapi.responses import Response, StreamingResponse
from sqlmodel import select
from .algorithm import create_model_mappings
from .auth import (
ReservationSnapshot,
get_reservation_snapshot,
pay_for_request,
revert_pay_for_request,
validate_bearer_key,
)
from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key
from .core import get_logger
from .core.db import (
ApiKey,
@@ -35,7 +29,6 @@ from .payment.helpers import (
)
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.request_correction import correct_request, extract_error_message
@@ -43,9 +36,10 @@ logger = get_logger(__name__)
proxy_router = APIRouter()
_upstreams: list[BaseUpstreamProvider] = []
_model_instances: dict[str, Model] = {} # All aliases -> Model
_provider_map: dict[
str, list[tuple[Model, BaseUpstreamProvider]]
] = {} # All aliases -> sorted [(candidate Model, its Provider)]
str, list[BaseUpstreamProvider]
] = {} # All aliases -> List[Provider]
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
@@ -77,44 +71,32 @@ def get_upstreams() -> list[BaseUpstreamProvider]:
return _upstreams
def get_candidates(
model_id: str,
) -> list[tuple[Model, BaseUpstreamProvider]] | None:
"""Get the sorted (model, provider) candidate list for a model ID.
Each provider is paired with its own model for the alias, so routing can
forward and bill the candidate that actually serves. Version suffixes
(e.g. ``-20251222``) are stripped as a retry when the exact ID is
unknown, since upstreams may return a specific version of a base model
we track.
"""
def get_model_instance(model_id: str) -> Model | None:
"""Get Model instance by ID from global cache."""
if not model_id:
return None
model_id_lower = model_id.lower()
if candidates := _provider_map.get(model_id_lower):
return candidates
# Try exact match first
if model := _model_instances.get(model_id_lower):
return model
# Try stripping common version suffixes (e.g., -20251222)
# This handles cases where upstream returns a specific version
# but we only track the base model name.
import re
base_model_id = re.sub(r"-\d{8}$", "", model_id_lower)
if base_model_id != model_id_lower:
if candidates := _provider_map.get(base_model_id):
return candidates
if model := _model_instances.get(base_model_id):
return model
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)
return candidates[0][0] if candidates else None
def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None:
"""Get the sorted UpstreamProvider list for a model ID."""
candidates = get_candidates(model_id)
return [provider for _, provider in candidates] if candidates else None
"""Get UpstreamProvider list for model ID from global cache."""
return _provider_map.get(model_id.lower())
def get_unique_models() -> list[Model]:
@@ -122,39 +104,11 @@ def get_unique_models() -> list[Model]:
return list(_unique_models.values())
def _is_tinfoil_attestation_path(path: str) -> bool:
"""Return True for exact Tinfoil attestation routes, with optional slash."""
return path in {
"attestation",
"attestation/",
"tee/attestation",
"tee/attestation/",
}
def _select_unauthenticated_get_upstreams(
path: str, upstreams: list[BaseUpstreamProvider]
) -> list[BaseUpstreamProvider]:
"""Select upstream candidates for unauthenticated GET bypass paths.
Tinfoil attestation endpoints are provider-specific. Trying every enabled
upstream can return an unrelated provider's 404 before Tinfoil is reached,
so route those paths only to Tinfoil providers.
"""
if _is_tinfoil_attestation_path(path):
return [
upstream
for upstream in upstreams
if getattr(upstream, "provider_type", None) == "tinfoil"
]
return upstreams
async def refresh_model_maps() -> None:
"""Refresh global model and provider maps using the cost-based algorithm."""
from sqlalchemy.orm import selectinload
global _provider_map, _unique_models
global _model_instances, _provider_map, _unique_models
async with create_session() as session:
# Fetch all providers with their models in a single logical operation
@@ -164,23 +118,22 @@ async def refresh_model_maps() -> None:
result = await session.exec(query)
provider_rows = result.all()
overrides_by_key: dict[tuple[str, int], tuple[ModelRow, float]] = {}
disabled_model_keys: set[tuple[str, int]] = set()
overrides_by_id: dict[str, tuple[ModelRow, float]] = {}
disabled_model_ids: set[str] = set()
for provider in provider_rows:
if not provider.enabled:
continue
for model in provider.models:
model_key = (model.id.lower(), model.upstream_provider_id)
if model.enabled:
overrides_by_key[model_key] = (model, provider.provider_fee)
overrides_by_id[model.id] = (model, provider.provider_fee)
else:
disabled_model_keys.add(model_key)
disabled_model_ids.add(model.id)
_, _provider_map, _unique_models = create_model_mappings(
_model_instances, _provider_map, _unique_models = create_model_mappings(
upstreams=_upstreams,
overrides_by_key=overrides_by_key,
disabled_model_keys=disabled_model_keys,
overrides_by_id=overrides_by_id,
disabled_model_ids=disabled_model_ids,
)
@@ -213,7 +166,6 @@ _API_PATH_PREFIXES = (
"moderations",
"providers",
"tee/",
"attestation",
)
@@ -232,54 +184,20 @@ async def proxy(
is_responses_api = path.startswith("v1/responses") or path.startswith("responses")
request_body = await request.body()
request_body_dict = parse_request_body_json(request_body, path)
# EHBP (Encrypted HTTP Body Protocol) requests carry an Ehbp-Encapsulated-Key
# header and a binary HPKE-sealed body. The proxy cannot parse the body to
# 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", "")
if not model_id:
return create_error_response(
"invalid_request",
"EHBP request missing X-Routstr-Model header",
400,
request=request,
)
else:
request_body_dict = parse_request_body_json(request_body, path)
if is_responses_api:
model_id = extract_model_from_responses_request(request_body_dict)
else:
model_id = request_body_dict.get("model", "unknown")
# Exact Tinfoil attestation GET routes don't map to models — forward
# 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):
selected_upstreams = _select_unauthenticated_get_upstreams(path, _upstreams)
if not selected_upstreams:
return create_error_response(
"upstream_error",
"No upstream available for unauthenticated GET path",
502,
request=request,
)
# /tee/* GET requests (e.g. attestation) don't map to models — just
# forward to all enabled upstreams without model/cost/auth lookups.
if request.method == "GET" and path.startswith("tee/"):
all_upstreams = _upstreams
last_error_response = None
for i, upstream in enumerate(selected_upstreams):
for i, upstream in enumerate(all_upstreams):
try:
headers = upstream.prepare_headers(dict(request.headers))
response = await upstream.forward_get_request(request, path, headers)
if (
response.status_code in [502, 429]
and i < len(selected_upstreams) - 1
):
if response.status_code in [502, 429] and i < len(all_upstreams) - 1:
logger.warning(
"Upstream %s returned %s for unauthenticated GET %s, trying next",
"Upstream %s returned %s for tee GET %s, trying next",
upstream.provider_type,
response.status_code,
path,
@@ -288,43 +206,42 @@ async def proxy(
return response
except UpstreamError as e:
logger.warning(
"Upstream %s failed for unauthenticated GET %s: %s",
"Upstream %s failed for tee GET %s: %s",
upstream.provider_type,
path,
e,
)
if i == len(selected_upstreams) - 1:
if i == len(all_upstreams) - 1:
last_error_response = create_upstream_error_response(e, request)
continue
return last_error_response or create_error_response(
"upstream_error", "All upstreams failed", 502, request=request
)
candidates = get_candidates(model_id)
if is_responses_api:
model_id = extract_model_from_responses_request(request_body_dict)
else:
model_id = request_body_dict.get("model", "unknown")
if not candidates:
model_obj = get_model_instance(model_id)
if not model_obj:
return create_error_response(
"invalid_model", f"Model '{model_id}' not found", 400, request=request
)
if is_ehbp:
candidates = [
(model, upstream)
for model, upstream in candidates
if upstream.supports_ehbp
]
if not candidates:
return create_error_response(
"unsupported_request",
f"No EHBP-capable provider found for model '{model_id}'",
400,
request=request,
)
upstreams = get_provider_for_model(model_id)
if not upstreams:
return create_error_response(
"invalid_model",
f"No provider found for model '{model_id}'",
400,
request=request,
)
# Reserve/max-cost checks use the best-ranked candidate; the failover loop
# below rebinds (model_obj, upstream) per candidate so forwarding and
# settlement always use the model of the provider actually being tried.
model_obj = candidates[0][0]
# todo figure out cost calculation since fallback provider is usually not the same price
# Use first provider for initial checks/cost calculation
# primary_upstream = upstreams[0]
_max_cost_for_model = await get_max_cost_for_model(
model=model_id, session=session, model_obj=model_obj
@@ -339,25 +256,9 @@ async def proxy(
if x_cashu := headers.get("x-cashu", None):
last_error = None
for i, (model_obj, upstream) in enumerate(candidates):
for i, upstream in enumerate(upstreams):
try:
if is_ehbp:
if not upstream.supports_ehbp:
logger.warning(
"Upstream %s does not support EHBP for model=%s",
upstream.provider_type,
model_id,
)
continue
return await forward_ehbp_x_cashu_request(
request=request,
x_cashu_token=x_cashu,
path=path,
max_cost_for_model=max_cost_for_model,
model_obj=model_obj,
upstream=upstream,
)
elif is_responses_api:
if is_responses_api:
return await upstream.handle_x_cashu_responses(
request, x_cashu, path, max_cost_for_model, model_obj
)
@@ -377,7 +278,7 @@ async def proxy(
"status_code": e.status_code,
},
)
if i == len(candidates) - 1:
if i == len(upstreams) - 1:
last_error = e
continue
@@ -404,12 +305,12 @@ async def proxy(
logger.debug("Processing unauthenticated GET request", extra={"path": path})
last_error_response = None
for i, (_, upstream) in enumerate(candidates):
for i, upstream in enumerate(upstreams):
try:
headers = upstream.prepare_headers(dict(request.headers))
response = await upstream.forward_get_request(request, path, headers)
if response.status_code in [502, 429] and i < len(candidates) - 1:
if response.status_code in [502, 429] and i < len(upstreams) - 1:
error_message = ""
try:
if hasattr(response, "body"):
@@ -439,81 +340,28 @@ async def proxy(
return response
except UpstreamError as e:
logger.warning(f"Upstream {upstream.provider_type} failed (GET): {e}")
if i == len(candidates) - 1:
if i == len(upstreams) - 1:
last_error_response = create_upstream_error_response(e, request)
continue
return last_error_response or create_error_response(
"upstream_error", "All upstreams failed", 502, request=request
)
reservation_snapshot: ReservationSnapshot | None = None
if is_ehbp or request_body_dict:
if request_body_dict:
await pay_for_request(key, max_cost_for_model, session)
reservation_snapshot = await get_reservation_snapshot(key, session)
# Tracks request params already removed in response to upstream rejections,
# shared across providers so a stripped param stays stripped on failover and
# the reactive retry can never loop unboundedly.
already_stripped: set[str] = set()
for i, (model_obj, upstream) in enumerate(candidates):
if i > 0 and request_body_dict:
# The reservation was sized to the previous candidate's envelope;
# settlement bills the serving candidate, so a pricier fallback
# must be re-reserved at its own max cost before it is tried. A
# candidate whose envelope the key cannot cover is rejected, just
# as it would be had it been ranked first.
candidate_max = await get_max_cost_for_model(
model=model_id, session=session, model_obj=model_obj
)
candidate_max = await calculate_discounted_max_cost(
candidate_max, request_body_dict, model_obj=model_obj
)
candidate_max = max(candidate_max, settings.min_request_msat)
if candidate_max > max_cost_for_model:
await revert_pay_for_request(
key, session, max_cost_for_model, reservation_snapshot
)
try:
await pay_for_request(key, candidate_max, session)
except HTTPException:
if i == len(candidates) - 1:
raise
await pay_for_request(key, max_cost_for_model, session)
reservation_snapshot = await get_reservation_snapshot(key, session)
continue
reservation_snapshot = await get_reservation_snapshot(key, session)
max_cost_for_model = candidate_max
for i, upstream in enumerate(upstreams):
headers = upstream.prepare_headers(dict(request.headers))
try:
while True:
try:
if is_ehbp:
if not upstream.supports_ehbp:
logger.warning(
"Upstream %s does not support EHBP for model=%s",
upstream.provider_type,
model_id,
)
raise UpstreamError(
f"Provider {upstream.provider_type} does not support EHBP",
status_code=400,
)
response = await forward_ehbp_request(
request=request,
path=path,
headers=headers,
request_body=request_body,
upstream=upstream,
key=key,
max_cost_for_model=max_cost_for_model,
session=session,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
)
elif is_responses_api:
if is_responses_api:
response = await upstream.forward_responses_request(
request,
path,
@@ -523,7 +371,6 @@ async def proxy(
max_cost_for_model,
session,
model_obj,
reservation_snapshot,
)
else:
response = await upstream.forward_request(
@@ -535,7 +382,6 @@ async def proxy(
max_cost_for_model,
session,
model_obj,
reservation_snapshot,
)
except UpstreamError:
# Let the outer UpstreamError handler manage retry/revert
@@ -552,9 +398,7 @@ async def proxy(
"max_cost_for_model": max_cost_for_model,
},
)
await revert_pay_for_request(
key, session, max_cost_for_model, reservation_snapshot
)
await revert_pay_for_request(key, session, max_cost_for_model)
raise
# Reactive recovery: some models reject one specific request
@@ -562,7 +406,7 @@ async def proxy(
# 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.
if response.status_code == 400 and not is_ehbp:
if response.status_code == 400:
correction = correct_request(
request_body,
extract_error_message(response),
@@ -590,7 +434,7 @@ async def proxy(
if response.status_code != 200:
# Check if we should retry (502 Upstream Error or 429 Rate Limit)
should_retry = response.status_code in [502, 429, 400, 401, 403, 404]
if should_retry and i < len(candidates) - 1:
if should_retry and i < len(upstreams) - 1:
error_message = ""
try:
if hasattr(response, "body"):
@@ -623,9 +467,7 @@ async def proxy(
continue
# 4xx error (user error), or other non-retryable error, or last provider failed
await revert_pay_for_request(
key, session, max_cost_for_model, reservation_snapshot
)
await revert_pay_for_request(key, session, max_cost_for_model)
logger.warning(
"Upstream request failed, revert payment "
"(provider=%s model=%s status=%s path=%s)",
@@ -657,10 +499,8 @@ async def proxy(
"max_cost_for_model": max_cost_for_model,
},
)
# The cancellation has been caught, so complete exact cleanup in
# this task before the request-scoped session can be torn down.
await revert_pay_for_request(
key, session, max_cost_for_model, reservation_snapshot
await asyncio.shield(
revert_pay_for_request(key, session, max_cost_for_model)
)
raise
@@ -674,15 +514,13 @@ async def proxy(
"provider": upstream.provider_type,
"model": model_id,
"status_code": e.status_code,
"retry": i < len(candidates) - 1,
"retry": i < len(upstreams) - 1,
},
)
# If this was the last provider
if i == len(candidates) - 1:
await revert_pay_for_request(
key, session, max_cost_for_model, reservation_snapshot
)
if i == len(upstreams) - 1:
await revert_pay_for_request(key, session, max_cost_for_model)
return create_upstream_error_response(e, request)
# Otherwise loop continues to next provider
-2
View File
@@ -11,7 +11,6 @@ from .openrouter import OpenRouterUpstreamProvider
from .perplexity import PerplexityUpstreamProvider
from .ppqai import PPQAIUpstreamProvider
from .routstr import RoutstrUpstreamProvider
from .tinfoil import TinfoilUpstreamProvider
from .xai import XAIUpstreamProvider
upstream_provider_classes: list[type[BaseUpstreamProvider]] = [
@@ -27,7 +26,6 @@ upstream_provider_classes: list[type[BaseUpstreamProvider]] = [
PerplexityUpstreamProvider,
PPQAIUpstreamProvider,
RoutstrUpstreamProvider,
TinfoilUpstreamProvider,
XAIUpstreamProvider,
]
"""List of all upstream classes"""
+1 -1
View File
@@ -24,7 +24,7 @@ class AnthropicUpstreamProvider(BaseUpstreamProvider):
)
@classmethod
def _build_from_row(
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
) -> "AnthropicUpstreamProvider":
return cls(
+2 -47
View File
@@ -4,14 +4,7 @@ import json
from sqlmodel import select
from ..core import get_logger
from ..core.db import (
CashuTransaction,
UpstreamProviderRow,
create_session,
)
from ..core.db import (
store_cashu_transaction_with_retry as store_cashu_transaction,
)
from ..core.db import UpstreamProviderRow, create_session
from ..wallet import send_token
from .routstr import RoutstrUpstreamProvider
@@ -104,8 +97,6 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
# Instantiate provider and check balance
provider = RoutstrUpstreamProvider.from_db_row(row)
if provider is None:
return
balance = await provider.get_balance()
if balance is None:
@@ -130,6 +121,7 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
},
)
print(amount, mint_url)
try:
token = await send_token(amount, "sat", mint_url)
except Exception as e:
@@ -144,23 +136,6 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
)
return
try:
await store_cashu_transaction(
token=token,
amount=amount,
unit="sat",
mint_url=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": mint_url},
)
return
result = await provider.topup(token)
if "error" in result:
@@ -172,26 +147,6 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
},
)
else:
async with create_session() as session:
transaction = (
await session.exec(
select(CashuTransaction).where(
CashuTransaction.token == token,
CashuTransaction.type == "out",
CashuTransaction.source == "auto_topup",
)
)
).first()
if transaction is None:
logger.critical(
"Completed auto top-up transaction is missing from the database",
extra={"provider_id": row.id, "mint_url": mint_url},
)
else:
transaction.collected = True
session.add(transaction)
await session.commit()
logger.info(
"Auto top-up completed successfully",
extra={
+1 -1
View File
@@ -38,7 +38,7 @@ class AzureUpstreamProvider(BaseUpstreamProvider):
self.api_version = api_version
@classmethod
def _build_from_row(
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
) -> "AzureUpstreamProvider | None":
if not provider_row.api_version:
+183 -458
View File
File diff suppressed because it is too large Load Diff
+9 -16
View File
@@ -3,15 +3,10 @@
litellm's bundled cost map does not yet ship ``deepseek-v4-flash`` /
``deepseek-v4-pro``. Without an entry, ``backfill_cache_pricing`` cannot find a
``cache_read_input_token_cost`` and cache reads fall back to the full input
rate a large overcharge on cache hits (DeepSeek V4 hits are ~0.008-0.02x
input, i.e. cached tokens cost 50-120x less than regular input).
rate a ~60% overcharge on cache hits (DeepSeek hits are ~0.2x input).
This module injects the missing entries into ``litellm.model_cost`` at startup
so the existing backfill path resolves them. Rates mirror the canonical
``deepseek`` provider entries now in litellm's ``model_prices`` map
(``input_cost_per_token`` is the cache-*miss* rate;
``cache_read_input_token_cost`` is the cache-*hit* rate), sourced from
https://api-docs.deepseek.com/quick_start/pricing via
so the existing backfill path resolves them. Rates mirror the open upstream PR
https://github.com/BerriAI/litellm/pull/26380 (issue
https://github.com/BerriAI/litellm/issues/30430).
@@ -27,23 +22,21 @@ from ..core import get_logger
logger = get_logger(__name__)
# USD per token. Mirrors the canonical ``deepseek`` provider entries in
# litellm's model_prices map (source: DeepSeek API pricing docs). Keep these in
# sync with ``litellm.model_cost["deepseek/deepseek-v4-*"]``.
# USD per token. Source: BerriAI/litellm PR #26380.
_DEEPSEEK_V4_RATES: dict[str, dict[str, float]] = {
"deepseek-v4-flash": {
"input_cost_per_token": 1.4e-07,
"output_cost_per_token": 2.8e-07,
"cache_read_input_token_cost": 2.8e-09,
"cache_read_input_token_cost": 2.8e-08,
"cache_creation_input_token_cost": 0.0,
"input_cost_per_token_cache_hit": 2.8e-09,
"input_cost_per_token_cache_hit": 2.8e-08,
},
"deepseek-v4-pro": {
"input_cost_per_token": 4.35e-07,
"output_cost_per_token": 8.7e-07,
"cache_read_input_token_cost": 3.625e-09,
"input_cost_per_token": 1.74e-06,
"output_cost_per_token": 3.48e-06,
"cache_read_input_token_cost": 1.4e-07,
"cache_creation_input_token_cost": 0.0,
"input_cost_per_token_cache_hit": 3.625e-09,
"input_cost_per_token_cache_hit": 1.4e-07,
},
}
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -20,7 +20,7 @@ class FireworksUpstreamProvider(BaseUpstreamProvider):
)
@classmethod
def _build_from_row(
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
) -> "FireworksUpstreamProvider":
return cls(
+1 -1
View File
@@ -51,7 +51,7 @@ class GeminiUpstreamProvider(BaseUpstreamProvider):
return self._client
@classmethod
def _build_from_row(
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
) -> "GeminiUpstreamProvider":
return cls(
+45 -92
View File
@@ -5,12 +5,6 @@ from typing import TYPE_CHECKING
import httpx
from .base import BaseUpstreamProvider
from .pricing_resolver import (
FallbackPricingResolver,
ResolvedPricing,
_as_float,
estimate_context_length,
)
if TYPE_CHECKING:
from ..core.db import UpstreamProviderRow
@@ -51,7 +45,7 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
)
@classmethod
def _build_from_row(
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
) -> "GenericUpstreamProvider":
return cls(
@@ -70,40 +64,6 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
"platform_url": cls.platform_url,
}
def _native_pricing(
self, model_id: str, model_spec: dict
) -> ResolvedPricing | None:
"""Read pricing/metadata from Venice's bespoke ``model_spec`` schema.
Returns ``None`` when the upstream reported no *usable* native price
absent, non-numeric, negative, or both-zero so the caller falls
through to the shared resolution chain instead of fabricating a number
or trusting a bogus one. This mirrors the money-safety guards the
litellm and OpenRouter rungs already apply: a both-zero price would
serve the model free, a negative one would credit the caller, and a
non-numeric string would otherwise throw and drop the whole catalog.
"""
pricing_info = model_spec.get("pricing", {})
input_usd = _as_float(pricing_info.get("input", {}).get("usd"))
output_usd = _as_float(pricing_info.get("output", {}).get("usd"))
if input_usd is None or output_usd is None:
return None
if input_usd < 0 or output_usd < 0 or (input_usd == 0 and output_usd == 0):
return None
capabilities = model_spec.get("capabilities", {})
input_modalities = ["text"]
if capabilities.get("supportsVision", False):
input_modalities.append("image")
return ResolvedPricing(
prompt=input_usd / 1_000_000,
completion=output_usd / 1_000_000,
context_length=model_spec.get("availableContextTokens"),
source="native",
input_modalities=input_modalities,
)
async def fetch_models(self) -> list[Model]:
"""Fetch models from upstream API using /models endpoint."""
from ..payment.models import Architecture, Model, Pricing, TopProvider
@@ -118,7 +78,6 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
response.raise_for_status()
data = response.json()
resolver = FallbackPricingResolver()
models_list = []
for model_data in data.get("data", []):
model_id = model_data.get("id", "")
@@ -130,44 +89,41 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
owned_by = model_data.get("owned_by", "unknown")
model_spec = model_data.get("model_spec", {})
resolved = self._native_pricing(model_id, model_spec)
if resolved is None:
resolved = await resolver.resolve(model_id)
context_length = 4096
if model_spec.get("availableContextTokens"):
context_length = model_spec["availableContextTokens"]
elif any(
pattern in model_id.lower() for pattern in ["32k", "32000"]
):
context_length = 32768
elif any(
pattern in model_id.lower() for pattern in ["16k", "16000"]
):
context_length = 16384
elif any(pattern in model_id.lower() for pattern in ["8k", "8000"]):
context_length = 8192
elif "gpt-4" in model_id.lower():
context_length = 8192
elif "claude" in model_id.lower():
context_length = 200000
if resolved is None:
# Fail closed: never invent a price. Import the model
# disabled with a warning so the operator can price it
# (the admin UI surfaces disabled remote models).
logger.warning(
f"No pricing source resolved for '{model_id}' from "
f"{self.upstream_name}; importing it disabled",
extra={"model_id": model_id, "base_url": self.base_url},
)
resolved = ResolvedPricing(
prompt=0.0,
completion=0.0,
context_length=None,
source="unresolved",
)
enabled = False
else:
enabled = True
pricing_info = model_spec.get("pricing", {})
input_pricing = pricing_info.get("input", {})
output_pricing = pricing_info.get("output", {})
# Prefer the source's own modality string (OpenRouter ships
# one, e.g. "text+image->text"); otherwise derive it from the
# captured input/output modalities in the same "in->out" shape
# rather than flattening vision models to "text->text".
modality = resolved.modality or (
f"{'+'.join(resolved.input_modalities)}"
f"->{'+'.join(resolved.output_modalities)}"
)
prompt_price = input_pricing.get("usd", 0.001) / 1000000
completion_price = output_pricing.get("usd", 0.001) / 1000000
# A source can carry a price but no context (e.g. a litellm
# entry missing max_input_tokens); fall back to an id-based
# estimate so we never persist a zero-length window.
context_length = resolved.context_length or estimate_context_length(
model_id
)
capabilities = model_spec.get("capabilities", {})
input_modalities = ["text"]
output_modalities = ["text"]
if capabilities.get("supportsVision", False):
input_modalities.append("image")
modality = "text"
if capabilities.get("supportsVision", False):
modality = "text->text"
spec_name = model_spec.get("name", model_name)
description = f"{spec_name}"
@@ -183,33 +139,30 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
context_length=context_length,
architecture=Architecture(
modality=modality,
input_modalities=resolved.input_modalities,
output_modalities=resolved.output_modalities,
tokenizer=resolved.tokenizer,
instruct_type=resolved.instruct_type,
input_modalities=input_modalities,
output_modalities=output_modalities,
tokenizer="unknown",
instruct_type=None,
),
pricing=Pricing(
prompt=resolved.prompt,
completion=resolved.completion,
prompt=prompt_price,
completion=completion_price,
request=0.0,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
input_cache_read=resolved.input_cache_read,
input_cache_write=resolved.input_cache_write,
max_prompt_cost=0.001,
max_completion_cost=0.001,
max_cost=0.001,
),
sats_pricing=None,
per_request_limits=None,
top_provider=TopProvider(
context_length=context_length,
max_completion_tokens=(
resolved.max_completion_tokens
if resolved.max_completion_tokens is not None
else context_length // 2
),
is_moderated=bool(resolved.is_moderated),
max_completion_tokens=context_length // 2,
is_moderated=False,
),
enabled=enabled,
enabled=True,
upstream_provider_id=None,
canonical_slug=None,
)
+1 -1
View File
@@ -20,7 +20,7 @@ class GroqUpstreamProvider(BaseUpstreamProvider):
)
@classmethod
def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "GroqUpstreamProvider":
def from_db_row(cls, provider_row: "UpstreamProviderRow") -> "GroqUpstreamProvider":
return cls(
api_key=provider_row.api_key,
provider_fee=provider_row.provider_fee,
+17 -44
View File
@@ -12,7 +12,6 @@ from sqlmodel import select
from ..core import get_logger
from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session
from ..core.provider_slugs import allocate_unique_provider_slug
from ..payment.models import Model
from .base import BaseUpstreamProvider
@@ -94,10 +93,12 @@ async def get_all_models_with_overrides(
provider_result = await session.exec(select(UpstreamProviderRow))
providers_by_id = {p.id: p for p in provider_result.all()}
overrides_by_key: dict[tuple[str, int], tuple[ModelRow, float]] = {
(row.id.lower(), row.upstream_provider_id): (
overrides_by_id: dict[str, tuple[ModelRow, float]] = {
row.id: (
row,
providers_by_id[row.upstream_provider_id].provider_fee,
providers_by_id[row.upstream_provider_id].provider_fee
if row.upstream_provider_id in providers_by_id
else 1.01,
)
for row in override_rows
if row.upstream_provider_id is not None
@@ -105,28 +106,17 @@ async def get_all_models_with_overrides(
and providers_by_id[row.upstream_provider_id].enabled
}
all_models: dict[tuple[str, str], Model] = {}
all_models: dict[str, Model] = {}
for upstream in upstreams:
upstream_db_id = getattr(upstream, "db_id", None)
provider_key = (
f"db:{upstream_db_id}"
if isinstance(upstream_db_id, int)
else f"{getattr(upstream, 'provider_type', '')}|{getattr(upstream, 'base_url', '')}"
)
for model in upstream.get_cached_models():
model_key = (
(model.id.lower(), upstream_db_id)
if isinstance(upstream_db_id, int)
else None
)
if model_key is not None and model_key in overrides_by_key:
override_row, provider_fee = overrides_by_key[model_key]
all_models[(model.id.lower(), provider_key)] = _row_to_model(
if model.id in overrides_by_id:
override_row, provider_fee = overrides_by_id[model.id]
all_models[model.id] = _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee
)
elif model.enabled:
all_models[(model.id.lower(), provider_key)] = model
all_models[model.id] = model
return list(all_models.values())
@@ -225,6 +215,9 @@ async def init_upstreams() -> list[BaseUpstreamProvider]:
provider = _instantiate_provider(provider_row)
if provider:
# Keep provider DB id on runtime instance so model mapping can
# bind DB overrides to the correct upstream.
setattr(provider, "db_id", provider_row.id)
await provider.refresh_models_cache()
logger.debug(
f"Initialized {provider_row.provider_type} provider",
@@ -257,7 +250,6 @@ async def _seed_providers_from_settings(
providers_to_add: list[UpstreamProviderRow] = []
seeded_provider_keys: set[tuple[str, str]] = set()
reserved_slugs: set[str] = set()
provider_classes_by_type = {
cls.provider_type: cls
@@ -272,7 +264,6 @@ async def _seed_providers_from_settings(
("PERPLEXITY_API_KEY", "perplexity", None, None),
("FIREWORKS_API_KEY", "fireworks", None, None),
("XAI_API_KEY", "xai", None, None),
("TINFOIL_API_KEY", "tinfoil", None, None),
]
for env_key, provider_type, _, _ in env_mappings:
@@ -288,13 +279,8 @@ async def _seed_providers_from_settings(
)
)
if not result.first():
slug = await allocate_unique_provider_slug(
session, provider_type, reserved_slugs
)
reserved_slugs.add(slug)
providers_to_add.append(
UpstreamProviderRow(
slug=slug,
provider_type=provider_type,
base_url=base_url,
api_key=api_key,
@@ -313,13 +299,8 @@ async def _seed_providers_from_settings(
)
)
if not result.first():
slug = await allocate_unique_provider_slug(
session, "ollama", reserved_slugs
)
reserved_slugs.add(slug)
providers_to_add.append(
UpstreamProviderRow(
slug=slug,
provider_type="ollama",
base_url=ollama_base_url,
api_key=ollama_api_key,
@@ -339,13 +320,8 @@ async def _seed_providers_from_settings(
)
)
if not result.first():
slug = await allocate_unique_provider_slug(
session, "azure", reserved_slugs
)
reserved_slugs.add(slug)
providers_to_add.append(
UpstreamProviderRow(
slug=slug,
provider_type="azure",
base_url=base_url,
api_key=api_key,
@@ -366,13 +342,8 @@ async def _seed_providers_from_settings(
)
)
if not result.first():
slug = await allocate_unique_provider_slug(
session, "custom", reserved_slugs
)
reserved_slugs.add(slug)
providers_to_add.append(
UpstreamProviderRow(
slug=slug,
provider_type="custom",
base_url=base_url,
api_key=api_key,
@@ -385,7 +356,7 @@ async def _seed_providers_from_settings(
session.add(provider)
logger.info(
f"Seeding {provider.provider_type} provider", # type: ignore[str-format]
extra={"base_url": provider.base_url, "slug": provider.slug},
extra={"base_url": provider.base_url},
)
@@ -420,7 +391,9 @@ def _instantiate_provider(
return provider
if provider_row.provider_type == "custom":
return BaseUpstreamProvider.from_db_row(provider_row)
return BaseUpstreamProvider(
provider_row.base_url, provider_row.api_key, provider_row.provider_fee
)
logger.error(
f"Unknown provider type: {provider_row.provider_type}",
+1 -1
View File
@@ -43,7 +43,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
)
@classmethod
def _build_from_row(
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
) -> "OllamaUpstreamProvider":
return cls(
+1 -1
View File
@@ -20,7 +20,7 @@ class OpenAIUpstreamProvider(BaseUpstreamProvider):
)
@classmethod
def _build_from_row(
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
) -> "OpenAIUpstreamProvider":
return cls(
+34 -1
View File
@@ -1,3 +1,4 @@
import json
from typing import TYPE_CHECKING
import httpx
@@ -18,6 +19,38 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider):
supports_anthropic_messages = True
litellm_provider_prefix = "openrouter/"
def prepare_request_body(
self, body: bytes | None, model_obj: Model
) -> bytes | None:
"""Set provider.require_parameters on tool-use requests.
Without it OpenRouter can route a tool call to an endpoint that doesn't
support function calling and 404 with "No endpoints found that support
tool use". We leave a client-supplied value untouched.
"""
body = super().prepare_request_body(body, model_obj)
if not body:
return body
try:
data = json.loads(body)
except json.JSONDecodeError:
return body
if not isinstance(data, dict) or not data.get("tools"):
return body
provider = data.get("provider")
if not isinstance(provider, dict):
provider = {}
if "require_parameters" in provider:
return body
provider["require_parameters"] = True
data["provider"] = provider
return json.dumps(data).encode()
def _apply_provider_field(self, response_json: object) -> None:
"""Stamp the ``provider`` field for OpenRouter responses.
@@ -57,7 +90,7 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider):
)
@classmethod
def _build_from_row(
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
) -> "OpenRouterUpstreamProvider":
return cls(
+1 -1
View File
@@ -23,7 +23,7 @@ class PerplexityUpstreamProvider(BaseUpstreamProvider):
)
@classmethod
def _build_from_row(
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
) -> "PerplexityUpstreamProvider":
return cls(
+8 -24
View File
@@ -8,7 +8,6 @@ from pydantic.v1 import BaseModel, Field
from ..core.logging import get_logger
from ..payment.models import Architecture, Model, Pricing, async_fetch_openrouter_models
from .base import BaseUpstreamProvider, TopupData
from .ehbp import EHBPForwardingTarget
if TYPE_CHECKING:
from ..core.db import UpstreamProviderRow
@@ -40,10 +39,6 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
default_base_url = "https://api.ppq.ai"
platform_url = "https://ppq.ai/api-docs"
IGNORED_MODEL_IDS: list[str] = ["auto"]
# PPQ.AI has a private encrypted endpoint, but this proxy currently has no
# provider-attested usage extractor/model binding for it. Keep EHBP disabled
# until a ConfidentialInferenceProfile can bill it without max-cost fallback.
supports_ehbp = False
def __init__(self, api_key: str, provider_fee: float = 1.0):
super().__init__(
@@ -51,7 +46,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
)
@classmethod
def _build_from_row(
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
) -> "PPQAIUpstreamProvider":
return cls(
@@ -75,20 +70,6 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
def transform_model_name(self, model_id: str) -> str:
return model_id
def get_ehbp_forwarding_target(
self, path: str, model_obj: Model
) -> EHBPForwardingTarget:
"""Return the PPQ.AI private enclave target for EHBP requests.
PPQ.AI exposes EHBP-aware inference under /private/v1/... separate
from the public /v1/... endpoint. The encrypted body remains opaque to
Routstr, so PPQ.AI also needs X-Private-Model for routing/billing.
"""
return EHBPForwardingTarget(
url=f"{self.base_url.rstrip('/')}/private/{path.lstrip('/')}",
headers={"X-Private-Model": model_obj.forwarded_model_id or model_obj.id},
)
@classmethod
async def create_account_static(cls) -> dict[str, object]:
"""Create a new PPQ.AI account without requiring an instance.
@@ -248,14 +229,17 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
f"Disabling PPQ.AI provider ({self.base_url}) due to insufficient balance",
extra={"error": error_message},
)
from sqlmodel import select
from ..core.db import UpstreamProviderRow, create_session
async with create_session() as session:
provider = (
await session.get(UpstreamProviderRow, self.db_id)
if self.db_id is not None
else None
statement = select(UpstreamProviderRow).where(
UpstreamProviderRow.base_url == self.base_url,
UpstreamProviderRow.api_key == self.api_key,
)
result = await session.exec(statement)
provider = result.first()
if provider:
provider.enabled = False
-203
View File
@@ -1,203 +0,0 @@
"""Shared price/metadata resolution chain for upstream model discovery.
Most OpenAI-compatible ``/models`` responses carry no pricing. Rather than let
a provider fabricate one, this module resolves a model through decreasingly
trustworthy sources litellm's bundled cost map (curated list prices, mirrors
provider docs), then the OpenRouter feed (resale prices, broader coverage)
and returns ``None`` when none of them know the model, so the caller can fail
closed instead of inventing a number.
Provider-native pricing (a gateway's own ``/models`` schema, e.g. Venice's
``model_spec``) is authoritative and handled by the provider before this chain
is consulted; only the shared fallback lives here so a later refactor can hoist
it into the base provider unchanged.
"""
from __future__ import annotations
from dataclasses import dataclass, field
@dataclass
class ResolvedPricing:
"""Per-token pricing plus whatever metadata the answering source carried.
Prices are USD per token. ``source`` records provenance
(``native``/``litellm``/``openrouter``/``unresolved``) so later work can
surface where each price came from.
"""
prompt: float
completion: float
context_length: int | None
source: str
modality: str | None = None
max_completion_tokens: int | None = None
input_cache_read: float = 0.0
input_cache_write: float = 0.0
input_modalities: list[str] = field(default_factory=lambda: ["text"])
output_modalities: list[str] = field(default_factory=lambda: ["text"])
tokenizer: str = "unknown"
instruct_type: str | None = None
is_moderated: bool | None = None
def estimate_context_length(model_id: str) -> int:
"""Best-effort context window from a model id when no source reports one.
The last rung of the fallback chain, reached only for a model whose price
resolved but whose context did not (or that imported disabled). Context is
not a billing input, so a rough id-based guess is acceptable here where a
guessed *price* never would be.
"""
lowered = model_id.lower()
if any(pattern in lowered for pattern in ["32k", "32000"]):
return 32768
if any(pattern in lowered for pattern in ["16k", "16000"]):
return 16384
if any(pattern in lowered for pattern in ["8k", "8000"]):
return 8192
if "gpt-4" in lowered:
return 8192
if "claude" in lowered:
return 200000
return 4096
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
def _as_int(value: object) -> int | None:
"""Coerce an already-numeric token count to ``int``, else ``None``."""
return int(value) if isinstance(value, (int, float)) else 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.
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)):
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):
return None
input_modalities = ["text"]
if info.get("supports_vision"):
input_modalities.append("image")
return ResolvedPricing(
prompt=float(prompt),
completion=float(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
# estimate downstream rather than borrowing the output cap.
context_length=_as_int(info.get("max_input_tokens")),
source="litellm",
max_completion_tokens=_as_int(info.get("max_output_tokens")),
input_cache_read=float(info.get("cache_read_input_token_cost") or 0.0),
input_cache_write=float(info.get("cache_creation_input_token_cost") or 0.0),
input_modalities=input_modalities,
)
def _match_openrouter(model_id: str, feed: list[dict]) -> dict | None:
"""Find ``model_id`` in the OpenRouter feed, exact id before bare tail.
Bare-tail matching (``deepseek-chat`` ``deepseek/deepseek-chat``) is a
looser, lower-trust match OpenRouter fans a model out across resellers
so an exact id match always wins first. When several entries share the bare
tail, the one with the highest *combined* (prompt + completion) per-token
cost wins: the choice must be deterministic (not feed-order-dependent) and
money-safe whichever way traffic leans, since undercharging is the hazard.
Ranking on prompt alone could pick an entry that is cheap on input but dear
on output. The live feed has no such collisions today; this only governs
the latent case.
"""
bare = model_id.split("/", 1)[-1]
exact = next((m for m in feed if m.get("id") == model_id), None)
if exact is not None:
return exact
matches = [m for m in feed if m.get("id", "").split("/", 1)[-1] == bare]
if not matches:
return None
def _combined_cost(m: dict) -> float:
pricing = m.get("pricing", {})
return (_as_float(pricing.get("prompt")) or 0.0) + (
_as_float(pricing.get("completion")) or 0.0
)
return max(matches, key=_combined_cost)
def _from_openrouter(model_id: str, feed: list[dict]) -> ResolvedPricing | None:
entry = _match_openrouter(model_id, feed)
if entry is None:
return None
pricing = entry.get("pricing", {})
prompt = _as_float(pricing.get("prompt"))
completion = _as_float(pricing.get("completion"))
if prompt is None or completion is None:
return None
architecture = entry.get("architecture", {})
top_provider = entry.get("top_provider", {})
return ResolvedPricing(
prompt=prompt,
completion=completion,
context_length=_as_int(entry.get("context_length")),
source="openrouter",
modality=architecture.get("modality"),
max_completion_tokens=_as_int(top_provider.get("max_completion_tokens")),
input_cache_read=_as_float(pricing.get("input_cache_read")) or 0.0,
input_cache_write=_as_float(pricing.get("input_cache_write")) or 0.0,
input_modalities=architecture.get("input_modalities") or ["text"],
output_modalities=architecture.get("output_modalities") or ["text"],
tokenizer=architecture.get("tokenizer") or "unknown",
instruct_type=architecture.get("instruct_type"),
is_moderated=top_provider.get("is_moderated"),
)
class FallbackPricingResolver:
"""Resolves models via litellm → OpenRouter for one discovery pass.
The OpenRouter catalog is fetched at most once and only when a model
actually misses litellm, so a provider full of litellm-known models never
touches the network. Instantiate one per ``fetch_models`` call.
"""
def __init__(self) -> None:
self._openrouter_feed: list[dict] | None = None
async def resolve(self, model_id: str) -> ResolvedPricing | None:
"""Resolve ``model_id``; ``None`` if no source knows it."""
resolved = _from_litellm(model_id)
if resolved is not None:
return resolved
if self._openrouter_feed is None:
# Lazy import so tests can patch the feed at its source.
from ..payment.models import async_fetch_openrouter_models
self._openrouter_feed = await async_fetch_openrouter_models()
return _from_openrouter(model_id, self._openrouter_feed)
+1 -1
View File
@@ -55,7 +55,7 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider):
return path.lstrip("/")
@classmethod
def _build_from_row(
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
) -> "RoutstrUpstreamProvider":
import json
-236
View File
@@ -1,236 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import httpx
from fastapi import Request
from fastapi.responses import Response, StreamingResponse
from pydantic.v1 import BaseModel
from ..core.exceptions import UpstreamError
from ..core.logging import get_logger
from ..payment.models import Architecture, Model, Pricing
from .base import BaseUpstreamProvider
from .ehbp import (
_ENCLAVE_URL_HEADER,
_PROXY_ONLY_HEADERS,
_RESPONSE_USAGE_HEADER,
ConfidentialInferenceProfile,
EHBPForwardingTarget,
)
if TYPE_CHECKING:
from ..core.db import UpstreamProviderRow
logger = get_logger(__name__)
class TinfoilModelPricing(BaseModel):
inputTokenPricePer1M: float = 0.0
outputTokenPricePer1M: float = 0.0
requestPrice: float = 0.0
class TinfoilModel(BaseModel):
id: str
context_window: int = 0
created: int = 0
multimodal: bool = False
reasoning: bool = False
tool_calling: bool = False
type: str = "chat"
pricing: TinfoilModelPricing = TinfoilModelPricing()
endpoints: list[str] = []
class TinfoilUpstreamProvider(BaseUpstreamProvider):
"""Direct upstream provider for the Tinfoil inference API.
Tinfoil hosts open-source models inside attested secure enclaves and exposes
an OpenAI-compatible API at ``https://inference.tinfoil.sh``. Request and
response bodies are encrypted end-to-end with EHBP (HPKE), so Routstr acts
as a blind relay: it forwards the opaque encrypted body, never sees
plaintext, and bills from the ``X-Tinfoil-Usage-Metrics`` header that
Tinfoil returns outside the encrypted body when
``X-Tinfoil-Request-Usage-Metrics: true`` is set.
"""
provider_type = "tinfoil"
default_base_url = "https://inference.tinfoil.sh"
platform_url = "https://docs.tinfoil.sh"
supports_ehbp = True
confidential_inference_profile = ConfidentialInferenceProfile(
usage_response_header=_RESPONSE_USAGE_HEADER,
client_target_url_header=_ENCLAVE_URL_HEADER,
allow_client_target_override=True,
proxy_only_headers=_PROXY_ONLY_HEADERS,
)
def __init__(self, api_key: str, provider_fee: float = 1.0):
super().__init__(
base_url=self.default_base_url,
api_key=api_key,
provider_fee=provider_fee,
)
@classmethod
def _build_from_row(
cls, provider_row: "UpstreamProviderRow"
) -> "TinfoilUpstreamProvider":
return cls(
api_key=provider_row.api_key,
provider_fee=provider_row.provider_fee,
)
@classmethod
def get_provider_metadata(cls) -> dict[str, object]:
return {
"id": cls.provider_type,
"name": "Tinfoil",
"default_base_url": cls.default_base_url,
"fixed_base_url": True,
"platform_url": cls.platform_url,
"can_create_account": False,
"can_topup": False,
"can_show_balance": False,
}
def transform_model_name(self, model_id: str) -> str:
return model_id.removeprefix("tinfoil/")
def get_confidential_inference_profile(self) -> ConfidentialInferenceProfile:
return self.confidential_inference_profile
async def forward_get_request(
self,
request: Request,
path: str,
headers: dict,
) -> Response | StreamingResponse:
"""Handle Tinfoil-specific GET endpoints.
* ``/attestation`` (or ``/tee/attestation``): proxy to the Tinfoil ATC
(attestation bundle proxy) at ``https://atc.tinfoil.sh/attestation``.
* Other GETs: forward to the provider base URL
(``https://inference.tinfoil.sh``). ``X-Tinfoil-Enclave-Url`` is an
EHBP-only header used for encrypted POST requests and is not honored
for unencrypted GET requests.
"""
clean_path = path.removeprefix("tee/").rstrip("/")
if clean_path == "attestation":
return await self._proxy_attestation(headers)
return await super().forward_get_request(request, path, headers)
async def _proxy_attestation(self, headers: dict) -> Response:
url = "https://atc.tinfoil.sh/attestation"
async with httpx.AsyncClient(
transport=httpx.AsyncHTTPTransport(retries=1),
timeout=30.0,
) as client:
try:
resp = await client.get(
url,
headers={
"Accept": headers.get("accept", "application/json"),
},
)
response_headers = dict(resp.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=resp.content,
status_code=resp.status_code,
headers=response_headers,
)
except Exception as exc:
raise UpstreamError(
f"Error fetching Tinfoil attestation: {type(exc).__name__}",
status_code=502,
) from exc
def get_ehbp_forwarding_target(
self, path: str, model_obj: Model
) -> EHBPForwardingTarget:
"""Return the Tinfoil enclave target for EHBP requests.
Requests usage metrics from the enclave so Routstr can bill exactly
without decrypting the response body. The actual forwarding URL is
overridden at dispatch time by ``X-Tinfoil-Enclave-Url`` when the SDK
sends it (see ``routstr/upstream/ehbp.py``).
"""
return EHBPForwardingTarget(
url=f"{self.base_url.rstrip('/')}/{path.lstrip('/')}",
headers={"X-Tinfoil-Request-Usage-Metrics": "true"},
profile=self.confidential_inference_profile,
)
async def fetch_models(self) -> list[Model]:
"""Fetch models from the public Tinfoil models endpoint.
``GET /v1/models`` is unauthenticated and returns all available models
with their pricing in USD per 1M tokens.
"""
url = f"{self.base_url}/v1/models"
try:
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.get(url)
response.raise_for_status()
data = response.json()
models_data = data.get("data", [])
models: list[Model] = []
for model_data in models_data:
try:
tf = TinfoilModel.parse_obj(model_data)
input_price = tf.pricing.inputTokenPricePer1M
output_price = tf.pricing.outputTokenPricePer1M
request_price = tf.pricing.requestPrice
modality = "text->text"
input_modalities = ["text"]
output_modalities = ["text"]
if tf.multimodal:
modality = "text->text+image"
input_modalities = ["text", "image"]
models.append(
Model(
id=tf.id,
name=tf.id,
created=tf.created,
description=f"Tinfoil {tf.type} model",
context_length=tf.context_window,
architecture=Architecture(
modality=modality,
input_modalities=input_modalities,
output_modalities=output_modalities,
tokenizer="Unknown",
instruct_type=None,
),
pricing=Pricing(
prompt=input_price / 1_000_000,
completion=output_price / 1_000_000,
request=request_price,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
),
)
)
except Exception as e:
logger.warning(
"Failed to parse Tinfoil 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 Tinfoil",
extra={"error": str(e), "error_type": type(e).__name__},
)
return []
-189
View File
@@ -1,189 +0,0 @@
"""h11-based HTTP client for EHBP requests that captures HTTP trailers.
httpx/httpcore silently discard HTTP trailers during chunked transfer
decoding. Tinfoil returns ``X-Tinfoil-Usage-Metrics`` as a trailer on
streaming responses, so we need a lower-level HTTP client that preserves
trailers from the h11 ``EndOfMessage`` event.
Because EHBP response bodies are opaque encrypted blobs, buffering the full
response is acceptable the client decrypts the complete body regardless of
whether it arrived streamed or buffered.
"""
from __future__ import annotations
import asyncio
import ssl
from dataclasses import dataclass, field
from urllib.parse import urlsplit
import h11
from ..core import get_logger
logger = get_logger(__name__)
_READ_BUFSIZE = 65536
_DEFAULT_TIMEOUT_SECONDS = 30.0
_DEFAULT_CLOSE_TIMEOUT_SECONDS = 1.0
_DEFAULT_MAX_RESPONSE_BYTES = 25 * 1024 * 1024
_HOP_BY_HOP_HEADERS = {
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailer",
"transfer-encoding",
"upgrade",
}
@dataclass
class TrailerResponse:
"""Buffered HTTP response with optional trailer headers."""
status_code: int
headers: list[tuple[str, str]]
body: bytes
trailers: list[tuple[str, str]] = field(default_factory=list)
def _get_header(headers: list[tuple[str, str]], name: str) -> str | None:
name_lower = name.lower()
for k, v in headers:
if k.lower() == name_lower:
return v
return None
def _strip_hop_by_hop_headers(headers: dict[str, str]) -> dict[str, str]:
"""Remove connection-specific headers before serializing a new request."""
connection_tokens: set[str] = set()
for key, value in headers.items():
if key.lower() == "connection":
connection_tokens.update(
token.strip().lower() for token in value.split(",") if token.strip()
)
excluded = _HOP_BY_HOP_HEADERS | connection_tokens
return {key: value for key, value in headers.items() if key.lower() not in excluded}
async def forward_with_trailer(
*,
method: str,
url: str,
headers: dict[str, str],
body: bytes,
timeout_seconds: float = _DEFAULT_TIMEOUT_SECONDS,
max_response_bytes: int = _DEFAULT_MAX_RESPONSE_BYTES,
close_timeout_seconds: float = _DEFAULT_CLOSE_TIMEOUT_SECONDS,
) -> TrailerResponse:
"""Send an HTTP/1.1 request via h11 and capture HTTP trailers.
Returns a :class:`TrailerResponse` with the full buffered body and any
trailer headers from the ``EndOfMessage`` event.
"""
parsed = urlsplit(url)
host = parsed.hostname
if not host:
raise ValueError(f"Invalid URL (no hostname): {url}")
port = parsed.port or 443
path = parsed.path or "/"
if parsed.query:
path = f"{path}?{parsed.query}"
# FastAPI has already decoded the incoming request body. Do not carry the
# original connection's framing or other hop-by-hop metadata into the new
# upstream connection.
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:
# Build HTTP/1.1 request
header_lines = [f"{method} {path} HTTP/1.1"]
has_host = any(k.lower() == "host" for k in headers)
if not has_host:
header_lines.append(f"Host: {host}")
header_lines.append("Connection: close")
for key, value in headers.items():
if key.lower() == "host":
continue
header_lines.append(f"{key}: {value}")
if body and not any(k.lower() == "content-length" for k in headers):
header_lines.append(f"Content-Length: {len(body)}")
request_data = "\r\n".join(header_lines).encode() + b"\r\n\r\n"
if body:
request_data += body
writer.write(request_data)
await asyncio.wait_for(writer.drain(), timeout=timeout_seconds)
# Parse response with h11
conn = h11.Connection(h11.CLIENT)
status_code = 0
resp_headers: list[tuple[str, str]] = []
body_chunks: list[bytes] = []
body_size = 0
trailers: list[tuple[str, str]] = []
while True:
event = conn.next_event()
if event is h11.NEED_DATA:
data = await asyncio.wait_for(
reader.read(_READ_BUFSIZE),
timeout=timeout_seconds,
)
conn.receive_data(data if data else b"")
continue
if isinstance(event, h11.Response):
status_code = event.status_code
resp_headers = [(k.decode(), v.decode()) for k, v in event.headers]
elif isinstance(event, h11.Data):
body_size += len(event.data)
if body_size > max_response_bytes:
raise ValueError(
f"EHBP response exceeded {max_response_bytes} bytes"
)
body_chunks.append(event.data)
elif isinstance(event, h11.EndOfMessage):
for k, v in event.headers:
trailers.append((k.decode(), v.decode()))
break
elif isinstance(event, h11.PAUSED):
# Shouldn't happen for simple request/response, but break safely
logger.warning("h11 PAUSED event during EHBP response parsing")
break
elif isinstance(event, h11.ConnectionClosed):
break
return TrailerResponse(
status_code=status_code,
headers=resp_headers,
body=b"".join(body_chunks),
trailers=trailers,
)
finally:
writer.close()
if close_timeout_seconds > 0:
try:
await asyncio.wait_for(
writer.wait_closed(), timeout=close_timeout_seconds
)
except Exception:
pass
+1 -1
View File
@@ -21,7 +21,7 @@ class XAIUpstreamProvider(BaseUpstreamProvider):
)
@classmethod
def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "XAIUpstreamProvider":
def from_db_row(cls, provider_row: "UpstreamProviderRow") -> "XAIUpstreamProvider":
return cls(
api_key=provider_row.api_key,
provider_fee=provider_row.provider_fee,
+113 -356
View File
@@ -1,11 +1,9 @@
import asyncio
import re
import socket
import time
import typing
from typing import TypedDict
import httpx
from cashu.core.base import Proof, Token
from cashu.core.mint_info import MintInfo as _CashuMintInfo
from cashu.wallet.helpers import deserialize_token_from_string
@@ -14,7 +12,7 @@ from pydantic_core import PydanticUndefined
from sqlmodel import col, select, update
from .core import db, get_logger
from .core.db import store_cashu_transaction_with_retry as store_cashu_transaction
from .core.db import store_cashu_transaction
from .core.settings import settings
from .payment.lnurl import raw_send_to_lnurl
@@ -33,149 +31,6 @@ _CashuMintInfo.model_rebuild(force=True)
logger = get_logger(__name__)
class MintConnectionError(Exception):
"""The mint could not be reached (network transport failure).
Maps to a 503, not a 4xx: the token is fine, the mint is just unavailable.
"""
class TokenConsumedError(Exception):
"""A failure that happened AFTER the token's proofs were spent (melt
succeeded, or redemption already returned) e.g. minting on the primary
mint or the DB credit then failed.
Non-retryable: the same token will not work again. Seals the cause chain so
a transport error underneath is never re-surfaced as a retryable
mint_unreachable.
"""
# httpx base classes cover their subclasses. HTTPStatusError is excluded on
# purpose — that means the mint answered, just with an error status.
_TRANSPORT_EXC_TYPES: tuple[type[BaseException], ...] = (
httpx.NetworkError,
httpx.TimeoutException,
ConnectionError, # refused/reset/aborted
socket.gaierror, # DNS failure
asyncio.TimeoutError,
)
def is_mint_connection_error(error: BaseException) -> bool:
"""True if ``error`` (or anything in its cause/context chain) is a mint
transport failure. Walks the chain because some sites re-raise transport
errors wrapped in ValueError/MintConnectionError; matches on TYPE, not text.
"""
seen: set[int] = set()
current: BaseException | None = error
while current is not None and id(current) not in seen:
seen.add(id(current))
if isinstance(current, TokenConsumedError):
# Sealed: the token was already spent, so whatever transport error
# sits underneath must not make this look retryable.
return False
if isinstance(current, MintConnectionError):
return True
if isinstance(current, _TRANSPORT_EXC_TYPES):
return True
current = current.__cause__ or current.__context__
return False
# Redemption ``code`` values whose token is spent/consumed/unusable — the
# X-Cashu path must NOT echo the original token for these (echoing invites a
# retry with a token that can never succeed again).
SPENT_TOKEN_CODES: frozenset[str] = frozenset(
{
"cashu_token_already_spent",
"cashu_token_consumed",
"cashu_token_zero_value",
"internal_error",
}
)
def classify_redemption_error(
error: Exception,
) -> tuple[str, int, str, str] | None:
"""Map a token-redemption failure to ``(type, status, message, code)``.
Single source of truth for every endpoint that redeems a token (bearer,
X-Cashu, top-up) so the same failure yields the same taxonomy everywhere.
``type`` and ``code`` are stable client contract; ``message`` is sanitized
(raw error text stays in logs). Returns None for an unclassified internal
fault the caller emits a generic 500.
"""
if isinstance(error, TokenConsumedError):
return (
"token_consumed",
500,
"Token was redeemed but could not be credited; do not retry",
"cashu_token_consumed",
)
if is_mint_connection_error(error):
return (
"mint_unreachable",
503,
"Cashu mint is unreachable",
"cashu_mint_unreachable",
)
lowered = str(error).lower()
if "already spent" in lowered:
return (
"token_already_spent",
400,
"Cashu token already spent",
"cashu_token_already_spent",
)
if (
"insufficient" in lowered
or "melt fee" in lowered
or "exceed token amount" in lowered
or "estimate fees" in lowered
):
return (
"mint_error",
422,
"Token value is too small to cover swap fees",
"cashu_token_swap_fees_exceed_amount",
)
if "failed to melt" in lowered:
return (
"mint_error",
422,
"Failed to swap token from foreign mint",
"cashu_foreign_mint_swap_failed",
)
if ("invalid" in lowered or "decode" in lowered) and "token" in lowered:
# Anchored to "token" so internal faults whose text merely contains
# "invalid"/"decode" fall through to the 500 branch, not a token error.
return (
"invalid_token",
400,
"Invalid Cashu token",
"invalid_cashu_token",
)
if "must be positive" in lowered or "yielded no value" in lowered:
# Redeemed to <= 0 (empty/dust token, or value fully consumed by fees).
# Consumed, so non-retryable, but its own code — not the generic bucket.
return (
"cashu_error",
400,
"Failed to redeem Cashu token: token yielded no value",
"cashu_token_zero_value",
)
if isinstance(error, ValueError):
return (
"cashu_error",
400,
"Failed to redeem Cashu token",
"cashu_token_redemption_failed",
)
return None
async def get_balance(unit: str) -> int:
wallet = await get_wallet(settings.primary_mint, unit)
return wallet.available_balance.amount
@@ -219,12 +74,8 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int
"""Internal send function - returns amount and serialized token"""
effective_mint_url = mint_url or settings.primary_mint
wallet: Wallet = await get_wallet(effective_mint_url, unit)
all_proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
proofs = [proof for proof in all_proofs if not proof.reserved]
# Fallback must compare the requested amount with liquid proofs only. Counting
# reserved proofs here can suppress fallback even though they cannot be sent.
proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
proofs_for_mint = sum(p.amount for p in proofs)
reserved_for_mint = sum(p.amount for p in all_proofs if p.reserved)
# Fallback: proofs from untrusted source mints are swapped to primary_mint
# during receive, so the user's preferred refund_mint_url may have no proofs
@@ -236,10 +87,8 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int
)
effective_mint_url = settings.primary_mint
wallet = await get_wallet(effective_mint_url, unit)
all_proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
proofs = [proof for proof in all_proofs if not proof.reserved]
proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
proofs_for_mint = sum(p.amount for p in proofs)
reserved_for_mint = sum(p.amount for p in all_proofs if p.reserved)
all_mint_urls = list({k.mint_url for k in wallet.keysets.values()})
proof_summary = {
@@ -253,8 +102,7 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int
raw_proofs_by_keyset[p.id] = raw_proofs_by_keyset.get(p.id, 0) + p.amount
logger.info(
f"send: proof inventory | mint={effective_mint_url} unit={unit} amount={amount} "
f"primary_mint={settings.primary_mint} liquid_proofs_for_mint={proofs_for_mint} "
f"reserved_proofs_for_mint={reserved_for_mint} "
f"primary_mint={settings.primary_mint} proofs_for_mint={proofs_for_mint} "
f"all_mints={all_mint_urls} by_keyset={proof_summary} "
f"raw_proofs_by_keyset_id={raw_proofs_by_keyset} "
f"total_wallet_proofs={sum(p.amount for p in wallet.proofs)}"
@@ -410,8 +258,6 @@ async def _calculate_swap_amount(
"swap_to_primary_mint: fee estimation failed",
extra={"error": str(e)},
)
if is_mint_connection_error(e):
raise MintConnectionError("Cashu mint is unreachable") from e
raise ValueError(f"Failed to estimate fees: {e}") from e
@@ -537,13 +383,6 @@ async def swap_to_primary_mint(
quote_id=melt_quote.quote,
)
except Exception as e:
# A down mint won't fix itself by retrying with a smaller amount.
if is_mint_connection_error(e):
logger.error(
"swap_to_primary_mint: melt failed — mint unreachable",
extra={"error": str(e), "foreign_mint": token_obj.mint},
)
raise MintConnectionError("Cashu mint is unreachable") from e
shortfall = _melt_insufficient_shortfall(e)
recomputed = 0
if shortfall is not None:
@@ -619,21 +458,21 @@ async def swap_to_primary_mint(
# Recovery scan ran but did NOT restore the orphaned proofs
# (mint reports them as spent — they're stuck). Refuse to
# credit the API key balance for proofs we don't actually hold.
raise TokenConsumedError(
raise ValueError(
f"Swap recovery failed: mint signed outputs but proofs are "
f"unrecoverable (mint reports them spent). "
f"Expected {minted_amount}, recovered {balance_gained}. "
f"Local wallet DB ('.wallet/') state is corrupted — "
f"the counter for keyset is stuck at a bad index range."
)
except TokenConsumedError:
except ValueError:
raise
except Exception as recovery_err:
logger.error(
"swap_to_primary_mint: recovery failed",
extra={"error": str(recovery_err)},
)
raise TokenConsumedError(
raise ValueError(
f"Mint on primary failed and recovery unsuccessful: {e}"
) from e
else:
@@ -646,10 +485,7 @@ async def swap_to_primary_mint(
"mint_quote_id": mint_quote.quote,
},
)
# Foreign proofs already melted (spent) — non-retryable.
raise TokenConsumedError(
"Mint on primary failed after successful melt"
) from e
raise
logger.info(
"swap_to_primary_mint: completed successfully",
@@ -707,33 +543,19 @@ async def credit_balance(
extra={"old_balance": key.balance, "credit_amount": amount},
)
# The token is already redeemed (spent) here, so any crediting failure
# is post-redemption and non-retryable — surface it as TokenConsumedError
# (a key that vanished mid-flight, or an unexpected DB fault), never a
# retryable/token-error taxonomy.
try:
# Atomic UPDATE to prevent race conditions during concurrent topups.
stmt = (
update(db.ApiKey)
.where(col(db.ApiKey.hashed_key) == key.hashed_key)
.values(balance=(db.ApiKey.balance) + amount)
)
result = await session.exec(stmt) # type: ignore[call-overload]
# If pruning removed this key after redemption, do not commit a no-op
# balance update and pretend the top-up succeeded.
if (getattr(result, "rowcount", 0) or 0) == 0:
raise TokenConsumedError(
"Token redeemed but the API key disappeared before the "
"credit could be recorded"
)
await session.commit()
await session.refresh(key)
except TokenConsumedError:
raise
except Exception as db_error:
raise TokenConsumedError(
"Token redeemed but crediting the balance failed"
) from db_error
# Use atomic SQL UPDATE to prevent race conditions during concurrent topups
stmt = (
update(db.ApiKey)
.where(col(db.ApiKey.hashed_key) == key.hashed_key)
.values(balance=(db.ApiKey.balance) + amount)
)
result = await session.exec(stmt) # type: ignore[call-overload]
# If pruning removed this key after redemption, do not commit a no-op
# balance update and pretend the top-up succeeded.
if (getattr(result, "rowcount", 0) or 0) == 0:
raise ValueError("API key disappeared before credit could be recorded")
await session.commit()
await session.refresh(key)
logger.info(
"credit_balance: Balance updated successfully",
@@ -752,11 +574,11 @@ async def credit_balance(
)
except Exception:
pass
else:
logger.debug(
"Cashu token successfully redeemed and stored",
extra={"amount": amount, "unit": unit, "mint_url": mint_url},
)
logger.debug(
"Cashu token successfully redeemed and stored",
extra={"amount": amount, "unit": unit, "mint_url": mint_url},
)
return amount
except Exception as e:
logger.error(
@@ -855,7 +677,7 @@ async def fetch_all_balances(
"unit": unit,
"wallet_balance": proofs_balance,
"user_balance": user_balance,
"owner_balance": proofs_balance - user_balance,
"owner_balance": proofs_balance - user_balance if proofs_balance != 0 else 0,
}
return result
except Exception as e:
@@ -911,7 +733,9 @@ async def fetch_all_balances(
total_wallet_balance_sats += proofs_balance_sats
total_user_balance_sats += user_balance_sats
owner_balance = total_wallet_balance_sats - total_user_balance_sats
owner_balance = 0
if total_wallet_balance_sats != 0:
owner_balance = total_wallet_balance_sats - total_user_balance_sats
return (
balance_details,
@@ -924,132 +748,104 @@ async def fetch_all_balances(
async def periodic_payout() -> None:
while True:
await asyncio.sleep(settings.payout_interval_seconds)
print(settings.payout_interval_seconds)
if not settings.receive_ln_address:
continue
try:
if not settings.receive_ln_address:
continue
# Include the primary mint even if it is not listed in cashu_mints,
# matching fetch_all_balances(); otherwise primary-mint funds never
# auto-payout.
mint_urls: list[str] = list(settings.cashu_mints)
if settings.primary_mint and settings.primary_mint not in mint_urls:
mint_urls.append(settings.primary_mint)
async with db.create_session() as session:
for mint_url in mint_urls:
for mint_url in settings.cashu_mints:
for unit in ["sat", "msat"]:
# Isolate failures per mint/unit so one slow or failing
# mint does not abort payout for every other mint/unit.
try:
wallet = await get_wallet(mint_url, unit)
proofs = get_proofs_per_mint_and_unit(
wallet, mint_url, unit, not_reserved=True
wallet = await get_wallet(mint_url, unit)
proofs = get_proofs_per_mint_and_unit(
wallet, mint_url, unit, not_reserved=True
)
proofs = await slow_filter_spend_proofs(proofs, wallet)
await asyncio.sleep(5)
user_balance = await db.balances_for_mint_and_unit(
session, mint_url, unit
)
if unit == "sat":
user_balance = user_balance // 1000
proofs_balance = sum(proof.amount for proof in proofs)
available_balance = proofs_balance - user_balance
# Threshold is configured in sats; convert for msat wallets.
min_amount = (
settings.min_payout_sat
if unit == "sat"
else settings.min_payout_sat * 1000
)
if available_balance > min_amount:
amount_received = await raw_send_to_lnurl(
wallet,
proofs,
settings.receive_ln_address,
unit,
amount=available_balance,
)
proofs = await slow_filter_spend_proofs(proofs, wallet)
await asyncio.sleep(5)
user_balance = await db.balances_for_mint_and_unit(
session, mint_url, unit
)
if unit == "sat":
user_balance = user_balance // 1000
proofs_balance = sum(proof.amount for proof in proofs)
available_balance = proofs_balance - user_balance
# Threshold is configured in sats; convert for msat wallets.
min_amount = (
settings.min_payout_sat
if unit == "sat"
else settings.min_payout_sat * 1000
)
if available_balance > min_amount:
amount_received = await raw_send_to_lnurl(
wallet,
proofs,
settings.receive_ln_address,
unit,
amount=available_balance,
)
logger.info(
"Payout sent successfully",
extra={
"mint_url": mint_url,
"unit": unit,
"balance": available_balance,
"amount_received": amount_received,
},
)
except Exception as e:
logger.error(
f"Error sending payout: {type(e).__name__}",
logger.info(
"Payout sent successfully",
extra={
"error": str(e),
"mint_url": mint_url,
"unit": unit,
"balance": available_balance,
"amount_received": amount_received,
},
)
except Exception as e:
logger.error(
f"Error in periodic payout cycle: {type(e).__name__}",
f"Error sending payout: {type(e).__name__}",
extra={"error": str(e)},
)
async def _refund_sweep_once(cutoff: int) -> None:
async with db.create_session() as session:
stmt = select(db.CashuTransaction).where(
db.CashuTransaction.type == "out",
db.CashuTransaction.collected == False, # noqa: E712
db.CashuTransaction.swept == False, # noqa: E712
db.CashuTransaction.created_at < cutoff,
)
results = await session.exec(stmt)
refunds = results.all()
for refund in refunds:
try:
await recieve_token(refund.token)
refund.swept = True
session.add(refund)
logger.info(
"Swept uncollected refund",
extra={
"id": refund.id,
"amount": refund.amount,
"unit": refund.unit,
},
)
except Exception as e:
error_msg = str(e).lower()
if "already spent" in error_msg:
refund.collected = True
session.add(refund)
logger.info(
"Refund already spent (client collected), marking swept",
extra={
"id": refund.id,
},
)
else:
logger.warning(
"Failed to sweep refund",
extra={
"id": refund.id,
"error": str(e),
},
)
await session.commit()
async def refund_sweep_once() -> None:
"""Sweep eligible uncollected refund tokens once."""
cutoff = int(time.time()) - settings.refund_sweep_ttl_seconds
await _refund_sweep_once(cutoff)
async def periodic_refund_sweep() -> None:
while True:
await asyncio.sleep(60 * 60) # every hour
try:
await refund_sweep_once()
cutoff = int(time.time()) - settings.refund_sweep_ttl_seconds
async with db.create_session() as session:
stmt = select(db.CashuTransaction).where(
db.CashuTransaction.type == "out",
db.CashuTransaction.collected == False, # noqa: E712
db.CashuTransaction.swept == False, # noqa: E712
db.CashuTransaction.created_at < cutoff,
)
results = await session.exec(stmt)
refunds = results.all()
for refund in refunds:
try:
await recieve_token(refund.token)
refund.swept = True
session.add(refund)
logger.info(
"Swept uncollected refund",
extra={
"id": refund.id,
"amount": refund.amount,
"unit": refund.unit,
},
)
except Exception as e:
error_msg = str(e).lower()
if "already spent" in error_msg:
refund.collected = True
session.add(refund)
logger.info(
"Refund already spent (client collected), marking swept",
extra={
"id": refund.id,
},
)
else:
logger.warning(
"Failed to sweep refund",
extra={
"id": refund.id,
"error": str(e),
},
)
await session.commit()
except Exception as e:
logger.error(
"Error in periodic refund sweep",
@@ -1072,56 +868,17 @@ async def periodic_routstr_fee_payout() -> None:
try:
async with db.create_session() as session:
fee = await db.get_routstr_fee(session)
if fee.payout_in_progress_msats:
logger.critical(
"Routstr fee payout requires manual reconciliation",
extra={
"payout_in_progress_msats": fee.payout_in_progress_msats,
"payout_started_at": fee.payout_started_at,
},
)
continue
accumulated_sats = fee.accumulated_msats // 1000
if accumulated_sats >= ROUTSTR_FEE_DEFAULT_PAYOUT:
wallet = await get_wallet(settings.primary_mint, "sat")
proofs = get_proofs_per_mint_and_unit(
wallet, settings.primary_mint, "sat", not_reserved=True
)
amount_received = await raw_send_to_lnurl(
wallet, proofs, ROUTSTR_LN_ADDRESS, "sat", amount=accumulated_sats
)
paid_msats = accumulated_sats * 1000
payout_checkpointed = await db.reset_routstr_fee(
session, paid_msats
)
if not payout_checkpointed:
logger.warning("Routstr fee payout was already claimed")
continue
try:
amount_received = await raw_send_to_lnurl(
wallet,
proofs,
ROUTSTR_LN_ADDRESS,
"sat",
amount=accumulated_sats,
)
except Exception:
logger.critical(
"Routstr fee payout outcome is unknown; manual reconciliation required",
extra={"payout_in_progress_msats": paid_msats},
exc_info=True,
)
continue
payout_completed = await db.complete_routstr_fee_payout(
session, paid_msats
)
if not payout_completed:
logger.critical(
"Routstr fee payout sent but checkpoint was not completed",
extra={"payout_in_progress_msats": paid_msats},
)
continue
await db.reset_routstr_fee(session, paid_msats)
logger.info(
"Routstr fee payout sent",
extra={
-1
View File
@@ -1 +0,0 @@
"""Operational and recovery scripts for routstr-core (not a shipped package)."""
-104
View File
@@ -1,104 +0,0 @@
"""Recover admin access by resetting the stored admin password (issue #553).
The lockout escape hatch for an operator who has lost the admin password. It
talks to the ``secrets`` table directly and deliberately does *not* require
``ROUTSTR_SECRET_KEY``: the admin password is scrypt-hashed (key-independent),
so recovery works even when the encryption key is missing or has changed.
Two explicit, mutually exclusive actions running with no arguments only prints
help, so the password can't be reset by accident:
python scripts/reset_admin_password.py --password <new-password>
Hash and store <new-password> now.
python scripts/reset_admin_password.py --regenerate
Clear the stored hash; the next node startup generates a fresh random
password and logs it once (with the /admin URL).
"""
import argparse
import asyncio
import sys
import time
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import create_session, get_secret, set_admin_password
from routstr.core.vault import MIN_PASSWORD_LENGTH
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
prog="reset_admin_password",
description="Reset the node's admin password (recovery from lockout).",
)
action = parser.add_mutually_exclusive_group()
action.add_argument(
"--password",
metavar="NEW_PASSWORD",
help=f"set this as the new admin password (min {MIN_PASSWORD_LENGTH} chars)",
)
action.add_argument(
"--regenerate",
action="store_true",
help="clear the password so the next startup generates and logs a new one",
)
return parser
async def apply_reset(
session: AsyncSession,
*,
password: str | None = None,
regenerate: bool = False,
) -> str:
"""Perform the requested reset against ``session``; return a status message."""
if password is not None:
if len(password) < MIN_PASSWORD_LENGTH:
raise ValueError(
f"New password must be at least {MIN_PASSWORD_LENGTH} characters"
)
await set_admin_password(session, password)
return "Admin password updated."
if regenerate:
secret = await get_secret(session)
secret.admin_password_hash = None
secret.updated_at = int(time.time())
session.add(secret)
await session.commit()
return (
"Admin password cleared. The next node startup will generate a new "
"one and log it once with the /admin URL."
)
return ""
async def _run(password: str | None, regenerate: bool) -> str:
async with create_session() as session:
return await apply_reset(
session, password=password, regenerate=regenerate
)
def main(argv: list[str] | None = None) -> int:
parser = build_parser()
args = parser.parse_args(argv)
if args.password is None and not args.regenerate:
parser.print_help()
return 0
try:
message = asyncio.run(_run(args.password, args.regenerate))
except ValueError as exc:
print(f"error: {exc}", file=sys.stderr)
return 2
print(message)
return 0
if __name__ == "__main__":
raise SystemExit(main())
-15
View File
@@ -1,15 +0,0 @@
"""Shared pytest configuration for the whole suite.
A fixed, valid ``ROUTSTR_SECRET_KEY`` is set before any app import so that
secret encryption is deterministic across the suite and the mandatory-key
fail-fast does not break app-boot tests. Tests that need a different key (or an
absent one) override this per-test via ``monkeypatch``.
"""
import os
# 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)
-121
View File
@@ -1,121 +0,0 @@
"""Tests for admin password auth backed by the hashed Secret store (issue #553).
Login and password change verify against the one-way ``Secret.admin_password_hash``
(scrypt) instead of a plaintext settings field, which also closes the old ``!=``
timing-attack comparison. The first hash is written by ``bootstrap_secrets`` at
startup (generated, or migrated from a legacy password); here it is seeded
directly into the store, then login checks against it and a password change
re-hashes so the old password stops working and the new one starts.
"""
from __future__ import annotations
import pytest
from httpx import AsyncClient, Response
from routstr.core.db import create_session, set_admin_password
async def _seed_password(password: str) -> None:
# bootstrap_secrets owns first-password creation at startup; tests seed the
# hash straight into the store the same way, then exercise login/change.
async with create_session() as session:
await set_admin_password(session, password)
async def _login(client: AsyncClient, password: str) -> Response:
return await client.post("/admin/api/login", json={"password": password})
# --- login -----------------------------------------------------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_login_500_when_no_password_configured(
integration_client: AsyncClient,
) -> None:
resp = await _login(integration_client, "anything")
assert resp.status_code == 500
@pytest.mark.integration
@pytest.mark.asyncio
async def test_login_succeeds_with_correct_password(
integration_client: AsyncClient,
) -> None:
await _seed_password("correct horse")
resp = await _login(integration_client, "correct horse")
assert resp.status_code == 200
body = resp.json()
assert body["ok"] is True
assert isinstance(body["token"], str) and body["token"]
@pytest.mark.integration
@pytest.mark.asyncio
async def test_login_rejects_wrong_password(
integration_client: AsyncClient,
) -> None:
await _seed_password("correct horse")
resp = await _login(integration_client, "wrong horse")
assert resp.status_code == 401
# --- password change -------------------------------------------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_update_password_rehashes_so_only_new_works(
integration_client: AsyncClient,
) -> None:
await _seed_password("old password")
login = await _login(integration_client, "old password")
token = login.json()["token"]
integration_client.headers["Authorization"] = f"Bearer {token}"
resp = await integration_client.patch(
"/admin/api/password",
json={"current_password": "old password", "new_password": "new password"},
)
assert resp.status_code == 200, resp.text
# Drop admin auth so the login calls aren't treated as authenticated noise.
integration_client.headers.pop("Authorization", None)
assert (await _login(integration_client, "old password")).status_code == 401
assert (await _login(integration_client, "new password")).status_code == 200
@pytest.mark.integration
@pytest.mark.asyncio
async def test_update_password_rejects_wrong_current(
integration_client: AsyncClient,
) -> None:
await _seed_password("old password")
login = await _login(integration_client, "old password")
token = login.json()["token"]
integration_client.headers["Authorization"] = f"Bearer {token}"
resp = await integration_client.patch(
"/admin/api/password",
json={"current_password": "not the password", "new_password": "new password"},
)
assert resp.status_code == 401
@pytest.mark.integration
@pytest.mark.asyncio
async def test_update_password_rejects_short_new(
integration_client: AsyncClient,
) -> None:
await _seed_password("old password")
login = await _login(integration_client, "old password")
token = login.json()["token"]
integration_client.headers["Authorization"] = f"Bearer {token}"
resp = await integration_client.patch(
"/admin/api/password",
json={"current_password": "old password", "new_password": "x"},
)
assert resp.status_code == 400
@@ -1,229 +0,0 @@
import json
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.payment.cost_calculation import CostData, calculate_cost
from routstr.proxy import get_model_instance, reinitialize_upstreams
def _admin_headers() -> dict[str, str]:
token = "test-admin-cache-pricing-token"
admin_sessions[token] = int(
(datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp()
)
return {"Authorization": f"Bearer {token}"}
def _model_payload(
provider_id: int,
*,
cache_read: float,
cache_write: float,
) -> dict[str, object]:
return {
"id": "custom-cache-model",
"name": "Custom Cache Model",
"description": "custom model with explicit cache pricing",
"created": 0,
"context_length": 128000,
"architecture": {
"modality": "text",
"input_modalities": ["text"],
"output_modalities": ["text"],
"tokenizer": "unknown",
"instruct_type": None,
},
"pricing": {
"prompt": 1.4e-7,
"completion": 2.8e-7,
"input_cache_read": cache_read,
"input_cache_write": cache_write,
"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": True,
"forwarded_model_id": "custom-cache-model",
}
@pytest.mark.integration
@pytest.mark.asyncio
async def test_admin_provider_model_api_persists_cache_pricing_on_create_and_update(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
provider = UpstreamProviderRow(
provider_type="generic",
base_url="https://custom-upstream.example/v1",
api_key="test-key",
provider_fee=1.0,
)
integration_session.add(provider)
await integration_session.commit()
await integration_session.refresh(provider)
assert provider.id is not None
await reinitialize_upstreams()
headers = _admin_headers()
create_payload = _model_payload(
provider.id,
cache_read=2.8e-9,
cache_write=3.5e-9,
)
with patch("routstr.payment.models.sats_usd_price", return_value=1e-6):
create_response = await integration_client.post(
f"/admin/api/upstream-providers/{provider.id}/models",
headers=headers,
json=create_payload,
)
assert create_response.status_code == 200
create_body = create_response.json()
assert create_body["pricing"]["input_cache_read"] == pytest.approx(2.8e-9)
assert create_body["pricing"]["input_cache_write"] == pytest.approx(3.5e-9)
row = await integration_session.get(ModelRow, ("custom-cache-model", provider.id))
assert row is not None
stored_pricing = json.loads(row.pricing)
assert stored_pricing["input_cache_read"] == pytest.approx(2.8e-9)
assert stored_pricing["input_cache_write"] == pytest.approx(3.5e-9)
update_payload = _model_payload(
provider.id,
cache_read=1.25e-9,
cache_write=4.5e-9,
)
with patch("routstr.payment.models.sats_usd_price", return_value=1e-6):
update_response = await integration_client.post(
f"/admin/api/upstream-providers/{provider.id}/models",
headers=headers,
json=update_payload,
)
assert update_response.status_code == 200
update_body = update_response.json()
assert update_body["pricing"]["input_cache_read"] == pytest.approx(1.25e-9)
assert update_body["pricing"]["input_cache_write"] == pytest.approx(4.5e-9)
await integration_session.refresh(row)
updated_pricing = json.loads(row.pricing)
assert updated_pricing["input_cache_read"] == pytest.approx(1.25e-9)
assert updated_pricing["input_cache_write"] == pytest.approx(4.5e-9)
model = get_model_instance("custom-cache-model")
assert model is not None
assert model.sats_pricing is not None
assert model.sats_pricing.input_cache_read == pytest.approx(0.00125)
assert model.sats_pricing.input_cache_write == pytest.approx(0.0045)
with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=1e-6):
cost = await calculate_cost(
{
"model": "custom-cache-model",
"usage": {
"prompt_tokens": 1000,
"completion_tokens": 100,
"prompt_tokens_details": {"cached_tokens": 800},
},
},
max_cost=1_000_000,
)
assert isinstance(cost, CostData)
assert cost.input_tokens == 200
assert cost.cache_read_input_tokens == 800
assert cost.cache_read_msats == 1000
assert cost.output_msats == 28000
assert cost.input_msats == 29000
assert cost.total_msats == 57000
@pytest.mark.integration
@pytest.mark.asyncio
async def test_upstream_response_cost_uses_model_cache_pricing(
integration_session: AsyncSession,
patched_db_engine: None,
) -> None:
"""A model's configured cache price must discount upstream cached-token usage."""
provider = UpstreamProviderRow(
provider_type="generic",
base_url="https://cache-priced-upstream.example/v1",
api_key="test-key",
provider_fee=1.0,
)
integration_session.add(provider)
await integration_session.commit()
await integration_session.refresh(provider)
assert provider.id is not None
row = ModelRow(
id="cache-priced-model",
name="Cache Priced Model",
description="model seeded with explicit cache pricing",
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": 1.4e-7,
"completion": 2.8e-7,
"input_cache_read": 1.25e-9,
"input_cache_write": 4.5e-9,
}
),
upstream_provider_id=provider.id,
enabled=True,
forwarded_model_id="cache-priced-model",
)
integration_session.add(row)
await integration_session.commit()
with patch("routstr.payment.models.sats_usd_price", return_value=1e-6):
await reinitialize_upstreams()
with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=1e-6):
cost = await calculate_cost(
{
"model": "cache-priced-model",
"usage": {
"prompt_tokens": 1000,
"completion_tokens": 100,
"prompt_tokens_details": {"cached_tokens": 800},
},
},
max_cost=1_000_000,
)
assert isinstance(cost, CostData)
# prompt 1.4e-7 USD/token -> 0.14 sats/token -> 140 msats/token
# cache 1.25e-9 USD/token -> 0.00125 sats/token -> 1.25 msats/token
# completion 2.8e-7 USD/token -> 0.28 sats/token -> 280 msats/token
assert cost.input_tokens == 200
assert cost.cache_read_input_tokens == 800
assert cost.cache_read_msats == 1000
assert cost.input_msats == 29000
assert cost.output_msats == 28000
assert cost.total_msats == 57000
@@ -1,113 +0,0 @@
"""Tests for the admin nsec rotation endpoint (issue #553).
The Nostr identity is a secret: it lives encrypted in the Secret store, never in
the settings blob, so it cannot be set through the general settings PATCH. This
dedicated endpoint is the supported way to set/rotate/clear it it encrypts the
key at rest, updates the live runtime identity (so signing picks it up without a
restart), and derives the npub. Invalid keys are rejected.
"""
from __future__ import annotations
import secrets
import time
from collections.abc import AsyncGenerator
import pytest
import pytest_asyncio
from httpx import AsyncClient
from routstr.core import vault
from routstr.core.admin import admin_sessions
from routstr.core.db import AsyncSession, get_secret
from routstr.core.settings import derive_npub_from_nsec, settings
# A valid 64-char hex private key (accepted by nsec_to_keypair, as in bootstrap).
NSEC_HEX = "1" * 64
@pytest_asyncio.fixture
async def admin_client(
integration_client: AsyncClient,
) -> AsyncGenerator[AsyncClient, None]:
"""An integration_client pre-authenticated with an admin session token."""
token = secrets.token_urlsafe(24)
admin_sessions[token] = int(time.time()) + 3600
integration_client.headers["Authorization"] = f"Bearer {token}"
yield integration_client
admin_sessions.pop(token, None)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_update_nsec_stores_encrypted_and_derives_npub(
admin_client: AsyncClient,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "nsec", "")
monkeypatch.setattr(settings, "npub", "")
resp = await admin_client.patch("/admin/api/nsec", json={"nsec": NSEC_HEX})
assert resp.status_code == 200
expected_npub = derive_npub_from_nsec(NSEC_HEX)
assert resp.json() == {"ok": True, "npub": expected_npub}
# Stored encrypted at rest, decryptable back to the original key.
integration_session.expunge_all()
secret = await get_secret(integration_session)
assert secret.encrypted_nsec is not None
assert vault.is_encrypted(secret.encrypted_nsec)
assert vault.decrypt(secret.encrypted_nsec) == NSEC_HEX
# Live runtime identity updated so Nostr signing reflects it without restart.
assert settings.nsec == NSEC_HEX
assert settings.npub == expected_npub
@pytest.mark.integration
@pytest.mark.asyncio
async def test_update_nsec_rejects_invalid_key(
admin_client: AsyncClient,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "nsec", "")
monkeypatch.setattr(settings, "npub", "")
resp = await admin_client.patch(
"/admin/api/nsec", json={"nsec": "not-a-real-nsec"}
)
assert resp.status_code == 400
# Nothing stored, live identity untouched.
integration_session.expunge_all()
secret = await get_secret(integration_session)
assert secret.encrypted_nsec is None
assert settings.nsec == ""
@pytest.mark.integration
@pytest.mark.asyncio
async def test_update_nsec_clears_identity_with_empty_value(
admin_client: AsyncClient,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
# Start from a node that has an identity...
monkeypatch.setattr(settings, "nsec", "")
monkeypatch.setattr(settings, "npub", "")
set_resp = await admin_client.patch("/admin/api/nsec", json={"nsec": NSEC_HEX})
assert set_resp.status_code == 200
# ...then clear it.
clear_resp = await admin_client.patch("/admin/api/nsec", json={"nsec": ""})
assert clear_resp.status_code == 200
assert clear_resp.json() == {"ok": True, "npub": ""}
integration_session.expunge_all()
secret = await get_secret(integration_session)
assert secret.encrypted_nsec is None
assert settings.nsec == ""
assert settings.npub == ""
@@ -1,82 +0,0 @@
"""Tests for the admin settings endpoint's handling of secrets (issue #553).
``admin_password`` is no longer a settings field (it lives only as a one-way
hash in the Secret store), so it must never appear in the GET/PATCH payloads.
``nsec`` (in-memory at runtime) and ``upstream_api_key`` (still in the settings
blob) are both redacted on read and ignored on write they cannot be set
through the general settings endpoint, only through their dedicated paths.
"""
from __future__ import annotations
import secrets
import time
from collections.abc import AsyncGenerator
import pytest
import pytest_asyncio
from httpx import AsyncClient
from routstr.core.admin import admin_sessions
from routstr.core.db import AsyncSession
from routstr.core.settings import SettingsService, settings
@pytest_asyncio.fixture
async def admin_client(
integration_client: AsyncClient,
) -> AsyncGenerator[AsyncClient, None]:
"""An integration_client pre-authenticated with an admin session token."""
token = secrets.token_urlsafe(24)
admin_sessions[token] = int(time.time()) + 3600
integration_client.headers["Authorization"] = f"Bearer {token}"
yield integration_client
admin_sessions.pop(token, None)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_get_settings_omits_admin_password_and_redacts_secrets(
admin_client: AsyncClient, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(settings, "nsec", "nsec-secret")
monkeypatch.setattr(settings, "upstream_api_key", "sk-secret")
resp = await admin_client.get("/admin/api/settings")
assert resp.status_code == 200
data = resp.json()
assert "admin_password" not in data
assert data["nsec"] == "[REDACTED]"
assert data["upstream_api_key"] == "[REDACTED]"
@pytest.mark.integration
@pytest.mark.asyncio
async def test_patch_settings_ignores_secret_fields(
admin_client: AsyncClient,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
# The PATCH path persists through SettingsService, which needs an
# initialized current snapshot and a settings row in the shared test DB.
await SettingsService.initialize(integration_session)
monkeypatch.setattr(settings, "nsec", "original-nsec")
resp = await admin_client.patch(
"/admin/api/settings",
json={
"name": "Renamed",
"nsec": "attacker-nsec",
"upstream_api_key": "attacker-key",
"admin_password": "attacker-pw",
},
)
assert resp.status_code == 200
data = resp.json()
assert data["name"] == "Renamed"
assert "admin_password" not in data
assert data["nsec"] == "[REDACTED]"
# The live secret was not overwritten through the general settings endpoint.
assert settings.nsec == "original-nsec"
@@ -15,7 +15,6 @@ from unittest.mock import patch
import pytest
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.auth import ReservationSnapshot
from routstr.core.db import ApiKey
from routstr.payment.cost_calculation import CostData
@@ -24,7 +23,7 @@ def _make_key(balance: int, reserved: int) -> ApiKey:
return ApiKey(
hashed_key=f"test_{uuid.uuid4().hex}",
balance=balance,
reserved_balance=0,
reserved_balance=reserved,
total_spent=0,
total_requests=1,
)
@@ -76,8 +75,6 @@ async def test_balance_never_negative_when_cost_exceeds_reservation(
key = _make_key(balance=deducted_max_cost, reserved=deducted_max_cost)
integration_session.add(key)
await integration_session.commit()
from routstr.auth import pay_for_request
await pay_for_request(key, deducted_max_cost, integration_session)
response_data = {"model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}}
@@ -85,9 +82,7 @@ async def test_balance_never_negative_when_cost_exceeds_reservation(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
key, response_data, integration_session, deducted_max_cost, None, None
)
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
await _refresh(integration_session, key)
@@ -114,8 +109,6 @@ async def test_balance_floor_at_zero_on_overrun(
key = _make_key(balance=500, reserved=500)
integration_session.add(key)
await integration_session.commit()
from routstr.auth import pay_for_request
await pay_for_request(key, deducted_max_cost, integration_session)
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}}
@@ -123,9 +116,7 @@ async def test_balance_floor_at_zero_on_overrun(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
key, response_data, integration_session, deducted_max_cost, None, None
)
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
await _refresh(integration_session, key)
@@ -157,8 +148,6 @@ async def test_full_cost_charged_when_balance_sufficient_for_overrun(
key = _make_key(balance=2000, reserved=990)
integration_session.add(key)
await integration_session.commit()
from routstr.auth import pay_for_request
await pay_for_request(key, deducted_max_cost, integration_session)
response_data = {"model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}}
@@ -166,9 +155,7 @@ async def test_full_cost_charged_when_balance_sufficient_for_overrun(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
key, response_data, integration_session, deducted_max_cost, None, None
)
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
await _refresh(integration_session, key)
@@ -197,11 +184,7 @@ async def test_concurrent_cost_overruns_never_negative(
"""Concurrent finalization with cost overruns must never produce negative balance."""
import asyncio
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
from routstr.auth import adjust_payment_for_tokens, pay_for_request
from routstr.core.db import create_session
deducted_max_cost = 990
@@ -227,14 +210,12 @@ async def test_concurrent_cost_overruns_never_negative(
async with create_session() as session:
key_to_reserve = await session.get(ApiKey, key_hash)
assert key_to_reserve is not None
reservations = []
for _ in range(n_requests):
await pay_for_request(key_to_reserve, deducted_max_cost, session)
reservations.append(await get_reservation_snapshot(key_to_reserve, session))
await session.refresh(key_to_reserve)
# Now finalize all concurrently with cost overrun
async def finalize(reservation: ReservationSnapshot) -> None:
async def finalize() -> None:
response_data = {
"model": "test-model",
"usage": {"prompt_tokens": 100, "completion_tokens": 100},
@@ -242,22 +223,15 @@ async def test_concurrent_cost_overruns_never_negative(
async with create_session() as session:
fresh_key = await session.get(ApiKey, key_hash)
assert fresh_key is not None
await adjust_payment_for_tokens(
fresh_key,
response_data,
session,
deducted_max_cost,
reservation_snapshot=reservation,
)
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
fresh_key, response_data, session, deducted_max_cost
)
# Patch once around the gather: entering the same patch target from
# concurrent tasks un-patches in the wrong order and leaks the mock into
# every later test in the session.
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await asyncio.gather(*(finalize(r) for r in reservations))
await asyncio.gather(*[finalize() for _ in range(n_requests)])
async with create_session() as session:
final_key = await session.get(ApiKey, key_hash)
@@ -298,8 +272,6 @@ async def test_zero_free_balance_overrun_is_safe(
key = _make_key(balance=1000, reserved=1000)
integration_session.add(key)
await integration_session.commit()
from routstr.auth import pay_for_request
await pay_for_request(key, deducted_max_cost, integration_session)
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 100}}
@@ -307,9 +279,7 @@ async def test_zero_free_balance_overrun_is_safe(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
key, response_data, integration_session, deducted_max_cost, None, None
)
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
await _refresh(integration_session, key)
@@ -338,11 +308,7 @@ async def test_parallel_requests_no_free_inference(
"""Second parallel finalization must be charged even when first depleted free balance."""
import asyncio
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
from routstr.auth import adjust_payment_for_tokens
from routstr.core.db import create_session
deducted_max_cost = 100
@@ -363,18 +329,14 @@ async def test_parallel_requests_no_free_inference(
key = ApiKey(
hashed_key=key_hash,
balance=starting_balance,
reserved_balance=0,
reserved_balance=deducted_max_cost * 2, # both slots pre-reserved
total_spent=0,
total_requests=2,
)
session.add(key)
await session.commit()
reservations = []
for _ in range(2):
await pay_for_request(key, deducted_max_cost, session)
reservations.append(await get_reservation_snapshot(key, session))
async def finalize(reservation: ReservationSnapshot) -> None:
async def finalize() -> None:
response_data = {
"model": "test-model",
"usage": {"prompt_tokens": 50, "completion_tokens": 100},
@@ -382,22 +344,15 @@ async def test_parallel_requests_no_free_inference(
async with create_session() as session:
fresh_key = await session.get(ApiKey, key_hash)
assert fresh_key is not None
await adjust_payment_for_tokens(
fresh_key,
response_data,
session,
deducted_max_cost,
reservation_snapshot=reservation,
)
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
fresh_key, response_data, session, deducted_max_cost
)
# Patch once around the gather: entering the same patch target from two
# concurrent tasks un-patches in the wrong order and leaks the mock into
# every later test in the session.
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await asyncio.gather(*(finalize(r) for r in reservations))
await asyncio.gather(finalize(), finalize())
async with create_session() as session:
final_key = await session.get(ApiKey, key_hash)
+1 -1
View File
@@ -77,7 +77,7 @@ async def test_child_key_flow(integration_session: AsyncSession) -> None:
try:
adjustment = await adjust_payment_for_tokens(
child_key_db, response_data, integration_session, 500, None, None
child_key_db, response_data, integration_session, 500
)
assert adjustment["total_msats"] == 400
-666
View File
@@ -1,666 +0,0 @@
"""Failover requests are billed and forwarded as the provider that served them.
Covers the whole-system settlement path when two enabled providers expose the
same model under different spellings and prices: the routing winner fails with
a 502, the fallback provider serves, and the response must be billed at the
fallback's configured rate, carry the fallback's model id in the forwarded
request body, and echo the fallback's model id to the client.
"""
import json
from typing import Any, AsyncGenerator
from unittest.mock import patch
import httpx
import pytest
from httpx import AsyncClient
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ApiKey, ReservationRelease
from routstr.payment.models import Architecture, Model, Pricing
from routstr.proxy import refresh_model_maps
from routstr.upstream.base import BaseUpstreamProvider
CHEAP_BASE_URL = "https://cheap.example.com/v1"
EXPENSIVE_BASE_URL = "https://expensive.example.com/v1"
THIRD_BASE_URL = "https://third.example.com/v1"
def _make_model(
model_id: str,
prompt_sats: float,
completion_sats: float,
max_cost: float = 50.0,
) -> Model:
"""Build a model whose USD and sats pricing rank consistently."""
return Model(
id=model_id,
name=model_id,
created=1,
description="test model",
context_length=8192,
architecture=Architecture(
modality="text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="gpt",
instruct_type=None,
),
pricing=Pricing(
prompt=prompt_sats, completion=completion_sats, max_cost=max_cost
),
sats_pricing=Pricing(
prompt=prompt_sats, completion=completion_sats, max_cost=max_cost
),
)
class _StaticProvider(BaseUpstreamProvider):
"""Upstream provider with a fixed model catalog and no remote refresh."""
def __init__(self, base_url: str, api_key: str, fee: float, model: Model) -> None:
super().__init__(base_url, api_key, fee)
self.provider_type = "custom"
self._static_model = model
def get_cached_models(self) -> list[Model]:
return [self._static_model]
async def refresh_models_cache(self) -> None:
pass
async def _install_providers(
providers: list[_StaticProvider],
) -> AsyncGenerator[None, None]:
"""Install providers into the routing maps, restoring the originals after."""
from routstr import proxy
original_upstreams = proxy.get_upstreams()
with patch("routstr.proxy._upstreams", providers):
await refresh_model_maps()
yield
with patch("routstr.proxy._upstreams", original_upstreams):
await refresh_model_maps()
@pytest.fixture
async def dual_provider_maps(
patched_db_engine: None,
) -> AsyncGenerator[tuple[_StaticProvider, _StaticProvider], None]:
"""Two same-tail providers under different spellings and prices."""
cheap = _StaticProvider(
CHEAP_BASE_URL,
"key-cheap",
1.0,
_make_model("prova/dual-model", 0.001, 0.002),
)
expensive = _StaticProvider(
EXPENSIVE_BASE_URL,
"key-expensive",
1.0,
_make_model("provb/dual-model", 0.005, 0.010, max_cost=100.0),
)
async for _ in _install_providers([cheap, expensive]):
yield cheap, expensive
def _upstream_response(request: httpx.Request) -> httpx.Response:
"""502 from the cheap (winning) provider; a served completion elsewhere."""
if request.url.host == "cheap.example.com":
return httpx.Response(
502,
content=json.dumps({"error": {"message": "bad gateway"}}).encode(),
headers={"content-type": "application/json"},
)
body = {
"id": "chatcmpl-served",
"object": "chat.completion",
"created": 1,
"model": "dual-model",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "hi"},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 1000,
"completion_tokens": 500,
"total_tokens": 1500,
},
}
return httpx.Response(
200,
content=json.dumps(body).encode(),
headers={"content-type": "application/json"},
)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_failover_serve_billed_at_serving_providers_rate(
authenticated_client: AsyncClient,
dual_provider_maps: tuple[_StaticProvider, _StaticProvider],
integration_session: AsyncSession,
) -> None:
"""A fallback serve is billed at the fallback's price, not the winner's.
The cheap provider ranks first for the shared tail; it 502s and the
expensive provider serves 1000 input + 500 output tokens. At the serving
provider's sats pricing (0.005/0.010 sats per token) that is 10_000 msats;
at the winner's (0.001/0.002) it would be 2_000 msats.
"""
sent_requests: list[httpx.Request] = []
# Patch the network transport (not AsyncClient.send) so the in-process
# ASGI test client is untouched and only the proxy's upstream hop is mocked.
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,
),
# cost_calculation binds sats_usd_price at import time, so the price
# patch in the app fixture does not reach it; patch its own binding.
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
payload = response.json()
# Both providers were attempted, cheapest first.
assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"expensive.example.com",
]
# The fallback must be asked for ITS OWN model spelling, not the winner's.
forwarded_body = json.loads(sent_requests[1].content)
assert forwarded_body["model"] == "provb/dual-model"
# The response echo names the model that actually served.
assert payload["model"] == "provb/dual-model"
# Billed at the serving provider's rate: 1000/1000*5000 + 500/1000*10000.
assert payload["cost"]["total_msats"] == 10_000
# The fallback's larger max-cost envelope requires a replacement
# reservation. The failed candidate is released, the serving candidate is
# charged, and no request-owned reservation remains active.
key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined]
records = (
await integration_session.exec(
select(ReservationRelease).where(ReservationRelease.key_hash == key_hash)
)
).all()
assert sorted(record.status for record in records) == ["charged", "released"]
@pytest.fixture
async def same_id_provider_maps(
patched_db_engine: None,
) -> AsyncGenerator[None, None]:
"""Two providers exposing the IDENTICAL model id at different prices."""
cheap = _StaticProvider(
CHEAP_BASE_URL,
"key-cheap",
1.0,
_make_model("dual-model", 0.001, 0.002),
)
expensive = _StaticProvider(
EXPENSIVE_BASE_URL,
"key-expensive",
1.0,
_make_model("dual-model", 0.005, 0.010),
)
async for _ in _install_providers([cheap, expensive]):
yield
@pytest.mark.integration
@pytest.mark.asyncio
async def test_same_id_failover_settles_at_serving_price(
authenticated_client: AsyncClient,
same_id_provider_maps: None,
) -> None:
"""Settlement must not re-derive pricing from the response's model string.
Both providers expose the exact same model id, so the forwarded body is
identical either way the only observable difference is the settled
amount. The response's model string resolves to the alias winner (cheap),
but the expensive provider served, so the bill must be 10_000 msats, not
the winner's 2_000.
"""
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
assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"expensive.example.com",
]
assert response.json()["cost"]["total_msats"] == 10_000
@pytest.mark.integration
@pytest.mark.asyncio
async def test_version_suffixed_model_id_routes(
authenticated_client: AsyncClient,
same_id_provider_maps: None,
) -> None:
"""A version-suffixed request (``…-YYYYMMDD``) routes to the base model.
Model resolution stripped the suffix but the provider lookup did not, so
such requests resolved a model yet found no provider and 400'd. With the
unified candidate lookup the strip applies to both.
"""
async def fake_transport(
request: httpx.Request, *args: Any, **kwargs: Any
) -> httpx.Response:
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-20260101",
"messages": [{"role": "user", "content": "hello"}],
},
)
assert response.status_code == 200
@pytest.fixture
async def fee_split_provider_maps(
patched_db_engine: None,
) -> AsyncGenerator[None, None]:
"""Same-tail providers whose fees differ; the serving one charges 1.5x."""
cheap = _StaticProvider(
CHEAP_BASE_URL,
"key-cheap",
1.0,
_make_model("dual-model", 0.001, 0.002),
)
expensive = _StaticProvider(
EXPENSIVE_BASE_URL,
"key-expensive",
1.5,
_make_model("dual-model", 0.005, 0.010),
)
async for _ in _install_providers([cheap, expensive]):
yield
@pytest.mark.integration
@pytest.mark.asyncio
async def test_usd_cost_serve_carries_serving_providers_fee(
authenticated_client: AsyncClient,
fee_split_provider_maps: None,
) -> None:
"""The USD-cost billing path applies the SERVING provider's fee.
The upstream that serves reports ``usage.cost`` in USD, so billing goes
through the USD-cost path where the provider fee is applied explicitly.
The serving provider's fee is 1.5; the alias winner's is 1.0. At 0.001 USD
reported cost and 0.0005 USD/sat: 0.001 * 1.5 / 0.0005 = 3 sats = 3000
msats (fee 1.0 would give 2000).
"""
sent_requests: list[httpx.Request] = []
def usd_cost_response(request: httpx.Request) -> httpx.Response:
if request.url.host == "cheap.example.com":
return httpx.Response(
502,
content=json.dumps({"error": {"message": "bad gateway"}}).encode(),
headers={"content-type": "application/json"},
)
body = {
"id": "chatcmpl-usd",
"object": "chat.completion",
"created": 1,
"model": "dual-model",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "hi"},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50,
"total_tokens": 150,
"cost": 0.001,
},
}
return httpx.Response(
200,
content=json.dumps(body).encode(),
headers={"content-type": "application/json"},
)
async def fake_transport(
request: httpx.Request, *args: Any, **kwargs: Any
) -> httpx.Response:
sent_requests.append(request)
return usd_cost_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
assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"expensive.example.com",
]
assert response.json()["cost"]["total_msats"] == 3_000
@pytest.fixture
async def envelope_split_provider_maps(
patched_db_engine: None,
) -> AsyncGenerator[None, None]:
"""Same-id providers where the fallback's max cost dwarfs the key balance."""
cheap = _StaticProvider(
CHEAP_BASE_URL,
"key-cheap",
1.0,
_make_model("dual-model", 0.001, 0.002, max_cost=50.0),
)
expensive = _StaticProvider(
EXPENSIVE_BASE_URL,
"key-expensive",
1.0,
_make_model("dual-model", 0.005, 0.010, max_cost=20_000.0),
)
async for _ in _install_providers([cheap, expensive]):
yield
@pytest.mark.integration
@pytest.mark.asyncio
async def test_failover_beyond_balance_envelope_is_rejected(
authenticated_client: AsyncClient,
envelope_split_provider_maps: None,
) -> None:
"""A fallback whose max-cost envelope exceeds the balance is not served.
Admission and reservation are sized to the best-ranked candidate's max
cost. When that candidate fails and the next one's envelope exceeds the
key's balance, serving it could settle far beyond what admission allowed,
so the request must be rejected (as it would be if the pricier candidate
were ranked first) instead of forwarded.
"""
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"}],
},
)
# The 20_000-sat envelope exceeds the key's 10_000-sat balance: the
# 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,
) -> AsyncGenerator[None, None]:
"""Same-id providers where the fallback needs a larger, affordable reserve."""
cheap = _StaticProvider(
CHEAP_BASE_URL,
"key-cheap",
1.0,
_make_model("dual-model", 0.001, 0.002, max_cost=50.0),
)
expensive = _StaticProvider(
EXPENSIVE_BASE_URL,
"key-expensive",
1.0,
_make_model("dual-model", 0.005, 0.010, max_cost=100.0),
)
async for _ in _install_providers([cheap, expensive]):
yield
@pytest.mark.integration
@pytest.mark.asyncio
async def test_failover_reserves_serving_candidates_envelope(
authenticated_client: AsyncClient,
raised_envelope_provider_maps: None,
integration_session: AsyncSession,
) -> None:
"""An affordable pricier fallback is re-reserved, served, and billed.
The fallback's max cost (100 sats) exceeds the winner's (50 sats) but fits
the key's balance, so the reservation is raised to the serving candidate's
envelope and the request completes, billed at the serving rate with the
unused reserve refunded.
"""
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
assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"expensive.example.com",
]
assert response.json()["cost"]["total_msats"] == 10_000
key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined]
records = (
await integration_session.exec(
select(ReservationRelease).where(ReservationRelease.key_hash == key_hash)
)
).all()
assert len(records) == 2
released = next(record for record in records if record.status == "released")
charged = next(record for record in records if record.status == "charged")
assert charged.reserved_msats > released.reserved_msats
assert all(record.status != "active" for record in records)
@@ -38,7 +38,7 @@ async def test_overrun_charges_after_reservation_swept(
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
from routstr.auth import adjust_payment_for_tokens
deducted_max_cost = 990 # discounted reservation
actual_token_cost = 1000 # actual cost overruns the reservation
@@ -47,10 +47,6 @@ async def test_overrun_charges_after_reservation_swept(
key = _make_key(balance=1000, reserved=0)
integration_session.add(key)
await integration_session.commit()
await pay_for_request(key, deducted_max_cost, integration_session)
key.reserved_balance = 0
integration_session.add(key)
await integration_session.commit()
response_data = {
"model": "test-model",
@@ -62,7 +58,7 @@ async def test_overrun_charges_after_reservation_swept(
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
key, response_data, integration_session, deducted_max_cost, None, None
key, response_data, integration_session, deducted_max_cost
)
await integration_session.refresh(key)
@@ -83,16 +79,8 @@ async def test_free_response_path_closed_end_to_end(
patched_db_engine: None,
) -> None:
"""A reservation released by the real sweeper must not yield a free response."""
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
from routstr.core.db import (
ReservationRelease,
create_session,
release_stale_reservations,
)
from routstr.auth import adjust_payment_for_tokens, pay_for_request
from routstr.core.db import create_session, release_stale_reservations
deducted_max_cost = 990
actual_token_cost = 1000
@@ -116,15 +104,10 @@ async def test_free_response_path_closed_end_to_end(
key = await session.get(ApiKey, key_hash)
assert key is not None
await pay_for_request(key, deducted_max_cost, session)
snapshot = await get_reservation_snapshot(key, session)
await session.refresh(key)
assert key.reserved_balance == deducted_max_cost
key.reserved_at = int(time.time()) - 10_000
record = await session.get(ReservationRelease, snapshot.release_id)
assert record is not None
record.created_at = int(time.time()) - 10_000
session.add(key)
session.add(record)
await session.commit()
# Sweeper releases the stale reservation without charging.
@@ -146,20 +129,18 @@ async def test_free_response_path_closed_end_to_end(
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
key,
response_data,
session,
deducted_max_cost,
reservation_snapshot=snapshot,
key, response_data, session, deducted_max_cost
)
async with create_session() as session:
final = await session.get(ApiKey, key_hash)
assert final is not None
# Stale release is terminal for this reservation. A late finalizer must not
# charge aggregate balance that may now belong to a newer request.
assert final.total_spent == 0
assert final.balance == 1000
assert final.total_spent == actual_token_cost, (
f"Free response: total_spent={final.total_spent}, expected {actual_token_cost}"
)
assert final.balance == 1000 - actual_token_cost, (
f"Balance not charged after sweep: {final.balance}"
)
assert final.balance >= 0
assert final.reserved_balance == 0
@@ -241,10 +241,8 @@ async def test_http_402_response_shape_on_insufficient_balance(
mock_upstream.prepare_headers = MagicMock(return_value={})
with (
patch(
"routstr.proxy.get_candidates",
return_value=[(mock_model, mock_upstream)],
),
patch("routstr.proxy.get_model_instance", return_value=mock_model),
patch("routstr.proxy.get_provider_for_model", return_value=[mock_upstream]),
# Patch where it is used (proxy imports it at module level)
patch(
"routstr.proxy.get_max_cost_for_model",
@@ -1,128 +0,0 @@
"""A live upstream provider resolves its OWN database row by stable identity
(its primary key), not by its mutable/secret ``api_key``.
Today ``from_db_row`` drops ``provider_row.id`` and the two self-referential
paths PPQ.AI's insufficient-balance self-disable and the base
``refresh_models_cache`` re-find their own row with
``WHERE base_url == self.base_url AND api_key == self.api_key``. That uses a
rotatable secret as a self-handle: the moment the row's key changes underneath a
live object (a rotation racing an in-flight request), the object can no longer
find itself. These tests pin the invariant that a provider looks itself up by
identity, so the lookup survives a key change (and, later, key encryption).
"""
from unittest.mock import AsyncMock, patch
import pytest
from routstr.core.db import UpstreamProviderRow
from routstr.upstream.ppqai import PPQAIUpstreamProvider
@pytest.mark.integration
@pytest.mark.asyncio
async def test_provider_object_carries_its_persistent_identity(
integration_session: object,
patched_db_engine: None,
) -> None:
"""``from_db_row`` gives the in-memory object its row's identity (``db_id``)."""
row = UpstreamProviderRow(
provider_type="ppqai",
base_url="https://api.ppq.ai",
api_key="sk-original",
enabled=True,
provider_fee=1.0,
)
integration_session.add(row) # type: ignore[attr-defined]
await integration_session.commit() # type: ignore[attr-defined]
await integration_session.refresh(row) # type: ignore[attr-defined]
provider = PPQAIUpstreamProvider.from_db_row(row)
assert provider is not None
assert provider.db_id == row.id
@pytest.mark.integration
@pytest.mark.asyncio
async def test_self_disable_targets_own_row_after_key_rotation(
integration_session: object,
patched_db_engine: None,
) -> None:
"""PPQ.AI self-disable must disable *its* row even after the key rotated.
RED (current): the object holds the pre-rotation key, so the
``(base_url, api_key)`` lookup misses the row the provider is never
disabled. GREEN: lookup by ``id`` finds it and disables it.
"""
row = UpstreamProviderRow(
provider_type="ppqai",
base_url="https://api.ppq.ai",
api_key="sk-original",
enabled=True,
provider_fee=1.0,
)
integration_session.add(row) # type: ignore[attr-defined]
await integration_session.commit() # type: ignore[attr-defined]
await integration_session.refresh(row) # type: ignore[attr-defined]
provider = PPQAIUpstreamProvider.from_db_row(row) # captures sk-original
assert provider is not None
# Key is rotated in the DB while `provider` is still live.
row.api_key = "sk-rotated"
integration_session.add(row) # type: ignore[attr-defined]
await integration_session.commit() # type: ignore[attr-defined]
with patch("routstr.proxy.reinitialize_upstreams", new=AsyncMock()):
await provider.on_upstream_error_redirect(402, "Insufficient balance")
await integration_session.refresh(row) # type: ignore[attr-defined]
assert row.enabled is False
@pytest.mark.integration
@pytest.mark.asyncio
async def test_refresh_models_cache_finds_own_row_after_key_rotation(
integration_session: object,
patched_db_engine: None,
) -> None:
"""``refresh_models_cache`` must resolve its own row after a key rotation.
``refresh_models_cache`` swallows every exception (it only logs), so the
observable proof it found its row is that it reaches ``list_models`` which
is called with the row's ``id`` only *after* the row is resolved. RED
(current): the stale-key ``(base_url, api_key)`` lookup returns nothing, the
method raises ``404`` internally and returns before ``list_models`` is ever
called. GREEN: lookup by ``id`` finds the row and ``list_models`` runs for
that ``id``.
"""
row = UpstreamProviderRow(
provider_type="ppqai",
base_url="https://api.ppq.ai",
api_key="sk-original",
enabled=True,
provider_fee=1.0,
)
integration_session.add(row) # type: ignore[attr-defined]
await integration_session.commit() # type: ignore[attr-defined]
await integration_session.refresh(row) # type: ignore[attr-defined]
row_id = row.id
provider = PPQAIUpstreamProvider.from_db_row(row) # captures sk-original
assert provider is not None
row.api_key = "sk-rotated"
integration_session.add(row) # type: ignore[attr-defined]
await integration_session.commit() # type: ignore[attr-defined]
list_models_mock = AsyncMock(return_value=[])
with (
patch.object(provider, "fetch_models", new=AsyncMock(return_value=[])),
patch("routstr.upstream.base.list_models", new=list_models_mock),
):
await provider.refresh_models_cache()
list_models_mock.assert_awaited_once()
assert list_models_mock.await_args is not None
assert list_models_mock.await_args.kwargs["upstream_id"] == row_id
@@ -120,9 +120,7 @@ async def test_finalise_releases_reservation_and_charges_balance(
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}}
with patch("routstr.auth.calculate_cost", return_value=cost_data):
await adjust_payment_for_tokens(
key, response_data, integration_session, cost, None, None
)
await adjust_payment_for_tokens(key, response_data, integration_session, cost)
await integration_session.refresh(key)
@@ -141,7 +141,7 @@ async def test_revert_with_zero_reserved_balance_is_noop(
Previously this would drive reserved_balance negative. With the floor guard,
it should return False and leave reserved_balance at 0.
"""
from routstr.auth import pay_for_request, revert_pay_for_request
from routstr.auth import revert_pay_for_request
unique_key = f"test_revert_key_{uuid.uuid4().hex[:8]}"
test_key = ApiKey(
@@ -151,12 +151,8 @@ async def test_revert_with_zero_reserved_balance_is_noop(
)
integration_session.add(test_key)
await integration_session.commit()
await pay_for_request(test_key, 100, integration_session)
test_key.reserved_balance = 0
integration_session.add(test_key)
await integration_session.commit()
# A stale cleanup already released the aggregate reservation.
# Try to revert more than available — should be a no-op
result = await revert_pay_for_request(test_key, integration_session, 100)
await integration_session.refresh(test_key)
@@ -165,8 +161,8 @@ async def test_revert_with_zero_reserved_balance_is_noop(
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 test_key.total_requests == 0, (
f"Total requests should remain 0, got: {test_key.total_requests}"
)
@@ -175,18 +171,17 @@ async def test_revert_with_sufficient_reserved_balance_succeeds(
integration_session: AsyncSession,
) -> None:
"""Test that revert_pay_for_request works correctly when there is enough reserved balance."""
from routstr.auth import pay_for_request, revert_pay_for_request
from routstr.auth import revert_pay_for_request
unique_key = f"test_revert_ok_{uuid.uuid4().hex[:8]}"
test_key = ApiKey(
hashed_key=unique_key,
balance=5000,
reserved_balance=0,
total_requests=2,
reserved_balance=500,
total_requests=3,
)
integration_session.add(test_key)
await integration_session.commit()
await pay_for_request(test_key, 500, integration_session)
result = await revert_pay_for_request(test_key, integration_session, 500)
@@ -207,21 +202,17 @@ async def test_revert_partial_reserved_balance_is_noop(
integration_session: AsyncSession,
) -> None:
"""Test that reverting more than the current reserved_balance is a no-op."""
from routstr.auth import pay_for_request, revert_pay_for_request
from routstr.auth import revert_pay_for_request
unique_key = f"test_revert_partial_{uuid.uuid4().hex[:8]}"
test_key = ApiKey(
hashed_key=unique_key,
balance=5000,
reserved_balance=0,
total_requests=0,
reserved_balance=50,
total_requests=1,
)
integration_session.add(test_key)
await integration_session.commit()
await pay_for_request(test_key, 500, integration_session)
test_key.reserved_balance = 50
integration_session.add(test_key)
await integration_session.commit()
# Try to revert 500 when only 50 is reserved — should be no-op
result = await revert_pay_for_request(test_key, integration_session, 500)
@@ -246,28 +237,20 @@ async def test_double_revert_prevented(
This simulates the double-revert scenario where both upstream/base.py
and proxy.py attempt to revert the same reservation.
"""
from routstr.auth import (
get_reservation_snapshot,
pay_for_request,
revert_pay_for_request,
)
from routstr.auth import revert_pay_for_request
unique_key = f"test_double_revert_{uuid.uuid4().hex[:8]}"
test_key = ApiKey(
hashed_key=unique_key,
balance=10000,
reserved_balance=0,
total_requests=4,
reserved_balance=500,
total_requests=5,
)
integration_session.add(test_key)
await integration_session.commit()
await pay_for_request(test_key, 500, integration_session)
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)
await integration_session.refresh(test_key)
assert result1 is True
@@ -275,9 +258,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)
await integration_session.refresh(test_key)
assert result2 is False, "Second revert should be a no-op"
@@ -298,30 +279,22 @@ async def test_sequential_reverts_never_go_negative(
Simulates the double-revert scenario where multiple code paths
attempt to revert the same reservation.
"""
from routstr.auth import (
get_reservation_snapshot,
pay_for_request,
revert_pay_for_request,
)
from routstr.auth import revert_pay_for_request
unique_key = f"test_multi_revert_{uuid.uuid4().hex[:8]}"
test_key = ApiKey(
hashed_key=unique_key,
balance=10000,
reserved_balance=0,
total_requests=4,
reserved_balance=500,
total_requests=5,
)
integration_session.add(test_key)
await integration_session.commit()
await pay_for_request(test_key, 500, integration_session)
snapshot = await get_reservation_snapshot(test_key, integration_session)
# 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)
results.append(r)
await integration_session.refresh(test_key)
@@ -344,11 +317,7 @@ 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,
)
from routstr.auth import revert_pay_for_request
parent_key_hash = f"test_parent_{uuid.uuid4().hex[:8]}"
child_key_hash = f"test_child_{uuid.uuid4().hex[:8]}"
@@ -356,26 +325,22 @@ async def test_child_key_revert_floor_guard(
parent_key = ApiKey(
hashed_key=parent_key_hash,
balance=10000,
reserved_balance=0,
total_requests=2,
reserved_balance=500,
total_requests=3,
)
child_key = ApiKey(
hashed_key=child_key_hash,
balance=0,
reserved_balance=0,
total_requests=2,
reserved_balance=500,
total_requests=3,
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
)
result1 = await revert_pay_for_request(child_key, integration_session, 500)
await integration_session.refresh(parent_key)
await integration_session.refresh(child_key)
@@ -384,9 +349,7 @@ async def test_child_key_revert_floor_guard(
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
)
result2 = await revert_pay_for_request(child_key, integration_session, 500)
await integration_session.refresh(parent_key)
await integration_session.refresh(child_key)
@@ -1,76 +0,0 @@
"""Tests for the ``reset_admin_password`` recovery script (issue #553).
The script is the lockout escape hatch: it works without ``ROUTSTR_SECRET_KEY``
(scrypt hashing is key-independent). Two explicit, mutually exclusive actions
``--password`` sets a new hash now, ``--regenerate`` clears the hash so the next
boot generates and logs a fresh one. A bare invocation is informational only and
must never touch the database (so nobody resets their password by accident).
"""
import pytest
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core import vault
from routstr.core.db import get_secret, set_admin_password
from scripts.reset_admin_password import apply_reset, build_parser, main
@pytest.mark.asyncio
async def test_password_sets_a_verifiable_hash(
integration_session: AsyncSession,
) -> None:
await apply_reset(integration_session, password="recover-me-123")
secret = await get_secret(integration_session)
assert secret.admin_password_hash is not None
assert vault.verify_password("recover-me-123", secret.admin_password_hash) is True
assert secret.updated_at is not None
@pytest.mark.asyncio
async def test_regenerate_clears_the_hash(
integration_session: AsyncSession,
) -> None:
# Start from a node that already has an admin password set.
await set_admin_password(integration_session, "old-password-9")
assert (await get_secret(integration_session)).admin_password_hash is not None
await apply_reset(integration_session, regenerate=True)
secret = await get_secret(integration_session)
# Cleared -> the next boot's bootstrap_secrets generates and logs a new one.
assert secret.admin_password_hash is None
assert secret.updated_at is not None
@pytest.mark.asyncio
async def test_password_below_min_length_is_rejected(
integration_session: AsyncSession,
) -> None:
await set_admin_password(integration_session, "old-password-9")
with pytest.raises(ValueError, match="8 characters"):
await apply_reset(integration_session, password="short")
# The existing password is untouched by the rejected reset.
secret = await get_secret(integration_session)
assert vault.verify_password("old-password-9", secret.admin_password_hash or "")
def test_password_and_regenerate_are_mutually_exclusive() -> None:
parser = build_parser()
with pytest.raises(SystemExit):
parser.parse_args(["--password", "abcd1234", "--regenerate"])
def test_no_args_prints_help_and_never_opens_a_session(
capsys: pytest.CaptureFixture[str],
monkeypatch: pytest.MonkeyPatch,
) -> None:
def _fail() -> None:
raise AssertionError("a bare invocation must not touch the database")
monkeypatch.setattr("scripts.reset_admin_password.create_session", _fail)
assert main([]) == 0
assert "usage" in capsys.readouterr().out.lower()
-484
View File
@@ -1,484 +0,0 @@
"""Tests for ``bootstrap_secrets`` — moving node secrets into the Secret store.
Specifies the per-secret bootstrap that runs at startup (issue #553). For both
the admin password and the nsec it follows the same three branches: use the
column if already set, otherwise migrate any legacy plaintext (env first, then
the old settings blob), otherwise admin password only generate and log one.
A column written under a different ROUTSTR_SECRET_KEY fails fast rather than
silently corrupting state.
"""
import json
from contextlib import asynccontextmanager
from pathlib import Path
from typing import Any, AsyncGenerator
import pytest
from sqlmodel import text
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core import vault
from routstr.core.db import NsecState, get_secret, set_nsec
from routstr.core.settings import (
SettingsService,
bootstrap_secrets,
derive_npub_from_nsec,
settings,
)
# Valid Fernet keys; must match the suite default in tests/conftest.py.
TEST_SECRET_KEY = "l_Tkp-7xmjcQ-IFhr6qhILrU8HPRbEmYMrfSbo_5srU="
TEST_SECRET_KEY_ALT = "_Teyrky_iToeDK51Tj1FsI9MJ340_cqKGmeher-a7MQ="
NSEC_HEX = "1" * 64
# A different key, standing in for a stale value left behind in env/blob after
# the vault has taken ownership of the real one.
STALE_NSEC_HEX = "2" * 64
@pytest.fixture
def clean_secret_env(monkeypatch: pytest.MonkeyPatch) -> None:
"""No ambient legacy secrets, and a known in-memory settings baseline."""
monkeypatch.delenv("ADMIN_PASSWORD", raising=False)
monkeypatch.delenv("NSEC", raising=False)
monkeypatch.setenv("ROUTSTR_SECRET_KEY", TEST_SECRET_KEY)
monkeypatch.setattr(settings, "nsec", "")
monkeypatch.setattr(settings, "npub", "")
monkeypatch.setattr(settings, "http_url", "")
async def _create_settings_blob(session: AsyncSession, data: dict) -> None:
await session.exec( # type: ignore
text(
"CREATE TABLE IF NOT EXISTS settings "
"(id INTEGER PRIMARY KEY, data TEXT NOT NULL, "
"updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP)"
)
)
await session.exec( # type: ignore
text("INSERT INTO settings (id, data) VALUES (1, :data)").bindparams(
data=json.dumps(data)
)
)
await session.commit()
# --- admin password --------------------------------------------------------
@pytest.mark.asyncio
async def test_generates_admin_password_when_none(
clean_secret_env: None, integration_session: AsyncSession
) -> None:
await bootstrap_secrets(integration_session)
secret = await get_secret(integration_session)
assert secret.admin_password_hash is not None
assert secret.admin_password_hash.startswith("scrypt:")
@pytest.mark.asyncio
async def test_admin_password_generation_is_idempotent(
clean_secret_env: None, integration_session: AsyncSession
) -> None:
await bootstrap_secrets(integration_session)
first = (await get_secret(integration_session)).admin_password_hash
await bootstrap_secrets(integration_session)
second = (await get_secret(integration_session)).admin_password_hash
assert first is not None and first == second
@pytest.mark.asyncio
async def test_hashes_legacy_admin_password_from_env(
clean_secret_env: None,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("ADMIN_PASSWORD", "hunter2")
await bootstrap_secrets(integration_session)
secret = await get_secret(integration_session)
assert secret.admin_password_hash is not None
assert vault.verify_password("hunter2", secret.admin_password_hash) is True
@pytest.mark.asyncio
async def test_hashes_legacy_admin_password_from_blob(
clean_secret_env: None, integration_session: AsyncSession
) -> None:
# No ADMIN_PASSWORD in env, but the old settings blob carries one.
await _create_settings_blob(integration_session, {"admin_password": "blobpw"})
await bootstrap_secrets(integration_session)
secret = await get_secret(integration_session)
assert vault.verify_password("blobpw", secret.admin_password_hash or "") is True
@pytest.mark.asyncio
async def test_admin_password_race_adopts_winner_without_clobber(
clean_secret_env: None,
integration_engine: Any,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
capsys: pytest.CaptureFixture[str],
) -> None:
# Two workers boot against one shared DB and both read a null admin password.
# The first to commit "wins" and shows the operator its generated password. A
# worker that read null but lost the race must NOT overwrite the winner's hash
# (which the operator may already be logging in with) and must NOT print a
# second password that will never work.
#
# The race window is forced deterministically: a hook fires inside bootstrap's
# generate branch (so it only runs once this worker has committed to
# generating) and commits the winner's password on a separate connection
# before this worker writes its own.
import sqlite3
from routstr.core import settings as settings_mod
db_file = integration_engine.url.database
winner_hash = vault.hash_password("winner-password-123")
real_token = settings_mod.secrets.token_urlsafe
def commit_winner_then_generate(nbytes: int) -> str:
conn = sqlite3.connect(db_file)
conn.execute(
"UPDATE secrets SET admin_password_hash = ? WHERE id = 1", (winner_hash,)
)
conn.commit()
conn.close()
return real_token(nbytes)
monkeypatch.setattr(
settings_mod.secrets, "token_urlsafe", commit_winner_then_generate
)
await get_secret(integration_session) # row exists, password still null
capsys.readouterr() # drop anything emitted before the race resolves
await bootstrap_secrets(integration_session)
secret = await get_secret(integration_session)
assert secret.admin_password_hash is not None
# The winner's password survives and still verifies — no clobber.
assert vault.verify_password("winner-password-123", secret.admin_password_hash)
# The losing worker stayed silent — no second generated password was leaked.
assert "generated a temporary" not in capsys.readouterr().out
# --- nsec ------------------------------------------------------------------
@pytest.mark.asyncio
async def test_encrypts_legacy_nsec_from_env_and_derives_npub(
clean_secret_env: None,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("NSEC", NSEC_HEX)
await bootstrap_secrets(integration_session)
secret = await get_secret(integration_session)
assert secret.encrypted_nsec is not None
assert vault.is_encrypted(secret.encrypted_nsec) is True
assert vault.decrypt(secret.encrypted_nsec) == NSEC_HEX
# In-memory runtime value is the decrypted nsec, and npub is derived from it.
assert settings.nsec == NSEC_HEX
assert settings.npub == derive_npub_from_nsec(NSEC_HEX)
@pytest.mark.asyncio
async def test_decrypts_existing_nsec_column(
clean_secret_env: None, integration_session: AsyncSession
) -> None:
secret = await get_secret(integration_session)
secret.encrypted_nsec = vault.encrypt(NSEC_HEX)
secret.nsec_state = NsecState.encrypted
integration_session.add(secret)
await integration_session.commit()
stored = secret.encrypted_nsec
await bootstrap_secrets(integration_session)
reloaded = await get_secret(integration_session)
assert settings.nsec == NSEC_HEX
# The column is reused, not re-encrypted.
assert reloaded.encrypted_nsec == stored
@pytest.mark.asyncio
async def test_fail_fast_when_nsec_encrypted_with_different_key(
clean_secret_env: None,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
# Encrypt the column under the alternate key, then bootstrap under the
# suite key -> the value cannot be decrypted -> clear startup failure.
monkeypatch.setenv("ROUTSTR_SECRET_KEY", TEST_SECRET_KEY_ALT)
secret = await get_secret(integration_session)
secret.encrypted_nsec = vault.encrypt(NSEC_HEX)
secret.nsec_state = NsecState.encrypted
integration_session.add(secret)
await integration_session.commit()
monkeypatch.setenv("ROUTSTR_SECRET_KEY", TEST_SECRET_KEY)
with pytest.raises(RuntimeError, match="ROUTSTR_SECRET_KEY"):
await bootstrap_secrets(integration_session)
# --- encryption is mandatory, key custody is not: upgrade without a key --------
@pytest.mark.asyncio
async def test_legacy_nsec_without_secret_key_generates_and_encrypts(
clean_secret_env: None,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
capsys: pytest.CaptureFixture[str],
) -> None:
# A node upgrading with a legacy plaintext NSEC but no ROUTSTR_SECRET_KEY must
# NOT break. Encryption at rest stays mandatory (the nsec is never persisted
# in plaintext), but the key custody is flexible: bootstrap generates a master
# key, persists it to the key file, warns loudly, and encrypts the identity —
# so the node keeps running instead of refusing to boot.
monkeypatch.delenv("ROUTSTR_SECRET_KEY", raising=False)
key_file = tmp_path / "routstr_secret.key"
monkeypatch.setenv("ROUTSTR_SECRET_KEY_FILE", str(key_file))
monkeypatch.setenv("NSEC", NSEC_HEX)
await bootstrap_secrets(integration_session)
# A master key was generated and persisted...
assert key_file.exists()
# ...the nsec is encrypted at rest under it, never stored in plaintext...
secret = await get_secret(integration_session)
assert secret.encrypted_nsec is not None
assert vault.is_encrypted(secret.encrypted_nsec) is True
assert vault.decrypt(secret.encrypted_nsec) == NSEC_HEX
assert secret.nsec_state == NsecState.encrypted
# ...the node holds the live identity (npub derived from it)...
assert settings.nsec == NSEC_HEX
assert settings.npub == derive_npub_from_nsec(NSEC_HEX)
# ...and the operator is loudly told a key was generated and must be backed up
# (path shown, but never the key value) so an upgrade cannot silently create
# an unbacked key nor leak the key into captured stdout / aggregated logs.
out = capsys.readouterr().out
assert str(key_file) in out
assert key_file.read_text().strip() not in out
assert "BACK UP" in out.upper()
# --- boot ordering: rescue legacy blob secrets before they are stripped ----
@pytest.mark.asyncio
async def test_blob_only_nsec_is_migrated_before_blob_is_stripped(
clean_secret_env: None, integration_session: AsyncSession
) -> None:
# Legacy node whose nsec lives ONLY in the settings blob (never in env).
# bootstrap_secrets must run *before* SettingsService.initialize strips the
# blob, or the only copy of the secret would be lost.
await _create_settings_blob(
integration_session, {"nsec": NSEC_HEX, "name": "LegacyNode"}
)
await bootstrap_secrets(integration_session)
await SettingsService.initialize(integration_session)
secret = await get_secret(integration_session)
# The plaintext nsec has been moved into the encrypted Secret store...
assert secret.encrypted_nsec is not None
assert vault.decrypt(secret.encrypted_nsec) == NSEC_HEX
assert settings.nsec == NSEC_HEX
# ...and stripped from the persisted settings blob.
row = await integration_session.exec( # type: ignore
text("SELECT data FROM settings WHERE id = 1")
)
blob = json.loads(row.first()[0])
assert "nsec" not in blob
assert blob["name"] == "LegacyNode"
@pytest.mark.asyncio
async def test_initialize_does_not_clobber_store_only_nsec(
clean_secret_env: None, integration_session: AsyncSession
) -> None:
# Steady state after migration: the nsec lives ONLY in the encrypted Secret
# store (NSEC removed from env, blob already stripped on a previous boot).
# bootstrap decrypts it into memory; initialize then re-derives settings from
# the secret-free blob and must NOT wipe the live nsec back to empty (or the
# node would silently stop signing Nostr announcements).
await _create_settings_blob(integration_session, {"name": "LegacyNode"})
secret = await get_secret(integration_session)
secret.encrypted_nsec = vault.encrypt(NSEC_HEX)
secret.nsec_state = NsecState.encrypted
integration_session.add(secret)
await integration_session.commit()
await bootstrap_secrets(integration_session)
assert settings.nsec == NSEC_HEX # bootstrap decrypted it into memory
await SettingsService.initialize(integration_session)
# The live secret survives initialize even though no env/blob carries it...
assert settings.nsec == NSEC_HEX
# ...and is still never written back to the persisted blob.
row = await integration_session.exec( # type: ignore
text("SELECT data FROM settings WHERE id = 1")
)
assert "nsec" not in json.loads(row.first()[0])
@pytest.mark.asyncio
async def test_stale_env_nsec_does_not_override_vault_nsec(
clean_secret_env: None,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
# The vault owns the nsec, but a stale NSEC (e.g. the operator rotated the
# key in the UI yet left the old value in .env) is still in the environment.
# bootstrap decrypts the store value; initialize must NOT let the stale env
# value clobber it, or a restart silently reverts to the old identity.
await set_nsec(integration_session, NSEC_HEX)
monkeypatch.setenv("NSEC", STALE_NSEC_HEX)
await _create_settings_blob(integration_session, {"name": "LegacyNode"})
await bootstrap_secrets(integration_session)
await SettingsService.initialize(integration_session)
# The vault value wins; the stale env value is ignored.
assert settings.nsec == NSEC_HEX
@pytest.mark.asyncio
async def test_stale_env_nsec_does_not_split_npub_from_vault_nsec(
clean_secret_env: None,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
# As above, the vault owns the nsec while a stale NSEC lingers in env. The
# private key correctly comes from the vault, but the npub must too: if
# initialize derives the public key from the stale env nsec, the node ends up
# with a private key from the vault and a public key from the old env value,
# and anything reading settings.npub announces the wrong Nostr identity.
expected_npub = derive_npub_from_nsec(NSEC_HEX)
stale_npub = derive_npub_from_nsec(STALE_NSEC_HEX)
assert expected_npub and stale_npub and expected_npub != stale_npub # guard
await set_nsec(integration_session, NSEC_HEX)
monkeypatch.setenv("NSEC", STALE_NSEC_HEX)
await _create_settings_blob(integration_session, {"name": "LegacyNode"})
await bootstrap_secrets(integration_session)
await SettingsService.initialize(integration_session)
assert settings.nsec == NSEC_HEX
assert settings.npub == expected_npub
@pytest.mark.asyncio
async def test_cleared_nsec_stays_cleared_across_reboot(
clean_secret_env: None,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
# An identity was imported from env, then the operator cleared it via the
# admin API. The old NSEC is still in env. On the NEXT PROCESS the cleared
# identity must stay cleared, not get resurrected from the stale env value.
monkeypatch.setenv("NSEC", NSEC_HEX)
await bootstrap_secrets(integration_session)
assert settings.nsec == NSEC_HEX
# Clear via the admin path (store empty, vault owns it).
await set_nsec(integration_session, "")
# Simulate a fresh process rather than pre-clearing the live singleton: the
# pydantic settings global reloads the (still-stale) NSEC from env and derives
# its npub, which is exactly the in-memory state a new boot starts from before
# bootstrap runs. The cleared store must win over this stale live value.
monkeypatch.setattr(settings, "nsec", NSEC_HEX)
monkeypatch.setattr(settings, "npub", derive_npub_from_nsec(NSEC_HEX))
await bootstrap_secrets(integration_session)
reloaded = await get_secret(integration_session)
assert reloaded.nsec_state == NsecState.cleared
assert reloaded.encrypted_nsec is None # not re-imported
assert settings.nsec == "" # stays cleared
assert settings.npub == "" # and no derived public identity survives
@pytest.mark.asyncio
async def test_initialize_keeps_npub_matching_store_only_nsec(
clean_secret_env: None, integration_session: AsyncSession
) -> None:
# Steady state with mandatory encryption: the nsec lives ONLY in the
# encrypted Secret store (env carries no NSEC) and the blob has no npub.
# bootstrap decrypts the nsec and derives the npub into memory; initialize
# then re-derives settings from the npub-less blob and must NOT wipe the npub
# back to empty, or the node holds a private key with no matching public key
# and silently stops announcing a usable Nostr identity.
expected_npub = derive_npub_from_nsec(NSEC_HEX)
assert expected_npub # guard: the test key must yield a real npub
await _create_settings_blob(integration_session, {"name": "LegacyNode"})
secret = await get_secret(integration_session)
secret.encrypted_nsec = vault.encrypt(NSEC_HEX)
secret.nsec_state = NsecState.encrypted
integration_session.add(secret)
await integration_session.commit()
await bootstrap_secrets(integration_session)
assert settings.npub == expected_npub # bootstrap derived it
await SettingsService.initialize(integration_session)
# The npub still matches the live nsec...
assert settings.nsec == NSEC_HEX
assert settings.npub == expected_npub
# ...and is persisted to the blob (it is public, not a stripped secret).
row = await integration_session.exec( # type: ignore
text("SELECT data FROM settings WHERE id = 1")
)
assert json.loads(row.first()[0])["npub"] == expected_npub
@pytest.mark.asyncio
async def test_startup_runs_bootstrap_before_settings_initialize(
monkeypatch: pytest.MonkeyPatch,
) -> None:
# The two tests above prove the migration outcome *given* the call order;
# they hardcode that order themselves. This one guards the order at its real
# call site — the application lifespan — so a reorder in main.py (which would
# strip a blob-only secret before bootstrap could rescue it) is caught.
import routstr.core.main as main
order: list[str] = []
class _Abort(Exception):
pass
@asynccontextmanager
async def fake_create_session() -> AsyncGenerator[None, None]:
yield None
async def fake_bootstrap(session: Any) -> None:
order.append("bootstrap")
async def fake_initialize(session: Any) -> None:
order.append("initialize")
# Stop startup here, before the background-task fan-out (prices, nostr,
# upstreams) that we don't want to run in a unit test.
raise _Abort()
async def noop_init_db() -> None:
return None
monkeypatch.setattr(main, "configure_litellm", lambda: None)
monkeypatch.setattr(main, "register_deepseek_v4_pricing", lambda: None)
monkeypatch.setattr(main, "run_migrations", lambda: None)
monkeypatch.setattr(main, "init_db", noop_init_db)
monkeypatch.setattr(main, "create_session", fake_create_session)
monkeypatch.setattr(main, "bootstrap_secrets", fake_bootstrap)
monkeypatch.setattr(main.SettingsService, "initialize", fake_initialize)
with pytest.raises(_Abort):
async with main.lifespan(main.app):
pass
assert order == ["bootstrap", "initialize"]
-93
View File
@@ -1,93 +0,0 @@
"""Tests for the ``Secret`` singleton model (issue #553).
Specifies the node-level secret store: a single row (``id=1``, like
``RoutstrFee``) holding the one-way admin-password hash and the encrypted nsec.
``get_secret`` is get-or-create, so callers always get the singleton without
worrying whether it has been initialised yet. Encoding of the values themselves
lives in ``routstr.core.vault``; here we only assert the row persists and stays
a singleton.
"""
import time
from typing import Any
import pytest
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import Secret, get_secret
@pytest.mark.asyncio
async def test_get_secret_creates_singleton(
integration_session: AsyncSession,
) -> None:
secret = await get_secret(integration_session)
assert secret.id == 1
# Fresh row carries no secret material yet.
assert secret.admin_password_hash is None
assert secret.encrypted_nsec is None
assert secret.updated_at is None
@pytest.mark.asyncio
async def test_get_secret_is_idempotent(
integration_session: AsyncSession,
) -> None:
first = await get_secret(integration_session)
second = await get_secret(integration_session)
assert first.id == second.id == 1
rows = (await integration_session.exec(select(Secret))).all()
assert len(rows) == 1
@pytest.mark.asyncio
async def test_secret_fields_round_trip(
integration_session: AsyncSession,
) -> None:
secret = await get_secret(integration_session)
secret.admin_password_hash = "scrypt:16384:8:1:c2FsdA==:aGFzaA=="
secret.encrypted_nsec = "fernet:v1:gAAAAA"
secret.updated_at = int(time.time())
integration_session.add(secret)
await integration_session.commit()
integration_session.expunge_all()
reloaded = await get_secret(integration_session)
assert reloaded.admin_password_hash == "scrypt:16384:8:1:c2FsdA==:aGFzaA=="
assert reloaded.encrypted_nsec == "fernet:v1:gAAAAA"
assert reloaded.updated_at is not None
@pytest.mark.asyncio
async def test_get_secret_tolerates_concurrent_first_insert(
integration_engine: Any,
integration_session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
# A second worker wins the race and commits the singleton row first.
async with AsyncSession(integration_engine, expire_on_commit=False) as other:
other.add(Secret(id=1, admin_password_hash="scrypt:from-other-worker"))
await other.commit()
# Reproduce the race window: our session's first read still sees no row, so
# it attempts to INSERT a duplicate id=1. The real IntegrityError that follows
# must be recovered (roll back, re-read) rather than crashing startup.
real_get = integration_session.get
calls = {"n": 0}
async def stale_first_read(model: Any, pk: Any) -> Any:
calls["n"] += 1
if calls["n"] == 1:
return None
return await real_get(model, pk)
monkeypatch.setattr(integration_session, "get", stale_first_read)
secret = await get_secret(integration_session)
# Recovered the other worker's row; no crash, still a single row.
assert secret.id == 1
assert secret.admin_password_hash == "scrypt:from-other-worker"
rows = (await integration_session.exec(select(Secret))).all()
assert len(rows) == 1
+4 -4
View File
@@ -174,12 +174,12 @@ async def test_topup_retries_when_quote_fee_exceeds_estimate(
@pytest.mark.integration
@pytest.mark.asyncio
async def test_topup_returns_422_when_retries_exhausted(
async def test_topup_returns_400_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."""
exhausts the retry budget: clean 400 with an actionable message, melt never
executed."""
mock_token, token_wallet, primary_wallet = _make_swap_mocks(
1000, fee_reserves=[1, 10, 25, 50]
)
@@ -188,7 +188,7 @@ async def test_topup_returns_422_when_retries_exhausted(
authenticated_client, mock_token, token_wallet, primary_wallet
)
assert response.status_code == 422
assert response.status_code == 400
assert "too small to cover swap fees" in response.json()["detail"]
assert token_wallet.melt_quote.call_count == 4 # estimation + 3 attempts
token_wallet.melt.assert_not_called()
+14 -13
View File
@@ -12,7 +12,10 @@ from sqlmodel import select
from routstr.core.db import ApiKey
from .utils import ResponseValidator
from .utils import (
CashuTokenGenerator,
ResponseValidator,
)
@pytest.mark.integration
@@ -76,31 +79,29 @@ async def test_api_key_generation_invalid_token(
# Capture initial state
await db_snapshot.capture()
# Non-Cashu bearer values are invalid API keys (401). Malformed values that
# look like Cashu tokens use the shared Cashu taxonomy (400 invalid_token).
invalid_tokens: list[tuple[str, int, str | None]] = [
("not-a-cashu-token", 401, None),
("sk-not-a-real-api-key", 401, None),
("cashuA", 400, "invalid_cashu_token"),
("cashuA" + "x" * 1000, 400, "invalid_cashu_token"),
# Test various invalid tokens
invalid_tokens = [
CashuTokenGenerator.generate_invalid_token(), # Malformed token
"not-a-cashu-token", # Wrong format
"cashuA", # Empty token
"cashuA" + "x" * 1000, # Invalid base64
]
for invalid_token, expected_status, expected_code in invalid_tokens:
for invalid_token in invalid_tokens:
integration_client.headers["Authorization"] = f"Bearer {invalid_token}"
response = await integration_client.get("/v1/wallet/info")
assert response.status_code == expected_status, (
# Should fail with 401
assert response.status_code == 401, (
f"Token {invalid_token[:20]}... should be invalid"
)
# Validate error response
validator = ResponseValidator()
error_validation = validator.validate_error_response(
response, expected_status=expected_status, expected_error_key="detail"
response, expected_status=401, expected_error_key="detail"
)
assert error_validation["valid"]
if expected_code is not None:
assert response.json()["detail"]["error"]["code"] == expected_code
# Verify no database changes
diff = await db_snapshot.diff()
+14 -8
View File
@@ -14,7 +14,6 @@ from httpx import AsyncClient
from sqlmodel import select
from routstr.core.db import ApiKey, CashuTransaction
from routstr.wallet import MintConnectionError
@pytest.mark.integration
@@ -504,19 +503,26 @@ async def test_mint_unavailability_handling(
# The global mock in conftest.py is already in place,
# so we need to temporarily modify it
raw_error = "Mint unavailable: Connection refused"
from unittest.mock import patch
# Make the send_token method raise a typed mint connection exception.
# Make the send_token method raise an exception
with patch(
"routstr.balance.send_token",
side_effect=MintConnectionError(raw_error),
side_effect=Exception("Mint unavailable: Connection refused"),
):
response = await authenticated_client.post("/v1/wallet/refund")
assert response.status_code == 503
assert response.json()["detail"] == "Mint service unavailable"
assert raw_error not in response.text
# The exception should propagate as a 503 error (Service Unavailable)
# But we need to handle it properly
try:
response = await authenticated_client.post("/v1/wallet/refund")
# If we get here, check the status code
assert response.status_code == 503
assert "Mint service unavailable" in response.json()["detail"]
except Exception as e:
# If the exception propagates, that's also a failure scenario
assert "Mint unavailable" in str(e)
# Balance should remain unchanged (transaction should roll back)
# Note: Current implementation might not handle this perfectly
wallet_response = await authenticated_client.get("/v1/wallet/")
assert wallet_response.status_code == 200
assert wallet_response.json()["balance"] == 10_000_000
-82
View File
@@ -1,82 +0,0 @@
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
from routstr.core import admin
@pytest.mark.asyncio
@pytest.mark.parametrize("requested_mint", [None, "https://secondary.example"])
async def test_withdraw_uses_effective_mint_and_records_outgoing_transaction(
monkeypatch: pytest.MonkeyPatch, requested_mint: str | None
) -> None:
primary_mint = "https://primary.example"
effective_mint = requested_mint or primary_mint
wallet = object()
proofs = [SimpleNamespace(amount=40), SimpleNamespace(amount=60)]
token = "cashuBoutgoing"
get_wallet = AsyncMock(return_value=wallet)
get_proofs = Mock(return_value=proofs)
filter_proofs = AsyncMock(return_value=proofs)
send_token = AsyncMock(return_value=token)
store_transaction = AsyncMock(return_value=True)
monkeypatch.setattr(admin, "get_wallet", get_wallet)
monkeypatch.setattr(admin, "get_proofs_per_mint_and_unit", get_proofs)
monkeypatch.setattr(admin, "slow_filter_spend_proofs", filter_proofs)
monkeypatch.setattr(admin, "send_token", send_token)
monkeypatch.setattr(admin, "store_cashu_transaction", store_transaction)
monkeypatch.setattr(admin.settings, "primary_mint", primary_mint)
result = await admin.withdraw(
Mock(),
admin.WithdrawRequest(amount=75, mint_url=requested_mint, unit="sat"),
)
assert result == {"token": token}
get_wallet.assert_awaited_once_with(effective_mint, "sat")
get_proofs.assert_called_once_with(wallet, effective_mint, "sat", not_reserved=True)
filter_proofs.assert_awaited_once_with(proofs, wallet)
send_token.assert_awaited_once_with(75, "sat", effective_mint)
store_transaction.assert_awaited_once_with(
token=token,
amount=75,
unit="sat",
mint_url=effective_mint,
typ="out",
collected=False,
source="admin",
)
@pytest.mark.asyncio
async def test_withdraw_returns_issued_token_when_audit_storage_fails(
monkeypatch: pytest.MonkeyPatch,
) -> None:
mint = "https://primary.example"
proofs = [SimpleNamespace(amount=100)]
token = "cashuBrecoverable"
monkeypatch.setattr(admin, "get_wallet", AsyncMock(return_value=object()))
monkeypatch.setattr(
admin, "get_proofs_per_mint_and_unit", Mock(return_value=proofs)
)
monkeypatch.setattr(
admin, "slow_filter_spend_proofs", AsyncMock(return_value=proofs)
)
monkeypatch.setattr(admin, "send_token", AsyncMock(return_value=token))
monkeypatch.setattr(
admin,
"store_cashu_transaction",
AsyncMock(side_effect=RuntimeError("database unavailable")),
)
critical = Mock()
monkeypatch.setattr(admin.logger, "critical", critical)
monkeypatch.setattr(admin.settings, "primary_mint", mint)
result = await admin.withdraw(Mock(), admin.WithdrawRequest(amount=75))
assert result == {"token": token}
critical.assert_called_once()
+6 -68
View File
@@ -137,12 +137,12 @@ def test_create_model_mappings_includes_db_override_for_missing_cached_model(
model_instances, provider_map, unique_models = create_model_mappings(
upstreams=[provider],
overrides_by_key={("azure/gpt-4o", 7): (override_row, 1.01)},
disabled_model_keys=set(),
overrides_by_id={"azure/gpt-4o": (override_row, 1.01)},
disabled_model_ids=set(),
)
assert "azure/gpt-4o" in model_instances
assert [p for _, p in provider_map["azure/gpt-4o"]] == [provider]
assert provider_map["azure/gpt-4o"] == [provider]
assert "gpt-4o" in unique_models
@@ -182,73 +182,11 @@ def test_create_model_mappings_dedupes_with_provider_identity_not_provider_type(
_, provider_map, _ = create_model_mappings(
upstreams=[provider_a, provider_b],
overrides_by_key={("azure/gpt-4o", 2): (override_row, 1.01)},
disabled_model_keys=set(),
overrides_by_id={"azure/gpt-4o": (override_row, 1.01)},
disabled_model_ids=set(),
)
providers_for_alias = [p for _, p in provider_map["azure/gpt-4o"]]
providers_for_alias = provider_map["azure/gpt-4o"]
assert provider_a in providers_for_alias
assert provider_b in providers_for_alias
assert len(providers_for_alias) == 2
def test_create_model_mappings_applies_override_only_to_matching_provider(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Same-id overrides must not add provider-specific aliases to other providers."""
provider_a_model = create_test_model("same-id", prompt_price=0.01)
provider_a = create_test_provider(
"provider-a",
"https://provider-a.example/v1",
db_id=1,
models=[provider_a_model],
)
provider_b_model = create_test_model("same-id", prompt_price=0.02)
provider_b = create_test_provider(
"provider-b",
"https://provider-b.example/v1",
db_id=2,
models=[provider_b_model],
)
override_model = create_test_model("same-id", prompt_price=0.001)
override_model.alias_ids = ["provider-b-only"]
override_row = SimpleNamespace(id="same-id", upstream_provider_id=2, enabled=True)
def fake_row_to_model(*args, **kwargs) -> Model: # type: ignore[no-untyped-def]
return override_model
monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model)
_, provider_map, _ = create_model_mappings(
upstreams=[provider_a, provider_b],
overrides_by_key={("same-id", 2): (override_row, 1.01)},
disabled_model_keys=set(),
)
assert [p for _, p in provider_map["provider-b-only"]] == [provider_b]
assert {p for _, p in provider_map["same-id"]} == {provider_a, provider_b}
def test_create_model_mappings_disables_only_matching_provider() -> None:
"""Disabled overrides are scoped to the provider row, not the shared model id."""
provider_a = create_test_provider(
"provider-a",
"https://provider-a.example/v1",
db_id=1,
models=[create_test_model("same-id")],
)
provider_b = create_test_provider(
"provider-b",
"https://provider-b.example/v1",
db_id=2,
models=[create_test_model("same-id")],
)
_, provider_map, _ = create_model_mappings(
upstreams=[provider_a, provider_b],
overrides_by_key={},
disabled_model_keys={("same-id", 2)},
)
assert [p for _, p in provider_map["same-id"]] == [provider_a]
+1 -235
View File
@@ -1,9 +1,8 @@
import hashlib
from types import SimpleNamespace
from typing import AsyncGenerator, cast
from typing import AsyncGenerator
from unittest.mock import AsyncMock, patch
import httpx
import pytest
from fastapi import HTTPException
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
@@ -13,19 +12,6 @@ from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.auth import validate_bearer_key
from routstr.core.db import ApiKey
from routstr.wallet import MintConnectionError
def _value_error_wrapping_transport() -> ValueError:
"""A ValueError re-raised ``from`` a real httpx transport error, mirroring
``wallet.py`` wrapping a connection failure. The sanitized classifier must
still see the mint-unreachable signal through the ``__cause__`` chain."""
try:
raise httpx.ConnectError("All connection attempts failed")
except httpx.ConnectError as exc:
err = ValueError("Failed to estimate fees: connection failed")
err.__cause__ = exc
return err
def _make_engine() -> AsyncEngine:
@@ -71,223 +57,3 @@ async def test_failed_first_cashu_redemption_rolls_back_empty_api_key(
await validate_bearer_key(token, session)
assert await session.get(ApiKey, hashed_key) is None
@pytest.mark.parametrize(
("error", "expected_status", "expected_type", "expected_message", "expected_code"),
[
(
ValueError("Mint Error: Token already spent. (Code: 11001)"),
400,
"token_already_spent",
"Cashu token already spent",
"cashu_token_already_spent",
),
(
# Raw httpx transport error propagated unwrapped from cashu.
httpx.ConnectError("All connection attempts failed"),
503,
"mint_unreachable",
"Cashu mint is unreachable",
"cashu_mint_unreachable",
),
(
# Typed error raised by wallet.py at a wrap site.
MintConnectionError("connect to http://mint:3338 refused"),
503,
"mint_unreachable",
"Cashu mint is unreachable",
"cashu_mint_unreachable",
),
(
# ValueError wrapping the httpx error in its __cause__ chain.
_value_error_wrapping_transport(),
503,
"mint_unreachable",
"Cashu mint is unreachable",
"cashu_mint_unreachable",
),
(
# asyncio.TimeoutError is builtin TimeoutError on 3.11+.
TimeoutError("Timed out connecting to Cashu mint http://mint:3338"),
503,
"mint_unreachable",
"Cashu mint is unreachable",
"cashu_mint_unreachable",
),
(
ValueError(
"Token amount (5 sat) is insufficient to cover melt fees. "
"Needed: 7 sat (amount: 5 + fee: 1 + input_fees: 1)"
),
422,
"mint_error",
"Token value is too small to cover swap fees",
"cashu_token_swap_fees_exceed_amount",
),
(
ValueError(
"Failed to estimate fees: Fees (7 sat) exceed token amount (5 sat)"
),
422,
"mint_error",
"Token value is too small to cover swap fees",
"cashu_token_swap_fees_exceed_amount",
),
(
ValueError(
"Failed to melt token from foreign mint http://foreign:3338: boom"
),
422,
"mint_error",
"Failed to swap token from foreign mint",
"cashu_foreign_mint_swap_failed",
),
(
ValueError("could not decode token"),
400,
"invalid_token",
"Invalid Cashu token",
"invalid_cashu_token",
),
(
ValueError("some unexpected wallet condition"),
400,
"cashu_error",
"Failed to redeem Cashu token",
"cashu_token_redemption_failed",
),
(
ValueError("Redeemed token amount must be positive, got 0 msats"),
400,
"cashu_error",
"Failed to redeem Cashu token: token yielded no value",
"cashu_token_zero_value",
),
],
)
@pytest.mark.asyncio
async def test_redemption_failure_returns_sanitized_error(
session: AsyncSession,
error: Exception,
expected_status: int,
expected_type: str,
expected_message: str,
expected_code: str,
) -> None:
"""Redemption failures reuse the shared X-Cashu taxonomy (carried in
``type``), expose stable sanitized messages and granular machine-readable
``code`` values, and leave no orphan ApiKey row."""
token = "cashuAredemption_fails_with_specific_error"
hashed_key = hashlib.sha256(token.encode()).hexdigest()
token_obj = SimpleNamespace(mint="http://mint:3338", unit="sat")
from routstr.core.settings import settings
with (
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
patch("routstr.auth.deserialize_token_from_string", return_value=token_obj),
patch(
"routstr.auth.credit_balance",
new=AsyncMock(side_effect=error),
),
):
with pytest.raises(HTTPException) as exc_info:
await validate_bearer_key(token, session)
assert exc_info.value.status_code == expected_status
detail = cast(dict[str, dict[str, object]], exc_info.value.detail)
error_detail = detail["error"]
assert error_detail["type"] == expected_type
assert error_detail["code"] == expected_code
assert error_detail["message"] == expected_message
assert str(error) not in cast(str, error_detail["message"])
assert await session.get(ApiKey, hashed_key) is None
@pytest.mark.asyncio
async def test_unexpected_redemption_error_returns_internal_error(
session: AsyncSession,
) -> None:
"""Unexpected (non-wallet) failures surface as generic 500s without
leaking internal details, instead of masquerading as token errors."""
token = "cashuAredemption_fails_with_internal_error"
hashed_key = hashlib.sha256(token.encode()).hexdigest()
token_obj = SimpleNamespace(mint="http://mint:3338", unit="sat")
from routstr.core.settings import settings
with (
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
patch("routstr.auth.deserialize_token_from_string", return_value=token_obj),
patch(
"routstr.auth.credit_balance",
new=AsyncMock(side_effect=RuntimeError("db exploded at /var/lib/secret")),
),
):
with pytest.raises(HTTPException) as exc_info:
await validate_bearer_key(token, session)
assert exc_info.value.status_code == 500
detail = cast(dict[str, dict[str, str]], exc_info.value.detail)
error_detail = detail["error"]
assert error_detail["code"] == "internal_error"
assert "/var/lib/secret" not in error_detail["message"]
assert await session.get(ApiKey, hashed_key) is None
@pytest.mark.asyncio
async def test_internal_error_with_invalid_keyword_does_not_masquerade(
session: AsyncSession,
) -> None:
"""A non-wallet fault whose text merely contains "invalid" (but not
"token") must fall through to a generic 500, not a 401 token error.
Guards the anchored `"invalid"/"decode"` + `"token"` gate against stdlib/
driver strings like "Invalid isoformat string" leaking as token errors."""
token = "cashuAinternal_fault_mentions_invalid"
hashed_key = hashlib.sha256(token.encode()).hexdigest()
token_obj = SimpleNamespace(mint="http://mint:3338", unit="sat")
from routstr.core.settings import settings
with (
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
patch("routstr.auth.deserialize_token_from_string", return_value=token_obj),
patch(
"routstr.auth.credit_balance",
new=AsyncMock(
side_effect=RuntimeError("Invalid isoformat string: '2020-13-99'")
),
),
):
with pytest.raises(HTTPException) as exc_info:
await validate_bearer_key(token, session)
assert exc_info.value.status_code == 500
detail = cast(dict[str, dict[str, str]], exc_info.value.detail)
assert detail["error"]["code"] == "internal_error"
assert await session.get(ApiKey, hashed_key) is None
@pytest.mark.asyncio
async def test_malformed_cashu_token_returns_400_invalid_token(
session: AsyncSession,
) -> None:
"""A malformed 'cashu...' token that fails to decode maps to 400
invalid_cashu_token (shared taxonomy), not the generic 401 invalid_api_key."""
token = "cashuAthis_is_not_a_valid_token"
with patch(
"routstr.auth.deserialize_token_from_string",
side_effect=ValueError("unable to decode token: bad base64"),
):
with pytest.raises(HTTPException) as exc_info:
await validate_bearer_key(token, session)
assert exc_info.value.status_code == 400
detail = cast(dict[str, dict[str, str]], exc_info.value.detail)
assert detail["error"]["type"] == "invalid_token"
assert detail["error"]["code"] == "invalid_cashu_token"
# Raw decoder text must not leak to the client.
assert "base64" not in detail["error"]["message"]
-143
View File
@@ -1,143 +0,0 @@
import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from routstr.core.db import CashuTransaction
from routstr.upstream.auto_topup import _check_and_topup
def _row() -> MagicMock:
row = MagicMock()
row.id = "provider-1"
row.base_url = "https://provider.test"
row.api_key = "secret"
row.provider_settings = json.dumps(
{
"auto_topup": True,
"topup_threshold": 100,
"topup_amount_limit": 50,
"topup_mint_url": "https://mint.test",
}
)
return row
class _Session:
def __init__(self, transaction: CashuTransaction) -> None:
self.transaction = transaction
self.commit = AsyncMock()
async def __aenter__(self) -> "_Session":
return self
async def __aexit__(self, *args: object) -> None:
return None
async def exec(self, query: object) -> MagicMock:
result = MagicMock()
result.first.return_value = self.transaction
return result
def add(self, transaction: CashuTransaction) -> None:
self.transaction = transaction
@pytest.mark.asyncio
async def test_auto_topup_persists_before_sending_and_marks_success_collected() -> None:
provider = MagicMock()
provider.get_balance = AsyncMock(return_value=0)
provider.topup = AsyncMock(return_value={"balance": 50})
transaction = CashuTransaction(
token="cashu-token", amount=50, unit="sat", source="auto_topup"
)
session = _Session(transaction)
with (
patch(
"routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row",
return_value=provider,
),
patch(
"routstr.upstream.auto_topup.send_token",
AsyncMock(return_value="cashu-token"),
),
patch(
"routstr.upstream.auto_topup.store_cashu_transaction",
AsyncMock(return_value=True),
) as store,
patch("routstr.upstream.auto_topup.create_session", return_value=session),
):
await _check_and_topup(_row())
store.assert_awaited_once_with(
token="cashu-token",
amount=50,
unit="sat",
mint_url="https://mint.test",
typ="out",
collected=False,
source="auto_topup",
)
provider.topup.assert_awaited_once_with("cashu-token")
assert transaction.collected is True
session.commit.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("outcome", [{"error": "rejected"}, RuntimeError("network")])
async def test_auto_topup_failure_leaves_persisted_token_uncollected(
outcome: object,
) -> None:
provider = MagicMock()
provider.get_balance = AsyncMock(return_value=0)
provider.topup = AsyncMock(
side_effect=outcome if isinstance(outcome, Exception) else None,
return_value=outcome,
)
with (
patch(
"routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row",
return_value=provider,
),
patch(
"routstr.upstream.auto_topup.send_token",
AsyncMock(return_value="cashu-token"),
),
patch(
"routstr.upstream.auto_topup.store_cashu_transaction",
AsyncMock(return_value=True),
),
patch("routstr.upstream.auto_topup.create_session") as create_session,
):
if isinstance(outcome, Exception):
with pytest.raises(RuntimeError):
await _check_and_topup(_row())
else:
await _check_and_topup(_row())
create_session.assert_not_called()
@pytest.mark.asyncio
async def test_auto_topup_does_not_send_untracked_token() -> None:
provider = MagicMock()
provider.get_balance = AsyncMock(return_value=0)
provider.topup = AsyncMock()
with (
patch(
"routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row",
return_value=provider,
),
patch(
"routstr.upstream.auto_topup.send_token",
AsyncMock(return_value="cashu-token"),
),
patch(
"routstr.upstream.auto_topup.store_cashu_transaction",
AsyncMock(side_effect=RuntimeError("database unavailable")),
),
):
await _check_and_topup(_row())
provider.topup.assert_not_awaited()
+3 -288
View File
@@ -1,13 +1,12 @@
import json
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi.responses import JSONResponse
from routstr.balance import refund_wallet_endpoint, topup_wallet_endpoint
from routstr.balance import refund_wallet_endpoint
from routstr.core.db import ApiKey, CashuTransaction
from routstr.wallet import MintConnectionError, credit_balance
from routstr.wallet import credit_balance
def _make_cashu_tx(
@@ -104,64 +103,6 @@ async def test_refund_x_cashu_not_found_raises_404() -> None:
assert exc_info.value.status_code == 404
@pytest.mark.asyncio
async def test_refund_x_cashu_pending_raises_425() -> None:
"""in row exists with a request_id but out row not yet created → 425.
This is the race condition where /v1/wallet/refund is polled while the
upstream request is still in flight. The endpoint must signal "retry"
rather than a permanent 404.
"""
from fastapi import HTTPException
x_cashu_token = "cashuApending_token"
in_tx = _make_cashu_tx(
token=x_cashu_token, amount=0, unit="msat", type="in", request_id="req-pending"
)
session = MagicMock()
session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(None)])
session.add = MagicMock()
session.commit = AsyncMock()
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization="Bearer sk-somekey",
x_cashu=x_cashu_token,
session=session,
)
assert exc_info.value.status_code == 425
assert exc_info.value.headers == {"Retry-After": "2"}
@pytest.mark.asyncio
async def test_refund_x_cashu_in_tx_without_request_id_raises_404() -> None:
"""in row exists but has no request_id (cannot link to a refund) → 404.
This is a genuine "no refund will ever exist" case, distinct from the
pending 425 path.
"""
from fastapi import HTTPException
x_cashu_token = "cashuAnoreqid_token"
in_tx = _make_cashu_tx(
token=x_cashu_token, amount=0, unit="msat", type="in", request_id=None
)
session = MagicMock()
session.exec = AsyncMock(side_effect=[_exec_result(in_tx)])
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization="Bearer sk-somekey",
x_cashu=x_cashu_token,
session=session,
)
assert exc_info.value.status_code == 404
@pytest.mark.asyncio
async def test_refund_x_cashu_swept_raises_410() -> None:
from fastapi import HTTPException
@@ -398,10 +339,7 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch(
"routstr.balance.send_token",
AsyncMock(side_effect=MintConnectionError("raw mint outage detail")),
),
patch("routstr.balance.send_token", AsyncMock(side_effect=Exception("mint down"))),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
@@ -415,46 +353,10 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
)
assert exc_info.value.status_code == 503
assert exc_info.value.detail == "Mint service unavailable"
assert "raw mint outage detail" not in exc_info.value.detail
# Verify two exec calls: debit + restore
assert session.exec.await_count == 2
@pytest.mark.asyncio
async def test_apikey_refund_generic_failure_is_sanitized_500() -> None:
"""Unexpected send-side failures restore balance without leaking exception text."""
from fastapi import HTTPException
key = _make_api_key(balance=5000, refund_currency="sat")
raw_error = "database secret token raw-mint-response"
session = MagicMock()
session.get = AsyncMock(return_value=key)
session.exec = AsyncMock(side_effect=[_update_result(1), _update_result(1)])
session.commit = AsyncMock()
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(side_effect=RuntimeError(raw_error))),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch("routstr.balance.logger"),
):
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization="Bearer sk-testhash",
x_cashu=None,
session=session,
)
assert exc_info.value.status_code == 500
assert exc_info.value.detail == "Refund failed"
assert raw_error not in exc_info.value.detail
assert session.exec.await_count == 2
# ---------------------------------------------------------------------------
# no-create guarantee: fresh Cashu/unknown sk- tokens must not create API keys
# ---------------------------------------------------------------------------
@@ -499,190 +401,3 @@ async def test_refund_unknown_sk_bearer_returns_401() -> None:
assert exc_info.value.status_code == 401
session.get.assert_awaited_once()
# --- Topup redemption error taxonomy (POST /v1/wallet/topup) ------------------
@pytest.mark.asyncio
@pytest.mark.parametrize(
"error",
[
httpx.ConnectError("All connection attempts failed"),
MintConnectionError("connect to mint refused"),
TimeoutError("timed out connecting to mint"),
],
)
async def test_topup_mint_unreachable_returns_503(error: Exception) -> None:
"""A down mint must surface 503 (retryable), not 400 or 500 — the token is
fine, so the client should retry once the mint recovers."""
from fastapi import HTTPException
key = _make_api_key(balance=1000)
session = MagicMock()
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)),
):
with pytest.raises(HTTPException) as exc_info:
await topup_wallet_endpoint(
cashu_token="cashuAtoken", key=key, session=session
)
assert exc_info.value.status_code == 503
assert exc_info.value.detail == "Cashu mint is unreachable"
@pytest.mark.asyncio
async def test_topup_already_spent_still_returns_400() -> None:
"""Regression: the mint-unreachable short-circuit must not swallow the
existing ValueError substring buckets."""
from fastapi import HTTPException
key = _make_api_key(balance=1000)
session = MagicMock()
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch(
"routstr.balance.credit_balance",
AsyncMock(side_effect=ValueError("Token already spent")),
),
):
with pytest.raises(HTTPException) as exc_info:
await topup_wallet_endpoint(
cashu_token="cashuAtoken", key=key, session=session
)
assert exc_info.value.status_code == 400
assert exc_info.value.detail == "Cashu token already spent"
@pytest.mark.asyncio
async def test_topup_zero_value_returns_400_zero_value_message() -> None:
"""A dust/zero redemption maps to the documented zero-value message, not the
generic redemption-failed one."""
from fastapi import HTTPException
key = _make_api_key(balance=1000)
session = MagicMock()
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch(
"routstr.balance.credit_balance",
AsyncMock(
side_effect=ValueError("Redeemed token amount must be positive, got 0 msats")
),
),
):
with pytest.raises(HTTPException) as exc_info:
await topup_wallet_endpoint(
cashu_token="cashuAtoken", key=key, session=session
)
assert exc_info.value.status_code == 400
assert exc_info.value.detail == "Failed to redeem Cashu token: token yielded no value"
@pytest.mark.asyncio
async def test_topup_token_consumed_returns_500() -> None:
"""A post-redemption crediting failure (token spent) is a non-retryable 500,
not a 4xx that invites a retry."""
from fastapi import HTTPException
from routstr.wallet import TokenConsumedError
key = _make_api_key(balance=1000)
session = MagicMock()
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch(
"routstr.balance.credit_balance",
AsyncMock(side_effect=TokenConsumedError("credit failed")),
),
):
with pytest.raises(HTTPException) as exc_info:
await topup_wallet_endpoint(
cashu_token="cashuAtoken", key=key, session=session
)
assert exc_info.value.status_code == 500
assert exc_info.value.detail == (
"Token was redeemed but could not be credited; do not retry"
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("error", "expected_status", "expected_detail"),
[
(
ValueError(
"Failed to estimate fees: Fees (7 sat) exceed token amount (5 sat)"
),
422,
"Token value is too small to cover swap fees",
),
(
ValueError(
"Token amount (5 sat) is insufficient to cover melt fees."
),
422,
"Token value is too small to cover swap fees",
),
(
ValueError("Failed to melt token from foreign mint http://m: boom"),
422,
"Failed to swap token from foreign mint",
),
],
)
async def test_topup_fee_and_swap_failures_return_422(
error: Exception, expected_status: int, expected_detail: str
) -> None:
"""Fee/swap failures map to 422 (shared taxonomy), matching the bearer and
X-Cashu paths previously top-up flattened these to 400."""
from fastapi import HTTPException
key = _make_api_key(balance=1000)
session = MagicMock()
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)),
):
with pytest.raises(HTTPException) as exc_info:
await topup_wallet_endpoint(
cashu_token="cashuAtoken", key=key, session=session
)
assert exc_info.value.status_code == expected_status
assert exc_info.value.detail == expected_detail
@pytest.mark.asyncio
async def test_topup_unexpected_non_valueerror_returns_500() -> None:
"""A non-ValueError, non-transport fault is an internal error (500), not a
sanitized 400 the merged except must preserve this."""
from fastapi import HTTPException
key = _make_api_key(balance=1000)
session = MagicMock()
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch(
"routstr.balance.credit_balance",
AsyncMock(side_effect=RuntimeError("db exploded")),
),
):
with pytest.raises(HTTPException) as exc_info:
await topup_wallet_endpoint(
cashu_token="cashuAtoken", key=key, session=session
)
assert exc_info.value.status_code == 500
assert exc_info.value.detail == "Internal server error"
-102
View File
@@ -81,21 +81,6 @@ def test_backfill_strips_vendor_prefix_for_litellm_lookup() -> None:
assert result.input_cache_read == expected
def test_backfill_case_insensitive_lookup() -> None:
"""A generic upstream may report a mixed-case id
(deepseek-ai/DeepSeek-V4-Flash); litellm keys are lowercase. The
case-insensitive fallback still resolves the cache rate."""
pricing = Pricing(prompt=1.4e-07, completion=2.8e-07)
result = backfill_cache_pricing("deepseek-ai/DeepSeek-V4-Flash", pricing)
expected = litellm.model_cost["deepseek-v4-flash"][
"cache_read_input_token_cost"
]
assert result.input_cache_read == expected
assert result.input_cache_read < pricing.prompt # sanity: it's a discount
def test_backfill_fills_cache_write_rate() -> None:
"""Anthropic cache writes cost more than input (1.25x); billing them at
the input rate undercharges. litellm carries the write rate."""
@@ -151,93 +136,6 @@ def test_provider_fee_applies_to_backfilled_cache_rates() -> None:
assert adjusted.pricing.prompt == pytest.approx(2.8e-07 * 2.0)
def test_row_to_model_backfills_cache_rate() -> None:
"""The DB-override path (admin-configured providers, e.g. a generic
upstream) stores pricing without cache rates. ``_row_to_model`` must
backfill them from litellm just like ``_apply_provider_fee_to_model``,
otherwise cache reads bill at the full input rate."""
import json
from routstr.core.db import ModelRow
from routstr.payment.models import _row_to_model
row = ModelRow(
id="deepseek-v4-flash",
name="deepseek-v4-flash",
created=0,
description="",
context_length=1000000,
architecture=json.dumps(
{
"modality": "text",
"input_modalities": ["text"],
"output_modalities": ["text"],
"tokenizer": "unknown",
"instruct_type": None,
}
),
# Stored pricing omits input_cache_read (generic provider never sets it).
pricing=json.dumps({"prompt": 1.4e-07, "completion": 2.8e-07}),
enabled=True,
upstream_provider_id=1,
)
with patch(
"routstr.payment.models.sats_usd_price", return_value=5.0e-5
):
model = _row_to_model(row, apply_provider_fee=True, provider_fee=1.0)
litellm_read = litellm.model_cost["deepseek-v4-flash"][
"cache_read_input_token_cost"
]
assert model.pricing.input_cache_read == pytest.approx(litellm_read)
assert model.pricing.input_cache_read < model.pricing.prompt # a discount
assert model.sats_pricing is not None
assert model.sats_pricing.input_cache_read > 0
def test_row_to_model_backfills_via_forwarded_model_id() -> None:
"""An alias row (id != forwarded_model_id) must backfill cache rates from
the *forwarded* model name the real upstream model litellm prices
not the alias id, which litellm doesn't know."""
import json
from routstr.core.db import ModelRow
from routstr.payment.models import _row_to_model
row = ModelRow(
id="local-alias", # litellm has no such key
name="local-alias",
created=0,
description="",
context_length=1000000,
architecture=json.dumps(
{
"modality": "text",
"input_modalities": ["text"],
"output_modalities": ["text"],
"tokenizer": "unknown",
"instruct_type": None,
}
),
pricing=json.dumps({"prompt": 1.4e-07, "completion": 2.8e-07}),
enabled=True,
upstream_provider_id=1,
forwarded_model_id="deepseek-v4-flash",
)
with patch(
"routstr.payment.models.sats_usd_price", return_value=5.0e-5
):
model = _row_to_model(row, apply_provider_fee=True, provider_fee=1.0)
litellm_read = litellm.model_cost["deepseek-v4-flash"][
"cache_read_input_token_cost"
]
assert model.pricing.input_cache_read == pytest.approx(litellm_read)
assert model.pricing.input_cache_read < model.pricing.prompt
# ============================================================================
# calculate_cost — cached tokens billed at cache rates
# ============================================================================
@@ -1,56 +0,0 @@
from unittest.mock import AsyncMock, patch
import pytest
from routstr.core.db import store_cashu_transaction
@pytest.mark.asyncio
@pytest.mark.parametrize(
"error",
[
OSError("disk full"),
RuntimeError("connection lost"),
ConnectionRefusedError("database unavailable"),
],
)
async def test_store_cashu_transaction_propagates_commit_errors(
error: Exception,
) -> None:
session = AsyncMock()
session.commit.side_effect = error
session.__aenter__.return_value = session
session.__aexit__.return_value = None
with (
patch("routstr.core.db.create_session", return_value=session),
patch("routstr.core.db.logger.critical") as critical,
):
with pytest.raises(type(error), match=str(error)):
await store_cashu_transaction(
token="cashuAtest",
amount=1_000,
unit="sat",
mint_url="https://mint.example",
typ="out",
request_id="request-1",
)
critical.assert_called_once()
@pytest.mark.asyncio
async def test_store_cashu_transaction_returns_true_after_commit() -> None:
session = AsyncMock()
session.__aenter__.return_value = session
session.__aexit__.return_value = None
with patch("routstr.core.db.create_session", return_value=session):
stored = await store_cashu_transaction(
token="cashuAtest",
amount=1_000,
unit="sat",
)
assert stored is True
session.commit.assert_awaited_once()
@@ -1,91 +0,0 @@
from typing import Any
from unittest.mock import AsyncMock, patch
import pytest
from sqlalchemy.ext.asyncio import create_async_engine
from sqlmodel import SQLModel, select
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core import db
@pytest.mark.asyncio
async def test_cashu_transaction_storage_retries_then_succeeds() -> None:
store = AsyncMock(side_effect=[OSError("database locked"), True])
sleep = AsyncMock()
with (
patch("routstr.core.db.store_cashu_transaction", store),
patch("routstr.core.db.asyncio.sleep", sleep),
):
stored = await db.store_cashu_transaction_with_retry(
token="cashuAretry",
amount=100,
unit="sat",
)
assert stored is True
assert store.await_count == 2
sleep.assert_awaited_once_with(0.25)
@pytest.mark.asyncio
async def test_cashu_transaction_retry_is_idempotent_after_ambiguous_commit() -> None:
engine = create_async_engine("sqlite+aiosqlite://")
async with engine.begin() as connection:
await connection.run_sync(SQLModel.metadata.create_all)
original_store = db.store_cashu_transaction
attempts = 0
async def ambiguous_store(**kwargs: Any) -> bool:
nonlocal attempts
attempts += 1
stored = await original_store(**kwargs)
if attempts == 1:
raise OSError("connection dropped after commit")
return stored
with (
patch.object(db, "engine", engine),
patch("routstr.core.db.store_cashu_transaction", ambiguous_store),
patch("routstr.core.db.asyncio.sleep", AsyncMock()),
):
stored = await db.store_cashu_transaction_with_retry(
token="cashuAambiguous",
amount=100,
unit="sat",
)
async with AsyncSession(engine) as session:
result = await session.exec(select(db.CashuTransaction))
transactions = result.all()
assert stored is True
assert attempts == 2
assert len(transactions) == 1
await engine.dispose()
@pytest.mark.asyncio
async def test_cashu_transaction_storage_raises_after_bounded_retries() -> None:
error = OSError("database unavailable")
store = AsyncMock(side_effect=error)
sleep = AsyncMock()
with (
patch("routstr.core.db.store_cashu_transaction", store),
patch("routstr.core.db.asyncio.sleep", sleep),
patch("routstr.core.db.logger.critical") as critical,
):
with pytest.raises(OSError, match="database unavailable"):
await db.store_cashu_transaction_with_retry(
token="cashuAfail",
amount=100,
unit="sat",
max_attempts=3,
)
assert store.await_count == 3
assert [call.args[0] for call in sleep.await_args_list] == [0.25, 0.5]
critical.assert_called_once()
-231
View File
@@ -404,240 +404,9 @@ async def test_deepseek_malformed_hit_tokens_coerce_to_zero() -> None:
assert result.cache_read_input_tokens == 0
# ============================================================================
# Truly-empty response with a non-zero USD cost → full refund
#
# When an upstream reports a USD cost but the response carries NO tokens at all
# (input, output, cache-read and cache-creation all zero), billing the
# USD-derived cost charges the user for nothing. Refund in full. The gate is
# tightened relative to PR #489: a cache-read/-creation-only turn legitimately
# reports zero prompt/completion tokens with a real cost and must still bill.
# ============================================================================
@pytest.mark.asyncio
async def test_truly_empty_usd_cost_response_is_refunded(
mock_fixed_pricing: None,
) -> None:
"""0 input + 0 output + 0 cache tokens with a non-zero USD cost → refund."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 0,
"completion_tokens": 0,
"total_cost": 0.01, # non-zero USD cost despite no tokens
},
}
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.total_msats == 0 # full refund
assert result.input_msats == 0
assert result.output_msats == 0
assert result.total_usd == 0.0
assert result.input_tokens == 0
assert result.output_tokens == 0
assert result.cache_read_input_tokens == 0
assert result.cache_creation_input_tokens == 0
@pytest.mark.asyncio
async def test_cache_read_only_usd_cost_response_is_billed(
mock_fixed_pricing: None,
) -> None:
"""Cache-read-only turn (0 prompt/completion, non-zero cost) still bills."""
response = {
"model": "claude-3-5-sonnet",
"usage": {
"input_tokens": 0,
"output_tokens": 0,
"cache_read_input_tokens": 1000, # real cached usage
"cache_creation_input_tokens": 0,
"total_cost": 0.01, # non-zero USD cost
},
}
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
# NOT refunded — the USD cost is billed in full. Pinning the exact value
# guards against any future regression that would over-refund a cache-only
# turn (the bug in PR #489, which refunded whenever prompt+completion == 0).
assert result.total_msats == 200000
assert result.total_usd == 0.01
assert result.cache_read_input_tokens == 1000
@pytest.mark.asyncio
@pytest.mark.parametrize(
("total_cost", "input_cost", "output_cost", "expected_msats"),
[
(0.000471, 0.00023451, 0.00023649, 9420),
(0.00000004, 0.00000002, 0.00000002, 1),
],
)
async def test_small_usd_cost_components_sum_to_rounded_total(
total_cost: float,
input_cost: float,
output_cost: float,
expected_msats: int,
) -> None:
"""Small USD component costs must retain every billed millisatoshi."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 1,
"completion_tokens": 1,
"cost_details": {
"total_cost": total_cost,
"input_cost": input_cost,
"output_cost": output_cost,
},
},
}
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.total_msats == expected_msats
assert result.input_msats + result.output_msats == result.total_msats
@pytest.mark.asyncio
async def test_openrouter_upstream_inference_cost_components_are_used() -> None:
"""OpenRouter component aliases must determine the input/output split."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 375,
"completion_tokens": 158,
"total_tokens": 533,
"cost": 0.00022354,
"is_byok": False,
"prompt_tokens_details": {
"cached_tokens": 286,
"cache_write_tokens": 0,
},
"cost_details": {
"upstream_inference_cost": 0.00022354,
"upstream_inference_prompt_cost": 0.00004974,
"upstream_inference_completions_cost": 0.0001738,
},
"completion_tokens_details": {"reasoning_tokens": 17},
},
}
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.input_msats == 994
assert result.output_msats == 3477
assert result.input_msats + result.output_msats == result.total_msats == 4471
# ============================================================================
# PPQ.AI BYOK: upstream_inference_cost + BYOK fee billing
#
# PPQ.AI (bring-your-own-key) returns a small ~5 % BYOK routing fee in
# ``usage.cost`` and the real inference cost in
# ``cost_details.upstream_inference_cost``. The old code billed only the fee,
# under-charging by ~20×. The fix bills ``upstream_inference_cost + byok_fee``
# — what PPQ actually deducts from the balance.
# Payload numbers are from a live ``glm-5.2-fast`` request (GitHub issue #615).
# ============================================================================
@pytest.mark.asyncio
async def test_ppq_byok_bills_upstream_inference_cost_plus_fee() -> None:
"""PPQ.AI BYOK must bill upstream_inference_cost + byok_fee, not just the
fee. Mirrors the live request from GitHub issue #615."""
response = {
"model": "glm-5.2-fast",
"usage": {
"prompt_tokens": 164371, # includes 159301 cached
"completion_tokens": 99,
"cost": 0.002260057305, # ~5% BYOK routing fee
"is_byok": True,
"prompt_tokens_details": {"cached_tokens": 159301},
"cost_details": {
"upstream_inference_cost": 0.04475361,
"upstream_inference_prompt_cost": 0.04410021,
"upstream_inference_completions_cost": 0.0006534,
},
},
}
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
# The fix bills upstream_inference_cost + byok_fee (~0.047 USD → ~940k
# msats), not the fee alone (~0.0023 USD → ~45k msats). ~20× correction.
assert result.total_msats == 940274
assert result.input_msats + result.output_msats == result.total_msats
assert result.input_msats == 926546
assert result.output_msats == 13728
assert result.total_usd == pytest.approx(0.047013667305)
# Token normalisation (OpenAI dialect: cached included in prompt_tokens)
assert result.input_tokens == 5070 # 164371 - 159301
assert result.cache_read_input_tokens == 159301
assert result.output_tokens == 99
@pytest.mark.asyncio
async def test_ppq_byok_fee_only_would_undercharge() -> None:
"""Sanity check: billing only usage.cost (the BYOK fee) under-charges by
~20×. This documents the regression the fix prevents."""
response = {
"model": "glm-5.2-fast",
"usage": {
"prompt_tokens": 164371,
"completion_tokens": 99,
"cost": 0.002260057305, # BYOK fee only — no upstream_inference_cost
"is_byok": True,
"prompt_tokens_details": {"cached_tokens": 159301},
},
}
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
# Without upstream_inference_cost, only the fee is billed — the old bug.
assert result.total_msats == 45202
# ============================================================================
# Test 13: Missing Usage Block
# ============================================================================
@pytest.mark.asyncio
async def test_missing_upstream_cost_uses_litellm_model_pricing(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Token usage without an upstream cost is priced from LiteLLM."""
monkeypatch.setattr(settings, "fixed_pricing", True)
monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", 0)
monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", 0)
monkeypatch.setattr(
"routstr.payment.models.litellm_cost_entry",
lambda model: {
"input_cost_per_token": 0.000001,
"output_cost_per_token": 0.000002,
},
)
response = {
"model": "priced-by-litellm",
"usage": {
"prompt_tokens": 90,
"completion_tokens": 80,
"total_tokens": 170,
"prompt_tokens_details": {"cached_tokens": 0},
"completion_tokens_details": {"reasoning_tokens": 74},
"prompt_cache_hit_tokens": 0,
"prompt_cache_miss_tokens": 90,
},
}
result = await calculate_cost(response, max_cost=10000)
assert isinstance(result, CostData)
assert not isinstance(result, MaxCostData)
assert result.input_msats == 1800
assert result.output_msats == 3200
assert result.input_msats + result.output_msats == result.total_msats == 5000
@pytest.mark.asyncio
async def test_missing_usage_block(mock_fixed_pricing: None) -> None:
"""When usage is missing, return MaxCostData with zero tokens."""
-173
View File
@@ -1,173 +0,0 @@
"""Coverage tests for admin.py (currently 35%).
Tests admin endpoints that are testable without full app setup:
withdraw validation, authentication guards, and slug validation.
"""
from unittest.mock import Mock, patch
import pytest
from fastapi import HTTPException, Request
# ===========================================================================
# withdraw — validation and edge cases
# ===========================================================================
@pytest.mark.asyncio
async def test_withdraw_rejects_zero_amount() -> None:
"""withdraw validation rejects amount <= 0."""
from routstr.core.admin import WithdrawRequest, withdraw
request = Request(scope={"type": "http", "method": "POST"})
with pytest.raises(HTTPException) as exc_info:
await withdraw(request, WithdrawRequest(amount=0, unit="sat"))
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_withdraw_rejects_negative_amount() -> None:
"""withdraw validation rejects negative amounts."""
from routstr.core.admin import WithdrawRequest, withdraw
request = Request(scope={"type": "http", "method": "POST"})
with pytest.raises(HTTPException) as exc_info:
await withdraw(request, WithdrawRequest(amount=-100, unit="sat"))
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_withdraw_rejects_insufficient_balance() -> None:
"""withdraw returns 400 when wallet balance is insufficient."""
from routstr.core.admin import WithdrawRequest, withdraw
request = Request(scope={"type": "http", "method": "POST"})
with patch("routstr.core.admin.get_wallet") as mock_wallet, \
patch("routstr.core.admin.get_proofs_per_mint_and_unit") as mock_proofs, \
patch("routstr.core.admin.slow_filter_spend_proofs") as mock_filter:
mock_w = Mock()
mock_w.keysets = {}
mock_w.proofs = []
mock_wallet.return_value = mock_w
mock_proofs.return_value = []
mock_filter.return_value = []
with pytest.raises(HTTPException) as exc_info:
await withdraw(request, WithdrawRequest(amount=1000000, unit="sat"))
assert exc_info.value.status_code == 400
assert "Insufficient" in str(exc_info.value.detail)
# ===========================================================================
# update_password — validation
# ===========================================================================
@pytest.mark.asyncio
async def test_update_password_rejects_empty_new() -> None:
"""update_password rejects empty new password."""
from routstr.core.admin import PasswordUpdate, update_password
request = Request(scope={"type": "http", "method": "POST"})
with pytest.raises(HTTPException) as exc_info:
await update_password(
request,
PasswordUpdate(current_password="old", new_password=""),
)
# Returns 500 (no admin password configured) or 400 (validation)
assert exc_info.value.status_code in (400, 500, 422)
@pytest.mark.asyncio
async def test_update_password_rejects_short_new() -> None:
"""update_password rejects short passwords."""
from routstr.core.admin import PasswordUpdate, update_password
request = Request(scope={"type": "http", "method": "POST"})
with pytest.raises(HTTPException) as exc_info:
await update_password(
request,
PasswordUpdate(current_password="old", new_password="ab"),
)
assert exc_info.value.status_code in (400, 500, 422)
# ===========================================================================
# require_admin_api guard
# ===========================================================================
@pytest.mark.asyncio
async def test_require_admin_rejects_no_session() -> None:
"""require_admin_api rejects requests without admin session cookie."""
from routstr.core.admin import require_admin_api
request = Request(scope={
"type": "http",
"method": "GET",
"headers": [],
})
with pytest.raises(HTTPException) as exc_info:
await require_admin_api(request)
# 401 or 403 depending on auth configuration
assert exc_info.value.status_code in (401, 403)
# ===========================================================================
# _validate_slug
# ===========================================================================
def test_validate_slug_accepts_valid() -> None:
"""Valid slugs pass validation."""
from routstr.core.admin import _validate_slug
assert _validate_slug("valid-slug") == "valid-slug"
assert _validate_slug("valid123") == "valid123"
assert _validate_slug("my-provider") == "my-provider"
def test_validate_slug_rejects_spaces() -> None:
"""Slugs with spaces are rejected."""
from fastapi import HTTPException
from routstr.core.admin import _validate_slug
with pytest.raises(HTTPException):
_validate_slug("invalid slug")
def test_validate_slug_rejects_too_short() -> None:
"""Slugs shorter than 3 chars are rejected."""
from fastapi import HTTPException
from routstr.core.admin import _validate_slug
with pytest.raises(HTTPException):
_validate_slug("ab")
# ===========================================================================
# admin login endpoint
# ===========================================================================
@pytest.mark.asyncio
async def test_admin_login_requires_payload() -> None:
"""admin_login requires a payload — verify it exists."""
# Verify the function signature
import inspect
from routstr.core.admin import admin_login
sig = inspect.signature(admin_login)
params = list(sig.parameters.keys())
assert "request" in params
assert "payload" in params or len(params) >= 2
-204
View File
@@ -1,204 +0,0 @@
"""Coverage tests for base.py (currently 41%).
Tests preparers, builders, accessors, and model cache methods.
"""
from unittest.mock import Mock
import pytest
from routstr.upstream.base import BaseUpstreamProvider
# ===========================================================================
# prepare_headers
# ===========================================================================
def test_prepare_headers_adds_auth() -> None:
"""API key is added as Bearer token."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
headers = p.prepare_headers({})
assert "Authorization" in headers
assert headers["Authorization"] == "Bearer sk-test-key"
def test_prepare_headers_preserves_existing() -> None:
"""Existing headers are preserved."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
headers = p.prepare_headers({"X-Custom": "value", "Content-Type": "application/json"})
assert headers["X-Custom"] == "value"
assert headers["Content-Type"] == "application/json"
def test_prepare_headers_auth_header_passthrough() -> None:
"""Authorization header is handled — verify current behaviour."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
headers = p.prepare_headers({"Authorization": "Bearer user-key"})
# Currently provider key is used (may be intentional for proxy pattern)
assert "Authorization" in headers
# ===========================================================================
# prepare_params
# ===========================================================================
@pytest.mark.asyncio
async def test_prepare_params_passes_through() -> None:
"""Query params are preserved by default."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
params = p.prepare_params("/v1/chat/completions", {"temperature": "0.7"})
assert params["temperature"] == "0.7"
# ===========================================================================
# transform_model_name / normalize_request_path / get_request_base_url
# ===========================================================================
def test_transform_model_name_default_passthrough() -> None:
"""Default returns model_id unchanged."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
assert p.transform_model_name("gpt-4") == "gpt-4"
assert p.transform_model_name("") == ""
def test_normalize_request_path_passthrough() -> None:
"""Default returns path unchanged."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
assert p.normalize_request_path("/v1/chat/completions") == "/v1/chat/completions"
def test_get_request_base_url_default() -> None:
"""Default returns the provider's base_url."""
p = BaseUpstreamProvider("https://api.test.com/v1", "sk-test-key")
url = p.get_request_base_url("/v1/chat/completions")
assert url == "https://api.test.com/v1"
# ===========================================================================
# build_request_url
# ===========================================================================
def test_build_request_url_combines_base_and_path() -> None:
"""Combines base_url and path."""
p = BaseUpstreamProvider("https://api.test.com/v1", "sk-test-key")
url = p.build_request_url("/chat/completions")
assert "api.test.com" in url
assert "/chat/completions" in url
# ===========================================================================
# get_litellm_provider_prefix / get_provider_metadata
# ===========================================================================
def test_get_litellm_provider_prefix_default() -> None:
"""Default returns a string prefix."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
prefix = p.get_litellm_provider_prefix()
assert isinstance(prefix, str)
def test_get_provider_metadata_returns_dict() -> None:
"""Default metadata has name and capabilities."""
metadata = BaseUpstreamProvider.get_provider_metadata()
assert isinstance(metadata, dict)
assert "name" in metadata
# ===========================================================================
# from_db_row
# ===========================================================================
@pytest.mark.asyncio
async def test_from_db_row_returns_provider() -> None:
"""from_db_row constructs a provider from a valid row."""
mock_row = Mock()
mock_row.base_url = "https://api.test.com"
mock_row.api_key = "sk-test-key"
mock_row.slug = "test-slug"
mock_row.provider_fee = 1.0
mock_row.field_overrides = None
mock_row.name = "Test"
result = BaseUpstreamProvider.from_db_row(mock_row)
assert result is not None
# ===========================================================================
# prepare_request_body
# ===========================================================================
def test_prepare_request_body_with_model() -> None:
"""prepare_request_body takes bytes body and Model object."""
mock_model = Mock()
mock_model.id = "gpt-4"
mock_model.forwarded_model_id = None
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
# None body returns None
result = p.prepare_request_body(None, mock_model)
assert result is None
# ===========================================================================
# prepare_responses_request_body
# ===========================================================================
def test_prepare_responses_request_body_none() -> None:
"""None body returns None."""
model_obj = Mock()
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
result = p.prepare_responses_request_body(None, model_obj)
assert result is None
# ===========================================================================
# _upstream_accepts_cache_control
# ===========================================================================
def test_upstream_accepts_cache_control_default() -> None:
"""Default: upstream does NOT accept cache-control."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
assert p._upstream_accepts_cache_control() is False
# ===========================================================================
# inject_cost_metadata
# ===========================================================================
def test_inject_cost_metadata_adds_metadata() -> None:
"""Cost metadata is injected into the response dict."""
mock_key = Mock()
mock_key.balance_msat = 500000
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
data = {"model": "gpt-4", "usage": {"prompt_tokens": 100}}
cost_data = {
"base_msats": 200000,
"input_msats": 100000,
"output_msats": 100000,
"total_msats": 200000,
"total_usd": 0.01,
"input_tokens": 100,
"output_tokens": 50,
}
p.inject_cost_metadata(data, cost_data, mock_key)
# Metadata is nested under metadata.routstr.cost
assert "metadata" in data or "routstr_cost" in data or "cost" in data
# ===========================================================================
# _apply_provider_field
# ===========================================================================
def test_apply_provider_field_adds_to_response() -> None:
"""Provider field is added to response JSON."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
data = {"id": "chatcmpl-123"}
p._apply_provider_field(data)
assert "provider" in data
-217
View File
@@ -1,217 +0,0 @@
"""Additional coverage tests for base.py (41% → target 50%+).
Tests error message extraction, static helpers, model cache, and cost hooks.
These test existing correct behavior all should PASS.
"""
import json
from unittest.mock import Mock
import pytest
from routstr.upstream.base import BaseUpstreamProvider
# ===========================================================================
# _extract_upstream_error_message
# ===========================================================================
def test_extract_error_from_json_body() -> None:
"""Error message is extracted from JSON upstream error response."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
body = json.dumps({"error": {"message": "Model not found", "type": "not_found"}}).encode()
msg, error_type = p._extract_upstream_error_message(body)
assert "Model not found" in msg
assert error_type == "not_found"
def test_extract_error_from_simple_json() -> None:
"""Simple JSON error with direct message key."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
body = json.dumps({"message": "Rate limit exceeded"}).encode()
msg, error_type = p._extract_upstream_error_message(body)
assert "Rate limit" in msg
def test_extract_error_from_text_body() -> None:
"""Non-JSON text body is returned as-is."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
msg, error_type = p._extract_upstream_error_message(b"Internal Server Error")
assert "Internal Server Error" in msg
def test_extract_error_empty_body() -> None:
"""Empty body returns a generic message."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
msg, error_type = p._extract_upstream_error_message(b"")
assert isinstance(msg, str)
assert len(msg) > 0
def test_extract_error_simple_error_string_not_parsed() -> None:
"""JSON error as plain string (not dict) falls through to generic message."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
body = json.dumps({"error": "Invalid API key"}).encode()
msg, error_type = p._extract_upstream_error_message(body)
# Simple error strings not nested in a dict object use generic message
assert "Upstream request failed" in msg or "Invalid" in msg
# ===========================================================================
# on_upstream_error_redirect
# ===========================================================================
@pytest.mark.asyncio
async def test_on_upstream_error_redirect_noop() -> None:
"""Default implementation is a no-op for non-redirect statuses."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
await p.on_upstream_error_redirect(402, "Insufficient balance")
@pytest.mark.asyncio
async def test_on_upstream_error_redirect_429() -> None:
"""429 rate limit passes through (subclasses may override)."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
await p.on_upstream_error_redirect(429, "Rate limited")
# ===========================================================================
# _fold_cache_into_input_tokens (static method)
# ===========================================================================
def test_fold_cache_no_cache_data() -> None:
"""Usage without cache details is unchanged."""
from routstr.upstream.base import BaseUpstreamProvider
usage = Mock()
usage.prompt_tokens = 100
del usage.prompt_tokens_details # No cache details
BaseUpstreamProvider._fold_cache_into_input_tokens(usage)
# Should not modify the usage object when no cache exists
def test_fold_cache_preserves_total() -> None:
"""Total prompt tokens remain the same after folding cache."""
from routstr.upstream.base import BaseUpstreamProvider
usage = Mock()
usage.prompt_tokens = 100
details = Mock()
details.cached_tokens = 30
usage.prompt_tokens_details = details
BaseUpstreamProvider._fold_cache_into_input_tokens(usage)
# prompt_tokens should still be 100 (total unchanged)
assert usage.prompt_tokens == 100
# ===========================================================================
# get_cached_models / get_cached_model_by_id
# ===========================================================================
def test_get_cached_models_returns_list() -> None:
"""get_cached_models always returns a list."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
models = p.get_cached_models()
assert isinstance(models, list)
def test_get_cached_model_by_id_unknown_returns_none() -> None:
"""Unknown model ID returns None."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
result = p.get_cached_model_by_id("nonexistent-model-xyz-12345")
assert result is None
# ===========================================================================
# get_x_cashu_cost
# ===========================================================================
def test_get_x_cashu_cost_with_usage() -> None:
"""Cost is calculated from response data with usage info."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
response_data = {
"model": "gpt-4",
"usage": {"prompt_tokens": 100, "completion_tokens": 50},
}
result = p.get_x_cashu_cost(response_data, 100000, None)
# Either returns None (needs more data) or a cost object
assert result is not None
def test_get_x_cashu_cost_no_usage() -> None:
"""Response without usage returns MaxCostData."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
response_data = {"model": "gpt-4"}
result = p.get_x_cashu_cost(response_data, 100000, None)
# Without usage, uses max_cost
assert result is not None
# ===========================================================================
# get_balance
# ===========================================================================
@pytest.mark.asyncio
async def test_get_balance_raises_not_implemented() -> None:
"""Default get_balance raises NotImplementedError (no account support)."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
with pytest.raises(NotImplementedError):
await p.get_balance()
# ===========================================================================
# refresh_models_cache
# ===========================================================================
@pytest.mark.asyncio
async def test_refresh_models_cache_no_providers() -> None:
"""refresh_models_cache handles empty provider list gracefully."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
# Default implementation may be a no-op or raise
try:
await p.refresh_models_cache()
except Exception:
pass # May fail without DB — that's fine
# ===========================================================================
# fetch_models
# ===========================================================================
@pytest.mark.asyncio
async def test_fetch_models_returns_list() -> None:
"""fetch_models returns a model list (or empty) for default provider."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
try:
result = await p.fetch_models()
assert isinstance(result, list)
except Exception:
pass # May fail without network
# ===========================================================================
# create_account
# ===========================================================================
@pytest.mark.asyncio
async def test_create_account_raises_not_implemented() -> None:
"""Default create_account raises NotImplementedError."""
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
with pytest.raises(NotImplementedError):
await p.create_account()
-150
View File
@@ -1,150 +0,0 @@
"""Coverage-filling tests for middleware.py (currently 38% coverage).
Only LoggingMiddleware and request_id_context exist on main.
ConcurrencyLimiterMiddleware + TimeoutMiddleware are on an unmerged branch.
"""
from fastapi import FastAPI, Request
from fastapi.testclient import TestClient
# ---------------------------------------------------------------------------
# LoggingMiddleware
# ---------------------------------------------------------------------------
def test_logging_middleware_adds_request_id() -> None:
"""Every request gets an x-routstr-request-id header."""
from routstr.core.middleware import LoggingMiddleware
app = FastAPI()
@app.get("/test")
async def test_endpoint(request: Request) -> dict:
assert hasattr(request.state, "request_id")
assert request.state.request_id is not None
return {"ok": True}
app.add_middleware(LoggingMiddleware)
client = TestClient(app)
response = client.get("/test")
assert response.status_code == 200
assert "x-routstr-request-id" in response.headers
assert len(response.headers["x-routstr-request-id"]) == 36 # UUID4 length
def test_logging_middleware_skips_head_requests() -> None:
"""HEAD requests are skipped by _should_log (health probes)."""
from routstr.core.middleware import LoggingMiddleware
app = FastAPI()
@app.head("/test")
async def test_endpoint(request: Request) -> dict:
return {"ok": True}
app.add_middleware(LoggingMiddleware)
client = TestClient(app)
response = client.head("/test")
assert response.status_code == 200
assert "x-routstr-request-id" in response.headers
def test_logging_middleware_skips_options_requests() -> None:
"""OPTIONS requests (CORS preflight) are skipped."""
from routstr.core.middleware import LoggingMiddleware
app = FastAPI()
@app.options("/test")
async def test_endpoint(request: Request) -> dict:
return {"ok": True}
app.add_middleware(LoggingMiddleware)
client = TestClient(app)
response = client.options("/test")
assert response.status_code == 200
assert "x-routstr-request-id" in response.headers
def test_should_log_rejects_admin_api_prefix() -> None:
"""Admin API polling paths are skipped."""
from routstr.core.middleware import _should_log
assert _should_log("GET", "/admin/api/balances") is False
assert _should_log("GET", "/admin/api/logs") is False
assert _should_log("GET", "/admin/api/providers") is False
def test_should_log_rejects_nextjs_chunks() -> None:
"""Next.js static chunks are skipped."""
from routstr.core.middleware import _should_log
assert _should_log("GET", "/_next/static/chunks/main.js") is False
assert _should_log("GET", "/_next/data/build-id/page.json") is False
def test_should_log_rejects_exact_paths() -> None:
"""Exact paths like /favicon.ico are skipped."""
from routstr.core.middleware import _should_log
assert _should_log("GET", "/favicon.ico") is False
assert _should_log("GET", "/v1/wallet/info") is False
assert _should_log("GET", "/index.txt") is False
assert _should_log("GET", "/login/index.txt") is False
def test_should_log_accepts_normal_paths() -> None:
"""Normal API paths are logged."""
from routstr.core.middleware import _should_log
assert _should_log("GET", "/v1/chat/completions") is True
assert _should_log("POST", "/v1/chat/completions") is True
assert _should_log("GET", "/v1/models") is True
assert _should_log("POST", "/api/some-endpoint") is True
def test_should_log_accepts_non_skipped_path() -> None:
"""Generic paths not in skip list are logged."""
from routstr.core.middleware import _should_log
assert _should_log("GET", "/some/random/path") is True
assert _should_log("POST", "/api/custom") is True
def test_request_id_context_is_contextvar() -> None:
"""request_id_context is a ContextVar[str | None] with no default value."""
from contextvars import ContextVar
from routstr.core.middleware import request_id_context
assert isinstance(request_id_context, ContextVar)
# ContextVar without a default raises LookupError when accessed without being set
try:
val = request_id_context.get()
# If it returns, it should be None
assert val is None
except LookupError:
# Expected: ContextVar with no default raises LookupError
pass
def test_middleware_exports() -> None:
"""Only LoggingMiddleware is exported on main."""
from routstr.core.middleware import LoggingMiddleware, request_id_context
assert LoggingMiddleware is not None
assert request_id_context is not None
def test_middleware_skips_health_probe_path() -> None:
"""Health probe paths pass through without logging."""
from routstr.core.middleware import _should_log
# HEAD method is always skipped regardless of path
assert _should_log("HEAD", "/v1/chat/completions") is False
assert _should_log("OPTIONS", "/v1/chat/completions") is False
-181
View File
@@ -1,181 +0,0 @@
"""Coverage-filling tests for payment/helpers.py (currently 52% coverage).
Tests the real public API: check_token_balance, get_max_cost_for_model,
estimate_tokens, create_error_response, etc.
"""
from unittest.mock import Mock, patch
import pytest
# ---------------------------------------------------------------------------
# check_token_balance
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_check_token_balance_x_cashu_present() -> None:
"""X-Cashu header triggers token deserialization and balance check."""
from routstr.payment.helpers import check_token_balance
headers = {"x-cashu": "cashuAtest_token"}
body = {"model": "gpt-4"}
with patch("routstr.payment.helpers.deserialize_token_from_string") as mock_deser:
mock_token = Mock()
mock_token.amount = 50000
mock_token.unit = "sat"
mock_deser.return_value = mock_token
# Should not raise — balance is sufficient
check_token_balance(headers, body, 1000)
@pytest.mark.asyncio
async def test_check_token_balance_no_x_cashu_raises() -> None:
"""Missing X-Cashu header raises HTTPException (401 on main)."""
from fastapi import HTTPException
from routstr.payment.helpers import check_token_balance
headers: dict[str, str] = {}
body = {"model": "gpt-4"}
with pytest.raises(HTTPException) as exc_info:
check_token_balance(headers, body, 1000)
assert exc_info.value.status_code == 401
@pytest.mark.asyncio
async def test_check_token_balance_insufficient_raises() -> None:
"""Token with insufficient balance raises HTTPException 402.
max_cost_for_model is in msat, so with amount=100 sat (=100,000 msat),
max_cost=200,000 msat triggers the insufficient balance check.
"""
from fastapi import HTTPException
from routstr.payment.helpers import check_token_balance
headers = {"x-cashu": "cashuAtest_token"}
body = {"model": "gpt-4"}
with patch("routstr.payment.helpers.deserialize_token_from_string") as mock_deser:
mock_token = Mock()
mock_token.amount = 100 # 100 sat
mock_token.unit = "sat"
mock_deser.return_value = mock_token
with pytest.raises(HTTPException) as exc_info:
# 200,000 msat > 100,000 msat (100 sat * 1000)
check_token_balance(headers, body, 200000)
assert exc_info.value.status_code == 402
# ---------------------------------------------------------------------------
# estimate_tokens
# ---------------------------------------------------------------------------
def test_estimate_tokens_empty_messages() -> None:
"""Empty message list returns 0 tokens."""
from routstr.payment.helpers import estimate_tokens
result = estimate_tokens([])
assert result == 0
def test_estimate_tokens_text_content() -> None:
"""Text messages are counted."""
from routstr.payment.helpers import estimate_tokens
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello, how are you?"},
]
result = estimate_tokens(messages)
assert result > 0
assert isinstance(result, int)
def test_estimate_tokens_long_text() -> None:
"""Longer messages produce higher token counts."""
from routstr.payment.helpers import estimate_tokens
short = estimate_tokens([{"role": "user", "content": "Hi"}])
long = estimate_tokens([{"role": "user", "content": "Hello " * 100}])
assert long > short
# ---------------------------------------------------------------------------
# create_error_response
# ---------------------------------------------------------------------------
def test_create_error_response_402() -> None:
"""402 Payment Required error is properly formatted."""
from fastapi import Request
from routstr.payment.helpers import create_error_response
request = Request(scope={"type": "http", "method": "GET"})
result = create_error_response("insufficient_funds", "Insufficient balance", 402, request)
assert result.status_code == 402
def test_create_error_response_500() -> None:
"""500 Internal Server Error is properly formatted."""
from fastapi import Request
from routstr.payment.helpers import create_error_response
request = Request(scope={"type": "http", "method": "GET"})
result = create_error_response("server_error", "Internal error", 500, request)
assert result.status_code == 500
# ---------------------------------------------------------------------------
# Image token estimation helpers
# ---------------------------------------------------------------------------
def test_image_dimensions_valid_png() -> None:
"""_get_image_dimensions returns width and height for a valid PNG."""
from routstr.payment.helpers import _get_image_dimensions
# A minimal 1x1 red PNG (valid minimal file)
png = (
b"\x89PNG\r\n\x1a\n"
b"\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01\x08\x02"
b"\x00\x00\x00\x90wS\xde"
b"\x00\x00\x00\x0cIDAT\x08\xd7c\xf8\x0f\x00\x00\x01\x01\x00\x05"
b"\x18\xd8N"
b"\x00\x00\x00\x00IEND\xaeB`\x82"
)
w, h = _get_image_dimensions(png)
assert w == 1
assert h == 1
def test_calculate_image_tokens_low_detail() -> None:
"""Low detail images are always 85 tokens."""
from routstr.payment.helpers import _calculate_image_tokens
tokens = _calculate_image_tokens(1024, 1024, "low")
assert tokens == 85
def test_calculate_image_tokens_high_detail() -> None:
"""High detail images are scaled and tile-based."""
from routstr.payment.helpers import _calculate_image_tokens
tokens = _calculate_image_tokens(1024, 1024, "high")
assert tokens > 85
assert isinstance(tokens, int)
-161
View File
@@ -1,161 +0,0 @@
"""Coverage tests for proxy.py (currently 47%).
Tests request parsing, model extraction, and routing helpers.
"""
import json
import pytest
from fastapi import HTTPException
# ===========================================================================
# parse_request_body_json
# ===========================================================================
def test_parse_json_valid_body() -> None:
"""Valid JSON body is parsed correctly for chat completions."""
from routstr.proxy import parse_request_body_json
body = json.dumps({"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]}).encode()
result = parse_request_body_json(body, "/v1/chat/completions")
assert result["model"] == "gpt-4"
assert result["messages"][0]["role"] == "user"
def test_parse_json_invalid_raises_400() -> None:
"""Invalid JSON raises HTTPException 400."""
from routstr.proxy import parse_request_body_json
with pytest.raises(HTTPException) as exc_info:
parse_request_body_json(b"not json", "/v1/chat/completions")
assert exc_info.value.status_code == 400
def test_parse_json_empty_body() -> None:
"""Empty body returns empty dict."""
from routstr.proxy import parse_request_body_json
result = parse_request_body_json(b"", "/v1/chat/completions")
assert isinstance(result, dict)
assert result == {}
def test_parse_json_responses_path() -> None:
"""Responses API path is handled."""
from routstr.proxy import parse_request_body_json
body = json.dumps({"model": "gpt-4", "input": "hello"}).encode()
result = parse_request_body_json(body, "/v1/responses")
assert "model" in result
def test_parse_json_rejects_non_integer_max_tokens() -> None:
"""max_tokens must be an integer."""
from routstr.proxy import parse_request_body_json
body = json.dumps({"model": "gpt-4", "max_tokens": "abc"}).encode()
with pytest.raises(HTTPException) as exc_info:
parse_request_body_json(body, "/v1/chat/completions")
assert exc_info.value.status_code == 400
# ===========================================================================
# extract_model_from_responses_request
# ===========================================================================
def test_extract_model_from_responses() -> None:
"""Model name is extracted from Responses API request."""
from routstr.proxy import extract_model_from_responses_request
body = {"model": "gpt-4o", "input": "test"}
model = extract_model_from_responses_request(body)
assert model == "gpt-4o"
def test_extract_model_returns_unknown_for_missing() -> None:
"""Missing model field returns 'unknown'."""
from routstr.proxy import extract_model_from_responses_request
body = {"input": "test"}
model = extract_model_from_responses_request(body)
assert model == "unknown"
def test_extract_model_empty_body_returns_unknown() -> None:
"""Empty body returns 'unknown'."""
from routstr.proxy import extract_model_from_responses_request
model = extract_model_from_responses_request({})
assert model == "unknown"
def test_extract_model_from_input_nested() -> None:
"""Model nested in input dict is found."""
from routstr.proxy import extract_model_from_responses_request
body = {"input": {"model": "claude-sonnet", "text": "hi"}}
model = extract_model_from_responses_request(body)
# The function checks input_data.get("model") for nested
assert model in ("claude-sonnet", "unknown")
# ===========================================================================
# get_model_instance / get_provider_for_model / get_unique_models
# ===========================================================================
def test_get_model_instance_unknown_returns_none() -> None:
"""Unknown model ID returns None."""
from routstr.proxy import get_model_instance
result = get_model_instance("nonexistent-model-xyz-12345")
assert result is None
def test_get_provider_for_model_unknown_returns_none() -> None:
"""Unknown model returns None."""
from routstr.proxy import get_provider_for_model
result = get_provider_for_model("nonexistent-model-xyz-12345")
assert result is None
def test_get_unique_models_returns_list() -> None:
"""get_unique_models always returns a list."""
from routstr.proxy import get_unique_models
result = get_unique_models()
assert isinstance(result, list)
def test_get_upstreams_returns_list() -> None:
"""get_upstreams returns a list of providers."""
from routstr.proxy import get_upstreams
result = get_upstreams()
assert isinstance(result, list)
# ===========================================================================
# parse_request_body_json — nested objects
# ===========================================================================
def test_parse_body_preserves_nested_objects() -> None:
"""Nested JSON objects are preserved during parsing."""
from routstr.proxy import parse_request_body_json
body = json.dumps({
"model": "claude-3",
"messages": [{"role": "system", "content": "You are helpful."}],
"temperature": 0.7,
"max_tokens": 1024,
}).encode()
result = parse_request_body_json(body, "/v1/chat/completions")
assert result["temperature"] == 0.7
assert result["max_tokens"] == 1024
assert len(result["messages"]) == 1
@@ -1,89 +0,0 @@
"""Tests asserting CORRECT behavior for DB persistence and payout safety.
RED tests FAIL against current main until bugs are fixed.
"""
import inspect
# ===========================================================================
# RED TESTS: Fee payout crash safety
# ===========================================================================
def test_fee_payout_pre_reset_or_lock_exists() -> None:
"""FIX REQUIRED: pay-then-reset must become lock-then-pay-then-unlock.
wallet.py:1076-1080 currently: raw_send_to_lnurl() THEN reset_routstr_fee().
A crash between these lines causes double payment.
Fix: set a lock flag BEFORE paying, clear it AFTER resetting.
On startup, reconcile any locked-but-not-reset payouts.
"""
from routstr import wallet
source = inspect.getsource(wallet.periodic_routstr_fee_payout)
pay_pos = source.find("raw_send_to_lnurl")
reset_pos = source.find("reset_routstr_fee")
assert pay_pos > 0 and reset_pos > 0, "Pay and reset both exist"
# After fix: lock/safeguard must exist BEFORE the pay call
pre_pay_section = source[:pay_pos]
has_pre_guard = any(
kw in pre_pay_section.lower()
for kw in ["lock", "payout_state", "is_paying", "in_progress",
"pre_reset", "reconcile", "checkpoint"]
)
assert has_pre_guard, (
"FIX REQUIRED: Fee payout pays before resetting with no crash guard. "
"A crash between pay and reset causes double payment. "
"Fix: add a DB lock/payout_state flag before paying."
)
# ===========================================================================
# RED TESTS: DB store resilience
# ===========================================================================
def test_retry_wrapper_exists() -> None:
"""FIX REQUIRED: A retry wrapper for critical DB writes must exist."""
from routstr.core import db
assert hasattr(db, "store_cashu_transaction_with_retry"), (
"FIX REQUIRED: No retry wrapper exists for critical money-path "
"DB writes. Was merged (#600) then reverted (#604). Must be "
"reinstated with CRITICAL logging on final failure."
)
# ===========================================================================
# Wallet caching mechanism (informational — not a bug on main)
# ===========================================================================
def test_wallet_cache_uses_global_dict() -> None:
"""get_wallet uses a global _wallets dict — verify mechanism."""
from routstr import wallet
source = inspect.getsource(wallet.get_wallet)
assert "_wallets" in source
assert "load_mint" in source
assert "load_proofs" in source
# ===========================================================================
# Mint rate limiter setting (informational)
# ===========================================================================
def test_mint_concurrency_setting_exists_or_documents_gap() -> None:
"""If mint_max_concurrency exists, it must NOT be 0.
0 disables 429 cooldown tracking on the PR #597 branch.
"""
from routstr.core.settings import settings
concurrency = getattr(settings, "mint_max_concurrency", None)
if concurrency is not None:
assert concurrency > 0, (
f"mint_max_concurrency = {concurrency}. 0 disables 429 cooldown."
)
-208
View File
@@ -1,208 +0,0 @@
from __future__ import annotations
from typing import Any, AsyncGenerator
from unittest.mock import AsyncMock, MagicMock
import pytest
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlalchemy.pool import StaticPool
from sqlmodel import SQLModel, select
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.auth import get_reservation_snapshot, pay_for_request
from routstr.core.db import ApiKey, ReservationRelease
from routstr.upstream.ehbp import (
finalize_ehbp_actual_cost_payment,
finalize_ehbp_max_cost_payment,
)
def _make_engine() -> AsyncEngine:
return create_async_engine(
"sqlite+aiosqlite://",
poolclass=StaticPool,
connect_args={"check_same_thread": False},
)
@pytest.fixture
async def session(monkeypatch: pytest.MonkeyPatch) -> AsyncGenerator[AsyncSession, None]:
monkeypatch.setattr("routstr.upstream.ehbp.ROUTSTR_FEE_PERCENT", 0)
engine = _make_engine()
async with engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.create_all)
db_session = AsyncSession(engine, expire_on_commit=False)
try:
yield db_session
finally:
await db_session.close()
await engine.dispose()
async def _api_key(session: AsyncSession, hashed_key: str) -> ApiKey | None:
return (
await session.exec(select(ApiKey).where(ApiKey.hashed_key == hashed_key))
).one_or_none()
def _fail_nth_api_key_update(
session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
target_update: int,
) -> None:
"""Return rowcount=0 for one API-key UPDATE without mutating the database."""
original_exec = session.exec
api_key_updates = 0
async def exec_with_failure(
statement: Any, *args: Any, **kwargs: Any
) -> Any:
nonlocal api_key_updates
table = getattr(statement, "table", None)
if getattr(table, "name", None) == "api_keys":
api_key_updates += 1
if api_key_updates == target_update:
return MagicMock(rowcount=0)
return await original_exec(statement, *args, **kwargs)
monkeypatch.setattr(session, "exec", exec_with_failure)
@pytest.mark.asyncio
async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve(
session: AsyncSession,
) -> None:
key = ApiKey(hashed_key="ehbp-actual", balance=10_000)
session.add(key)
await session.commit()
await pay_for_request(key, 3_000, session)
reservation = await get_reservation_snapshot(key, session)
await finalize_ehbp_actual_cost_payment(
key,
session,
reserved_cost_for_model=3_000,
model_id="tinfoil/model",
cost_info={
"total_msats": 1_200,
"input_tokens": 10,
"output_tokens": 20,
"input_msats": 500,
"output_msats": 700,
},
reservation_snapshot=reservation,
)
updated = await _api_key(session, "ehbp-actual")
assert updated is not None
assert updated.balance == 8_800
assert updated.reserved_balance == 0
assert updated.reserved_at is None
assert updated.total_spent == 1_200
@pytest.mark.asyncio
async def test_finalize_max_cost_payment_updates_parent_and_child_spend(
session: AsyncSession,
) -> None:
parent = ApiKey(hashed_key="ehbp-parent", balance=10_000)
child = ApiKey(
hashed_key="ehbp-child", balance=0, parent_key_hash="ehbp-parent"
)
session.add(parent)
session.add(child)
await session.commit()
await pay_for_request(child, 3_000, session)
reservation = await get_reservation_snapshot(child, session)
await finalize_ehbp_max_cost_payment(
child,
session,
max_cost_for_model=3_000,
model_id="tinfoil/model",
reservation_snapshot=reservation,
)
updated_parent = await _api_key(session, "ehbp-parent")
updated_child = await _api_key(session, "ehbp-child")
assert updated_parent is not None
assert updated_child is not None
assert updated_parent.balance == 7_000
assert updated_parent.reserved_balance == 0
assert updated_parent.reserved_at is None
assert updated_parent.total_spent == 3_000
assert updated_child.balance == 0
assert updated_child.reserved_balance == 0
assert updated_child.reserved_at is None
assert updated_child.total_spent == 3_000
@pytest.mark.asyncio
async def test_finalize_actual_cost_payment_rolls_back_when_parent_update_matches_no_rows(
session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
key = ApiKey(hashed_key="ehbp-missing-parent", balance=10_000)
session.add(key)
await session.commit()
await pay_for_request(key, 3_000, session)
reservation = await get_reservation_snapshot(key, session)
_fail_nth_api_key_update(session, monkeypatch, target_update=1)
rollback_spy = AsyncMock(wraps=session.rollback)
monkeypatch.setattr(session, "rollback", rollback_spy)
await finalize_ehbp_actual_cost_payment(
key,
session,
reserved_cost_for_model=3_000,
model_id="tinfoil/model",
cost_info={"total_msats": 1_200},
reservation_snapshot=reservation,
)
rollback_spy.assert_awaited_once()
updated = await _api_key(session, "ehbp-missing-parent")
assert updated is not None
assert updated.balance == 10_000
assert updated.reserved_balance == 3_000
assert updated.total_spent == 0
release = await session.get(ReservationRelease, reservation.release_id)
assert release is not None
assert release.status == "active"
@pytest.mark.asyncio
async def test_finalize_max_cost_payment_rolls_back_parent_when_child_update_matches_no_rows(
session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
parent = ApiKey(hashed_key="ehbp-rollback-parent", balance=10_000)
child = ApiKey(
hashed_key="ehbp-missing-child",
balance=0,
parent_key_hash="ehbp-rollback-parent",
)
session.add(parent)
session.add(child)
await session.commit()
await pay_for_request(child, 3_000, session)
reservation = await get_reservation_snapshot(child, session)
_fail_nth_api_key_update(session, monkeypatch, target_update=2)
await finalize_ehbp_max_cost_payment(
child,
session,
max_cost_for_model=3_000,
model_id="tinfoil/model",
reservation_snapshot=reservation,
)
updated_parent = await _api_key(session, "ehbp-rollback-parent")
assert updated_parent is not None
assert updated_parent.balance == 10_000
assert updated_parent.reserved_balance == 3_000
assert updated_parent.total_spent == 0
updated_child = await _api_key(session, "ehbp-missing-child")
assert updated_child is not None
assert updated_child.reserved_balance == 3_000
assert updated_child.total_spent == 0
@@ -1,215 +0,0 @@
"""Tests asserting CORRECT behavior for emergency refund and DB persistence.
These tests FAIL against current main because the code is buggy.
They serve as the "RED" phase of TDD once the bugs are fixed, they go green.
Correct behavior required:
1. store_cashu_transaction should raise on failure (not silently return False)
2. Emergency refund paths must NOT use try/except/pass for DB stores
3. A retry wrapper must exist for critical money-path DB writes
"""
from unittest.mock import AsyncMock, patch
import pytest
# ===========================================================================
# RED TESTS: store_cashu_transaction should RAISE on failure
# ===========================================================================
@pytest.mark.asyncio
async def test_store_cashu_raises_on_db_failure_not_returns_false() -> None:
"""FIX REQUIRED: store_cashu_transaction must raise on DB failure.
Currently returns False silently callers never detect the failure.
Correct behavior: raise an exception so callers can recover.
"""
from routstr.core.db import store_cashu_transaction
with patch("routstr.core.db.create_session") as mock_create:
mock_session = AsyncMock()
mock_session.commit = AsyncMock(side_effect=OSError("disk full"))
mock_session.__aenter__ = AsyncMock(return_value=mock_session)
mock_session.__aexit__ = AsyncMock(return_value=None)
mock_create.return_value = mock_session
with pytest.raises(Exception) as exc_info:
await store_cashu_transaction(
token="cashuAtest_refund_token",
amount=1000,
unit="sat",
mint_url="http://mint:3338",
typ="out",
request_id="req-123",
)
# Must raise a meaningful exception, not silently return False
# OSError or a custom DB error is acceptable
assert "disk full" in str(exc_info.value) or isinstance(
exc_info.value, (OSError, RuntimeError)
), (
f"Expected store to propagate the failure, got {type(exc_info.value).__name__}: "
f"{exc_info.value}"
)
@pytest.mark.asyncio
async def test_store_cashu_raises_on_any_error() -> None:
"""FIX REQUIRED: All DB errors must propagate, not just OSError."""
from routstr.core.db import store_cashu_transaction
errors = [
OSError("disk full"),
RuntimeError("connection lost"),
ConnectionRefusedError("db down"),
]
for error in errors:
with patch("routstr.core.db.create_session") as mock_create:
mock_session = AsyncMock()
mock_session.commit = AsyncMock(side_effect=error)
mock_session.__aenter__ = AsyncMock(return_value=mock_session)
mock_session.__aexit__ = AsyncMock(return_value=None)
mock_create.return_value = mock_session
with pytest.raises(Exception):
await store_cashu_transaction(
token="cashuAtest",
amount=1000,
unit="sat",
typ="out",
)
# ===========================================================================
# RED TESTS: Retry wrapper must exist
# ===========================================================================
def test_retry_wrapper_exists_for_critical_writes() -> None:
"""FIX REQUIRED: store_cashu_transaction_with_retry must exist.
Currently reverted (#600 → #604). All critical money-path DB writes
(after minting a token) need retry with backoff + CRITICAL logging.
"""
from routstr.core import db
assert hasattr(db, "store_cashu_transaction_with_retry"), (
"FIX REQUIRED: store_cashu_transaction_with_retry does not exist. "
"Was merged in PR #600, reverted in PR #604. "
"All post-mint DB writes need retry + backoff + CRITICAL logging."
)
@pytest.mark.asyncio
async def test_retry_wrapper_retries_on_transient_failure() -> None:
"""FIX REQUIRED: retry wrapper must retry, not fail on first attempt."""
from routstr.core import db
# Skip if the retry wrapper doesn't exist yet
if not hasattr(db, "store_cashu_transaction_with_retry"):
pytest.skip("store_cashu_transaction_with_retry does not exist yet")
with patch("routstr.core.db.store_cashu_transaction") as mock_store:
mock_store = AsyncMock()
mock_store.side_effect = [OSError("transient"), None] # 1st fails, 2nd succeeds
# We'd test that the wrapper retries, but it doesn't exist yet
# This test documents the expected behavior
# ===========================================================================
# RED TESTS: Emergency refund must not silently lose tokens
# ===========================================================================
def test_emergency_refund_no_try_except_pass() -> None:
"""FIX REQUIRED: Emergency refund paths must NOT use try/except/pass.
base.py:3643-3653 (chat) and base.py:4607-4617 (responses) both use
try/except/pass around store_cashu_transaction after minting a refund
token. If DB write fails, the token is permanently lost.
The fix: remove try/except/pass. Let the exception propagate so
the caller can detect failure and at minimum log the token.
"""
import inspect
from routstr.upstream.base import BaseUpstreamProvider
# Check chat emergency refund handler
chat_src = inspect.getsource(
BaseUpstreamProvider.handle_x_cashu_non_streaming_response
)
# Find the emergency refund section
emergency_start = chat_src.find("emergency_refund = amount")
assert emergency_start > 0, "Emergency refund path exists"
emergency_section = chat_src[emergency_start : emergency_start + 500]
# The try/except/pass around store_cashu_transaction must NOT exist
has_except_pass = "except Exception:" in emergency_section and "pass" in emergency_section
assert not has_except_pass, (
"FIX REQUIRED: Emergency refund (chat) uses try/except/pass around "
"store_cashu_transaction. A failed DB write silently loses the minted "
"token. Fix: let the exception propagate or log at CRITICAL with the "
"full token for manual recovery."
)
def test_emergency_refund_responses_api_no_silent_failure() -> None:
"""FIX REQUIRED: Responses API emergency refund same fix as chat."""
import inspect
from routstr.upstream.base import BaseUpstreamProvider
responses_src = inspect.getsource(
BaseUpstreamProvider.handle_x_cashu_non_streaming_responses_response
)
has_emergency = "emergency_refund = amount" in responses_src
if has_emergency:
emergency_start = responses_src.find("emergency_refund = amount")
emergency_section = responses_src[emergency_start : emergency_start + 500]
has_except_pass = (
"except Exception:" in emergency_section and "pass" in emergency_section
)
assert not has_except_pass, (
"FIX REQUIRED: Responses API emergency refund also uses "
"try/except/pass. Same fund-loss vulnerability as chat path."
)
# ===========================================================================
# RED TESTS: Fee payout crash safety
# ===========================================================================
def test_fee_payout_has_crash_guard() -> None:
"""FIX REQUIRED: Fee payout must have guard against double-pay on crash.
wallet.py:1076-1080 pays LNURL THEN resets the fee counter.
A crash between these steps causes double payment on restart.
Fix options:
1. Pre-reset the counter before paying (if pay fails, restore it)
2. Add a "payout_lock" DB flag that's set before pay and cleared after
3. Record payout in DB and reconcile on startup
"""
import inspect
from routstr import wallet
source = inspect.getsource(wallet.periodic_routstr_fee_payout)
# After the fix, the pay-then-reset pattern should be replaced
# with a safe sequence. Verify the guard exists.
has_guard = any(
kw in source.lower()
for kw in ["payout_lock", "is_paying", "payout_in_progress",
"pre_reset", "reset_before", "reconcile"]
)
assert has_guard, (
"FIX REQUIRED: Fee payout has no crash guard. Pay-then-reset "
"pattern in periodic_routstr_fee_payout can double-pay on "
"process restart."
)
-160
View File
@@ -1,160 +0,0 @@
import asyncio
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock, patch
import pytest
from sqlalchemy.ext.asyncio import create_async_engine
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr import wallet
from routstr.core import db
@asynccontextmanager
async def _session_context(session: Mock) -> AsyncIterator[Mock]:
yield session
@pytest.mark.asyncio
async def test_fee_payout_checkpoint_is_atomic_and_durable() -> None:
engine = create_async_engine("sqlite+aiosqlite://")
async with engine.begin() as connection:
await connection.run_sync(SQLModel.metadata.create_all)
async with AsyncSession(engine) as session:
session.add(db.RoutstrFee(id=1, accumulated_msats=5_000))
await session.commit()
assert await db.reset_routstr_fee(session, 5_000) is True
assert await db.reset_routstr_fee(session, 5_000) is False
fee = await db.get_routstr_fee(session)
await session.refresh(fee)
assert fee.accumulated_msats == 0
assert fee.payout_in_progress_msats == 5_000
assert fee.total_paid_msats == 0
assert await db.complete_routstr_fee_payout(session, 5_000) is True
await session.refresh(fee)
assert fee.payout_in_progress_msats == 0
assert fee.total_paid_msats == 5_000
assert fee.last_paid_at is not None
await engine.dispose()
@pytest.mark.asyncio
async def test_fee_payout_checkpoints_before_sending() -> None:
session = Mock()
fee = SimpleNamespace(
accumulated_msats=5_000,
payout_in_progress_msats=0,
payout_started_at=None,
)
payout_wallet = Mock()
events: list[str] = []
async def checkpoint(*_args: object) -> bool:
events.append("checkpoint")
return True
async def send(*_args: object, **_kwargs: object) -> int:
events.append("send")
return 5
async def complete(*_args: object) -> bool:
events.append("complete")
return True
with (
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
patch(
"routstr.wallet.asyncio.sleep",
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
),
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
patch("routstr.wallet.db.reset_routstr_fee", side_effect=checkpoint),
patch("routstr.wallet.db.complete_routstr_fee_payout", side_effect=complete),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=payout_wallet)),
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
patch("routstr.wallet.raw_send_to_lnurl", side_effect=send),
):
with pytest.raises(asyncio.CancelledError):
await wallet.periodic_routstr_fee_payout()
assert events == ["checkpoint", "send", "complete"]
@pytest.mark.asyncio
async def test_fee_payout_does_not_retry_an_unresolved_checkpoint() -> None:
session = Mock()
fee = SimpleNamespace(
accumulated_msats=10_000,
payout_in_progress_msats=5_000,
payout_started_at=123,
)
with (
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
patch(
"routstr.wallet.asyncio.sleep",
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
),
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock()) as checkpoint,
patch("routstr.wallet.get_wallet", AsyncMock()) as get_wallet,
patch("routstr.wallet.raw_send_to_lnurl", AsyncMock()) as send,
patch("routstr.wallet.logger.critical") as critical,
):
with pytest.raises(asyncio.CancelledError):
await wallet.periodic_routstr_fee_payout()
checkpoint.assert_not_awaited()
get_wallet.assert_not_awaited()
send.assert_not_awaited()
critical.assert_called_once()
@pytest.mark.asyncio
async def test_fee_payout_keeps_checkpoint_when_send_outcome_is_unknown() -> None:
session = Mock()
fee = SimpleNamespace(
accumulated_msats=5_000,
payout_in_progress_msats=0,
payout_started_at=None,
)
complete = AsyncMock()
with (
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
patch(
"routstr.wallet.asyncio.sleep",
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
),
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock(return_value=True)),
patch("routstr.wallet.db.complete_routstr_fee_payout", complete),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
patch(
"routstr.wallet.raw_send_to_lnurl",
AsyncMock(side_effect=TimeoutError("unknown outcome")),
),
patch("routstr.wallet.logger.critical") as critical,
):
with pytest.raises(asyncio.CancelledError):
await wallet.periodic_routstr_fee_payout()
complete.assert_not_awaited()
critical.assert_called_once()
-46
View File
@@ -1,46 +0,0 @@
import os
import sqlite3
import subprocess
import sys
from pathlib import Path
def _run_alembic(root: Path, database_url: str, revision: str) -> None:
env = os.environ.copy()
env["DATABASE_URL"] = database_url
subprocess.run(
[sys.executable, "-m", "alembic", "upgrade", revision],
cwd=root,
env=env,
check=True,
capture_output=True,
text=True,
)
def test_fee_payout_checkpoint_migration_preserves_existing_row(
tmp_path: Path,
) -> None:
root = Path(__file__).resolve().parents[2]
database_path = tmp_path / "migration.db"
database_url = f"sqlite+aiosqlite:///{database_path}"
_run_alembic(root, database_url, "c6d7e8f9a0b1")
with sqlite3.connect(database_path) as connection:
result = connection.execute(
"UPDATE routstr_fees SET accumulated_msats = 5000, "
"total_paid_msats = 1000, last_paid_at = 123 WHERE id = 1"
)
assert result.rowcount == 1
connection.commit()
_run_alembic(root, database_url, "head")
with sqlite3.connect(database_path) as connection:
row = connection.execute(
"SELECT accumulated_msats, total_paid_msats, last_paid_at, "
"payout_in_progress_msats, payout_started_at "
"FROM routstr_fees WHERE id = 1"
).fetchone()
assert row == (5000, 1000, 123, 0, None)
+2 -29
View File
@@ -11,9 +11,7 @@ async def _fake_session(): # type: ignore[no-untyped-def]
yield MagicMock()
def _patches( # type: ignore[no-untyped-def]
proof_amount: int = 1000, user_balance_msats: int = 0
):
def _patches(proof_amount: int = 1000): # type: ignore[no-untyped-def]
proof = MagicMock(amount=proof_amount)
return [
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
@@ -27,7 +25,7 @@ def _patches( # type: ignore[no-untyped-def]
),
patch(
"routstr.wallet.db.balances_for_mint_and_unit",
AsyncMock(return_value=user_balance_msats),
AsyncMock(return_value=0),
),
patch("routstr.wallet.db.create_session", _fake_session),
]
@@ -54,31 +52,6 @@ async def test_fetch_all_balances_falls_back_to_primary_mint() -> None:
assert total_wallet == 1000
@pytest.mark.asyncio
async def test_fetch_all_balances_reports_liability_when_wallet_is_empty() -> None:
"""An empty wallet must not hide outstanding user liabilities."""
from routstr.core.settings import settings
with patch.object(settings, "cashu_mints", []), patch.object(
settings, "primary_mint", "http://primary:3338"
):
for p in _patches(proof_amount=0, user_balance_msats=5000):
p.start()
try:
details, total_wallet, total_user, owner = await fetch_all_balances(
units=["sat"]
)
finally:
patch.stopall()
assert details[0]["wallet_balance"] == 0
assert details[0]["user_balance"] == 5
assert details[0]["owner_balance"] == -5
assert total_wallet == 0
assert total_user == 5
assert owner == -5
@pytest.mark.asyncio
async def test_fetch_all_balances_no_duplicate_primary_mint() -> None:
"""primary_mint already in cashu_mints is not inspected twice."""
-74
View File
@@ -1,74 +0,0 @@
"""raw_send_to_lnurl() must not hang forever on an unresponsive mint.
The Cashu library issues POST /v1/melt/bolt11 with timeout=None, so a hung
mint would block the melt (and the payout loop) indefinitely. raw_send_to_lnurl
now wraps wallet.melt() in asyncio.wait_for(MELT_TIMEOUT_SECONDS) and surfaces a
timeout as LNURLError instead of hanging.
"""
import asyncio
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from routstr.payment import lnurl
from routstr.payment.lnurl import LNURLError, raw_send_to_lnurl
@pytest.mark.asyncio
async def test_raw_send_to_lnurl_times_out_on_hung_melt() -> None:
proofs = [MagicMock(amount=1000)]
wallet = MagicMock()
wallet.melt_quote = AsyncMock(return_value=MagicMock(fee_reserve=1, quote="q"))
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
async def _hang(**kwargs: object) -> None:
await asyncio.sleep(5) # far longer than the patched timeout
wallet.melt = AsyncMock(side_effect=_hang)
lnurl_data = {
"callback_url": "https://ln.tld/cb",
"min_sendable": 1_000,
"max_sendable": 100_000_000,
}
with patch.object(lnurl, "MELT_TIMEOUT_SECONDS", 0.05), patch(
"routstr.payment.lnurl.get_lnurl_data", AsyncMock(return_value=lnurl_data)
), patch(
"routstr.payment.lnurl.get_lnurl_invoice",
AsyncMock(return_value=("lnbc1...", {})),
):
with pytest.raises(LNURLError, match="Melt timed out"):
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
@pytest.mark.asyncio
async def test_raw_send_to_lnurl_succeeds_within_timeout() -> None:
"""A prompt melt still returns the net amount, unaffected by the guard."""
proofs = [MagicMock(amount=1000)]
wallet = MagicMock()
wallet.melt_quote = AsyncMock(return_value=MagicMock(fee_reserve=1, quote="q"))
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
wallet.melt = AsyncMock(return_value=MagicMock())
lnurl_data = {
"callback_url": "https://ln.tld/cb",
"min_sendable": 1_000,
"max_sendable": 100_000_000,
}
with patch.object(lnurl, "MELT_TIMEOUT_SECONDS", 5), patch(
"routstr.payment.lnurl.get_lnurl_data", AsyncMock(return_value=lnurl_data)
), patch(
"routstr.payment.lnurl.get_lnurl_invoice",
AsyncMock(return_value=("lnbc1...", {})),
):
paid = await raw_send_to_lnurl(
wallet, proofs, "owner@ln.tld", "sat", amount=1000
)
assert paid > 0
wallet.melt.assert_awaited_once()
+2 -352
View File
@@ -11,19 +11,16 @@ import os
from typing import Any, AsyncIterator
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi.responses import Response, StreamingResponse
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
from routstr.auth import ReservationSnapshot # noqa: E402
from routstr.core.db import ApiKey # noqa: E402
from routstr.payment.cost_calculation import CostData # noqa: E402
from routstr.payment.models import Architecture, Model, Pricing # noqa: E402
from routstr.upstream.base import BaseUpstreamProvider # noqa: E402
from routstr.wallet import MintConnectionError, TokenConsumedError # noqa: E402
# ---------------------------------------------------------------------------
# Fixtures
@@ -499,27 +496,14 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None:
yield {"type": "message_stop"}
fake_cost = {"total_msats": 4321, "total_usd": 0.00015}
reservation = ReservationSnapshot(
release_id="messages-stream",
key_hash=key.hashed_key,
billing_key_hash=key.hashed_key,
reserved_msats=10_000,
)
captured_cost_call: dict[str, Any] = {}
async def fake_adjust(
fresh_key: Any,
combined_data: Any,
sess: Any,
max_cost: int,
model_obj: Any = None,
provider_fee: Any = None,
reservation_snapshot: Any = None,
fresh_key: Any, combined_data: Any, sess: Any, max_cost: int, usage: Any = None
) -> dict:
captured_cost_call["combined_data"] = combined_data
captured_cost_call["max_cost"] = max_cost
captured_cost_call["reservation_snapshot"] = reservation_snapshot
return fake_cost
fake_session = MagicMock()
@@ -553,7 +537,6 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None:
session=session,
max_cost_for_model=10_000,
model_obj=model,
reservation_snapshot=reservation,
)
assert isinstance(result, StreamingResponse)
@@ -576,7 +559,6 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None:
assert combined["usage"]["input_tokens"] == 5
assert combined["usage"]["output_tokens"] == 7
assert combined["model"] == "openai/gpt-4o-mini"
assert captured_cost_call["reservation_snapshot"] is reservation
@pytest.mark.asyncio
@@ -604,25 +586,12 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None:
yield b'event: message_stop\ndata: {"type":"message_stop"}\n\n'
fake_cost = {"total_msats": 999, "total_usd": 0.0001}
reservation = ReservationSnapshot(
release_id="messages-byte-stream",
key_hash=key.hashed_key,
billing_key_hash=key.hashed_key,
reserved_msats=10_000,
)
captured: dict[str, Any] = {}
async def fake_adjust(
fresh_key: Any,
combined_data: Any,
sess: Any,
max_cost: int,
model_obj: Any = None,
provider_fee: Any = None,
reservation_snapshot: Any = None,
fresh_key: Any, combined_data: Any, sess: Any, max_cost: int, usage: Any = None
) -> dict:
captured["combined_data"] = combined_data
captured["reservation_snapshot"] = reservation_snapshot
return fake_cost
fake_session = MagicMock()
@@ -655,7 +624,6 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None:
session=session,
max_cost_for_model=10_000,
model_obj=model,
reservation_snapshot=reservation,
)
assert isinstance(result, StreamingResponse)
@@ -679,7 +647,6 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None:
assert combined["usage"]["input_tokens"] == 3
assert combined["usage"]["output_tokens"] == 4
assert combined["model"] == "openai/gpt-4o-mini"
assert captured["reservation_snapshot"] is reservation
# ---------------------------------------------------------------------------
@@ -1341,320 +1308,3 @@ async def test_dispatch_uses_url_detected_prefix_for_fireworks_custom_row() -> N
"fireworks_ai/accounts/fireworks/models/glm-5"
)
assert captured_kwargs["api_base"] == "https://api.fireworks.ai/inference/v1"
# ---------------------------------------------------------------------------
# X-Cashu redemption error taxonomy (unreachable mint + string codes)
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
@pytest.mark.parametrize(
"handler_name",
["handle_x_cashu", "handle_x_cashu_responses"],
)
@pytest.mark.parametrize(
"error",
[
httpx.ConnectError("All connection attempts failed"),
MintConnectionError("Cashu mint is unreachable"),
TimeoutError("timed out connecting to mint"),
],
)
async def test_x_cashu_mint_unreachable_returns_503(
handler_name: str, error: Exception
) -> None:
"""Both X-Cashu entrypoints classify a down mint as 503 mint_unreachable,
not a generic 400 cashu_error."""
provider = _make_provider()
model = _make_model()
request = _make_request()
with patch(
"routstr.upstream.base.recieve_token", new=AsyncMock(side_effect=error)
):
handler = getattr(provider, handler_name)
response = await handler(
request=request,
x_cashu_token="cashuAtoken",
path="v1/chat/completions",
max_cost_for_model=10_000,
model_obj=model,
)
assert response.status_code == 503
body = json.loads(bytes(response.body))
assert body["error"]["type"] == "mint_unreachable"
assert body["error"]["message"] == "Cashu mint is unreachable"
assert body["error"]["code"] == "cashu_mint_unreachable"
if str(error) != body["error"]["message"]:
assert str(error) not in body["error"]["message"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"handler_name",
["handle_x_cashu", "handle_x_cashu_responses"],
)
@pytest.mark.parametrize(
(
"error",
"expected_status",
"expected_type",
"expected_message",
"expected_code",
),
[
(
ValueError("Mint Error: Token already spent"),
400,
"token_already_spent",
"Cashu token already spent",
"cashu_token_already_spent",
),
(
ValueError("invalid token: could not decode"),
400,
"invalid_token",
"Invalid Cashu token",
"invalid_cashu_token",
),
(
# Fee/swap failures now map to a granular 422 on the X-Cashu path,
# matching the bearer path (previously flattened to 400).
ValueError(
"Failed to estimate fees: Fees (7 sat) exceed token amount (5 sat)"
),
422,
"mint_error",
"Token value is too small to cover swap fees",
"cashu_token_swap_fees_exceed_amount",
),
(
ValueError("Failed to melt token from foreign mint http://m: boom"),
422,
"mint_error",
"Failed to swap token from foreign mint",
"cashu_foreign_mint_swap_failed",
),
(
ValueError("some unexpected wallet condition"),
400,
"cashu_error",
"Failed to redeem Cashu token",
"cashu_token_redemption_failed",
),
(
# Non-ValueError faults are internal errors (500), not token errors.
RuntimeError("db exploded"),
500,
"api_error",
"Internal error during token redemption",
"internal_error",
),
],
)
async def test_x_cashu_error_code_is_stable_string(
handler_name: str,
error: Exception,
expected_status: int,
expected_type: str,
expected_message: str,
expected_code: str,
) -> None:
"""X-Cashu emits a stable string ``code`` on every branch, matching the
bearer path's taxonomy instead of an int HTTP status."""
provider = _make_provider()
model = _make_model()
request = _make_request()
with patch(
"routstr.upstream.base.recieve_token", new=AsyncMock(side_effect=error)
):
handler = getattr(provider, handler_name)
response = await handler(
request=request,
x_cashu_token="cashuAtoken",
path="v1/chat/completions",
max_cost_for_model=10_000,
model_obj=model,
)
assert response.status_code == expected_status
body = json.loads(bytes(response.body))
assert body["error"]["type"] == expected_type
assert body["error"]["message"] == expected_message
assert body["error"]["code"] == expected_code
assert isinstance(body["error"]["code"], str)
if str(error) != expected_message:
assert str(error) not in body["error"]["message"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
("handler_name", "forward_attr"),
[
("handle_x_cashu", "forward_x_cashu_request"),
("handle_x_cashu_responses", "forward_x_cashu_responses_request"),
],
)
async def test_x_cashu_transport_error_after_redemption_is_not_retryable(
handler_name: str, forward_attr: str
) -> None:
"""A transport failure while forwarding (after the token is spent) maps to
502 upstream_error, never a retryable cashu_mint_unreachable."""
provider = _make_provider()
model = _make_model()
request = _make_request()
with (
patch(
"routstr.upstream.base.recieve_token",
new=AsyncMock(return_value=(5_000, "sat", "https://mint")),
),
patch("routstr.upstream.base.store_cashu_transaction", new=AsyncMock()),
patch.object(
provider,
forward_attr,
new=AsyncMock(side_effect=httpx.ConnectError("upstream down")),
),
):
handler = getattr(provider, handler_name)
response = await handler(
request=request,
x_cashu_token="cashuAtoken",
path="v1/chat/completions",
max_cost_for_model=10_000,
model_obj=model,
)
assert response.status_code == 502
body = json.loads(bytes(response.body))
assert body["error"]["type"] == "upstream_error"
assert body["error"]["code"] != "cashu_mint_unreachable"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"handler_name",
["handle_x_cashu", "handle_x_cashu_responses"],
)
async def test_x_cashu_token_consumed_returns_500_and_no_echo(
handler_name: str,
) -> None:
"""A post-redemption failure (token spent, crediting/minting failed) is a
non-retryable 500 token_consumed and must NOT echo the spent token back."""
provider = _make_provider()
model = _make_model()
request = _make_request()
with patch(
"routstr.upstream.base.recieve_token",
new=AsyncMock(side_effect=TokenConsumedError("credit failed")),
):
handler = getattr(provider, handler_name)
response = await handler(
request=request,
x_cashu_token="cashuAtoken",
path="v1/chat/completions",
max_cost_for_model=10_000,
model_obj=model,
)
assert response.status_code == 500
body = json.loads(bytes(response.body))
assert body["error"]["type"] == "token_consumed"
assert body["error"]["code"] == "cashu_token_consumed"
assert "X-Cashu" not in response.headers
@pytest.mark.asyncio
@pytest.mark.parametrize(
("handler_name", "error", "echoed"),
[
# Spent token: must NOT be echoed.
("handle_x_cashu", ValueError("Token already spent"), False),
("handle_x_cashu_responses", ValueError("Token already spent"), False),
# Unspent but unreachable mint: echo so the client can retry the token.
("handle_x_cashu", MintConnectionError("mint down"), True),
("handle_x_cashu_responses", MintConnectionError("mint down"), True),
# Consumed token (post-redemption): must NOT be echoed.
("handle_x_cashu", TokenConsumedError("credit failed"), False),
("handle_x_cashu_responses", TokenConsumedError("credit failed"), False),
],
)
async def test_x_cashu_echoes_token_only_when_recoverable(
handler_name: str, error: Exception, echoed: bool
) -> None:
provider = _make_provider()
model = _make_model()
request = _make_request()
with patch(
"routstr.upstream.base.recieve_token", new=AsyncMock(side_effect=error)
):
handler = getattr(provider, handler_name)
response = await handler(
request=request,
x_cashu_token="cashuAtoken",
path="v1/chat/completions",
max_cost_for_model=10_000,
model_obj=model,
)
if echoed:
assert response.headers.get("X-Cashu") == "cashuAtoken"
else:
assert "X-Cashu" not in response.headers
@pytest.mark.asyncio
@pytest.mark.parametrize(
("handler_name", "forward_attr"),
[
("handle_x_cashu", "forward_x_cashu_request"),
("handle_x_cashu_responses", "forward_x_cashu_responses_request"),
],
)
@pytest.mark.parametrize("amount", [0, -5])
async def test_x_cashu_zero_value_rejected_not_forwarded(
handler_name: str, forward_attr: str, amount: int
) -> None:
"""A token that redeems to <= 0 must be rejected as cashu_token_zero_value
(400) and NEVER forwarded as a free request the X-Cashu path lacked the
guard that credit_balance has."""
provider = _make_provider()
model = _make_model()
request = _make_request()
with (
patch(
"routstr.upstream.base.recieve_token",
new=AsyncMock(return_value=(amount, "sat", "https://mint")),
),
patch("routstr.upstream.base.store_cashu_transaction", new=AsyncMock()),
patch.object(
provider,
forward_attr,
new=AsyncMock(side_effect=AssertionError("must not forward a zero-value token")),
),
):
handler = getattr(provider, handler_name)
response = await handler(
request=request,
x_cashu_token="cashuAtoken",
path="v1/chat/completions",
max_cost_for_model=10_000,
model_obj=model,
)
assert response.status_code == 400
body = json.loads(bytes(response.body))
assert body["error"]["type"] == "cashu_error"
assert (
body["error"]["message"]
== "Failed to redeem Cashu token: token yielded no value"
)
assert body["error"]["code"] == "cashu_token_zero_value"
# Spent-to-zero token must not be echoed back for retry.
assert "X-Cashu" not in response.headers
-154
View File
@@ -1,154 +0,0 @@
"""Tests for periodic_payout() resilience fixes.
Covers two regressions from the auto-payout / primary-mint audit
(docs/auto-payout-primary-mint-failure-report.md):
1. periodic_payout() must include settings.primary_mint even when it is not
listed in settings.cashu_mints, matching fetch_all_balances(); otherwise
primary-mint funds never auto-payout.
2. A failure on one mint/unit must not abort payout for the remaining
mint/units in the same cycle (the try/except is now per mint/unit).
"""
from collections.abc import Callable, Coroutine
from contextlib import asynccontextmanager
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from routstr.wallet import periodic_payout
# Sentinel interval used to break the otherwise-infinite payout loop after
# exactly one full cycle.
_INTERVAL = 987
class _LoopBreak(Exception):
"""Raised via the patched sleep to stop periodic_payout after one cycle."""
@asynccontextmanager
async def _fake_session(): # type: ignore[no-untyped-def]
yield MagicMock()
def _one_cycle_sleep() -> Callable[[float], Coroutine[Any, Any, None]]:
"""Return an async sleep stub that lets exactly one payout cycle run.
The top-of-loop sleep uses the sentinel interval; the second time it is
seen (start of the second cycle) we raise to break out. The inner
``asyncio.sleep(5)`` pass-through is ignored.
"""
seen = {"interval": 0}
async def _sleep(seconds: float) -> None:
if seconds == _INTERVAL:
seen["interval"] += 1
if seen["interval"] >= 2:
raise _LoopBreak()
return _sleep
@pytest.mark.asyncio
async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> None:
"""primary_mint absent from cashu_mints is still paid out."""
from routstr.core.settings import settings
get_wallet = AsyncMock(return_value=MagicMock())
raw_send = AsyncMock(return_value=1000)
with patch.object(settings, "cashu_mints", []), patch.object(
settings, "primary_mint", "http://primary:3338"
), patch.object(settings, "receive_ln_address", "owner@ln.tld"), patch.object(
settings, "payout_interval_seconds", _INTERVAL
), patch.object(settings, "min_payout_sat", 10), patch(
"routstr.wallet.asyncio.sleep", _one_cycle_sleep()
), patch("routstr.wallet.db.create_session", _fake_session), patch(
"routstr.wallet.get_wallet", get_wallet
), patch(
"routstr.wallet.get_proofs_per_mint_and_unit",
MagicMock(return_value=[MagicMock(amount=100_000)]),
), patch(
"routstr.wallet.slow_filter_spend_proofs",
AsyncMock(side_effect=lambda proofs, wallet: proofs),
), patch(
"routstr.wallet.db.balances_for_mint_and_unit", AsyncMock(return_value=0)
), patch("routstr.wallet.raw_send_to_lnurl", raw_send):
with pytest.raises(_LoopBreak):
await periodic_payout()
processed = {call.args[0] for call in get_wallet.await_args_list}
assert processed == {"http://primary:3338"}
assert raw_send.await_count >= 1
@pytest.mark.asyncio
async def test_periodic_payout_isolates_failing_mint() -> None:
"""A failing mint does not prevent payout for the other mints."""
from routstr.core.settings import settings
async def _get_wallet(mint_url: str, unit: str) -> MagicMock:
if mint_url == "http://bad:3338":
raise RuntimeError("mint unreachable")
return MagicMock()
get_wallet = AsyncMock(side_effect=_get_wallet)
raw_send = AsyncMock(return_value=1000)
with patch.object(
settings, "cashu_mints", ["http://bad:3338", "http://good:3338"]
), patch.object(settings, "primary_mint", "http://good:3338"), patch.object(
settings, "receive_ln_address", "owner@ln.tld"
), patch.object(settings, "payout_interval_seconds", _INTERVAL), patch.object(
settings, "min_payout_sat", 10
), patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()), patch(
"routstr.wallet.db.create_session", _fake_session
), patch("routstr.wallet.get_wallet", get_wallet), patch(
"routstr.wallet.get_proofs_per_mint_and_unit",
MagicMock(return_value=[MagicMock(amount=100_000)]),
), patch(
"routstr.wallet.slow_filter_spend_proofs",
AsyncMock(side_effect=lambda proofs, wallet: proofs),
), patch(
"routstr.wallet.db.balances_for_mint_and_unit", AsyncMock(return_value=0)
), patch("routstr.wallet.raw_send_to_lnurl", raw_send):
with pytest.raises(_LoopBreak):
await periodic_payout()
# The bad mint raised on get_wallet for both units, yet the good mint was
# still reached and paid out for both units — failures are isolated.
good_calls = [
c for c in get_wallet.await_args_list if c.args[0] == "http://good:3338"
]
assert len(good_calls) == 2 # sat + msat
assert raw_send.await_count == 2 # good mint paid for both units
@pytest.mark.asyncio
async def test_periodic_payout_handles_session_creation_failure() -> None:
"""A db.create_session failure is logged and the payout loop continues."""
from routstr.core.settings import settings
create_session = MagicMock(side_effect=RuntimeError("db unavailable"))
logger = MagicMock()
with patch.object(settings, "cashu_mints", ["http://mint:3338"]), patch.object(
settings, "primary_mint", "http://mint:3338"
), patch.object(settings, "receive_ln_address", "owner@ln.tld"), patch.object(
settings, "payout_interval_seconds", _INTERVAL
), patch(
"routstr.wallet.asyncio.sleep", _one_cycle_sleep()
), patch(
"routstr.wallet.db.create_session", create_session
), patch("routstr.wallet.logger", logger):
with pytest.raises(_LoopBreak):
await periodic_payout()
create_session.assert_called_once()
logger.error.assert_called_once()
message = logger.error.call_args.args[0]
extra = logger.error.call_args.kwargs["extra"]
assert message == "Error in periodic payout cycle: RuntimeError"
assert extra == {"error": "db unavailable"}
-126
View File
@@ -1,126 +0,0 @@
from contextlib import asynccontextmanager
from typing import AsyncGenerator
from unittest.mock import AsyncMock
import pytest
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlalchemy.pool import StaticPool
from sqlmodel import SQLModel, select
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr import wallet
from routstr.core.db import CashuTransaction
def _make_engine() -> AsyncEngine:
return create_async_engine(
"sqlite+aiosqlite://",
poolclass=StaticPool,
connect_args={"check_same_thread": False},
)
@pytest.fixture
async def refund_db(
monkeypatch: pytest.MonkeyPatch,
) -> AsyncGenerator[AsyncEngine, None]:
engine = _make_engine()
async with engine.begin() as connection:
await connection.run_sync(SQLModel.metadata.create_all)
@asynccontextmanager
async def create_session() -> AsyncGenerator[AsyncSession, None]:
async with AsyncSession(engine, expire_on_commit=False) as session:
yield session
monkeypatch.setattr(wallet.db, "create_session", create_session)
try:
yield engine
finally:
await engine.dispose()
async def _store(
engine: AsyncEngine,
*,
token: str,
created_at: int,
typ: str = "out",
collected: bool = False,
swept: bool = False,
) -> None:
async with AsyncSession(engine, expire_on_commit=False) as session:
session.add(
CashuTransaction(
token=token,
amount=10,
unit="sat",
type=typ,
created_at=created_at,
collected=collected,
swept=swept,
)
)
await session.commit()
async def _transactions(engine: AsyncEngine) -> dict[str, CashuTransaction]:
async with AsyncSession(engine, expire_on_commit=False) as session:
results = await session.exec(select(CashuTransaction))
return {transaction.token: transaction for transaction in results.all()}
@pytest.mark.asyncio
async def test_refund_sweep_only_processes_eligible_outbound_transactions(
refund_db: AsyncEngine, monkeypatch: pytest.MonkeyPatch
) -> None:
cutoff = 1_000
await _store(refund_db, token="eligible", created_at=cutoff - 1)
await _store(refund_db, token="too-new", created_at=cutoff)
await _store(
refund_db, token="collected", created_at=cutoff - 1, collected=True
)
await _store(refund_db, token="swept", created_at=cutoff - 1, swept=True)
await _store(refund_db, token="inbound", created_at=cutoff - 1, typ="in")
receive_token = AsyncMock()
monkeypatch.setattr(wallet, "recieve_token", receive_token)
await wallet._refund_sweep_once(cutoff)
receive_token.assert_awaited_once_with("eligible")
transactions = await _transactions(refund_db)
assert transactions["eligible"].swept is True
assert transactions["too-new"].swept is False
assert transactions["collected"].collected is True
assert transactions["collected"].swept is False
assert transactions["swept"].swept is True
assert transactions["inbound"].swept is False
@pytest.mark.asyncio
async def test_refund_sweep_persists_success_and_isolates_token_failures(
refund_db: AsyncEngine, monkeypatch: pytest.MonkeyPatch
) -> None:
cutoff = 1_000
for token in ("success", "already-spent", "temporary-failure"):
await _store(refund_db, token=token, created_at=cutoff - 1)
async def receive_token(token: str) -> None:
if token == "already-spent":
raise RuntimeError("Token already spent")
if token == "temporary-failure":
raise RuntimeError("mint unavailable")
receive = AsyncMock(side_effect=receive_token)
monkeypatch.setattr(wallet, "recieve_token", receive)
await wallet._refund_sweep_once(cutoff)
assert receive.await_count == 3
transactions = await _transactions(refund_db)
assert transactions["success"].swept is True
assert transactions["success"].collected is False
assert transactions["already-spent"].collected is True
assert transactions["already-spent"].swept is False
assert transactions["temporary-failure"].collected is False
assert transactions["temporary-failure"].swept is False

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