Merge branch 'main' into feat/upstream-certification-harness

This commit is contained in:
9qeklajc
2026-10-02 19:41:33 +02:00
124 changed files with 14138 additions and 1634 deletions
+19
View File
@@ -45,6 +45,7 @@ ROUTSTR_SECRET_KEY=
# ONION_URL=http://mynode.onion (auto fetched from compose) # 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" # RELAYS="wss://relay.damus.io,wss://relay.nostr.band,wss://eden.nostr.land,wss://relay.routstr.com"
# ENABLE_ANALYTICS_SHARING=true # ENABLE_ANALYTICS_SHARING=true
# CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org" # this is the default
# CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org,https://ecashmint.otrta.me" # CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org,https://ecashmint.otrta.me"
# MINT_OPERATION_CONCURRENCY=4 # MINT_OPERATION_CONCURRENCY=4
# MINT_OPERATION_TIMEOUT_SECONDS=30 # MINT_OPERATION_TIMEOUT_SECONDS=30
@@ -64,6 +65,24 @@ ROUTSTR_SECRET_KEY=
# Network Configuration # Network Configuration
# CORS_ORIGINS=* # CORS_ORIGINS=*
# TOR_PROXY_URL=socks5://127.0.0.1:9050 # TOR_PROXY_URL=socks5://127.0.0.1:9050
# PROXY_EXTRA_ALLOWED_PATHS=
# Upstream Connection Pools (one pool per upstream origin)
# UPSTREAM_MAX_CONNECTIONS=200
# UPSTREAM_POOL_TIMEOUT=5
# UPSTREAM_READ_TIMEOUT=900
# Upstream Streaming Guards (0 disables; keep above reasoning models' think time)
# UPSTREAM_FIRST_TOKEN_TIMEOUT_SECONDS=0
# UPSTREAM_STREAM_IDLE_TIMEOUT_SECONDS=0
# UPSTREAM_ALLOWED_FAILS=3
# UPSTREAM_COOLDOWN_SECONDS=30
# Request and reservation lifetime limits (seconds)
# STALE_RESERVATION_TIMEOUT_SECONDS=300
# MAX_REQUEST_LIFETIME_SECONDS=1800
# DOWNSTREAM_SEND_TIMEOUT_SECONDS=60
# REQUEST_CLEANUP_TIMEOUT_SECONDS=30
# Logging # Logging
# LOG_LEVEL=INFO # LOG_LEVEL=INFO
+23 -5
View File
@@ -51,9 +51,25 @@ curl https://api.routstr.com/v1/chat/completions \
## Quick Start (Docker) ## Quick Start (Docker)
If you are a node runner, start a Routstr Core instance using Docker Compose: If you are a node runner, the recommended way to start Routstr Core is to clone
the repository at the latest release and run it with Docker Compose:
1. **Prepare your `.env`**: 1. **Clone the latest release**:
```bash
git clone https://github.com/Routstr/routstr-core.git
cd routstr-core
git checkout v0.4.7 # current release — see https://github.com/Routstr/routstr-core/releases/latest
```
Docker Compose builds the node and the admin dashboard from source, so there
is no image to pull.
2. **Prepare your `.env`**:
```bash
cp .env.example .env
```
Then edit it with your details:
```bash ```bash
# Optional: encrypts node secrets at rest. If unset, the node generates a key # 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 # on first start, writes it to routstr_secret.key, and prints it once — back
@@ -76,12 +92,14 @@ If you are a node runner, start a Routstr Core instance using Docker Compose:
uv run python -c "from cryptography.fernet import Fernet; print(Fernet.generate_key().decode())" uv run python -c "from cryptography.fernet import Fernet; print(Fernet.generate_key().decode())"
``` ```
2. **Start the services**: 3. **Start the services**:
```bash ```bash
docker compose up -d docker compose up -d
``` ```
3. **Get your admin password**: The first start builds both images (the dashboard build takes a few minutes).
4. **Get your admin password**:
On first start the node generates an admin password and logs it once with the On first start the node generates an admin password and logs it once with the
`/admin` URL. Read it from the logs: `/admin` URL. Read it from the logs:
```bash ```bash
@@ -89,7 +107,7 @@ If you are a node runner, start a Routstr Core instance using Docker Compose:
``` ```
(Lost it? Reset with `docker compose exec routstr /.venv/bin/python scripts/reset_admin_password.py --regenerate`.) (Lost it? Reset with `docker compose exec routstr /.venv/bin/python scripts/reset_admin_password.py --regenerate`.)
4. **Configure**: 5. **Configure**:
Open [http://localhost:8000/admin/](http://localhost:8000/admin/) to connect your AI providers and set pricing. 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/)**. For full instructions, see the **[Provider Quick Start Guide](https://docs.routstr.com/provider/quickstart/)**.
+3 -1
View File
@@ -264,7 +264,9 @@ Billing is input-token based (output tokens are free on Jev); the response's
- TypeSafe's `GET /v1/models` lists aliases only; the node additionally seeds - TypeSafe's `GET /v1/models` lists aliases only; the node additionally seeds
the known versioned ids so they can be requested directly. the known versioned ids so they can be requested directly.
- TypeSafe answers `429 Too Many Requests` and `529 Overloaded` when throttled. - TypeSafe answers `429 Too Many Requests` and `529 Overloaded` when throttled.
Both are forwarded as upstream errors; retry with exponential backoff. Both are forwarded as upstream errors; retry with exponential backoff. The
`429` keeps its status; the `529` is reported as `424`
(see [Upstream attribution](errors.md#upstream-attribution-424-failed-dependency)).
**Enabling the provider:** **Enabling the provider:**
+59 -9
View File
@@ -51,11 +51,50 @@ legacy status behavior.
| 403 | Forbidden | Access denied to resource | | 403 | Forbidden | Access denied to resource |
| 404 | Not Found | Endpoint or resource doesn't exist | | 404 | Not Found | Endpoint or resource doesn't exist |
| 422 | Unprocessable Entity | Validation errors | | 422 | Unprocessable Entity | Validation errors |
| 429 | Too Many Requests | Rate limit exceeded | | 424 | Failed Dependency | An upstream inference provider failed. This node is healthy — see [Upstream attribution](#upstream-attribution-424-failed-dependency) |
| 500 | Internal Server Error | Server-side error | | 429 | Too Many Requests | Rate limit exceeded (this node or an upstream provider) |
| 502 | Bad Gateway | Upstream API error | | 500 | Internal Server Error | Server-side error on this node |
| 502 | Bad Gateway | Gateway-level failure |
| 503 | Service Unavailable | Temporary outage | | 503 | Service Unavailable | Temporary outage |
### Upstream attribution (424 Failed Dependency)
When an upstream provider fails, this node is still healthy, so the failure is
reported as a **non-5xx** status. Clients should not mark the node down for it.
An upstream-attributable failure answers:
- **Status:** `424`
- **`error.code`:** `UPSTREAM_UNAVAILABLE`
- **Header:** `X-Routstr-Error-Scope: upstream`
- **`error.upstream_status`:** the provider's own status (e.g. `503`). Failures
built by the payment helpers carry it in `error.details.upstream_status`
instead
```http
HTTP/1.1 424 Failed Dependency
X-Routstr-Error-Scope: upstream
Content-Type: application/json
{
"error": {
"type": "upstream_error",
"message": "Service Unavailable",
"code": "UPSTREAM_UNAVAILABLE",
"upstream_status": 503
}
}
```
Two exceptions keep their own status:
- **Rate limits** answer `429` with `error.code = UPSTREAM_RATE_LIMIT`, even
when the provider wrapped them in a 5xx.
- **Provider-side 4xx** (`400`/`401`/`403`/`404`/`422`) passes through unchanged.
Node faults (unreachable mint, database failure, internal exception) still
answer `500` with **no** `X-Routstr-Error-Scope` header.
## Error Types ## Error Types
### Authentication Errors ### Authentication Errors
@@ -329,14 +368,18 @@ Retry-After: 45
### Upstream Errors ### Upstream Errors
#### Model Overloaded #### Upstream Unavailable
A provider returned a 5xx (overloaded, bad gateway, timeout, or a provider-side
outage). This node is healthy and your reservation has been reverted.
```json ```json
{ {
"error": { "error": {
"type": "upstream_error", "type": "upstream_error",
"message": "Model is currently overloaded", "message": "Model is currently overloaded",
"code": "model_overloaded", "code": "UPSTREAM_UNAVAILABLE",
"upstream_status": 503,
"details": { "details": {
"model": "gpt-4", "model": "gpt-4",
"retry_after": 5 "retry_after": 5
@@ -345,8 +388,11 @@ Retry-After: 45
} }
``` ```
**Status:** 503 **Status:** 424
**Resolution:** Retry request after delay **Header:** `X-Routstr-Error-Scope: upstream`
**Resolution:** Retry after a short backoff. If the node is configured with
alternative providers for the model, it already retried them before answering —
try another model or provider path if the failure persists.
#### Upstream Timeout #### Upstream Timeout
@@ -364,7 +410,8 @@ Retry-After: 45
} }
``` ```
**Status:** 504 **Status:** 424
**Header:** `X-Routstr-Error-Scope: upstream`
**Resolution:** Retry with shorter prompt or max_tokens **Resolution:** Retry with shorter prompt or max_tokens
### Content Policy ### Content Policy
@@ -416,7 +463,7 @@ def retry_with_backoff(
# Check if error is retryable # Check if error is retryable
if hasattr(e, 'status_code'): if hasattr(e, 'status_code'):
if e.status_code in [429, 502, 503, 504]: if e.status_code in [424, 429, 502, 503, 504]:
# Calculate delay with jitter # Calculate delay with jitter
delay = min( delay = min(
base_delay * (2 ** attempt) + random.uniform(0, 1), base_delay * (2 ** attempt) + random.uniform(0, 1),
@@ -441,6 +488,9 @@ Group errors for handling:
class ErrorHandler: class ErrorHandler:
# Errors that should be retried # Errors that should be retried
RETRYABLE_ERRORS = { RETRYABLE_ERRORS = {
'UPSTREAM_UNAVAILABLE', # upstream 5xx, reported as HTTP 424
'UPSTREAM_RATE_LIMIT', # HTTP 429
'UPSTREAM_TIMEOUT', # EHBP upstream timeout, reported as HTTP 424
'rate_limit', 'rate_limit',
'upstream_timeout', 'upstream_timeout',
'model_overloaded', 'model_overloaded',
+4 -3
View File
@@ -91,7 +91,7 @@ All errors follow a consistent format:
| `not_found` | 404 | Resource not found | | `not_found` | 404 | Resource not found |
| `rate_limit_exceeded` | 429 | Too many requests | | `rate_limit_exceeded` | 429 | Too many requests |
| `internal_error` | 500 | Server error | | `internal_error` | 500 | Server error |
| `upstream_error` | 502 | Upstream API error | | `upstream_error` | 424 | Upstream API error — the provider failed, this node is healthy. Carries `error.code = UPSTREAM_UNAVAILABLE`, the `X-Routstr-Error-Scope: upstream` header, and the provider's own status in `error.upstream_status`. Rate limits stay `429` + `UPSTREAM_RATE_LIMIT`. See [Error Handling](errors.md#upstream-attribution-424-failed-dependency) |
## Endpoint Categories ## Endpoint Categories
@@ -268,9 +268,10 @@ X-Webhook-Signature: sha256=...
| 402 | Payment required | | 402 | Payment required |
| 403 | Forbidden | | 403 | Forbidden |
| 404 | Not found | | 404 | Not found |
| 424 | Upstream provider failed (`X-Routstr-Error-Scope: upstream`) |
| 429 | Rate limited | | 429 | Rate limited |
| 500 | Server error | | 500 | Server error (no scope header) |
| 502 | Upstream error | | 502 | Gateway failure |
| 503 | Service unavailable | | 503 | Service unavailable |
## CORS Support ## CORS Support
+54 -1
View File
@@ -48,6 +48,29 @@ Connect to your AI provider(s):
| **Upstream URL** | API endpoint (e.g., `https://api.openai.com/v1`) | | **Upstream URL** | API endpoint (e.g., `https://api.openai.com/v1`) |
| **API Key** | Your provider's API key | | **API Key** | Your provider's API key |
### DeepSeek
Choose **DeepSeek** as the provider type and paste an API key from
[platform.deepseek.com](https://platform.deepseek.com/api_keys); the base URL
is fixed to `https://api.deepseek.com`. Setting `DEEPSEEK_API_KEY` seeds the
provider on startup instead.
Models are listed from DeepSeek's own `/models` and priced from a rate table
in `routstr/upstream/deepseek.py`, not from litellm or OpenRouter:
- **Peak rates only.** DeepSeek charges half price off-peak, but the node bills
one flat price per model, so it bills the peak rate. Clients overpay
off-peak; the node never bills below cost. Time-of-day pricing is planned.
- **Unknown models import disabled.** A model DeepSeek lists that the table
does not price shows up disabled in the Admin Dashboard. Enable it with a
manual price, or add it to the table.
- **Cache hits** bill at DeepSeek's cache-hit rate (about 2% of the input
rate on flash, about 3% on pro).
Thinking-mode `reasoning_content` is returned to clients unchanged in
responses, and forwarded unchanged when it appears in conversation history.
DeepSeek requires it on requests that carry `tools` and ignores it otherwise.
### PPQ Auto Top-up ### PPQ Auto Top-up
PPQ providers can automatically purchase more credits when their USD balance PPQ providers can automatically purchase more credits when their USD balance
@@ -139,6 +162,15 @@ Which mints to accept payments from:
| --------- | ------------------------------- | | --------- | ------------------------------- |
| **Mints** | List of trusted Cashu mint URLs | | **Mints** | List of trusted Cashu mint URLs |
A fresh node ships with two mints preconfigured:
- `https://mint.minibits.cash/Bitcoin`
- `https://mint.cubabitcoin.org`
Setting `CASHU_MINTS` (env) or editing the list in the dashboard replaces this
default entirely. An explicitly empty value leaves only the primary mint
trusted.
### Lightning Withdrawals ### Lightning Withdrawals
Automatic profit withdrawal: Automatic profit withdrawal:
@@ -197,13 +229,14 @@ Use environment variables for:
| `NPUB` | Nostr public key (bech32) | — | | `NPUB` | Nostr public key (bech32) | — |
| `NSEC` | Legacy seed for the Nostr private key (otherwise set from the admin UI) | — | | `NSEC` | Legacy seed for the Nostr private key (otherwise set from the admin UI) | — |
| `ENABLE_ANALYTICS_SHARING` | Enable usage analytics sharing to Nostr | `true` | | `ENABLE_ANALYTICS_SHARING` | Enable usage analytics sharing to Nostr | `true` |
| `CASHU_MINTS` | Comma-separated mint URLs | `https://mint.minibits.cash/Bitcoin` | | `CASHU_MINTS` | Comma-separated mint URLs | `https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org` |
| `MINT_OPERATION_CONCURRENCY` | Concurrent mint/unit balance reads | `4` | | `MINT_OPERATION_CONCURRENCY` | Concurrent mint/unit balance reads | `4` |
| `MINT_OPERATION_TIMEOUT_SECONDS` | Per-attempt timeout for mint network calls | `30` | | `MINT_OPERATION_TIMEOUT_SECONDS` | Per-attempt timeout for mint network calls | `30` |
| `MINT_MAX_CONCURRENCY` | Concurrent operations allowed per mint (`0` disables the limit) | `4` | | `MINT_MAX_CONCURRENCY` | Concurrent operations allowed per mint (`0` disables the limit) | `4` |
| `MINT_RETRY_MAX_ATTEMPTS` | Retries after a timeout or HTTP 429 (`0` disables retries) | `3` | | `MINT_RETRY_MAX_ATTEMPTS` | Retries after a timeout or HTTP 429 (`0` disables retries) | `3` |
| `RECEIVE_LN_ADDRESS` | Lightning address for withdrawals | — | | `RECEIVE_LN_ADDRESS` | Lightning address for withdrawals | — |
| `MIN_PAYOUT_SAT` | Min payout balance in sats (applies to all mints) | `210` | | `MIN_PAYOUT_SAT` | Min payout balance in sats (applies to all mints) | `210` |
| `MAX_PAYOUT_SAT` | Maximum gross budget per periodic payout in sats, including fees (all mints) | `250000` |
| `PAYOUT_INTERVAL_SECONDS` | Payout loop interval (seconds) | `900` | | `PAYOUT_INTERVAL_SECONDS` | Payout loop interval (seconds) | `900` |
| `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` | | `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` |
| `CORS_ORIGINS` | Allowed CORS origins | `*` | | `CORS_ORIGINS` | Allowed CORS origins | `*` |
@@ -216,6 +249,26 @@ Routstr's wallet mutation lock fail fast during that cooldown instead of waiting
while blocking every other wallet mutation. Callers receive an error and may retry while blocking every other wallet mutation. Callers receive an error and may retry
later; the current response does not include the cooldown duration. later; the current response does not include the cooldown duration.
Read-only `/v1/checkstate` requests start at the SDK request-model limit
(currently 1,000 proofs) and adapt downward on HTTP 413 or 500, down to one
proof. A 500 is a size hypothesis, not a confirmed limit. Successful reduced
sizes are cached per mint within each worker for 24 hours (and refreshed while
in use). HTTP 429 never reduces the batch
size. Scan deadline expiry opens a transport cooldown without shortening any
existing rate-limit cooldown. Invalid, incomplete, or failed scans do
not produce a partial spendable balance. Only explicit UNSPENT proofs qualify;
PENDING proofs are retained but excluded from payouts.
Each scan is bounded by a fixed 60-second deadline and a 128-request budget;
exhausting either aborts that scan safely. Automatic splitting applies only to
state checks, **not swaps or melts**. Their limits are independent, and ambiguous
mutation outcomes must be reconciled rather than retried with different inputs.
Periodic payouts reload local proofs without forcing a keyset refresh, skip
state checks at/below `MIN_PAYOUT_SAT`, and cap each gross payout budget at
`MAX_PAYOUT_SAT`. Oversized inputs receive enough change outputs to return the
excess; they are not automatically swapped. The cap is not a proof-count limit
or a guarantee of Lightning payment success.
### Priority ### Priority
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. 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.
+3 -1
View File
@@ -132,7 +132,9 @@ Connect to your AI provider:
### Cashu Mints ### Cashu Mints
Manage which mints you accept payments from: Manage which mints you accept payments from. A fresh node comes with two
default mints preconfigured (`https://mint.minibits.cash/Bitcoin` and
`https://mint.cubabitcoin.org`):
- **Add Mint** — Enter a mint URL - **Add Mint** — Enter a mint URL
- **Remove Mint** — Stop accepting from a mint - **Remove Mint** — Stop accepting from a mint
+79 -149
View File
@@ -2,179 +2,97 @@
Production deployment guide for Routstr Provider nodes. Production deployment guide for Routstr Provider nodes.
## All-in-One Docker Image (Preferred) ## Quick Start (Recommended)
The easiest way to deploy Routstr is using the all-in-one Docker image from Docker Hub, which includes both the FastAPI backend and the Next.js admin dashboard in a single container. The recommended way to run a provider node is to clone the repository at the
**latest release** and start the stack with Docker Compose. Compose builds both
### Quick Start the node and the admin dashboard from source, so there is no image to pull and no
dashboard build to keep in sync with the node.
```bash ```bash
docker run -d \ git clone https://github.com/Routstr/routstr-core.git
--name routstr \ cd routstr-core
-p 8000:8000 \
-v routstr-data:/app/data \
-e DATABASE_URL="sqlite:////app/data/routstr.db" \
9qeklajc/routstr:latest
```
Access your node: # Check out a release (v0.4.7 is current — see the releases page for the newest tag)
- **API & Admin Dashboard**: http://localhost:8000 git checkout v0.4.7
### Docker Compose Setup # Compose reads its configuration from .env
cp .env.example .env
Create `docker-compose.yml`:
```yaml
version: '3.8'
services:
routstr:
image: 9qeklajc/routstr:latest
container_name: routstr
restart: unless-stopped
ports:
- "8000:8000"
volumes:
- routstr-data:/app/data
environment:
DATABASE_URL: "sqlite:////app/data/routstr.db"
LOG_LEVEL: "info"
volumes:
routstr-data:
```
Start it:
```bash
docker compose up -d docker compose up -d
``` ```
--- Then open your node:
- **API & Admin Dashboard**: <http://localhost:8000>
## Docker Compose (Recommended) - **Admin login**: the password is generated and logged once on first start
For production, use Docker Compose with persistent storage and optional Tor support.
Use the included `compose.yml` for a flexible setup that handles both the UI and the node execution. This is useful for development or when you want to manage Tor as a separate service.
```bash ```bash
docker compose up -d docker compose logs routstr | grep -i admin
``` ```
This will: !!! note "The first start takes a few minutes"
1. **Build the UI**: Compiles the frontend and copies it to a shared volume. `docker compose up` builds both images locally, and the Next.js dashboard
2. **Start Routstr**: Runs the Python node, mounting the built UI. build is the slow part. Later starts reuse the built images.
3. **Start Tor**: Provides anonymous access via a `.onion` address.
!!! tip "Always tracking the newest release"
To check out whatever `releases/latest` currently points at, use:
```bash
git clone https://github.com/Routstr/routstr-core.git
cd routstr-core
git checkout "$(curl -sSL -o /dev/null -w '%{url_effective}' \
https://github.com/Routstr/routstr-core/releases/latest | sed 's|.*/tag/||')"
```
Omitting the `git checkout` entirely leaves you on `main` — newer, but not a
tested release.
--- ---
## With Tor (Anonymous Access) ## What Docker Compose Starts
Add Tor to serve your node as a hidden service—no port forwarding needed. `compose.yml` brings up three services:
```yaml 1. **ui** — builds the Next.js admin dashboard and copies the result into the
services: shared `./ui_out` volume.
routstr: 2. **routstr** — the Python node, serving the API and the dashboard built above.
image: ghcr.io/routstr/proxy:latest 3. **tor** — serves the node as a `.onion` hidden service, so no port forwarding
container_name: routstr is needed. See [Tor Support](tor.md) for how to read your `.onion` address.
restart: unless-stopped
ports:
- "8000:8000"
volumes:
- ./data:/app/data
- ./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
tor:
image: ghcr.io/hundehausen/tor-hidden-service:latest
container_name: tor
restart: unless-stopped
volumes:
- ./tor-data:/var/lib/tor
environment:
- HS_ROUTER=routstr:8000:80
```
After starting, find your `.onion` address:
```bash
docker exec tor cat /var/lib/tor/hidden_service/hostname
```
See [Tor Support](tor.md) for details.
--- ---
## Pre-Configuration (Optional) ## Pre-Configuration (Optional)
While everything can be configured via the dashboard, you can pre-configure settings with environment variables for automated deployments. Everything can be configured from the dashboard after first start, but you can
pre-configure a deployment by editing the `.env` file you created above:
### Using Environment Variables
```yaml
services:
routstr:
image: ghcr.io/routstr/proxy:latest
environment:
# Pre-configure upstream (optional)
- 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.
# Node identity
- NAME=My Provider Node
- DESCRIPTION=Fast GPT-4 access via Lightning
# 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
```
### Using an .env File
```yaml
services:
routstr:
image: ghcr.io/routstr/proxy:latest
env_file:
- .env
volumes:
- ./data:/app/data
```
Example `.env`:
```bash ```bash
# Upstream (optional — can also be set from the dashboard)
UPSTREAM_BASE_URL=https://api.openai.com/v1 UPSTREAM_BASE_URL=https://api.openai.com/v1
UPSTREAM_API_KEY=sk-proj-... 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 # 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 # your database (on the same volume) and its file is named once for backup. Set
# it explicitly to manage the key yourself. # it explicitly to manage the key yourself.
ROUTSTR_SECRET_KEY= ROUTSTR_SECRET_KEY=
# Node identity
NAME=My Provider Node NAME=My Provider Node
DESCRIPTION=Fast GPT-4 access via Lightning
# Lightning withdrawals
RECEIVE_LN_ADDRESS=me@walletofsatoshi.com RECEIVE_LN_ADDRESS=me@walletofsatoshi.com
``` ```
The admin password is generated and logged once on first start; set
`ADMIN_PASSWORD` only as a legacy seed for an existing deployment.
!!! note "Secret key persistence" !!! note "Secret key persistence"
If you leave `ROUTSTR_SECRET_KEY` unset, the node generates one and stores it 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 as `routstr_secret.key` **next to your database**, so it persists alongside
volume as your data — just include that volume in your backups. For stronger your data — just include that in your backups. For stronger isolation
isolation (keeping the key off the data volume), set `ROUTSTR_SECRET_KEY` from (keeping the key off the data volume), set `ROUTSTR_SECRET_KEY` from a
a secrets manager instead. secrets manager instead.
See [Configuration](configuration.md) for all available options. See [Configuration](configuration.md) for all available options.
@@ -182,17 +100,21 @@ See [Configuration](configuration.md) for all available options.
## Persistence ## Persistence
Point `DATABASE_URL` inside `/app/data` (as the examples above do) so everything With the default `compose.yml` the repository directory is mounted into the
Routstr persists lands on the mounted volume: container, so everything Routstr persists stays in the directory you cloned:
| Path | Contents | | Path | Contents |
|------|----------| |------|----------|
| `routstr.db` | SQLite database (settings, API keys, sessions) | | `keys.db` | SQLite database (settings, API keys, sessions) |
| `routstr_secret.key` | Auto-generated master key, written beside the database when `ROUTSTR_SECRET_KEY` is unset | | `routstr_secret.key` | Auto-generated master key, written beside the database when `ROUTSTR_SECRET_KEY` is unset |
| `.wallet/` | Cashu wallet data (your Bitcoin!) | | `.wallet/` | Cashu wallet data (your Bitcoin!) |
| `logs/` | Node logs |
!!! warning "Back Up Your Data" !!! warning "Back Up Your Data"
The `./data` volume contains your wallet. Losing it means losing funds. Back up regularly. Your cloned directory holds your wallet and your master key. Losing it means
losing funds. Back it up regularly — and don't delete the checkout to
"start fresh" without copying `keys.db`, `routstr_secret.key` and `.wallet/`
first.
--- ---
@@ -233,27 +155,35 @@ server {
## Updates ## Updates
Pull the latest image and restart: Check out the new release and rebuild:
```bash ```bash
docker compose pull git fetch --tags
docker compose up -d git checkout v0.4.7 # or the tag you are moving to
docker compose up -d --build
``` ```
`--build` is required: Compose reuses an existing image for a service unless you
ask it to rebuild.
!!! warning "Back up first"
Copy `keys.db`, `routstr_secret.key` and `.wallet/` before updating, and read
the release notes for the version you are moving to.
--- ---
## Building from Source ## Building Without Starting
### Using Docker Compose `docker compose up -d` already builds from source. To build the images
The easiest way to build everything from source: explicitly without starting them:
```bash ```bash
docker compose build docker compose build
``` ```
### Individual Components To build only the node image (the dashboard must already be built into
If you prefer building the node only (requires manual UI build first): `./ui_out`):
```bash ```bash
docker build -t routstr-node . docker build -t routstr-node .
``` ```
@@ -0,0 +1,42 @@
"""Immutable reservation start and absolute recovery deadline.
Revision ID: a73d19b6c204
Revises: e4c7a1b9d520
"""
import time
import sqlalchemy as sa
from alembic import op
revision = "a73d19b6c204"
down_revision = "e4c7a1b9d520"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"reservation_releases", sa.Column("started_at", sa.Integer(), nullable=True)
)
op.add_column(
"reservation_releases", sa.Column("expires_at", sa.Integer(), nullable=True)
)
op.create_index(
"ix_reservation_releases_expires_at", "reservation_releases", ["expires_at"]
)
# Original ages are unknowable for renewed legacy rows. Give them a finite
# migration grace period; deploy only after draining old workers.
op.execute(
sa.text(
"UPDATE reservation_releases SET expires_at = :expiry WHERE status = 'active'"
).bindparams(expiry=int(time.time()) + 1830)
)
def downgrade() -> None:
op.drop_index(
"ix_reservation_releases_expires_at", table_name="reservation_releases"
)
op.drop_column("reservation_releases", "expires_at")
op.drop_column("reservation_releases", "started_at")
@@ -0,0 +1,25 @@
"""Add direction to lightning_invoices
Revision ID: e4c7a1b9d520
Revises: 3a0fbd387f10
Create Date: 2026-09-20 00:00:00.000000
"""
import sqlalchemy as sa
from alembic import op
revision = "e4c7a1b9d520"
down_revision = "3a0fbd387f10"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"lightning_invoices",
sa.Column("direction", sa.String(), nullable=False, server_default="in"),
)
def downgrade() -> None:
op.drop_column("lightning_invoices", "direction")
+9 -2
View File
@@ -21,7 +21,9 @@ dependencies = [
"mdurl==0.1.2", "mdurl==0.1.2",
"pillow>=10", "pillow>=10",
"openai>=1.98.0", "openai>=1.98.0",
"litellm>=1.93.0,<1.94", # 1.93 is the first line supporting Python 3.14 "litellm>=1.101.2,<1.102",
"backoff>=2.2", # litellm's native Anthropic-messages streaming (e.g. deepseek/) imports litellm.proxy, which needs it
"orjson>=3.10",
] ]
[dependency-groups] [dependency-groups]
@@ -72,10 +74,12 @@ build-backend = "setuptools.build_meta"
[tool.setuptools] [tool.setuptools]
packages = ["routstr"] packages = ["routstr"]
[tool.ruff]
extend-exclude = ["examples"]
[tool.ruff.lint] [tool.ruff.lint]
select = ["E", "F", "I"] select = ["E", "F", "I"]
ignore = ["E501"] ignore = ["E501"]
exclude = ["examples"]
[tool.mypy] [tool.mypy]
python_version = "3.11" python_version = "3.11"
@@ -111,6 +115,9 @@ override-dependencies = [
# Transitive deps whose dependents allow the patched version but don't require # Transitive deps whose dependents allow the patched version but don't require
# it. Constraints raise the floor without bypassing any upstream pin. # it. Constraints raise the floor without bypassing any upstream pin.
constraint-dependencies = [ constraint-dependencies = [
"anyio>=4.14.2",
"pyjwt>=2.15.0",
"urllib3>=2.8.0",
"starlette>=1.3.1", "starlette>=1.3.1",
"httpcore>=1.0.9", # 1.0.8 caps h11<0.15 "httpcore>=1.0.9", # 1.0.8 caps h11<0.15
# 1.76 is the first grpcio-tools release with CPython 3.14 wheels. # 1.76 is the first grpcio-tools release with CPython 3.14 wheels.
+124 -8
View File
@@ -49,9 +49,7 @@ payments_logger = get_logger("routstr.payments")
# Routstr platform fee constants # Routstr platform fee constants
ROUTSTR_FEE_PERCENT: float = 2.1 ROUTSTR_FEE_PERCENT: float = 2.1
ROUTSTR_LN_ADDRESS: str = ( ROUTSTR_LN_ADDRESS: str = "routstr-fees@rizful.com"
"npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash"
)
ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS: int = 900 ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS: int = 900
ROUTSTR_FEE_DEFAULT_PAYOUT: int = 200 ROUTSTR_FEE_DEFAULT_PAYOUT: int = 200
@@ -552,8 +550,8 @@ async def _validate_bearer_key_locked(
async def pay_for_request( async def pay_for_request(
key: ApiKey, cost_per_request: int, session: AsyncSession key: ApiKey, cost_per_request: int, session: AsyncSession
) -> int: ) -> ReservationSnapshot:
"""Process payment for a request.""" """Reserve funds and return the durable identity for this request."""
# Ensure cost_per_request is at least the minimum allowed request cost # Ensure cost_per_request is at least the minimum allowed request cost
cost_per_request = max(cost_per_request, settings.min_request_msat) cost_per_request = max(cost_per_request, settings.min_request_msat)
@@ -637,6 +635,14 @@ async def pay_for_request(
) )
# Charge the base cost for the request atomically to avoid race conditions # Charge the base cost for the request atomically to avoid race conditions
from .core.lifecycle import request_lifetime
lifetime = request_lifetime.get()
remaining_lifetime = (
max(0, lifetime.deadline - asyncio.get_running_loop().time())
if lifetime is not None
else settings.max_request_lifetime_seconds
)
reserved_at_now = int(time.time()) reserved_at_now = int(time.time())
stmt = ( stmt = (
update(ApiKey) update(ApiKey)
@@ -686,6 +692,13 @@ async def pay_for_request(
billing_key_hash=reservation.billing_key_hash, billing_key_hash=reservation.billing_key_hash,
reserved_msats=reservation.reserved_msats, reserved_msats=reservation.reserved_msats,
status="active", status="active",
started_at=reserved_at_now,
# reserved_at_now floors to the second; add 1s margin so a
# finalizer finishing right at the nominal deadline isn't fenced
# out by truncation.
expires_at=reserved_at_now
+ math.ceil(remaining_lifetime + settings.request_cleanup_timeout_seconds)
+ 1,
) )
) )
# Publish the identity before commit. If the commit succeeds but its # Publish the identity before commit. If the commit succeeds but its
@@ -738,6 +751,53 @@ async def pay_for_request(
extra={"reservation_id": reservation.release_id}, extra={"reservation_id": reservation.release_id},
) )
try:
# Identity checks only: this call just committed the reservation, so the
# stale-reservation sweeper may legitimately have released it already.
# Release is a terminal state that settlement handles; it is not a
# mismatch between the record and the request.
await _validate_reservation_snapshot(
key, reservation, session, require_active=False
)
except BaseException:
released = False
try:
released = await _transition_reservation_to_released(
reservation,
session,
decrement_requests=True,
idempotent_success=True,
)
except BaseException:
try:
await session.rollback()
except BaseException:
pass
if not released:
try:
async with create_session() as cleanup_session:
released = await _transition_reservation_to_released(
reservation,
cleanup_session,
decrement_requests=True,
idempotent_success=True,
)
except BaseException:
logger.exception(
"Failed to release invalid billing reservation",
extra={"reservation_id": reservation.release_id},
)
if not released:
logger.error(
"Invalid billing reservation could not be released",
extra={"reservation_id": reservation.release_id},
)
await _stop_reservation_heartbeat(reservation.release_id)
_clear_current_reservation(reservation)
raise
logger.info( logger.info(
"Payment processed successfully", "Payment processed successfully",
extra={ extra={
@@ -762,7 +822,7 @@ async def pay_for_request(
}, },
) )
return cost_per_request return reservation
async def revert_pay_for_request( async def revert_pay_for_request(
@@ -828,6 +888,10 @@ async def renew_reservation(
update(ReservationRelease) update(ReservationRelease)
.where(col(ReservationRelease.id) == snapshot.release_id) .where(col(ReservationRelease.id) == snapshot.release_id)
.where(col(ReservationRelease.status) == "active") .where(col(ReservationRelease.status) == "active")
.where(
(col(ReservationRelease.expires_at).is_(None))
| (col(ReservationRelease.expires_at) > int(time.time()))
)
.values(created_at=int(time.time())) .values(created_at=int(time.time()))
) )
await session.commit() await session.commit()
@@ -855,12 +919,21 @@ def _start_reservation_heartbeat(snapshot: ReservationSnapshot) -> None:
""" """
interval = max(1, settings.stale_reservation_timeout_seconds // 3) interval = max(1, settings.stale_reservation_timeout_seconds // 3)
owner = asyncio.current_task() owner = asyncio.current_task()
from .core.lifecycle import request_lifetime
lifetime = request_lifetime.get()
deadline = asyncio.get_running_loop().time() + settings.max_request_lifetime_seconds
async def beat() -> None: async def beat() -> None:
try: try:
while True: while True:
await asyncio.sleep(interval) await asyncio.sleep(interval)
if owner is None or owner.done(): if (
owner is None
or owner.done()
or (lifetime is not None and lifetime.stopped)
or asyncio.get_running_loop().time() >= deadline
):
# Request control is gone; let the lease expire so the # Request control is gone; let the lease expire so the
# sweeper can release the reservation if no terminal # sweeper can release the reservation if no terminal
# transition ever ran. # transition ever ran.
@@ -1044,6 +1117,10 @@ async def _claim_reservation_for_charge(
update(ReservationRelease) update(ReservationRelease)
.where(col(ReservationRelease.id) == snapshot.release_id) .where(col(ReservationRelease.id) == snapshot.release_id)
.where(col(ReservationRelease.status) == "active") .where(col(ReservationRelease.status) == "active")
.where(
col(ReservationRelease.expires_at).is_(None)
| (col(ReservationRelease.expires_at) > int(time.time()))
)
.where(col(ReservationRelease.key_hash) == snapshot.key_hash) .where(col(ReservationRelease.key_hash) == snapshot.key_hash)
.where(col(ReservationRelease.billing_key_hash) == snapshot.billing_key_hash) .where(col(ReservationRelease.billing_key_hash) == snapshot.billing_key_hash)
.where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats) .where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats)
@@ -1104,7 +1181,7 @@ async def _charge_reservation_rows(
return True return True
async def adjust_payment_for_tokens( async def _adjust_payment_for_tokens(
key: ApiKey, key: ApiKey,
response_data: dict, response_data: dict,
session: AsyncSession, session: AsyncSession,
@@ -1540,6 +1617,45 @@ async def adjust_payment_for_tokens(
raise AssertionError("Unreachable: unhandled calculate_cost result") raise AssertionError("Unreachable: unhandled calculate_cost result")
async def adjust_payment_for_tokens(
key: ApiKey,
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:
"""Settle payment while exposing latency for every import path."""
started = time.perf_counter()
key_log_hash = key.hashed_key[:8] + "..."
succeeded = False
try:
result = await _adjust_payment_for_tokens(
key,
response_data,
session,
deducted_max_cost,
model_obj,
provider_fee,
reservation_snapshot,
)
succeeded = True
return result
finally:
logger.info(
"Payment settlement finished",
extra={
"key_hash": key_log_hash,
"model": response_data.get("model", "unknown"),
"settlement_duration_ms": round(
(time.perf_counter() - started) * 1000, 2
),
"settlement_succeeded": succeeded,
},
)
async def periodic_dead_key_prune() -> None: async def periodic_dead_key_prune() -> None:
"""Periodically prune dead API keys. Interval <= 0 disables it. """Periodically prune dead API keys. Interval <= 0 disables it.
+124
View File
@@ -0,0 +1,124 @@
"""Bounded adaptive batching for read-only NUT-07 requests, never mutations."""
import asyncio
import time
import httpx
from cashu.core.base import Proof, ProofSpentState, ProofState
from cashu.core.models import PostCheckStateRequest
from cashu.wallet.wallet import Wallet
from .core.logging import get_logger
from .mint import MINT_TRANSPORT_COOLDOWN_SECONDS, MintRateGuard, run_mint_operation
logger = get_logger(__name__)
_SDK_BATCH_LIMIT = PostCheckStateRequest.model_json_schema()["properties"]["Ys"][
"maxItems"
]
_LEARNED_TTL = 24 * 60 * 60
# Fixed scan bounds: adaptive halving makes the start size near irrelevant, and
# the deadline/request budget are safety limits, not tuning knobs.
_DEFAULT_BATCH_SIZE = _SDK_BATCH_LIMIT
_SCAN_TIMEOUT_SECONDS = 60
_MAX_REQUESTS = 128
_learned_sizes: dict[str, tuple[int, float]] = {}
async def filter_unspent_proofs(
proofs: list[Proof], wallet: Wallet, *, retry_on_rate_limit: bool = True
) -> list[Proof]:
if not proofs:
return []
mint_url = str(wallet.url)
key = mint_url.rstrip("/")
configured = _DEFAULT_BATCH_SIZE
learned, expires = _learned_sizes.get(key, (configured, 0.0))
batch_size = min(configured, learned) if expires > time.monotonic() else configured
unspent: list[Proof] = []
spent: list[Proof] = []
offset = 0
requests = 0
async def check_batch() -> tuple[list[Proof], list[ProofState]]:
nonlocal batch_size, requests
# Size fallback stays inside the rate guard's operation. A recoverable
# 500 during a cooldown probe must not open another cooldown first.
while True:
batch = proofs[offset : offset + batch_size]
if requests >= _MAX_REQUESTS:
raise ValueError("Proof-state request budget exhausted")
requests += 1
try:
response = await wallet.check_proof_state(batch)
except httpx.HTTPStatusError as error:
# A proxy 500 can mean a body limit (#761), but is not proof of
# one. Diagnostic retries are safe here because this is a read.
if error.response.status_code not in {413, 500} or len(batch) == 1:
logger.warning(
"Proof-state request failed; scan aborted",
extra={
"mint_url": mint_url,
"endpoint": "/v1/checkstate",
"status": error.response.status_code,
"content_type": error.response.headers.get("content-type"),
"request_bytes": error.request.headers.get(
"content-length"
),
"proof_count": len(batch),
"requests": requests,
},
)
raise
batch_size = max(1, len(batch) // 2)
logger.warning(
"Retrying proof-state check with a smaller batch",
extra={
"mint_url": mint_url,
"endpoint": "/v1/checkstate",
"status": error.response.status_code,
"content_type": error.response.headers.get("content-type"),
"request_bytes": error.request.headers.get("content-length"),
"proof_count": len(batch),
"next_batch_size": batch_size,
"requests": requests,
},
)
continue
states = response.states
if len(states) != len(batch) or any(
state.Y != proof.Y for proof, state in zip(batch, states)
):
raise ValueError("Invalid proof-state response: count or Y mismatch")
if any(state.state not in set(ProofSpentState) for state in states):
raise ValueError("Invalid proof-state response: unknown state")
return batch, states
# Bound the entire scan, including retries and cooldown waits.
deadline = asyncio.timeout(_SCAN_TIMEOUT_SECONDS)
try:
async with deadline:
while offset < len(proofs):
batch, states = await run_mint_operation(
check_batch,
op_name="check_proof_state",
mint_url=mint_url,
retry_on_rate_limit=retry_on_rate_limit,
)
if batch_size < configured:
_learned_sizes[key] = (batch_size, time.monotonic() + _LEARNED_TTL)
for proof, state in zip(batch, states):
if state.state == ProofSpentState.unspent:
unspent.append(proof)
elif state.state == ProofSpentState.spent:
spent.append(proof)
# Retain PENDING proofs without making them spendable.
offset += len(batch)
if spent:
await wallet.set_reserved_for_send(spent, reserved=True)
except TimeoutError:
if deadline.expired():
MintRateGuard.get(mint_url).apply_cooldown(
MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport"
)
raise
return unspent
+3
View File
@@ -2557,6 +2557,7 @@ async def get_transactions_api(
async def get_lightning_invoices_api( async def get_lightning_invoices_api(
status: str | None = None, status: str | None = None,
purpose: str | None = None, purpose: str | None = None,
direction: str | None = None,
search: str | None = None, search: str | None = None,
limit: int = 50, limit: int = 50,
offset: int = 0, offset: int = 0,
@@ -2569,6 +2570,8 @@ async def get_lightning_invoices_api(
base = base.where(LightningInvoice.status == status) base = base.where(LightningInvoice.status == status)
if purpose: if purpose:
base = base.where(LightningInvoice.purpose == purpose) base = base.where(LightningInvoice.purpose == purpose)
if direction:
base = base.where(LightningInvoice.direction == direction)
if search: if search:
pattern = f"%{search}%" pattern = f"%{search}%"
base = base.where( base = base.where(
+117 -4
View File
@@ -174,7 +174,10 @@ async def _transition_stale_reservation(
update(ReservationRelease) update(ReservationRelease)
.where(col(ReservationRelease.id) == reservation_id) .where(col(ReservationRelease.id) == reservation_id)
.where(col(ReservationRelease.status) == "active") .where(col(ReservationRelease.status) == "active")
.where(col(ReservationRelease.created_at) < cutoff) .where(
(col(ReservationRelease.created_at) < cutoff)
| (col(ReservationRelease.expires_at) <= int(time.time()))
)
.values(status="released") .values(status="released")
) )
return bool(transition.rowcount == 1) return bool(transition.rowcount == 1)
@@ -221,7 +224,10 @@ async def release_stale_reservations(
query = ( query = (
select(ReservationRelease) select(ReservationRelease)
.where(col(ReservationRelease.status) == "active") .where(col(ReservationRelease.status) == "active")
.where(col(ReservationRelease.created_at) < cutoff) .where(
(col(ReservationRelease.created_at) < cutoff)
| (col(ReservationRelease.expires_at) <= int(time.time()))
)
) )
if key_hash is not None: if key_hash is not None:
query = query.where( query = query.where(
@@ -509,14 +515,17 @@ class LightningInvoice(SQLModel, table=True): # type: ignore
status: str = Field( status: str = Field(
default="pending", default="pending",
description=( description=(
"pending, settlement_pending, paid, expired, cancelled, " "pending, settlement_pending, paid, failed, expired, cancelled, "
"reconciliation_required" "reconciliation_required"
), ),
) )
api_key_hash: str | None = Field( api_key_hash: str | None = Field(
default=None, description="Associated API key hash for topup operations" default=None, description="Associated API key hash for topup operations"
) )
purpose: str = Field(description="create or topup") direction: str = Field(
default="in", description="in for incoming invoices, out for payouts"
)
purpose: str = Field(description="create, topup or payout")
mint_url: str | None = Field( mint_url: str | None = Field(
default=None, default=None,
description="Mint URL where the quote was created (fallback tracking)", description="Mint URL where the quote was created (fallback tracking)",
@@ -781,6 +790,8 @@ class ReservationRelease(SQLModel, table=True): # type: ignore
key_hash: str = Field(index=True) key_hash: str = Field(index=True)
billing_key_hash: str = Field(index=True) billing_key_hash: str = Field(index=True)
reserved_msats: int reserved_msats: int
started_at: int | None = Field(default=None)
expires_at: int | None = Field(default=None, index=True)
status: str = Field(default="active") status: str = Field(default="active")
created_at: int = Field(default_factory=lambda: int(time.time())) created_at: int = Field(default_factory=lambda: int(time.time()))
@@ -1005,6 +1016,80 @@ async def complete_routstr_fee_payout(
return result.rowcount == 1 return result.rowcount == 1
async def record_lightning_payout(
session: AsyncSession,
*,
quote_id: str,
bolt11: str,
amount_sats: int,
mint_url: str,
destination: str,
) -> None:
"""Record a dispatched payout so it shows up in the Lightning history."""
session.add(
LightningInvoice(
id=uuid.uuid4().hex,
bolt11=bolt11,
amount_sats=amount_sats,
description=f"Payout to {destination}",
payment_hash=quote_id,
status="pending",
direction="out",
purpose="payout",
mint_url=mint_url,
# Payouts settle or fail at the mint; they never expire on our side.
expires_at=int(time.time()),
)
)
await session.commit()
async def settle_lightning_payout(
session: AsyncSession,
quote_id: str,
*,
status: str,
amount_sats: int | None = None,
) -> None:
result = await session.exec(
select(LightningInvoice)
.where(col(LightningInvoice.payment_hash) == quote_id)
.where(col(LightningInvoice.direction) == "out")
)
payout = result.first()
if payout is None:
logger.warning(
"No Lightning payout history row for quote",
extra={"quote_id": quote_id, "status": status},
)
return
payout.status = status
if status == "paid":
payout.paid_at = int(time.time())
if amount_sats is not None:
payout.amount_sats = amount_sats
session.add(payout)
await session.commit()
UNSETTLED_PAYOUT_STATUSES = ("pending", "reconciliation_required")
async def list_unsettled_lightning_payouts(
session: AsyncSession, mint_url: str, *, created_before: int
) -> list[LightningInvoice]:
"""Payout rows whose mint outcome was never written back to history."""
result = await session.exec(
select(LightningInvoice)
.where(col(LightningInvoice.direction) == "out")
.where(col(LightningInvoice.mint_url) == mint_url)
.where(col(LightningInvoice.status).in_(UNSETTLED_PAYOUT_STATUSES))
.where(col(LightningInvoice.created_at) < created_before)
.order_by(col(LightningInvoice.created_at))
)
return list(result.all())
async def total_user_liability(db_session: AsyncSession) -> int: async def total_user_liability(db_session: AsyncSession) -> int:
"""Return all outstanding user funds in millisatoshis. """Return all outstanding user funds in millisatoshis.
@@ -1021,6 +1106,34 @@ async def total_user_liability(db_session: AsyncSession) -> int:
return int(result.one() or 0) return int(result.one() or 0)
async def user_liability_for_mint_and_unit(
db_session: AsyncSession, mint_url: str, unit: str
) -> int:
"""Return outstanding user funds that refund from one mint and unit, in msats.
Single statement, for the same atomicity reason as ``total_user_liability``.
"""
key_balances = (
select(func.coalesce(func.sum(ApiKey.balance), 0))
.where(
col(ApiKey.refund_mint_url) == mint_url,
col(ApiKey.refund_currency) == unit,
)
.scalar_subquery()
)
unresolved_refunds = (
select(func.coalesce(func.sum(Refund.amount_msats), 0))
.where(
col(Refund.status).in_(REFUND_UNRESOLVED_STATUSES),
col(Refund.mint_url) == mint_url,
col(Refund.unit) == unit,
)
.scalar_subquery()
)
result = await db_session.exec(select(key_balances + unresolved_refunds))
return int(result.one() or 0)
async def balance_for_mint_and_unit( async def balance_for_mint_and_unit(
db_session: AsyncSession, mint_url: str, unit: str db_session: AsyncSession, mint_url: str, unit: str
) -> int: ) -> int:
+60
View File
@@ -0,0 +1,60 @@
"""Attribution scope for upstream-caused failures.
An upstream 5xx forwarded verbatim makes callers think this node is down.
Upstream failures are therefore reported as ``424`` with
``error.code = UPSTREAM_UNAVAILABLE``, the ``X-Routstr-Error-Scope: upstream``
header, and the provider's status in ``upstream_status``. Rate limits keep
``429``. Node faults keep ``500`` and carry no scope header.
"""
from __future__ import annotations
UPSTREAM_UNAVAILABLE = "UPSTREAM_UNAVAILABLE"
# Deliberately not 5xx: an upstream blip must not read as node health.
UPSTREAM_ERROR_STATUS = 424
ERROR_SCOPE_HEADER = "X-Routstr-Error-Scope"
ERROR_SCOPE_UPSTREAM = "upstream"
ERROR_SCOPE_NODE = "node"
def _is_rate_limit_code(code: object) -> bool:
# Lazy import: routstr.upstream imports this module.
from ..upstream.rate_limit import UPSTREAM_RATE_LIMIT
return code == UPSTREAM_RATE_LIMIT
def client_status_for_upstream_error(
status_code: int | None, code: object = None
) -> int:
"""Map an upstream status to the one the caller sees: 429 and 4xx pass
through, 5xx (or unknown) becomes :data:`UPSTREAM_ERROR_STATUS`."""
if _is_rate_limit_code(code):
return 429
if not status_code or status_code >= 500:
return UPSTREAM_ERROR_STATUS
return status_code
def client_code_for_upstream_error(
status_code: int | None, code: str | int | None
) -> str | int | None:
"""Return the ``error.code`` matching :func:`client_status_for_upstream_error`."""
if _is_rate_limit_code(code):
return code
if not status_code or status_code >= 500:
return UPSTREAM_UNAVAILABLE
return code
def upstream_status_details(
details: dict[str, object] | None, upstream_status: int | None
) -> dict[str, object] | None:
"""Add ``upstream_status`` to ``details`` when the caller sees a different status."""
merged: dict[str, object] = dict(details) if details else {}
if upstream_status and upstream_status != client_status_for_upstream_error(
upstream_status
):
merged["upstream_status"] = upstream_status
return merged or None
+48 -3
View File
@@ -5,6 +5,11 @@ from fastapi.encoders import jsonable_encoder
from fastapi.exceptions import RequestValidationError from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from .error_scope import (
ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS,
UPSTREAM_UNAVAILABLE,
)
from .logging import get_logger from .logging import get_logger
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -18,6 +23,19 @@ class UpstreamError(Exception):
string-matching the message. ``details`` holds optional structured, string-matching the message. ``details`` holds optional structured,
redaction-safe context. Both default to ``None`` for backwards redaction-safe context. Both default to ``None`` for backwards
compatibility. compatibility.
``from_upstream_response`` is True only when ``status_code`` is the status
the upstream itself answered with, as opposed to a status this proxy chose
for a transport failure, timeout or internal fault. Callers use it to
decide whether a status is safe to retry.
``scope`` is ``"upstream"`` for provider failures, reported to the caller
as ``424`` (see :mod:`routstr.core.error_scope`), or ``"node"`` for local
faults, which keep their status. ``status_code`` stays the provider's own
status whenever ``from_upstream_response`` is True; the caller-visible
mapping happens at response construction. Proxy-chosen statuses (transport
failure, timeout) may already be the caller-visible one — only read
``status_code`` as a provider status behind ``from_upstream_response``.
""" """
def __init__( def __init__(
@@ -26,11 +44,15 @@ class UpstreamError(Exception):
status_code: int = 502, status_code: int = 502,
code: str | None = None, code: str | None = None,
details: dict[str, object] | None = None, details: dict[str, object] | None = None,
from_upstream_response: bool = False,
scope: str = ERROR_SCOPE_UPSTREAM,
): ):
self.message = message self.message = message
self.status_code = status_code self.status_code = status_code
self.code = code self.code = code
self.details = details self.details = details
self.from_upstream_response = from_upstream_response
self.scope = scope
super().__init__(message) super().__init__(message)
@@ -38,8 +60,9 @@ class EhbpTimeoutError(UpstreamError):
"""Raised when an EHBP upstream times out waiting for a response. """Raised when an EHBP upstream times out waiting for a response.
Distinct from a generic :class:`UpstreamError` so callers can map the Distinct from a generic :class:`UpstreamError` so callers can map the
failure to a ``504 Gateway Timeout`` with a stable ``UPSTREAM_TIMEOUT`` failure to a stable ``UPSTREAM_TIMEOUT`` code instead of a misleading
code instead of a misleading ``500`` internal server error. ``500`` internal server error. Reported as ``424``: the timeout happened on
the provider hop, not this node.
``details`` carries optional structured, redaction-safe context and is ``details`` carries optional structured, redaction-safe context and is
forwarded to the client by ``create_upstream_error_response``. forwarded to the client by ``create_upstream_error_response``.
@@ -48,12 +71,34 @@ class EhbpTimeoutError(UpstreamError):
def __init__(self, message: str, details: dict[str, object] | None = None): def __init__(self, message: str, details: dict[str, object] | None = None):
super().__init__( super().__init__(
message, message,
status_code=504, status_code=UPSTREAM_ERROR_STATUS,
code="UPSTREAM_TIMEOUT", code="UPSTREAM_TIMEOUT",
details=details, details=details,
) )
class EhbpConnectionError(UpstreamError):
"""Raised when an EHBP upstream cannot be reached.
Covers transport failures while establishing the provider connection: DNS
resolution, TCP refused/reset, or a TLS error that is not a handshake
timeout. Distinct from a generic :class:`UpstreamError` so the failure is
attributed to the provider hop (``UPSTREAM_UNAVAILABLE``, reported as
``424``) instead of being flattened into a misleading node-scoped ``500``.
``details`` carries optional structured, redaction-safe context and is
forwarded to the client by ``create_upstream_error_response``.
"""
def __init__(self, message: str, details: dict[str, object] | None = None):
super().__init__(
message,
status_code=UPSTREAM_ERROR_STATUS,
code=UPSTREAM_UNAVAILABLE,
details=details,
)
def _error_message_from_detail(detail: object) -> str | None: def _error_message_from_detail(detail: object) -> str | None:
"""Extract a message from an HTTPException ``detail``, capped at 200 chars.""" """Extract a message from an HTTPException ``detail``, capped at 200 chars."""
if isinstance(detail, dict): if isinstance(detail, dict):
+134
View File
@@ -0,0 +1,134 @@
"""Supervise the real downstream connection, outside HTTP middleware wrappers."""
from __future__ import annotations
import asyncio
from contextvars import ContextVar
from dataclasses import dataclass
from starlette.types import ASGIApp, Message, Receive, Scope, Send
from . import get_logger
from .settings import settings
logger = get_logger(__name__)
class DownstreamTerminated(OSError):
"""Raised by downstream_send after disconnect; expected, not a server error."""
@dataclass
class RequestLifetime:
deadline: float = 0
stopped: bool = False
request_lifetime: ContextVar[RequestLifetime | None] = ContextVar(
"request_lifetime", default=None
)
class RequestLifecycleMiddleware:
def __init__(self, app: ASGIApp) -> None:
self.app = app
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
if scope["type"] != "http":
await self.app(scope, receive, send)
return
lifetime = RequestLifetime(
deadline=asyncio.get_running_loop().time()
+ settings.max_request_lifetime_seconds
)
token = request_lifetime.set(lifetime)
disconnected = asyncio.Event()
# One receive consumer. Backpressure uploads until consumed; after the
# final body message, continue listening independently of the app.
messages: asyncio.Queue[Message] = asyncio.Queue(maxsize=1)
response_started = False
async def pump() -> None:
while True:
message = await receive()
if message["type"] == "http.disconnect":
disconnected.set()
return
await messages.put(message)
async def downstream_receive() -> Message:
if disconnected.is_set():
return {"type": "http.disconnect"}
get = asyncio.create_task(messages.get())
gone = asyncio.create_task(disconnected.wait())
try:
await asyncio.wait((get, gone), return_when=asyncio.FIRST_COMPLETED)
if disconnected.is_set():
return {"type": "http.disconnect"}
return get.result()
finally:
for task in (get, gone):
task.cancel()
await asyncio.gather(get, gone, return_exceptions=True)
async def downstream_send(message: Message) -> None:
nonlocal response_started
if disconnected.is_set() or lifetime.stopped:
raise DownstreamTerminated("Downstream request terminated")
async with asyncio.timeout(settings.downstream_send_timeout_seconds):
await send(message)
if message["type"] == "http.response.start":
response_started = True
receiver = asyncio.create_task(pump())
work: asyncio.Future[None] = asyncio.ensure_future(
self.app(scope, downstream_receive, downstream_send)
)
gone = asyncio.create_task(disconnected.wait())
timed_out = False
try:
done, _ = await asyncio.wait(
(work, gone),
timeout=settings.max_request_lifetime_seconds,
return_when=asyncio.FIRST_COMPLETED,
)
if gone in done and not response_started and work not in done:
# A pre-response wallet or billing operation may have accepted
# funds already. Let it reach its own settlement before closing.
done, _ = await asyncio.wait(
(work,),
timeout=max(0, lifetime.deadline - asyncio.get_running_loop().time()),
)
if work in done:
try:
await work
except DownstreamTerminated:
if not disconnected.is_set():
raise
logger.debug("Client disconnected before response completed")
elif not disconnected.is_set():
timed_out = True
finally:
lifetime.stopped = True
for task in (receiver, gone, work):
task.cancel()
# Detached stream finalizers own settlement. The heartbeat stops
# with the request; durable expiry recovers any abandoned row.
done, pending = await asyncio.wait(
(receiver, gone, work), timeout=settings.request_cleanup_timeout_seconds
)
for task in done:
if not task.cancelled():
task.exception()
for task in pending:
task.cancel()
task.add_done_callback(
lambda t: t.exception() if not t.cancelled() else None
)
request_lifetime.reset(token)
if timed_out and work.done() and not disconnected.is_set() and not response_started:
async with asyncio.timeout(settings.downstream_send_timeout_seconds):
await send({"type": "http.response.start", "status": 504, "headers": []})
await send(
{"type": "http.response.body", "body": b"Request deadline exceeded"}
)
+223 -2
View File
@@ -28,7 +28,12 @@ DO NOT modify or remove these messages without updating the usage tracking logic
- The 'max_cost_for_model' field is extracted for refund calculation - The 'max_cost_for_model' field is extracted for refund calculation
- Must include 'max_cost_for_model' in extra dict - Must include 'max_cost_for_model' in extra dict
6. Any ERROR level logs with "upstream" in the message 6. "Payment settlement finished" (INFO) - routstr/auth.py and routstr/upstream/ehbp.py
- Emitted once per settlement attempt, including EHBP settlements
- Carries 'settlement_duration_ms' and 'settlement_succeeded'; the EHBP
emitter adds 'settlement_type'
7. Any ERROR level logs with "upstream" in the message
- Used to count upstream provider errors - Used to count upstream provider errors
- Helps identify service reliability issues - Helps identify service reliability issues
@@ -37,11 +42,15 @@ If you need to modify these messages, ensure you also update the parsing logic i
- routstr/core/log_manager.py - routstr/core/log_manager.py
""" """
import copy
import logging.config import logging.config
import logging.handlers import logging.handlers
import os import os
import queue
import re import re
import sys import sys
import threading
import time
import tomllib import tomllib
from datetime import datetime from datetime import datetime
from pathlib import Path from pathlib import Path
@@ -127,6 +136,218 @@ class DailyRotatingFileHandler(logging.handlers.TimedRotatingFileHandler):
pass pass
class QueuedDailyRotatingFileHandler(logging.Handler):
"""Move rotating-file I/O off request threads.
When both locks are needed, acquire the logging module lock before the
handler lock to match ``dictConfig``.
"""
_queue: queue.Queue[logging.LogRecord]
_target: DailyRotatingFileHandler
_listener: logging.handlers.QueueListener
_drain_timeout_seconds = 5.0
_reopen_backoff_seconds = 5.0
def __init__(self, filename: str, **kwargs: Any) -> None:
super().__init__()
self._filename = filename
self._kwargs = kwargs
self._stopped = True
self._next_open_attempt = 0.0
self._dropped_warned_at = -1.0
self._open()
def _open(self) -> None:
"""Attach a fresh rotating file handler and start draining it."""
# A new queue per listener: QueueListener's stop sentinel is a shared
# singleton, so two listeners on one queue would steal each other's.
record_queue: queue.Queue[logging.LogRecord] = queue.Queue()
target = DailyRotatingFileHandler(self._filename, **self._kwargs)
target.setFormatter(self.formatter)
listener = logging.handlers.QueueListener(record_queue, target)
try:
listener.start()
except Exception:
target.close()
raise
self._queue = record_queue
self._target = target
self._listener = listener
self._stopped = False
with getattr(logging, "_lock"):
handler_list = getattr(logging, "_handlerList")
# This wrapper owns the target's shutdown and lock ordering.
handler_list[:] = [
reference for reference in handler_list if reference() is not target
]
if not any(reference() is self for reference in handler_list):
getattr(logging, "_addHandlerRef")(self)
def _reopen_locked(self) -> bool:
"""Reopen using the lock order required by ``dictConfig``."""
with getattr(logging, "_lock"):
self.acquire()
try:
if not self._stopped:
return True
if time.monotonic() < self._next_open_attempt:
return False
try:
self._open()
except Exception:
self._next_open_attempt = (
time.monotonic() + self._reopen_backoff_seconds
)
raise
return True
finally:
self.release()
def setFormatter(self, fmt: logging.Formatter | None) -> None:
super().setFormatter(fmt)
self._target.setFormatter(fmt)
def handle(self, record: logging.LogRecord) -> bool:
if not self.filter(record):
return False
while True:
self.acquire()
try:
if not self._stopped:
self.emit(record)
return True
finally:
self.release()
if sys.is_finalizing():
# logging.shutdown() already ran; a new listener thread would
# never drain, so write the record synchronously instead.
self._emit_synchronously(record)
return False
try:
# Do not acquire the module lock while holding the handler lock.
if not self._reopen_locked():
self._warn_records_dropped()
return False
except Exception:
self.handleError(record)
return False
def _warn_records_dropped(self) -> None:
"""Report once per backoff window instead of dropping records silently."""
self.acquire()
try:
if self._dropped_warned_at >= self._next_open_attempt:
return
self._dropped_warned_at = self._next_open_attempt
finally:
self.release()
sys.stderr.write(
f"Logging listener for {self._filename} is unavailable; dropping "
f"records until the next reopen attempt in "
f"{self._reopen_backoff_seconds}s\n"
)
def _emit_synchronously(self, record: logging.LogRecord) -> None:
try:
sys.stderr.write(self.format(record) + "\n")
except Exception:
self.handleError(record)
def emit(self, record: logging.LogRecord) -> None:
try:
if not self._stopped:
self._queue.put_nowait(copy.copy(record))
except Exception:
# Handler.handle() does not catch exceptions raised by emit().
self.handleError(record)
def flush(self) -> None:
self.acquire()
try:
if self._stopped:
return
deadline = time.monotonic() + self._drain_timeout_seconds
with self._queue.all_tasks_done:
while self._queue.unfinished_tasks:
remaining = deadline - time.monotonic()
if remaining <= 0:
break
self._queue.all_tasks_done.wait(remaining)
pending = self._queue.unfinished_tasks
if pending:
sys.stderr.write(
f"Logging listener for {self._filename} still has {pending} "
f"record(s) queued after {self._drain_timeout_seconds}s flush\n"
)
self._target.flush()
finally:
self.release()
def _stop_listener(self) -> bool:
thread = self._listener._thread
if thread is None:
return True
self._listener.enqueue_sentinel()
thread.join(timeout=self._drain_timeout_seconds)
if thread.is_alive():
return False
self._listener._thread = None
return True
def _close_retired_listener(
self,
listener: logging.handlers.QueueListener,
target: DailyRotatingFileHandler,
) -> None:
def finish() -> None:
thread = listener._thread
if thread is not None:
thread.join()
listener._thread = None
try:
target.flush()
finally:
target.close()
threading.Thread(target=finish, daemon=True).start()
def close(self) -> None:
# Stop the listener under the handler lock, then close the target outside
# it because FileHandler.close() also takes the logging module lock.
self.acquire()
try:
target = None
retired = None
if not self._stopped:
listener = self._listener
current_target = self._target
stopped = self._stop_listener()
self._stopped = True
if stopped:
target = current_target
else:
retired = (listener, current_target)
sys.stderr.write(
f"Logging listener for {self._filename} did not stop "
f"within {self._drain_timeout_seconds}s; reopening on next record\n"
)
finally:
self.release()
if retired is not None:
self._close_retired_listener(*retired)
if target is not None:
target.flush()
target.close()
super().close()
def get_package_version() -> str: def get_package_version() -> str:
"""Read the package version from pyproject.toml.""" """Read the package version from pyproject.toml."""
try: try:
@@ -369,7 +590,7 @@ def setup_logging() -> None:
"handlers": { "handlers": {
"console": console_handler, "console": console_handler,
"file": { "file": {
"()": DailyRotatingFileHandler, "()": QueuedDailyRotatingFileHandler,
"level": log_level, "level": log_level,
"formatter": "json", "formatter": "json",
"filename": "logs/app.log", "filename": "logs/app.log",
+15 -6
View File
@@ -34,7 +34,7 @@ from ..payment.price import update_prices_periodically
from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically
from ..refund import periodic_refund_reconcile from ..refund import periodic_refund_reconcile
from ..upstream.auto_topup import periodic_auto_topup from ..upstream.auto_topup import periodic_auto_topup
from ..upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing from ..upstream.http_client import close_upstream_http_client
from ..upstream.litellm_routing import configure_litellm from ..upstream.litellm_routing import configure_litellm
from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout
from .admin import admin_router from .admin import admin_router
@@ -44,6 +44,7 @@ from .exceptions import (
http_exception_handler, http_exception_handler,
validation_exception_handler, validation_exception_handler,
) )
from .lifecycle import RequestLifecycleMiddleware
from .logging import get_logger, setup_logging from .logging import get_logger, setup_logging
from .middleware import LoggingMiddleware from .middleware import LoggingMiddleware
from .not_found import _NOT_FOUND_HTML, not_found_catch_all # noqa: F401 from .not_found import _NOT_FOUND_HTML, not_found_catch_all # noqa: F401
@@ -88,11 +89,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
# debug logging) before any upstream provider dispatches a request. # debug logging) before any upstream provider dispatches a request.
configure_litellm() configure_litellm()
# TEMPORARY: backfill DeepSeek V4 pricing missing from litellm's cost
# map (BerriAI/litellm#30430). Remove this call and
# deepseek_v4_pricing_shim.py once litellm ships these models.
register_deepseek_v4_pricing()
# Run database migrations on startup # Run database migrations on startup
run_migrations() run_migrations()
@@ -260,6 +256,14 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
"Error stopping background tasks", "Error stopping background tasks",
extra={"error": str(e), "error_type": type(e).__name__}, extra={"error": str(e), "error_type": type(e).__name__},
) )
finally:
try:
await close_upstream_http_client()
except Exception as e:
logger.error(
"Error closing upstream HTTP connection pools",
extra={"error": str(e), "error_type": type(e).__name__},
)
class _ImmutableStaticFiles(StaticFiles): class _ImmutableStaticFiles(StaticFiles):
@@ -289,6 +293,7 @@ app.add_middleware(
expose_headers=[ expose_headers=[
"x-routstr-request-id", "x-routstr-request-id",
"x-cashu", "x-cashu",
"x-routstr-error-scope",
"x-routstr-cost-msats", "x-routstr-cost-msats",
"x-routstr-cost-usd", "x-routstr-cost-usd",
"x-routstr-input-cost-msats", "x-routstr-input-cost-msats",
@@ -305,6 +310,10 @@ app.add_middleware(
# Add logging middleware # Add logging middleware
app.add_middleware(LoggingMiddleware) app.add_middleware(LoggingMiddleware)
# Outermost: observe the actual downstream connection, not middleware streams.
app.add_middleware(RequestLifecycleMiddleware)
# Add exception handlers # Add exception handlers
app.add_exception_handler(HTTPException, http_exception_handler) # type: ignore app.add_exception_handler(HTTPException, http_exception_handler) # type: ignore
app.add_exception_handler(RequestValidationError, validation_exception_handler) app.add_exception_handler(RequestValidationError, validation_exception_handler)
+172 -28
View File
@@ -1,7 +1,7 @@
import time import time
import uuid import uuid
from contextvars import ContextVar from contextvars import ContextVar
from typing import Callable from typing import AsyncIterator, Callable
from urllib.parse import urlsplit from urllib.parse import urlsplit
from fastapi import Request, Response from fastapi import Request, Response
@@ -9,6 +9,7 @@ from starlette.datastructures import Headers
from starlette.middleware.base import BaseHTTPMiddleware from starlette.middleware.base import BaseHTTPMiddleware
from .logging import get_logger from .logging import get_logger
from .settings import settings
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -86,20 +87,140 @@ _SKIP_LOG_EXACT: frozenset[str] = frozenset(
) )
def _should_log(method: str, path: str) -> bool: def _should_log(method: str, path: str, status_code: int | None = None) -> bool:
if method in _SKIP_LOG_METHODS: if method in _SKIP_LOG_METHODS:
return False return False
# Our own faults are never noise, whatever the path.
if status_code is not None and status_code >= 500:
return True
if path in _SKIP_LOG_EXACT: if path in _SKIP_LOG_EXACT:
return False # A 4xx storm on a UI-polled path is exactly what we need to see.
return status_code is not None and status_code >= 400
# Client errors on the skipped prefixes stay hidden: 404s under /_next/ are
# driven by whoever scans the node, and the admin UI's timer-driven polling
# turns one expired session into a 401 per poll.
return not any(path.startswith(prefix) for prefix in _SKIP_LOG_PREFIXES) return not any(path.startswith(prefix) for prefix in _SKIP_LOG_PREFIXES)
def _attribution(request: Request) -> dict[str, object]:
"""Model/provider fields, omitted rather than null on routes that resolve none."""
return {
field: value
for field in ("model", "provider")
if (value := getattr(request.state, field, None))
}
def mark(request: Request, name: str) -> None:
"""Record that stage ``name`` finished, for the completion log's timings."""
marks = getattr(request.state, "stage_marks", None)
if marks is not None:
marks[name] = time.monotonic()
def _request_content_length(headers: Headers) -> int | None:
"""Client-supplied length, dropped unless it is a plausible byte count."""
raw = headers.get("content-length")
if raw is None:
return None
try:
value = int(raw)
except ValueError:
return None
return value if value >= 0 else None
class LoggingMiddleware(BaseHTTPMiddleware): class LoggingMiddleware(BaseHTTPMiddleware):
"""Middleware to log proxy interactions and page navigation. """Middleware to log proxy interactions and page navigation.
Skips logging for static assets and Next.js chunks to avoid noise. Skips logging for static assets and Next.js chunks to avoid noise.
""" """
def _log_completion(
self,
*,
request: Request,
request_id: str,
path: str,
status_code: int,
duration: float,
headers_duration: float | None,
stage_start: float,
stage_marks: dict[str, float],
incoming_logged: bool,
) -> None:
if not _should_log(request.method, path, status_code):
return
extra: dict[str, object] = {
"request_id": request_id,
"method": request.method,
"path": path,
"status_code": status_code,
"duration_ms": round(duration * 1000, 2),
"content_length": _request_content_length(request.headers),
**_attribution(request),
}
if headers_duration is not None:
extra["time_to_headers_ms"] = round(headers_duration * 1000, 2)
if not incoming_logged:
# Tells log consumers that join on request_id why the matching
# "Incoming request" record is missing.
extra["incoming_suppressed"] = True
for name, marked_at in stage_marks.items():
extra[f"{name}_ms"] = round((marked_at - stage_start) * 1000, 2)
if status_code >= 400:
error_detail = getattr(request.state, "error_detail", None)
if isinstance(error_detail, dict):
extra["error_type"] = error_detail.get("error_type")
extra["error_code"] = error_detail.get("error_code")
extra["error_message"] = error_detail.get("error_message")
log = (
logger.warning
if duration > settings.slow_request_warn_seconds
else logger.info
)
log("Request completed", extra=extra)
async def _timed_body(
self,
body_iterator: AsyncIterator[bytes],
*,
request: Request,
request_id: str,
client_app: str,
path: str,
status_code: int,
stage_start: float,
stage_marks: dict[str, float],
headers_duration: float,
incoming_logged: bool,
) -> AsyncIterator[bytes]:
try:
async for chunk in body_iterator:
yield chunk
finally:
duration = time.monotonic() - stage_start
# dispatch() has already reset both context vars by now, and the
# logging filters read request_id/client_app from them.
request_token = request_id_context.set(request_id)
app_token = client_app_context.set(client_app)
try:
self._log_completion(
request=request,
request_id=request_id,
path=path,
status_code=status_code,
duration=duration,
headers_duration=headers_duration,
stage_start=stage_start,
stage_marks=stage_marks,
incoming_logged=incoming_logged,
)
finally:
request_id_context.reset(request_token)
client_app_context.reset(app_token)
async def dispatch(self, request: Request, call_next: Callable) -> Response: async def dispatch(self, request: Request, call_next: Callable) -> Response:
# Generate request ID # Generate request ID
request_id = str(uuid.uuid4()) request_id = str(uuid.uuid4())
@@ -108,15 +229,17 @@ class LoggingMiddleware(BaseHTTPMiddleware):
# Set request ID in context for logging # Set request ID in context for logging
token = request_id_context.set(request_id) token = request_id_context.set(request_id)
client_app_token = client_app_context.set( client_app = client_app_from_headers(request.headers)
client_app_from_headers(request.headers) client_app_token = client_app_context.set(client_app)
)
path = request.url.path path = request.url.path
should_log = _should_log(request.method, path) should_log = _should_log(request.method, path)
# Start timing # Start timing. Monotonic throughout: a wall-clock step would otherwise
start_time = time.time() # produce negative durations and bogus slow-request warnings.
stage_start = time.monotonic()
stage_marks: dict[str, float] = {}
request.state.stage_marks = stage_marks
if should_log: if should_log:
logger.info( logger.info(
@@ -135,33 +258,52 @@ class LoggingMiddleware(BaseHTTPMiddleware):
try: try:
response = await call_next(request) response = await call_next(request)
if should_log: headers_duration = time.monotonic() - stage_start
duration = time.time() - start_time
extra: dict[str, object] = {
"request_id": request_id,
"method": request.method,
"path": path,
"status_code": response.status_code,
"duration_ms": round(duration * 1000, 2),
}
if response.status_code >= 400:
error_detail = getattr(request.state, "error_detail", None)
if isinstance(error_detail, dict):
extra["error_type"] = error_detail.get("error_type")
extra["error_code"] = error_detail.get("error_code")
extra["error_message"] = error_detail.get("error_message")
logger.info(
"Request completed",
extra=extra,
)
if hasattr(response, "headers"): if hasattr(response, "headers"):
response.headers["x-routstr-request-id"] = request_id response.headers["x-routstr-request-id"] = request_id
# Headers are already on the wire before a streamed body ends,
# so this can only ever be time-to-headers.
response.headers["x-routstr-duration-ms"] = str(
round(headers_duration * 1000, 2)
)
body_iterator = getattr(response, "body_iterator", None)
if body_iterator is None:
self._log_completion(
request=request,
request_id=request_id,
path=path,
status_code=response.status_code,
duration=headers_duration,
headers_duration=None,
stage_start=stage_start,
stage_marks=stage_marks,
incoming_logged=should_log,
)
return response
# A StreamingResponse is barely started here: most of the time a
# slow completion spends in the node is spent relaying its body, so
# the completion log has to wait for the iterator to drain.
response.body_iterator = self._timed_body(
body_iterator,
request=request,
request_id=request_id,
client_app=client_app,
path=path,
status_code=response.status_code,
stage_start=stage_start,
stage_marks=stage_marks,
headers_duration=headers_duration,
incoming_logged=should_log,
)
return response return response
except Exception as e: except Exception as e:
# Always log failures, even for skipped paths, so we don't lose errors. # Always log failures, even for skipped paths, so we don't lose errors.
duration = time.time() - start_time duration = time.monotonic() - stage_start
logger.error( logger.error(
"Request failed", "Request failed",
extra={ extra={
@@ -171,6 +313,7 @@ class LoggingMiddleware(BaseHTTPMiddleware):
"duration_ms": round(duration * 1000, 2), "duration_ms": round(duration * 1000, 2),
"error": str(e), "error": str(e),
"error_type": type(e).__name__, "error_type": type(e).__name__,
**_attribution(request),
}, },
exc_info=True, exc_info=True,
) )
@@ -185,5 +328,6 @@ __all__ = [
"LoggingMiddleware", "LoggingMiddleware",
"UNKNOWN_CLIENT_APP", "UNKNOWN_CLIENT_APP",
"client_app_context", "client_app_context",
"mark",
"request_id_context", "request_id_context",
] ]
+71 -2
View File
@@ -11,6 +11,14 @@ from typing import Any
from pydantic.v1 import BaseModel, BaseSettings, Field from pydantic.v1 import BaseModel, BaseSettings, Field
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
# Mints a fresh node trusts out of the box, shared by the settings default and
# the primary-mint fallback. Defined before the Settings class because the
# default_factory lambda resolves it at class-definition time.
DEFAULT_CASHU_MINTS: list[str] = [
"https://mint.minibits.cash/Bitcoin",
"https://mint.cubabitcoin.org",
]
class Settings(BaseSettings): class Settings(BaseSettings):
class Config: class Config:
@@ -28,6 +36,26 @@ class Settings(BaseSettings):
# Core # Core
upstream_base_url: str = Field(default="", env="UPSTREAM_BASE_URL") upstream_base_url: str = Field(default="", env="UPSTREAM_BASE_URL")
upstream_api_key: str = Field(default="", env="UPSTREAM_API_KEY") upstream_api_key: str = Field(default="", env="UPSTREAM_API_KEY")
# Extra attempts against the same upstream on a transient 5xx. 0 disables.
upstream_5xx_retry_attempts: int = Field(
default=1, ge=0, env="UPSTREAM_5XX_RETRY_ATTEMPTS"
)
# Streaming guards, off by default (0). A stream that never produces a
# first chunk can still fail over; one that stalls later can only be
# aborted and billed for what it delivered. Reasoning models can stay
# silent for minutes, so set these above the longest expected think time.
upstream_first_token_timeout_seconds: float = Field(
default=0.0, ge=0, env="UPSTREAM_FIRST_TOKEN_TIMEOUT_SECONDS"
)
upstream_stream_idle_timeout_seconds: float = Field(
default=0.0, ge=0, env="UPSTREAM_STREAM_IDLE_TIMEOUT_SECONDS"
)
# Circuit breaker: timeouts/5xx per (provider, model) within a minute that
# take the pair out of candidate selection. 0 seconds disables it.
upstream_allowed_fails: int = Field(default=3, ge=1, env="UPSTREAM_ALLOWED_FAILS")
upstream_cooldown_seconds: float = Field(
default=30.0, ge=0, env="UPSTREAM_COOLDOWN_SECONDS"
)
# Node info # Node info
name: str = Field(default="ARoutstrNode", env="NAME") name: str = Field(default="ARoutstrNode", env="NAME")
@@ -37,7 +65,12 @@ class Settings(BaseSettings):
onion_url: str = Field(default="", env="ONION_URL") onion_url: str = Field(default="", env="ONION_URL")
# Cashu # Cashu
cashu_mints: list[str] = Field(default_factory=list, env="CASHU_MINTS") # Mints a fresh node trusts out of the box. Setting CASHU_MINTS (env or
# dashboard) replaces this list entirely; an explicitly empty value yields
# an empty list (no trusted mints beyond primary_mint).
cashu_mints: list[str] = Field(
default_factory=lambda: list(DEFAULT_CASHU_MINTS), env="CASHU_MINTS"
)
receive_ln_address: str = Field(default="", env="RECEIVE_LN_ADDRESS") receive_ln_address: str = Field(default="", env="RECEIVE_LN_ADDRESS")
primary_mint: str = Field(default="", env="PRIMARY_MINT_URL") primary_mint: str = Field(default="", env="PRIMARY_MINT_URL")
primary_mint_unit: str = Field(default="sat", env="PRIMARY_MINT_UNIT") primary_mint_unit: str = Field(default="sat", env="PRIMARY_MINT_UNIT")
@@ -49,6 +82,8 @@ class Settings(BaseSettings):
# Minimum available balance (in satoshis) before profit is paid out over # Minimum available balance (in satoshis) before profit is paid out over
# Lightning # Lightning
min_payout_sat: int = Field(default=210, gt=0, env="MIN_PAYOUT_SAT") min_payout_sat: int = Field(default=210, gt=0, env="MIN_PAYOUT_SAT")
# Gross payout budget in sats, including fees.
max_payout_sat: int = Field(default=250_000, gt=0, env="MAX_PAYOUT_SAT")
# Interval (seconds) between periodic payout attempts. Must be positive. # Interval (seconds) between periodic payout attempts. Must be positive.
payout_interval_seconds: int = Field( payout_interval_seconds: int = Field(
default=900, gt=0, env="PAYOUT_INTERVAL_SECONDS" default=900, gt=0, env="PAYOUT_INTERVAL_SECONDS"
@@ -98,6 +133,16 @@ class Settings(BaseSettings):
default=604_800, env="DEAD_KEY_MIN_AGE_SECONDS" default=604_800, env="DEAD_KEY_MIN_AGE_SECONDS"
) )
max_request_lifetime_seconds: float = Field(
default=1800, gt=0, env="MAX_REQUEST_LIFETIME_SECONDS"
)
downstream_send_timeout_seconds: float = Field(
default=60, gt=0, env="DOWNSTREAM_SEND_TIMEOUT_SECONDS"
)
request_cleanup_timeout_seconds: float = Field(
default=30, gt=0, env="REQUEST_CLEANUP_TIMEOUT_SECONDS"
)
# Network # Network
cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS") cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS")
# Comma-separated METHOD:path pairs adding to the proxy's canonical # Comma-separated METHOD:path pairs adding to the proxy's canonical
@@ -106,6 +151,14 @@ class Settings(BaseSettings):
# widens what the provider credential can be spent against, so wildcards # widens what the provider credential can be spent against, so wildcards
# and prefixes are not supported. # and prefixes are not supported.
proxy_extra_allowed_paths: str = Field(default="", env="PROXY_EXTRA_ALLOWED_PATHS") proxy_extra_allowed_paths: str = Field(default="", env="PROXY_EXTRA_ALLOWED_PATHS")
# Bound the client request body: a slow or oversized upload otherwise blocks
# the proxy before authentication and holds server resources for its duration.
request_body_timeout_seconds: float = Field(
default=30.0, gt=0, env="REQUEST_BODY_TIMEOUT_SECONDS"
)
max_request_body_bytes: int = Field(
default=20 * 1024 * 1024, gt=0, env="MAX_REQUEST_BODY_BYTES"
)
tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL") tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL")
providers_refresh_interval_seconds: int = Field( providers_refresh_interval_seconds: int = Field(
default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS" default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS"
@@ -158,9 +211,21 @@ class Settings(BaseSettings):
default=30.0, gt=0, env="DATABASE_BUSY_TIMEOUT" default=30.0, gt=0, env="DATABASE_BUSY_TIMEOUT"
) )
# Per-origin upstream connection pools. These fields are env-only below.
upstream_max_connections: int = Field(
default=200, ge=1, env="UPSTREAM_MAX_CONNECTIONS"
)
upstream_pool_timeout: float = Field(default=5.0, gt=0, env="UPSTREAM_POOL_TIMEOUT")
upstream_read_timeout: float = Field(
default=900.0, gt=0, env="UPSTREAM_READ_TIMEOUT"
)
# Logging # Logging
log_level: str = Field(default="INFO", env="LOG_LEVEL") log_level: str = Field(default="INFO", env="LOG_LEVEL")
enable_console_logging: bool = Field(default=True, env="ENABLE_CONSOLE_LOGGING") enable_console_logging: bool = Field(default=True, env="ENABLE_CONSOLE_LOGGING")
slow_request_warn_seconds: float = Field(
default=60.0, gt=0, env="SLOW_REQUEST_WARN_SECONDS"
)
# Other # Other
chat_completions_api_version: str = Field( chat_completions_api_version: str = Field(
@@ -219,6 +284,10 @@ ENV_ONLY_FIELDS = frozenset(
"database_pool_pre_ping", "database_pool_pre_ping",
"database_pool_hold_warn_seconds", "database_pool_hold_warn_seconds",
"database_busy_timeout", "database_busy_timeout",
# Reconfiguring a live pool would disrupt in-flight streams.
"upstream_max_connections",
"upstream_pool_timeout",
"upstream_read_timeout",
} }
) )
@@ -256,7 +325,7 @@ def _apply_to_live_settings(data: dict[str, Any]) -> None:
def _compute_primary_mint(cashu_mints: list[str]) -> str: def _compute_primary_mint(cashu_mints: list[str]) -> str:
return cashu_mints[0] if cashu_mints else "https://mint.minibits.cash/Bitcoin" return cashu_mints[0] if cashu_mints else DEFAULT_CASHU_MINTS[0]
def derive_npub_from_nsec(nsec: str) -> str | None: def derive_npub_from_nsec(nsec: str) -> str | None:
+12 -3
View File
@@ -443,7 +443,8 @@ async def get_invoice_status(
structured_errors: bool = Depends(_uses_v2_errors), structured_errors: bool = Depends(_uses_v2_errors),
) -> InvoiceStatusResponse: ) -> InvoiceStatusResponse:
invoice = await session.get(LightningInvoice, invoice_id) invoice = await session.get(LightningInvoice, invoice_id)
if not invoice: # Payout rows (direction="out") are operator history, never user invoices.
if not invoice or invoice.direction != "in":
raise _invoice_error( raise _invoice_error(
404, 404,
"Invoice not found", "Invoice not found",
@@ -486,7 +487,9 @@ async def recover_invoice(
structured_errors: bool = Depends(_uses_v2_errors), structured_errors: bool = Depends(_uses_v2_errors),
) -> InvoiceStatusResponse: ) -> InvoiceStatusResponse:
result = await session.exec( result = await session.exec(
select(LightningInvoice).where(LightningInvoice.bolt11 == request.bolt11) select(LightningInvoice)
.where(LightningInvoice.bolt11 == request.bolt11)
.where(col(LightningInvoice.direction) == "in")
) )
invoice = result.first() invoice = result.first()
@@ -942,6 +945,7 @@ async def _expire_overdue_invoices(now: int) -> int:
expired = await expiry_session.exec( # type: ignore[call-overload] expired = await expiry_session.exec( # type: ignore[call-overload]
update(LightningInvoice) update(LightningInvoice)
.where( .where(
col(LightningInvoice.direction) == "in",
col(LightningInvoice.status) == "pending", col(LightningInvoice.status) == "pending",
col(LightningInvoice.expires_at) < now, col(LightningInvoice.expires_at) < now,
) )
@@ -959,13 +963,17 @@ async def _process_invoice_watch_batch(session: AsyncSession, prev_now: int) ->
logger.info("Expired overdue invoices", extra={"invoice_count": swept}) logger.info("Expired overdue invoices", extra={"invoice_count": swept})
settling = await session.exec( settling = await session.exec(
select(LightningInvoice) select(LightningInvoice)
.where(col(LightningInvoice.status) == "settlement_pending") .where(
col(LightningInvoice.direction) == "in",
col(LightningInvoice.status) == "settlement_pending",
)
.order_by(col(LightningInvoice.created_at)) .order_by(col(LightningInvoice.created_at))
.limit(INVOICE_WATCH_BATCH_LIMIT // 2) .limit(INVOICE_WATCH_BATCH_LIMIT // 2)
) )
unpaid = await session.exec( unpaid = await session.exec(
select(LightningInvoice) select(LightningInvoice)
.where( .where(
col(LightningInvoice.direction) == "in",
col(LightningInvoice.status) == "pending", col(LightningInvoice.status) == "pending",
col(LightningInvoice.expires_at) >= now, col(LightningInvoice.expires_at) >= now,
) )
@@ -975,6 +983,7 @@ async def _process_invoice_watch_batch(session: AsyncSession, prev_now: int) ->
recoverable = await session.exec( recoverable = await session.exec(
select(LightningInvoice) select(LightningInvoice)
.where( .where(
col(LightningInvoice.direction) == "in",
col(LightningInvoice.status) == "expired", col(LightningInvoice.status) == "expired",
col(LightningInvoice.expires_at) > now - INVOICE_EXPIRY_GRACE_SECONDS, col(LightningInvoice.expires_at) > now - INVOICE_EXPIRY_GRACE_SECONDS,
) )
+38 -7
View File
@@ -15,6 +15,14 @@ from PIL import Image
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
from ..core import get_logger from ..core import get_logger
from ..core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_NODE,
ERROR_SCOPE_UPSTREAM,
client_code_for_upstream_error,
client_status_for_upstream_error,
upstream_status_details,
)
from ..core.exceptions import UpstreamError from ..core.exceptions import UpstreamError
from ..core.redaction import redact_org_ids from ..core.redaction import redact_org_ids
from ..core.settings import settings from ..core.settings import settings
@@ -654,13 +662,15 @@ def create_error_response(
token: str | None = None, token: str | None = None,
code: str | int | None = None, code: str | int | None = None,
details: dict[str, object] | None = None, details: dict[str, object] | None = None,
error_scope: str | None = None,
) -> Response: ) -> Response:
"""Create a standardized error response. """Create a standardized error response.
``code`` is a stable, machine-readable classification (e.g. ``code`` is a stable, machine-readable classification (e.g.
``UPSTREAM_RATE_LIMIT``); when omitted it defaults to the HTTP status code ``UPSTREAM_RATE_LIMIT``); when omitted it defaults to the HTTP status code
for backwards compatibility. ``details`` carries optional structured, for backwards compatibility. ``details`` carries optional structured,
redaction-safe context. redaction-safe context. ``error_scope`` is sent as the
:data:`ERROR_SCOPE_HEADER` response header.
""" """
error_obj: dict[str, object] = { error_obj: dict[str, object] = {
"message": redact_org_ids(message), "message": redact_org_ids(message),
@@ -669,6 +679,11 @@ def create_error_response(
} }
if details is not None: if details is not None:
error_obj["details"] = details error_obj["details"] = details
headers: dict[str, str] = {}
if token:
headers["X-Cashu"] = token
if error_scope is not None:
headers[ERROR_SCOPE_HEADER] = error_scope
return Response( return Response(
content=json.dumps( content=json.dumps(
{ {
@@ -678,7 +693,7 @@ def create_error_response(
), ),
status_code=status_code, status_code=status_code,
media_type="application/json", media_type="application/json",
headers={"X-Cashu": token} if token else {}, headers=headers,
) )
@@ -687,13 +702,29 @@ def create_upstream_error_response(
request: Request, request: Request,
fallback_status: int = 502, fallback_status: int = 502,
) -> Response: ) -> Response:
"""Build an error response from an :class:`UpstreamError`, preserving its """Build an error response from an :class:`UpstreamError`.
structured ``code``, ``details``, and original ``status_code``."""
Upstream-scoped errors are mapped via :mod:`routstr.core.error_scope`;
node-scoped errors keep their own status.
"""
status_code = error.status_code or fallback_status
code = getattr(error, "code", None)
details = getattr(error, "details", None)
if getattr(error, "scope", ERROR_SCOPE_UPSTREAM) == ERROR_SCOPE_NODE:
return create_error_response(
"upstream_error",
str(error),
status_code,
request=request,
code=code,
details=details,
)
return create_error_response( return create_error_response(
"upstream_error", "upstream_error",
str(error), str(error),
error.status_code or fallback_status, client_status_for_upstream_error(status_code, code),
request=request, request=request,
code=getattr(error, "code", None), code=client_code_for_upstream_error(status_code, code),
details=getattr(error, "details", None), details=upstream_status_details(details, status_code),
error_scope=ERROR_SCOPE_UPSTREAM,
) )
+5 -3
View File
@@ -325,7 +325,7 @@ async def raw_send_to_lnurl(
unit: str, unit: str,
amount: int | None = None, amount: int | None = None,
*, *,
on_melt_quote: Callable[[str], Awaitable[None]] | None = None, on_melt_quote: Callable[[str, str], Awaitable[None]] | None = None,
) -> int: ) -> int:
"""Send funds to an LNURL address. """Send funds to an LNURL address.
@@ -413,10 +413,12 @@ async def raw_send_to_lnurl(
raise LNURLError("Cashu melt fees exceed the requested gross amount") raise LNURLError("Cashu melt fees exceed the requested gross amount")
if on_melt_quote is not None: if on_melt_quote is not None:
await on_melt_quote(melt_quote_resp.quote) await on_melt_quote(melt_quote_resp.quote, bolt11_invoice)
assert selected_proofs is not None assert selected_proofs is not None
proofs = selected_proofs proofs = selected_proofs
# Cashu uses this argument only to size blank outputs, not set mint fees.
change_budget = sum(proof.amount for proof in proofs) - quoted_amount
await wallet.set_reserved_for_send(proofs, reserved=True) await wallet.set_reserved_for_send(proofs, reserved=True)
try: try:
@@ -424,7 +426,7 @@ async def raw_send_to_lnurl(
lambda: wallet.melt( lambda: wallet.melt(
proofs=proofs, proofs=proofs,
invoice=bolt11_invoice, invoice=bolt11_invoice,
fee_reserve_sat=melt_quote_resp.fee_reserve, fee_reserve_sat=change_budget,
quote_id=melt_quote_resp.quote, quote_id=melt_quote_resp.quote,
), ),
op_name="lnurl_melt", op_name="lnurl_melt",
+87 -50
View File
@@ -241,61 +241,98 @@ def _has_valid_pricing(model: dict) -> bool:
return True return True
async def async_fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: # OpenRouter occasionally answers /models with a truncated body, emptying the
"""Asynchronously fetch model information from OpenRouter API.""" # catalogue behind one log line. Retry, but keep 3 attempts within roughly the
# old single-attempt budget: this fetch blocks startup and the refresh loop.
OPENROUTER_MODELS_MAX_ATTEMPTS = 3
OPENROUTER_MODELS_TIMEOUT_SECONDS = 10
OPENROUTER_MODELS_RETRY_BACKOFF_SECONDS = 0.5
def _is_transient(error: BaseException) -> bool:
if isinstance(error, httpx.HTTPStatusError):
return error.response.status_code >= 500
return True
def _parse_models_response(response: httpx.Response | BaseException) -> list[dict]:
if isinstance(response, BaseException):
raise response
response.raise_for_status()
return [
model
for model in response.json().get("data", [])
if ":free" not in model.get("id", "").lower()
]
async def _fetch_openrouter_models_once(source_filter: str | None) -> list[dict]:
"""One attempt. Raises if /models is unusable; embeddings are best-effort."""
base_url = "https://openrouter.ai/api/v1" base_url = "https://openrouter.ai/api/v1"
timeout = OPENROUTER_MODELS_TIMEOUT_SECONDS
try: async with httpx.AsyncClient() as client:
async with httpx.AsyncClient() as client: models_response, embeddings_response = await asyncio.gather(
models_response, embeddings_response = await asyncio.gather( client.get(f"{base_url}/models", timeout=timeout),
client.get(f"{base_url}/models", timeout=30), client.get(f"{base_url}/embeddings/models", timeout=timeout),
client.get(f"{base_url}/embeddings/models", timeout=30), return_exceptions=True,
return_exceptions=True, )
)
def process_models_response( # Losing /models is what empties the node, so it fails the attempt and
response: httpx.Response | BaseException, # the caller retries. A missing embeddings half must not do the same.
) -> list[dict]: models_data = _parse_models_response(models_response)
if not isinstance(response, BaseException): try:
response.raise_for_status() models_data.extend(_parse_models_response(embeddings_response))
data = response.json() except Exception as e:
return [ logger.warning(f"Skipping OpenRouter embeddings models: {e}")
model
for model in data.get("data", []) # Apply source filter and exclusions
if ":free" not in model.get("id", "").lower() filtered_models = []
] for model in models_data:
model_id = model.get("id", "")
if source_filter:
source_prefix = f"{source_filter}/"
if not model_id.startswith(source_prefix):
continue
model = dict(model)
model["id"] = model_id[len(source_prefix) :]
model_id = model["id"]
if "(free)" in model.get("name", ""):
continue
if not _has_valid_pricing(model):
continue
filtered_models.append(model)
return filtered_models
async def async_fetch_openrouter_models(source_filter: str | None = None) -> list[dict]:
"""Fetch the OpenRouter catalogue; ``[]`` once every attempt has failed."""
for attempt in range(1, OPENROUTER_MODELS_MAX_ATTEMPTS + 1):
try:
return await _fetch_openrouter_models_once(source_filter)
except Exception as e:
last_attempt = attempt == OPENROUTER_MODELS_MAX_ATTEMPTS
if last_attempt or not _is_transient(e):
logger.error(
f"Error (async) fetching models from OpenRouter API "
f"after {attempt} attempt(s): {e}"
)
return [] return []
logger.warning(
f"OpenRouter models fetch attempt {attempt}/"
f"{OPENROUTER_MODELS_MAX_ATTEMPTS} failed: {e}; retrying"
)
# Jittered so nodes do not retry in lockstep.
backoff = OPENROUTER_MODELS_RETRY_BACKOFF_SECONDS * attempt
await asyncio.sleep(backoff * random.uniform(0.5, 1.5))
models_data: list[dict] = [] return []
models_data.extend(process_models_response(models_response))
models_data.extend(process_models_response(embeddings_response))
# Apply source filter and exclusions
filtered_models = []
for model in models_data:
model_id = model.get("id", "")
if source_filter:
source_prefix = f"{source_filter}/"
if not model_id.startswith(source_prefix):
continue
model = dict(model)
model["id"] = model_id[len(source_prefix) :]
model_id = model["id"]
if "(free)" in model.get("name", ""):
continue
if not _has_valid_pricing(model):
continue
filtered_models.append(model)
return filtered_models
except Exception as e:
logger.error(f"Error (async) fetching models from OpenRouter API: {e}")
return []
def _build_model_from_row( def _build_model_from_row(
+259 -41
View File
@@ -1,16 +1,16 @@
import asyncio import asyncio
import inspect import inspect
import json import json
import re
from typing import Any from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Request from fastapi import APIRouter, HTTPException, Request
from fastapi.responses import Response, StreamingResponse from fastapi.responses import Response, StreamingResponse
from sqlmodel import select from sqlmodel import select
from .algorithm import create_model_mappings from .algorithm import create_model_mappings
from .auth import ( from .auth import (
ReservationSnapshot, ReservationSnapshot,
get_reservation_snapshot,
pay_for_request, pay_for_request,
revert_pay_for_request, revert_pay_for_request,
validate_bearer_key, validate_bearer_key,
@@ -22,9 +22,15 @@ from .core.db import (
ModelRow, ModelRow,
UpstreamProviderRow, UpstreamProviderRow,
create_session, create_session,
get_session, )
from .core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS,
UPSTREAM_UNAVAILABLE,
) )
from .core.exceptions import UpstreamError from .core.exceptions import UpstreamError
from .core.middleware import mark
from .core.not_found import build_not_found_response from .core.not_found import build_not_found_response
from .core.settings import settings from .core.settings import settings
from .payment.helpers import ( from .payment.helpers import (
@@ -36,6 +42,12 @@ from .payment.helpers import (
) )
from .payment.models import Model from .payment.models import Model
from .upstream import BaseUpstreamProvider from .upstream import BaseUpstreamProvider
from .upstream.cooldown import (
candidate_model_identity,
is_cooling_down,
provider_identity,
record_failure,
)
from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request
from .upstream.helpers import init_upstreams from .upstream.helpers import init_upstreams
from .upstream.model_paths import ( from .upstream.model_paths import (
@@ -112,8 +124,6 @@ def get_candidates(
if candidates := _provider_map.get(model_id_lower): if candidates := _provider_map.get(model_id_lower):
return candidates return candidates
import re
base_model_id = re.sub(r"-\d{8}$", "", model_id_lower) base_model_id = re.sub(r"-\d{8}$", "", model_id_lower)
if base_model_id != model_id_lower: if base_model_id != model_id_lower:
if candidates := _provider_map.get(base_model_id): if candidates := _provider_map.get(base_model_id):
@@ -267,6 +277,10 @@ _ALLOWED_ENDPOINTS: dict[str, frozenset[str]] = {
"completions": frozenset({"POST"}), "completions": frozenset({"POST"}),
"responses": frozenset({"POST"}), "responses": frozenset({"POST"}),
"messages": frozenset({"POST"}), "messages": frozenset({"POST"}),
# Anthropic token-counting subroute; the proxy's allowlist is exact, so the
# "messages" entry above does not carry it. Clients (Claude Code, the
# Anthropic SDKs) call it before every request.
"messages/count_tokens": frozenset({"POST"}),
"embeddings": frozenset({"POST"}), "embeddings": frozenset({"POST"}),
# TypeSafe System One decision endpoint: POST {state, model, questions} # TypeSafe System One decision endpoint: POST {state, model, questions}
# -> {answers, usage}. Non-streaming, JSON in/out; billed from the # -> {answers, usage}. Non-streaming, JSON in/out; billed from the
@@ -395,23 +409,109 @@ def _forwarding_allowed(path: str, method: str) -> bool:
return method in _allowed_methods_for(_canonical_api_path(path)) return method in _allowed_methods_for(_canonical_api_path(path))
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) # Gateway conditions a retry usually clears. 500 is excluded: as likely to be a
async def proxy( # deterministic rejection that fails identically on the next attempt.
request: Request, path: str, session: AsyncSession = Depends(get_session) _RETRYABLE_UPSTREAM_5XX = frozenset({502, 503, 504})
) -> Response | StreamingResponse: _UPSTREAM_5XX_RETRY_BACKOFF_SECONDS = 0.5
"""Run proxy setup in a short request session, never across response streaming."""
def _counts_toward_cooldown(status_code: int) -> bool:
"""Provider faults and timeouts only — not client errors or rate limits."""
return status_code >= 500 or status_code == UPSTREAM_ERROR_STATUS
def _upstream_response_failure(response: Response) -> bool:
return (
_counts_toward_cooldown(response.status_code)
and response.headers.get(ERROR_SCOPE_HEADER) == ERROR_SCOPE_UPSTREAM
)
def _attribute_request(
request: Request, model_obj: Model, upstream: BaseUpstreamProvider
) -> None:
"""Attribute the completion log line to the candidate being tried.
Uses the provider's model id rather than the requested alias, so aliases
and cross-provider spellings resolve to the model that was forwarded.
"""
if model_obj.id:
request.state.model = model_obj.id
request.state.provider = upstream.provider_type
class _BodyLimitExceeded(Exception):
"""The client body is larger than ``max_request_body_bytes``."""
async def _read_bounded_body(request: Request) -> bytes | Response:
"""Read the request body under a size and time bound.
Returns the body, or the error response to send instead. Both bounds run
before any authentication or DB work, so an oversized or slowly uploaded
body cannot occupy the request for longer than the timeout.
"""
max_bytes = settings.max_request_body_bytes
timeout = settings.request_body_timeout_seconds
async def read() -> bytes:
declared = request.headers.get("content-length", "")
if declared.isdigit() and int(declared) > max_bytes:
raise _BodyLimitExceeded
body = bytearray()
async for chunk in request.stream():
body += chunk
# Chunked uploads declare no length, so the cap is enforced here.
if len(body) > max_bytes:
raise _BodyLimitExceeded
return bytes(body)
try: try:
return await _proxy(request, path, session) body = await asyncio.wait_for(read(), timeout)
finally: except _BodyLimitExceeded:
# FastAPI yield dependencies normally close after the response body is error_type, message, status = (
# sent. Close explicitly so a long stream cannot retain DB resources. "invalid_request",
close_result = session.close() f"Request body exceeds the {max_bytes} byte limit",
if inspect.isawaitable(close_result): 413,
await close_result )
except asyncio.TimeoutError:
error_type, message, status = (
"timeout",
f"Request body not received within {timeout} seconds",
408,
)
else:
# Draining the stream leaves Starlette unable to serve a second read.
# Cache the body so later readers (EHBP forwarding, upstream stream
# passthrough) get it instead of "Stream consumed".
request._body = body
return body
return create_error_response(error_type, message, status, request=request)
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
async def proxy(request: Request, path: str) -> Response | StreamingResponse:
"""Run proxy setup in a short request session, never across response streaming."""
# Read the body before opening a session: a slow uploader must not hold a
# DB connection while its request trickles in.
request_body = await _read_bounded_body(request)
if isinstance(request_body, Response):
return request_body
mark(request, "body_read")
async with create_session() as session:
try:
return await _proxy(request, path, session, request_body)
finally:
# Close explicitly so a long stream cannot retain DB resources
# while its response body is being sent.
close_result = session.close()
if inspect.isawaitable(close_result):
await close_result
async def _proxy( async def _proxy(
request: Request, path: str, session: AsyncSession request: Request, path: str, session: AsyncSession, request_body: bytes
) -> Response | StreamingResponse: ) -> Response | StreamingResponse:
# Screen the path before any routing decision: reject ambiguous spellings, # Screen the path before any routing decision: reject ambiguous spellings,
# then require a known API prefix so nothing unknown is forwarded with the # then require a known API prefix so nothing unknown is forwarded with the
@@ -426,7 +526,6 @@ async def _proxy(
return build_not_found_response(request, path) return build_not_found_response(request, path)
is_responses_api = path.startswith("v1/responses") or path.startswith("responses") is_responses_api = path.startswith("v1/responses") or path.startswith("responses")
request_body = await request.body()
# EHBP (Encrypted HTTP Body Protocol) requests carry an Ehbp-Encapsulated-Key # 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 # header and a binary HPKE-sealed body. The proxy cannot parse the body to
@@ -450,6 +549,12 @@ async def _proxy(
else: else:
model_id = request_body_dict.get("model", "unknown") model_id = request_body_dict.get("model", "unknown")
# Set before routing so the completion log is attributed even when the
# request fails before an upstream is chosen (400/401/402). "unknown" is
# the no-model sentinel, not a model.
if isinstance(model_id, str) and model_id and model_id != "unknown":
request.state.model = model_id
# Exact Tinfoil attestation GET routes don't map to models — forward # Exact Tinfoil attestation GET routes don't map to models — forward
# without model/cost/auth lookups. Do not prefix-match here: paths such as # without model/cost/auth lookups. Do not prefix-match here: paths such as
# /attestationjunk must continue through normal authentication. # /attestationjunk must continue through normal authentication.
@@ -472,11 +577,12 @@ async def _proxy(
last_error_response = None last_error_response = None
for i, upstream in enumerate(selected_upstreams): for i, upstream in enumerate(selected_upstreams):
request.state.provider = upstream.provider_type
try: try:
headers = upstream.prepare_headers(dict(request.headers)) headers = upstream.prepare_headers(dict(request.headers))
response = await upstream.forward_get_request(request, path, headers) response = await upstream.forward_get_request(request, path, headers)
if ( if (
response.status_code in [502, 429] response.status_code in [424, 502, 503, 429]
and i < len(selected_upstreams) - 1 and i < len(selected_upstreams) - 1
): ):
logger.warning( logger.warning(
@@ -498,7 +604,12 @@ async def _proxy(
last_error_response = create_upstream_error_response(e, request) last_error_response = create_upstream_error_response(e, request)
continue continue
return last_error_response or create_error_response( return last_error_response or create_error_response(
"upstream_error", "All upstreams failed", 502, request=request "upstream_error",
"All upstreams failed",
UPSTREAM_ERROR_STATUS,
request=request,
code=UPSTREAM_UNAVAILABLE,
error_scope=ERROR_SCOPE_UPSTREAM,
) )
selector: ModelPathSelector | None = None selector: ModelPathSelector | None = None
@@ -606,6 +717,20 @@ async def _proxy(
request=request, request=request,
) )
# A provider that just failed this model repeatedly is skipped while some
# other candidate can serve it. An explicit route is never rerouted.
if selector is None:
healthy = [
candidate
for candidate in candidates
if not is_cooling_down(
provider_identity(candidate[1]),
candidate_model_identity(candidate[0], model_id),
)
]
if healthy:
candidates = healthy
# Reserve/max-cost checks use the best-ranked candidate; the failover loop # Reserve/max-cost checks use the best-ranked candidate; the failover loop
# below rebinds (model_obj, upstream) per candidate so forwarding and # below rebinds (model_obj, upstream) per candidate so forwarding and
# settlement always use the model of the provider actually being tried. # settlement always use the model of the provider actually being tried.
@@ -623,6 +748,7 @@ async def _proxy(
if x_cashu := headers.get("x-cashu", None): if x_cashu := headers.get("x-cashu", None):
last_error = None last_error = None
for i, (model_obj, upstream) in enumerate(candidates): for i, (model_obj, upstream) in enumerate(candidates):
_attribute_request(request, model_obj, upstream)
try: try:
if is_ehbp: if is_ehbp:
if not upstream.supports_ehbp: if not upstream.supports_ehbp:
@@ -632,7 +758,7 @@ async def _proxy(
model_id, model_id,
) )
continue continue
return await forward_ehbp_x_cashu_request( response = await forward_ehbp_x_cashu_request(
request=request, request=request,
x_cashu_token=x_cashu, x_cashu_token=x_cashu,
path=path, path=path,
@@ -641,7 +767,7 @@ async def _proxy(
upstream=upstream, upstream=upstream,
) )
elif is_responses_api: elif is_responses_api:
return await upstream.handle_x_cashu_responses( response = await upstream.handle_x_cashu_responses(
request, request,
x_cashu, x_cashu,
path, path,
@@ -650,7 +776,7 @@ async def _proxy(
request_body=request_body, request_body=request_body,
) )
else: else:
return await upstream.handle_x_cashu( response = await upstream.handle_x_cashu(
request, request,
x_cashu, x_cashu,
path, path,
@@ -658,6 +784,12 @@ async def _proxy(
model_obj, model_obj,
request_body=request_body, request_body=request_body,
) )
if _upstream_response_failure(response):
record_failure(
provider_identity(upstream),
candidate_model_identity(model_obj, model_id),
)
return response
except UpstreamError as e: except UpstreamError as e:
logger.warning( logger.warning(
"Upstream %s failed (x-cashu) for model=%s: %s", "Upstream %s failed (x-cashu) for model=%s: %s",
@@ -670,6 +802,13 @@ async def _proxy(
"status_code": e.status_code, "status_code": e.status_code,
}, },
) )
if e.scope == ERROR_SCOPE_UPSTREAM and _counts_toward_cooldown(
e.status_code
):
record_failure(
provider_identity(upstream),
candidate_model_identity(model_obj, model_id),
)
if i == len(candidates) - 1: if i == len(candidates) - 1:
last_error = e last_error = e
continue continue
@@ -677,13 +816,19 @@ async def _proxy(
if last_error is not None: if last_error is not None:
return create_upstream_error_response(last_error, request) return create_upstream_error_response(last_error, request)
return create_error_response( return create_error_response(
"upstream_error", "All upstreams failed", 502, request=request "upstream_error",
"All upstreams failed",
UPSTREAM_ERROR_STATUS,
request=request,
code=UPSTREAM_UNAVAILABLE,
error_scope=ERROR_SCOPE_UPSTREAM,
) )
elif auth := headers.get("authorization", None): elif auth := headers.get("authorization", None):
key = await get_bearer_token_key( key = await get_bearer_token_key(
headers, path, session, auth, max_cost_for_model, model_id headers, path, session, auth, max_cost_for_model, model_id
) )
mark(request, "auth")
else: else:
if request.method not in ["GET"]: if request.method not in ["GET"]:
@@ -697,12 +842,16 @@ async def _proxy(
logger.debug("Processing unauthenticated GET request", extra={"path": path}) logger.debug("Processing unauthenticated GET request", extra={"path": path})
last_error_response = None last_error_response = None
for i, (_, upstream) in enumerate(candidates): for i, (model_obj, upstream) in enumerate(candidates):
_attribute_request(request, model_obj, upstream)
try: try:
headers = upstream.prepare_headers(dict(request.headers)) headers = upstream.prepare_headers(dict(request.headers))
response = await upstream.forward_get_request(request, path, 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 [424, 502, 503, 429]
and i < len(candidates) - 1
):
error_message = "" error_message = ""
try: try:
if hasattr(response, "body"): if hasattr(response, "body"):
@@ -736,14 +885,18 @@ async def _proxy(
last_error_response = create_upstream_error_response(e, request) last_error_response = create_upstream_error_response(e, request)
continue continue
return last_error_response or create_error_response( return last_error_response or create_error_response(
"upstream_error", "All upstreams failed", 502, request=request "upstream_error",
"All upstreams failed",
UPSTREAM_ERROR_STATUS,
request=request,
code=UPSTREAM_UNAVAILABLE,
error_scope=ERROR_SCOPE_UPSTREAM,
) )
reservation_snapshot: ReservationSnapshot | None = None reservation_snapshot: ReservationSnapshot | None = None
if is_ehbp or request_body_dict: if is_ehbp or request_body_dict:
await pay_for_request(key, max_cost_for_model, session) reservation_snapshot = await pay_for_request(key, max_cost_for_model, session)
reservation_snapshot = await get_reservation_snapshot(key, session) # pay_for_request refreshes the key after committing the reservation.
# Snapshot validation performs SELECTs after pay_for_request commits.
# End that read transaction before waiting on upstream response headers. # End that read transaction before waiting on upstream response headers.
await _finish_read_transaction(session) await _finish_read_transaction(session)
@@ -770,18 +923,25 @@ async def _proxy(
key, session, max_cost_for_model, reservation_snapshot key, session, max_cost_for_model, reservation_snapshot
) )
try: try:
await pay_for_request(key, candidate_max, session) reservation_snapshot = await pay_for_request(
key, candidate_max, session
)
except HTTPException: except HTTPException:
if i == len(candidates) - 1: if i == len(candidates) - 1:
raise raise
await pay_for_request(key, max_cost_for_model, session) reservation_snapshot = await pay_for_request(
reservation_snapshot = await get_reservation_snapshot(key, session) key, max_cost_for_model, session
)
await _finish_read_transaction(session) await _finish_read_transaction(session)
continue continue
reservation_snapshot = await get_reservation_snapshot(key, session)
await _finish_read_transaction(session) await _finish_read_transaction(session)
max_cost_for_model = candidate_max max_cost_for_model = candidate_max
# Only once the candidate is actually tried: a fallback skipped for its
# reservation must not take over the last attempted upstream's line.
_attribute_request(request, model_obj, upstream)
retries_left = settings.upstream_5xx_retry_attempts
retry_index = 0
headers = upstream.prepare_headers(dict(request.headers)) headers = upstream.prepare_headers(dict(request.headers))
try: try:
@@ -834,8 +994,39 @@ async def _proxy(
model_obj, model_obj,
reservation_snapshot, reservation_snapshot,
) )
except UpstreamError: except UpstreamError as e:
# Let the outer UpstreamError handler manage retry/revert # Only a gateway status the upstream itself answered with:
# re-sending the buffered body cannot double-bill. A 502 this
# proxy invented for a transport error or timeout is not
# retried — that request may already be running upstream.
if (
e.from_upstream_response
and e.status_code in _RETRYABLE_UPSTREAM_5XX
and retries_left > 0
):
retries_left -= 1
retry_index += 1
logger.warning(
"Upstream %s returned %s for model=%s; retrying same "
"upstream (attempt %s, %s retries left)",
upstream.provider_type,
e.status_code,
model_id,
retry_index + 1,
retries_left,
extra={
"provider": upstream.provider_type,
"model": model_id,
"status_code": e.status_code,
"path": path,
"retries_left": retries_left,
},
)
await asyncio.sleep(
_UPSTREAM_5XX_RETRY_BACKOFF_SECONDS * retry_index
)
continue
# Let the outer UpstreamError handler manage failover/revert
raise raise
except Exception as e: except Exception as e:
# Unexpected error (not an upstream failure) — revert and propagate # Unexpected error (not an upstream failure) — revert and propagate
@@ -873,7 +1064,7 @@ async def _proxy(
already_stripped.add(bad_param) already_stripped.add(bad_param)
logger.warning( logger.warning(
"Upstream %s rejected param '%s' for model=%s; " "Upstream %s rejected param '%s' for model=%s; "
"stripping and retrying same upstream", "correcting and retrying same upstream",
upstream.provider_type, upstream.provider_type,
bad_param, bad_param,
model_id, model_id,
@@ -888,8 +1079,23 @@ async def _proxy(
break break
if response.status_code != 200: if response.status_code != 200:
# Check if we should retry (502 Upstream Error or 429 Rate Limit) if _upstream_response_failure(response):
should_retry = response.status_code in [502, 429, 400, 401, 403, 404] record_failure(
provider_identity(upstream),
candidate_model_identity(model_obj, model_id),
)
# 424 is an upstream failure re-reported by error_scope.
# 502/503 are upstream errors, 429 rate limits.
should_retry = response.status_code in [
424,
502,
503,
429,
400,
401,
403,
404,
]
if should_retry and i < len(candidates) - 1: if should_retry and i < len(candidates) - 1:
error_message = "" error_message = ""
try: try:
@@ -965,6 +1171,13 @@ async def _proxy(
raise raise
except UpstreamError as e: except UpstreamError as e:
if e.scope == ERROR_SCOPE_UPSTREAM and _counts_toward_cooldown(
e.status_code
):
record_failure(
provider_identity(upstream),
candidate_model_identity(model_obj, model_id),
)
logger.warning( logger.warning(
"Upstream %s failed for model=%s: %s", "Upstream %s failed for model=%s: %s",
upstream.provider_type, upstream.provider_type,
@@ -990,7 +1203,12 @@ async def _proxy(
# Should not be reached given logic above # Should not be reached given logic above
return create_error_response( return create_error_response(
"upstream_error", "All upstreams failed", 502, request=request "upstream_error",
"All upstreams failed",
UPSTREAM_ERROR_STATUS,
request=request,
code=UPSTREAM_UNAVAILABLE,
error_scope=ERROR_SCOPE_UPSTREAM,
) )
+4
View File
@@ -1,6 +1,7 @@
from .anthropic import AnthropicUpstreamProvider from .anthropic import AnthropicUpstreamProvider
from .azure import AzureUpstreamProvider from .azure import AzureUpstreamProvider
from .base import BaseUpstreamProvider from .base import BaseUpstreamProvider
from .deepseek import DeepSeekUpstreamProvider
from .fireworks import FireworksUpstreamProvider from .fireworks import FireworksUpstreamProvider
from .gemini import GeminiUpstreamProvider from .gemini import GeminiUpstreamProvider
from .generic import GenericUpstreamProvider from .generic import GenericUpstreamProvider
@@ -13,11 +14,13 @@ from .ppqai import PPQAIUpstreamProvider
from .routstr import RoutstrUpstreamProvider from .routstr import RoutstrUpstreamProvider
from .tinfoil import TinfoilUpstreamProvider from .tinfoil import TinfoilUpstreamProvider
from .typesafe import TypeSafeUpstreamProvider from .typesafe import TypeSafeUpstreamProvider
from .venice import VeniceUpstreamProvider
from .xai import XAIUpstreamProvider from .xai import XAIUpstreamProvider
upstream_provider_classes: list[type[BaseUpstreamProvider]] = [ upstream_provider_classes: list[type[BaseUpstreamProvider]] = [
AnthropicUpstreamProvider, AnthropicUpstreamProvider,
AzureUpstreamProvider, AzureUpstreamProvider,
DeepSeekUpstreamProvider,
FireworksUpstreamProvider, FireworksUpstreamProvider,
GeminiUpstreamProvider, GeminiUpstreamProvider,
GenericUpstreamProvider, GenericUpstreamProvider,
@@ -30,6 +33,7 @@ upstream_provider_classes: list[type[BaseUpstreamProvider]] = [
RoutstrUpstreamProvider, RoutstrUpstreamProvider,
TinfoilUpstreamProvider, TinfoilUpstreamProvider,
TypeSafeUpstreamProvider, TypeSafeUpstreamProvider,
VeniceUpstreamProvider,
XAIUpstreamProvider, XAIUpstreamProvider,
] ]
"""List of all upstream classes""" """List of all upstream classes"""
+976 -705
View File
File diff suppressed because it is too large Load Diff
+79
View File
@@ -0,0 +1,79 @@
"""In-memory circuit breaker for a failing (provider, model) pair.
Process-local by design: each node observes its own upstream failures, and a
cooldown that outlives a restart would hide a provider that has recovered.
"""
from __future__ import annotations
import time
from typing import Any
from ..core import get_logger
from ..core.settings import settings
logger = get_logger(__name__)
_FAILURE_WINDOW_SECONDS = 60.0
_failures: dict[tuple[str, str], list[float]] = {}
_cooling_until: dict[tuple[str, str], float] = {}
def provider_identity(upstream: Any) -> str:
db_id = getattr(upstream, "db_id", None)
if isinstance(db_id, int):
return f"db:{db_id}"
return f"{upstream.provider_type.lower()}|{upstream.base_url.lower()}"
def model_identity(model_id: str) -> str:
return model_id.lower()
def candidate_model_identity(model: Any, requested_model_id: str) -> str:
model_id = getattr(model, "id", None)
return model_identity(
model_id if isinstance(model_id, str) and model_id else requested_model_id
)
def record_failure(provider_id: str, model_id: str) -> None:
"""Count a timeout or 5xx, opening a cooldown once too many land in a minute."""
if settings.upstream_cooldown_seconds <= 0:
return
pair = (provider_id, model_id)
now = time.monotonic()
recent = [t for t in _failures.get(pair, []) if now - t < _FAILURE_WINDOW_SECONDS]
recent.append(now)
if len(recent) >= settings.upstream_allowed_fails:
_failures.pop(pair, None)
_cooling_until[pair] = now + settings.upstream_cooldown_seconds
logger.warning(
"Upstream cooling down after repeated failures",
extra={
"provider": provider_id,
"model": model_id,
"cooldown_seconds": settings.upstream_cooldown_seconds,
},
)
else:
_failures[pair] = recent
def is_cooling_down(provider_id: str, model_id: str) -> bool:
pair = (provider_id, model_id)
until = _cooling_until.get(pair)
if until is None:
return False
if time.monotonic() >= until:
del _cooling_until[pair]
return False
return True
def reset_cooldowns() -> None:
_failures.clear()
_cooling_until.clear()
+106
View File
@@ -0,0 +1,106 @@
"""First-class upstream for the DeepSeek API.
Pricing comes from ``_PEAK_RATES`` below, not from litellm or OpenRouter:
litellm's bundled ``deepseek-v4-flash`` entry is stale (input, output and cache
rates alike), the OpenRouter feed carries resale prices below DeepSeek's own
peak rate, and neither the bundled map nor OpenRouter knows the current
``deepseek-flash`` id. A model DeepSeek lists that the table does not
cover is imported disabled rather than priced from those sources.
DeepSeek bills peak hours at twice the off-peak rate. The node has one flat
price per model, so the table holds the PEAK rates: a client may overpay
off-peak but the node never bills below its own cost.
Rates: https://api-docs.deepseek.com/quick_start/pricing (checked 2026-09-30).
"""
from __future__ import annotations
from typing import TYPE_CHECKING
from .base import BaseUpstreamProvider
from .generic import GenericUpstreamProvider
from .pricing_resolver import ResolvedPricing
if TYPE_CHECKING:
from ..core.db import UpstreamProviderRow
_CONTEXT_LENGTH = 1_000_000
_MAX_OUTPUT_TOKENS = 384_000
# USD per 1M tokens at DeepSeek's peak rate: (input cache miss, output, input
# cache hit). DeepSeek has no cache-write charge.
_FLASH = (0.30, 1.20, 0.006)
_PRO = (1.32, 3.96, 0.044)
_PEAK_RATES: dict[str, tuple[float, float, float]] = {
"deepseek-flash": _FLASH,
# Retired ids DeepSeek still accepts, served and billed as deepseek-flash.
"deepseek-v4-flash": _FLASH,
"deepseek-v4-flash-vision-exp": _FLASH,
"deepseek-v4-pro": _PRO,
}
# Pro is the only current model without vision support.
_TEXT_ONLY = {"deepseek-v4-pro"}
class DeepSeekUpstreamProvider(GenericUpstreamProvider):
"""Upstream provider specifically configured for the DeepSeek API."""
provider_type = "deepseek"
default_base_url = "https://api.deepseek.com"
platform_url = "https://platform.deepseek.com/api_keys"
litellm_provider_prefix = "deepseek/"
use_fallback_pricing = False
def __init__(self, api_key: str, provider_fee: float = 1.01):
super().__init__(
base_url=self.default_base_url,
api_key=api_key,
provider_fee=provider_fee,
upstream_name="DeepSeek",
)
@classmethod
def _build_from_row(
cls, provider_row: "UpstreamProviderRow"
) -> "DeepSeekUpstreamProvider":
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": "DeepSeek",
"default_base_url": cls.default_base_url,
"fixed_base_url": True,
"platform_url": cls.platform_url,
}
def _apply_provider_field(self, response_json: object) -> None:
# A first-party upstream: stamp "deepseek", not Generic's hostname.
BaseUpstreamProvider._apply_provider_field(self, response_json)
def transform_model_name(self, model_id: str) -> str:
"""Strip the 'deepseek/' prefix for DeepSeek API compatibility."""
return model_id.removeprefix("deepseek/")
def _native_pricing(
self, model_id: str, model_spec: dict
) -> ResolvedPricing | None:
"""Price ``model_id`` from the peak-rate table; ``None`` if absent."""
rates = _PEAK_RATES.get(model_id)
if rates is None:
return None
input_usd, output_usd, cache_hit_usd = rates
input_modalities = ["text"] if model_id in _TEXT_ONLY else ["text", "image"]
return ResolvedPricing(
prompt=input_usd / 1_000_000,
completion=output_usd / 1_000_000,
context_length=_CONTEXT_LENGTH,
source="native",
max_completion_tokens=_MAX_OUTPUT_TOKENS,
input_cache_read=cache_hit_usd / 1_000_000,
input_modalities=input_modalities,
)
@@ -1,73 +0,0 @@
"""TEMPORARY: local DeepSeek V4 pricing shim.
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).
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
https://github.com/BerriAI/litellm/pull/26380 (issue
https://github.com/BerriAI/litellm/issues/30430).
=== REMOVAL (once litellm ships these models) ===
Delete this file and the single ``register_deepseek_v4_pricing()`` call in
``routstr/core/main.py``. Nothing else depends on it. Entries are only added
when absent, so a stale shim is harmless after upstream lands — but remove it.
"""
import litellm
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-*"]``.
_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_creation_input_token_cost": 0.0,
"input_cost_per_token_cache_hit": 2.8e-09,
},
"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,
"cache_creation_input_token_cost": 0.0,
"input_cost_per_token_cache_hit": 3.625e-09,
},
}
def register_deepseek_v4_pricing() -> None:
"""Inject DeepSeek V4 pricing into ``litellm.model_cost`` if absent.
Idempotent and non-destructive: a key already present in the cost map
(e.g. once litellm ships it) is left untouched. Registers both the bare
(``deepseek-v4-flash``) and prefixed (``deepseek/deepseek-v4-flash``)
spellings since ``backfill_cache_pricing`` tries both.
"""
added = []
for bare, rates in _DEEPSEEK_V4_RATES.items():
for key in (bare, f"deepseek/{bare}"):
if key in litellm.model_cost:
continue
entry: dict[str, object] = dict(rates)
entry["litellm_provider"] = "deepseek"
entry["mode"] = "chat"
litellm.model_cost[key] = entry
added.append(key)
if added:
logger.info(
"Registered temporary DeepSeek V4 pricing shim",
extra={"models": added},
)
+85 -21
View File
@@ -5,7 +5,7 @@ import math
import time import time
import traceback import traceback
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import AsyncIterator, Mapping from typing import AsyncIterator, Awaitable, Mapping
from urllib.parse import urlsplit, urlunsplit from urllib.parse import urlsplit, urlunsplit
from fastapi import Request from fastapi import Request
@@ -31,6 +31,14 @@ from ..core.db import (
from ..core.db import ( from ..core.db import (
store_cashu_transaction_with_retry as store_cashu_transaction, store_cashu_transaction_with_retry as store_cashu_transaction,
) )
from ..core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_NODE,
ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS,
client_code_for_upstream_error,
client_status_for_upstream_error,
)
from ..core.exceptions import EhbpTimeoutError, UpstreamError from ..core.exceptions import EhbpTimeoutError, UpstreamError
from ..core.settings import settings from ..core.settings import settings
from ..payment.cost_calculation import ( from ..payment.cost_calculation import (
@@ -641,6 +649,37 @@ async def _release_failed_ehbp_charge(
) )
async def _record_ehbp_settlement(
operation: Awaitable[int],
*,
key: ApiKey,
model_id: str,
settlement_type: str,
) -> int:
"""Expose EHBP settlement latency alongside normal request settlement."""
started = time.perf_counter()
# A rollback can expire the ORM instance, so capture this before the operation.
key_log_hash = key.hashed_key[:8] + "..."
succeeded = False
try:
result = await operation
succeeded = True
return result
finally:
logger.info(
"Payment settlement finished",
extra={
"key_hash": key_log_hash,
"model": model_id,
"settlement_type": settlement_type,
"settlement_duration_ms": round(
(time.perf_counter() - started) * 1000, 2
),
"settlement_succeeded": succeeded,
},
)
async def finalize_ehbp_actual_cost_payment( async def finalize_ehbp_actual_cost_payment(
key: ApiKey, key: ApiKey,
session: AsyncSession, session: AsyncSession,
@@ -879,6 +918,7 @@ async def forward_ehbp_request(
f"EHBP upstream {provider_type} returned {resp.status_code} " f"EHBP upstream {provider_type} returned {resp.status_code} "
f"for model {model_obj.id}: {body_preview[:200] or '<empty>'}", f"for model {model_obj.id}: {body_preview[:200] or '<empty>'}",
status_code=resp.status_code, status_code=resp.status_code,
from_upstream_response=True,
) )
# Check for usage metrics in response headers (non-streaming) or # Check for usage metrics in response headers (non-streaming) or
@@ -928,13 +968,18 @@ async def forward_ehbp_request(
) )
billing_model = cost_info.pop("actual_model", None) or model_obj.id billing_model = cost_info.pop("actual_model", None) or model_obj.id
computed_msats = int(cost_info["total_msats"]) computed_msats = int(cost_info["total_msats"])
charged_msats = await finalize_ehbp_actual_cost_payment( charged_msats = await _record_ehbp_settlement(
key, finalize_ehbp_actual_cost_payment(
session, key,
max_cost_for_model, session,
billing_model, max_cost_for_model,
cost_info, billing_model,
reservation_snapshot, cost_info,
reservation_snapshot,
),
key=key,
model_id=billing_model,
settlement_type="ehbp_usage",
) )
cost_data = { cost_data = {
**cost_info, **cost_info,
@@ -954,12 +999,17 @@ async def forward_ehbp_request(
"key_hash": key.hashed_key[:8] + "...", "key_hash": key.hashed_key[:8] + "...",
}, },
) )
charged_msats = await finalize_ehbp_max_cost_payment( charged_msats = await _record_ehbp_settlement(
key, finalize_ehbp_max_cost_payment(
session, key,
max_cost_for_model, session,
model_obj.id, max_cost_for_model,
reservation_snapshot, model_obj.id,
reservation_snapshot,
),
key=key,
model_id=model_obj.id,
settlement_type="ehbp_unmeasured_release",
) )
cost_data = { cost_data = {
"total_msats": charged_msats, "total_msats": charged_msats,
@@ -1031,7 +1081,11 @@ async def forward_ehbp_request(
"traceback": tb, "traceback": tb,
}, },
) )
raise UpstreamError("An unexpected server error occurred", status_code=500) raise UpstreamError(
"An unexpected server error occurred",
status_code=500,
scope=ERROR_SCOPE_NODE,
)
async def forward_ehbp_x_cashu_request( async def forward_ehbp_x_cashu_request(
@@ -1127,15 +1181,21 @@ async def forward_ehbp_x_cashu_request(
"error": { "error": {
"message": "Error forwarding EHBP request to upstream", "message": "Error forwarding EHBP request to upstream",
"type": "upstream_error", "type": "upstream_error",
"code": resp.status_code, # Pass the status as the code so a provider 4xx
# keeps the legacy numeric ``code``.
"code": client_code_for_upstream_error(
resp.status_code, resp.status_code
),
"upstream_status": resp.status_code,
"refund_token": refund_token, "refund_token": refund_token,
} }
} }
), ),
status_code=resp.status_code, status_code=client_status_for_upstream_error(resp.status_code),
media_type="application/json", media_type="application/json",
) )
error_response.headers["X-Cashu"] = refund_token error_response.headers["X-Cashu"] = refund_token
error_response.headers[ERROR_SCOPE_HEADER] = ERROR_SCOPE_UPSTREAM
return error_response return error_response
# Compute refund from actual usage when available — check both # Compute refund from actual usage when available — check both
@@ -1241,9 +1301,10 @@ async def forward_ehbp_x_cashu_request(
error_response = create_error_response( error_response = create_error_response(
"upstream_timeout", "upstream_timeout",
str(e), str(e),
504, UPSTREAM_ERROR_STATUS,
request=request, request=request,
code="UPSTREAM_TIMEOUT", code="UPSTREAM_TIMEOUT",
error_scope=ERROR_SCOPE_UPSTREAM,
) )
error_response.headers["X-Cashu"] = refund_token error_response.headers["X-Cashu"] = refund_token
return error_response return error_response
@@ -1259,9 +1320,10 @@ async def forward_ehbp_x_cashu_request(
return create_error_response( return create_error_response(
"upstream_timeout", "upstream_timeout",
str(e), str(e),
504, UPSTREAM_ERROR_STATUS,
request=request, request=request,
code="UPSTREAM_TIMEOUT", code="UPSTREAM_TIMEOUT",
error_scope=ERROR_SCOPE_UPSTREAM,
) )
except Exception as e: except Exception as e:
@@ -1283,8 +1345,9 @@ async def forward_ehbp_x_cashu_request(
error_response = create_error_response( error_response = create_error_response(
"upstream_error", "upstream_error",
"EHBP request failed after token redemption; refunded token", "EHBP request failed after token redemption; refunded token",
502, UPSTREAM_ERROR_STATUS,
request=request, request=request,
error_scope=ERROR_SCOPE_UPSTREAM,
) )
error_response.headers["X-Cashu"] = refund_token error_response.headers["X-Cashu"] = refund_token
return error_response return error_response
@@ -1351,7 +1414,8 @@ async def forward_ehbp_x_cashu_request(
return create_error_response( return create_error_response(
"cashu_error" if not redeemed else "upstream_error", "cashu_error" if not redeemed else "upstream_error",
f"EHBP X-Cashu request failed: {error_message}", f"EHBP X-Cashu request failed: {error_message}",
400 if not redeemed else 502, 400 if not redeemed else UPSTREAM_ERROR_STATUS,
request=request, request=request,
token=x_cashu_token if not redeemed else None, token=x_cashu_token if not redeemed else None,
error_scope=None if not redeemed else ERROR_SCOPE_UPSTREAM,
) )
+76 -16
View File
@@ -44,6 +44,7 @@ Pipeline
from __future__ import annotations from __future__ import annotations
import asyncio
import json import json
import uuid import uuid
from collections.abc import AsyncGenerator, AsyncIterator from collections.abc import AsyncGenerator, AsyncIterator
@@ -52,8 +53,10 @@ from typing import Any, Callable
import httpx import httpx
from ..core import get_logger from ..core import get_logger
from ..core.error_scope import ERROR_SCOPE_NODE
from ..core.exceptions import UpstreamError from ..core.exceptions import UpstreamError
from ..payment.models import Model from ..payment.models import Model
from .http_client import acquire_upstream_http_client
from .messages_dispatch import ( from .messages_dispatch import (
ANTHROPIC_ONLY_FIELDS, ANTHROPIC_ONLY_FIELDS,
aggregate_anthropic_events_to_message, aggregate_anthropic_events_to_message,
@@ -63,6 +66,46 @@ logger = get_logger(__name__)
DUMMY_THOUGHT_SIGNATURE = "skip_thought_signature_validator" DUMMY_THOUGHT_SIGNATURE = "skip_thought_signature_validator"
class _ResponseOwnedIterator:
"""Close the upstream response even if iteration never starts."""
def __init__(
self, iterator: AsyncIterator[bytes], response: httpx.Response
) -> None:
self._iterator = iterator
self._response = response
self._cleanup_task: asyncio.Task[None] | None = None
def __aiter__(self) -> _ResponseOwnedIterator:
return self
async def __anext__(self) -> bytes:
try:
return await self._iterator.__anext__()
except StopAsyncIteration:
await self.aclose()
raise
except BaseException:
try:
await self.aclose()
finally:
raise
async def _cleanup(self) -> None:
try:
close = getattr(self._iterator, "aclose", None)
if close is not None:
await close()
finally:
await self._response.aclose()
async def aclose(self) -> None:
if self._cleanup_task is None:
self._cleanup_task = asyncio.create_task(self._cleanup())
await asyncio.shield(self._cleanup_task)
# Mapping: OpenAI finish_reason → Anthropic stop_reason # Mapping: OpenAI finish_reason → Anthropic stop_reason
_FINISH_TO_STOP = { _FINISH_TO_STOP = {
"stop": "end_turn", "stop": "end_turn",
@@ -112,6 +155,7 @@ def _translate_anthropic_to_openai(body: dict, model: str) -> dict:
raise UpstreamError( raise UpstreamError(
"Failed to translate Anthropic body to OpenAI format", "Failed to translate Anthropic body to OpenAI format",
status_code=500, status_code=500,
scope=ERROR_SCOPE_NODE,
) )
return dict(translated) return dict(translated)
@@ -297,17 +341,21 @@ async def _openai_chunks_to_anthropic_events(
yield _sse_event("message_stop", {"type": "message_stop"}) yield _sse_event("message_stop", {"type": "message_stop"})
GEMINI_STREAM_READ_TIMEOUT_SECONDS = 120.0
async def _post_and_stream( async def _post_and_stream(
base_url: str, base_url: str,
api_key: str, api_key: str,
payload: dict, payload: dict,
log_extra: dict[str, Any] | None, log_extra: dict[str, Any] | None,
) -> tuple[httpx.AsyncClient, httpx.Response]: ) -> httpx.Response:
"""POST to upstream chat-completions and return (client, response) for """POST to upstream chat-completions and return a streaming response."""
streaming. Caller is responsible for closing both."""
url = f"{base_url.rstrip('/')}/chat/completions" url = f"{base_url.rstrip('/')}/chat/completions"
client = httpx.AsyncClient(timeout=httpx.Timeout(120.0, read=120.0))
try: try:
client = acquire_upstream_http_client(url)
# HTTPX replaces rather than merges per-request timeout settings.
client_timeout = client.timeout
request = client.build_request( request = client.build_request(
"POST", "POST",
url, url,
@@ -317,10 +365,25 @@ async def _post_and_stream(
"Content-Type": "application/json", "Content-Type": "application/json",
"Accept": "text/event-stream", "Accept": "text/event-stream",
}, },
timeout=httpx.Timeout(
connect=client_timeout.connect,
read=GEMINI_STREAM_READ_TIMEOUT_SECONDS,
write=client_timeout.write,
pool=client_timeout.pool,
),
) )
response = await client.send(request, stream=True) response = await client.send(request, stream=True)
except UpstreamError:
raise
except httpx.PoolTimeout as exc:
logger.error(
"Gemini messages dispatch pool exhausted",
extra={"error": str(exc), "url": url, **(log_extra or {})},
)
raise UpstreamError(
"Upstream connection pool is busy", status_code=503
) from exc
except Exception as exc: except Exception as exc:
await client.aclose()
logger.error( logger.error(
"Gemini messages dispatch HTTP error", "Gemini messages dispatch HTTP error",
extra={"error": str(exc), "url": url, **(log_extra or {})}, extra={"error": str(exc), "url": url, **(log_extra or {})},
@@ -334,7 +397,6 @@ async def _post_and_stream(
body_bytes = await response.aread() body_bytes = await response.aread()
finally: finally:
await response.aclose() await response.aclose()
await client.aclose()
body_text = body_bytes.decode("utf-8", errors="replace") body_text = body_bytes.decode("utf-8", errors="replace")
logger.error( logger.error(
"Gemini messages dispatch upstream error", "Gemini messages dispatch upstream error",
@@ -348,9 +410,10 @@ async def _post_and_stream(
raise UpstreamError( raise UpstreamError(
f"Upstream error via gemini compat: {body_text}", f"Upstream error via gemini compat: {body_text}",
status_code=response.status_code, status_code=response.status_code,
from_upstream_response=True,
) )
return client, response return response
async def dispatch_gemini_messages( async def dispatch_gemini_messages(
@@ -371,9 +434,7 @@ async def dispatch_gemini_messages(
aggregates). aggregates).
""" """
if not request_body: if not request_body:
raise UpstreamError( raise UpstreamError("Missing request body for /v1/messages", status_code=400)
"Missing request body for /v1/messages", status_code=400
)
try: try:
body: dict = json.loads(request_body) body: dict = json.loads(request_body)
@@ -441,9 +502,7 @@ async def dispatch_gemini_messages(
}, },
) )
http_client, response = await _post_and_stream( response = await _post_and_stream(base_url, api_key, openai_kwargs, log_extra)
base_url, api_key, openai_kwargs, log_extra
)
async def line_iter() -> AsyncGenerator[str, None]: async def line_iter() -> AsyncGenerator[str, None]:
try: try:
@@ -451,10 +510,9 @@ async def dispatch_gemini_messages(
yield line yield line
finally: finally:
await response.aclose() await response.aclose()
await http_client.aclose()
anthropic_event_iter = _openai_chunks_to_anthropic_events( anthropic_event_iter = _ResponseOwnedIterator(
line_iter(), requested_model _openai_chunks_to_anthropic_events(line_iter(), requested_model), response
) )
if not client_stream: if not client_stream:
@@ -475,6 +533,8 @@ async def dispatch_gemini_messages(
f"Failed to aggregate upstream stream: {exc}", f"Failed to aggregate upstream stream: {exc}",
status_code=502, status_code=502,
) from exc ) from exc
finally:
await anthropic_event_iter.aclose()
return client_stream, aggregated, requested_model return client_stream, aggregated, requested_model
return client_stream, anthropic_event_iter, requested_model return client_stream, anthropic_event_iter, requested_model
+24 -3
View File
@@ -1,10 +1,12 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from urllib.parse import urlparse
import httpx import httpx
from .base import BaseUpstreamProvider from .base import BaseUpstreamProvider, _reported_provider
from .model_paths import public_provider_url
from .pricing_resolver import ( from .pricing_resolver import (
FallbackPricingResolver, FallbackPricingResolver,
ResolvedPricing, ResolvedPricing,
@@ -26,7 +28,11 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
provider_type = "generic" provider_type = "generic"
default_base_url = "http://localhost:8888" default_base_url = "http://localhost:8888"
platform_url = None platform_url: str | None = None
# Subclasses that own an authoritative price table set this False so a model
# the table misses imports disabled instead of taking a litellm/OpenRouter
# price that may undercut the upstream's own rate.
use_fallback_pricing = True
def __init__( def __init__(
self, self,
@@ -50,6 +56,21 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
provider_fee=provider_fee, provider_fee=provider_fee,
) )
def _apply_provider_field(self, response_json: object) -> None:
"""Stamp ``"generic:<upstream host>"`` unless the upstream named itself.
A generic upstream is not a router, so nothing identifies the serving
endpoint in the payload; the base URL host fills that role.
"""
if not isinstance(response_json, dict):
return
if _reported_provider(response_json) is None:
response_json["provider"] = (
urlparse(public_provider_url(self.base_url)).hostname
or self.upstream_name
)
super()._apply_provider_field(response_json)
@classmethod @classmethod
def _build_from_row( def _build_from_row(
cls, provider_row: "UpstreamProviderRow" cls, provider_row: "UpstreamProviderRow"
@@ -145,7 +166,7 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
model_spec = model_data.get("model_spec", {}) model_spec = model_data.get("model_spec", {})
resolved = self._native_pricing(model_id, model_spec) resolved = self._native_pricing(model_id, model_spec)
if resolved is None: if resolved is None and self.use_fallback_pricing:
resolved = await resolver.resolve(model_id) resolved = await resolver.resolve(model_id)
if resolved is None: if resolved is None:
+1
View File
@@ -272,6 +272,7 @@ async def _seed_providers_from_settings(
("PERPLEXITY_API_KEY", "perplexity", None, None), ("PERPLEXITY_API_KEY", "perplexity", None, None),
("FIREWORKS_API_KEY", "fireworks", None, None), ("FIREWORKS_API_KEY", "fireworks", None, None),
("XAI_API_KEY", "xai", None, None), ("XAI_API_KEY", "xai", None, None),
("DEEPSEEK_API_KEY", "deepseek", None, None),
("TINFOIL_API_KEY", "tinfoil", None, None), ("TINFOIL_API_KEY", "tinfoil", None, None),
("TYPESAFE_API_KEY", "typesafe", None, None), ("TYPESAFE_API_KEY", "typesafe", None, None),
] ]
+512
View File
@@ -0,0 +1,512 @@
"""Per-origin HTTP client pools with event-loop-aware shutdown."""
import asyncio
import concurrent.futures
import functools
import ipaddress
import ssl
import threading
import weakref
from dataclasses import dataclass, field
from typing import Any, cast
from urllib.parse import urlsplit
import httpx
from ..core import get_logger
from ..core.exceptions import UpstreamError
from ..core.settings import settings
logger = get_logger(__name__)
# Guards all module-level bookkeeping (_clients, _client_loop, _closing,
# _pending_closes, _failed_closes, _close_completed). Multiple event loops can
# live on different OS threads (tests and reload/shutdown paths exercise
# this), so compound read-modify-write sequences on these dicts need a real
# lock. Reentrant because _collect_completed_closes re-enters _schedule_close
# when rehoming clients. Never held across an await.
_state_lock = threading.RLock()
_clients: dict[str, httpx.AsyncClient] = {}
_client_loop: asyncio.AbstractEventLoop | None = None
_closing = False
UPSTREAM_MAX_KEEPALIVE_CONNECTIONS = 50
UPSTREAM_KEEPALIVE_EXPIRY = 60.0
UPSTREAM_CONNECT_TIMEOUT = 30.0
UPSTREAM_WRITE_TIMEOUT = 30.0
UPSTREAM_CONNECT_RETRIES = 1
@dataclass
class _CloseSubmission:
client: httpx.AsyncClient
completion: concurrent.futures.Future[None]
task: asyncio.Task[None] | None = None
retired: bool = False
settlement_lock: threading.Lock = field(default_factory=threading.Lock)
settled_outcome: tuple[str, object | None] | None = None
_pending_closes: dict[
asyncio.AbstractEventLoop,
dict[concurrent.futures.Future[None], _CloseSubmission],
] = {}
_failed_closes: dict[asyncio.AbstractEventLoop, set[httpx.AsyncClient]] = {}
_close_completed: weakref.WeakKeyDictionary[httpx.AsyncClient, bool] = (
weakref.WeakKeyDictionary()
)
class _StatelessCookies(httpx.Cookies):
"""Prevent response cookies from leaking between callers sharing a pool."""
def extract_cookies(self, response: httpx.Response) -> None:
return
def upstream_origin_key(url: str) -> str:
"""Return a canonical origin for an absolute HTTP(S) URL."""
error = "Upstream URL must be an absolute HTTP(S) URL with a valid authority"
if not isinstance(url, str):
raise ValueError(error)
try:
parts = urlsplit(url)
hostname = parts.hostname
port = parts.port
except ValueError as exc:
raise ValueError(error) from exc
scheme = parts.scheme.lower()
authority = parts.netloc.rsplit("@", 1)[-1]
if (
scheme not in {"http", "https"}
or not hostname
or "@" in parts.netloc
or any(character.isspace() for character in hostname)
or authority.endswith(":")
):
raise ValueError(error)
try:
address = ipaddress.ip_address(hostname)
except ValueError:
# HTTPX URL serialization applies the same IDNA normalization used for
# requests, so Unicode and punycode spellings share one pool key.
try:
normalized = httpx.URL(url).copy_with(
username=None,
password=None,
path="/",
query=None,
fragment=None,
)
except httpx.InvalidURL as exc:
raise ValueError(error) from exc
return str(normalized).rstrip("/")
canonical_host = address.compressed
if address.version == 6:
canonical_host = f"[{canonical_host}]"
default_port = 80 if scheme == "http" else 443
port_suffix = f":{port}" if port is not None and port != default_port else ""
return f"{scheme}://{canonical_host}{port_suffix}"
@functools.lru_cache(maxsize=1)
def _shared_ssl_context() -> ssl.SSLContext:
# Loading the CA bundle costs tens of milliseconds; do it once per process
# instead of once per origin pool.
return httpx.create_ssl_context()
def _build_client() -> httpx.AsyncClient:
limits = httpx.Limits(
max_connections=settings.upstream_max_connections,
max_keepalive_connections=UPSTREAM_MAX_KEEPALIVE_CONNECTIONS,
keepalive_expiry=UPSTREAM_KEEPALIVE_EXPIRY,
)
client = httpx.AsyncClient(
transport=httpx.AsyncHTTPTransport(
verify=_shared_ssl_context(),
limits=limits,
retries=UPSTREAM_CONNECT_RETRIES,
),
timeout=httpx.Timeout(
connect=UPSTREAM_CONNECT_TIMEOUT,
read=settings.upstream_read_timeout,
write=UPSTREAM_WRITE_TIMEOUT,
pool=settings.upstream_pool_timeout,
),
)
# AsyncClient's public setter copies into a concrete Cookies jar, so replace
# the backing jar directly to keep response cookies out of it.
client._cookies = _StatelessCookies()
return client
def _close_is_pending(client: httpx.AsyncClient) -> bool:
with _state_lock:
return any(
submission.client is client
for closes in _pending_closes.values()
for submission in closes.values()
)
def _forget_failed_client(client: httpx.AsyncClient) -> None:
with _state_lock:
for failed_loop, failed in list(_failed_closes.items()):
failed.discard(client)
if not failed:
_failed_closes.pop(failed_loop, None)
async def _close_client_resources(client: httpx.AsyncClient) -> None:
if not client.is_closed:
await client.aclose()
return
# HTTPX marks the client closed before awaiting its transports. A retry after
# cancellation or failure therefore has to resume at the transport boundary.
raw_client = cast(Any, client)
resources = [raw_client._transport]
resources.extend(
proxy for proxy in raw_client._mounts.values() if proxy is not None
)
seen: set[int] = set()
for resource in resources:
if id(resource) in seen:
continue
seen.add(id(resource))
await resource.aclose()
def _close_task_outcome(
completed: asyncio.Task[None],
) -> tuple[str, object | None]:
if completed.cancelled():
return ("cancelled", None)
exception = completed.exception()
if exception is not None:
return ("exception", exception)
return ("result", completed.result())
def _matching_close_outcomes(
first: tuple[str, object | None], second: tuple[str, object | None]
) -> bool:
if first[0] != second[0]:
return False
if first[0] == "cancelled":
return True
return first[1] is second[1]
def _settle_close_submission(
submission: _CloseSubmission, completed: asyncio.Task[None]
) -> None:
outcome = _close_task_outcome(completed)
with submission.settlement_lock:
if submission.settled_outcome is not None:
if _matching_close_outcomes(submission.settled_outcome, outcome):
return
raise RuntimeError("Close submission settled with conflicting outcomes")
if submission.completion.done():
raise RuntimeError("Close submission completion changed before settlement")
if outcome[0] == "result":
submission.completion.set_result(None)
elif outcome[0] == "exception":
submission.completion.set_exception(cast(BaseException, outcome[1]))
else:
submission.completion.set_exception(asyncio.CancelledError())
submission.settled_outcome = outcome
def _settle_submission_from_task(submission: _CloseSubmission) -> None:
task = submission.task
if task is not None and task.done():
_settle_close_submission(submission, task)
def _finish_close_submission(
submission: _CloseSubmission, completed: asyncio.Task[None]
) -> None:
if not submission.retired:
_settle_close_submission(submission, completed)
def _submit_close(
client: httpx.AsyncClient, loop: asyncio.AbstractEventLoop
) -> _CloseSubmission:
"""Submit a close without creating its coroutine until the loop runs it."""
submission = _CloseSubmission(client, concurrent.futures.Future())
def start() -> None:
if not submission.completion.set_running_or_notify_cancel():
return
submission.task = loop.create_task(_close_client_resources(client))
submission.task.add_done_callback(
lambda completed: _finish_close_submission(submission, completed)
)
loop.call_soon_threadsafe(start)
return submission
def _collect_completed_closes() -> None:
current_loop = asyncio.get_running_loop()
rehome: list[httpx.AsyncClient] = []
with _state_lock:
for loop, closes in list(_pending_closes.items()):
for future, submission in list(closes.items()):
client = submission.client
_settle_submission_from_task(submission)
if not future.done():
if submission.task is None and not loop.is_running():
submission.retired = True
future.cancel()
closes.pop(future)
rehome.append(client)
elif loop.is_closed():
# A task on a closed loop cannot resume, so it cannot
# race a retry at the owned transport boundary.
submission.retired = True
closes.pop(future)
rehome.append(client)
continue
closes.pop(future)
try:
future.result()
except concurrent.futures.CancelledError:
rehome.append(client)
except asyncio.CancelledError:
_close_completed.pop(client, None)
_failed_closes.setdefault(loop, set()).add(client)
except Exception as exc:
_close_completed.pop(client, None)
_failed_closes.setdefault(loop, set()).add(client)
logger.warning(
"Failed to close upstream HTTP client",
extra={"error": str(exc), "error_type": type(exc).__name__},
)
else:
_close_completed[client] = True
_forget_failed_client(client)
if not closes:
_pending_closes.pop(loop, None)
for client in rehome:
_schedule_close(client, current_loop)
def _resume_stopped_loop(
loop: asyncio.AbstractEventLoop,
tasks: list[asyncio.Task[None]],
timeout: float,
) -> bool:
if loop.is_closed() or loop.is_running():
return False
async def wait_for_tasks() -> None:
await asyncio.wait(tasks, timeout=timeout)
waiter = wait_for_tasks()
try:
loop.run_until_complete(waiter)
except RuntimeError:
waiter.close()
return False
return all(task.done() for task in tasks)
async def _drain_pending_closes(timeout: float = 5.0) -> None:
deadline = asyncio.get_running_loop().time() + timeout
while True:
_collect_completed_closes()
with _state_lock:
if not _pending_closes:
return
if all(loop.is_closed() for loop in _pending_closes):
return
pending_snapshot = [
(
owner_loop,
[
submission.task
for submission in closes.values()
if submission.task is not None and not submission.task.done()
],
)
for owner_loop, closes in _pending_closes.items()
]
remaining = deadline - asyncio.get_running_loop().time()
if remaining <= 0:
logger.error(
"Timed out draining upstream HTTP client closes; retaining them for retry"
)
return
resumed = False
for owner_loop, tasks in pending_snapshot:
if owner_loop.is_closed() or owner_loop.is_running():
continue
if not tasks:
continue
resumed = True
await asyncio.to_thread(
_resume_stopped_loop,
owner_loop,
tasks,
remaining,
)
_collect_completed_closes()
if not resumed:
await asyncio.sleep(min(0.01, remaining))
def _schedule_close(
client: httpx.AsyncClient, owner_loop: asyncio.AbstractEventLoop
) -> None:
"""Schedule closure on the owning loop, retaining unfinished work."""
with _state_lock:
if _close_completed.get(client, False):
_forget_failed_client(client)
return
if _close_is_pending(client):
return
current_loop = asyncio.get_running_loop()
execution_loop = owner_loop
if owner_loop.is_closed() or not owner_loop.is_running():
execution_loop = current_loop
logger.warning(
"Closing upstream HTTP client outside its inactive event loop"
)
try:
submission = _submit_close(client, execution_loop)
except RuntimeError:
if execution_loop is current_loop:
_failed_closes.setdefault(owner_loop, set()).add(client)
return
logger.warning("Upstream HTTP client event loop stopped during shutdown")
submission = _submit_close(client, current_loop)
execution_loop = current_loop
_forget_failed_client(client)
_pending_closes.setdefault(execution_loop, {})[submission.completion] = (
submission
)
def acquire_upstream_http_client(url: str) -> httpx.AsyncClient:
"""Return the pooled client for ``url``, mapping failures to ``UpstreamError``.
Shutdown becomes a 503 so callers can fail over; a malformed provider URL
becomes a 502 instead of an unhandled 500.
"""
try:
return get_upstream_http_client(url)
except RuntimeError as exc:
raise UpstreamError(str(exc), status_code=503) from exc
except ValueError as exc:
raise UpstreamError(str(exc), status_code=502) from exc
def get_upstream_http_client(url: str) -> httpx.AsyncClient:
"""Return the shared client for an absolute upstream URL's origin."""
global _client_loop
loop = asyncio.get_running_loop()
with _state_lock:
if _closing:
raise RuntimeError("Upstream HTTP client is shutting down")
_collect_completed_closes()
if _client_loop is not loop:
stale_clients = list(_clients.values())
stale_loop = _client_loop
_clients.clear()
_client_loop = loop
if stale_loop is not None:
for stale_client in stale_clients:
_schedule_close(stale_client, stale_loop)
key = upstream_origin_key(url)
client = _clients.get(key)
if client is not None and not client.is_closed:
return client
client = _build_client()
_clients[key] = client
logger.debug(
"Opened upstream HTTP connection pool",
extra={
"origin": key,
"max_connections": settings.upstream_max_connections,
"max_keepalive_connections": UPSTREAM_MAX_KEEPALIVE_CONNECTIONS,
"pool_timeout": settings.upstream_pool_timeout,
"read_timeout": settings.upstream_read_timeout,
},
)
return client
async def close_upstream_http_client() -> None:
"""Close every pool, using its owner loop while that loop remains active."""
global _client_loop, _closing
with _state_lock:
_collect_completed_closes()
clients = list(_clients.values())
owner_loop = _client_loop
failed_clients = [
(failed_loop, client)
for failed_loop, failed in _failed_closes.items()
for client in failed
]
if not clients and not failed_clients and not _pending_closes:
return
_closing = True
_clients.clear()
_client_loop = None
if owner_loop is not None:
for client in clients:
_schedule_close(client, owner_loop)
for failed_loop, client in failed_clients:
_schedule_close(client, failed_loop)
try:
await _drain_pending_closes()
finally:
with _state_lock:
_closing = False
def build_x_cashu_client() -> httpx.AsyncClient:
"""Build a per-request client for x-cashu forwarding.
This path intentionally bypasses the shared per-origin pools in this
module: the response and client are handed off to
``OwnedUpstreamStream``/``close_upstream_exchange`` in
``stream_ownership.py``, which close the client once the exchange
finishes. Closing a pooled client would tear down the shared pool for
every caller, so ownership stays per-request here at the cost of a fresh
connection per call.
"""
return httpx.AsyncClient(
transport=httpx.AsyncHTTPTransport(
verify=_shared_ssl_context(),
retries=UPSTREAM_CONNECT_RETRIES,
),
timeout=httpx.Timeout(
connect=UPSTREAM_CONNECT_TIMEOUT,
read=settings.upstream_read_timeout,
write=UPSTREAM_WRITE_TIMEOUT,
pool=settings.upstream_pool_timeout,
),
)
+29
View File
@@ -0,0 +1,29 @@
"""Fast JSON for per-chunk streaming paths, with stdlib fallback.
orjson rejects a few inputs the stdlib accepts (``NaN``/``Infinity`` on load,
non-string keys and integers beyond 64 bits on dump). Streaming must never break
on those, so each call falls back to :mod:`json` instead of raising.
"""
import json
import orjson
def loads(data: bytes | str) -> object | None:
"""Parse JSON, returning ``None`` when the payload is not valid JSON."""
try:
return orjson.loads(data)
except orjson.JSONDecodeError:
try:
return json.loads(data)
except ValueError:
return None
def dumps(obj: object) -> bytes:
"""Serialize to compact UTF-8 JSON bytes."""
try:
return orjson.dumps(obj)
except TypeError:
return json.dumps(obj).encode()
+170 -47
View File
@@ -36,6 +36,12 @@ from .reasoning_effort import adapt_messages_body_for_litellm
logger = get_logger(__name__) logger = get_logger(__name__)
# Sent in place of a blank upstream key. LiteLLM treats ``""`` as missing and
# falls back to the provider's env var (e.g. ``OPENAI_API_KEY``), failing with
# an AuthenticationError for keyless upstreams such as self-hosted
# OpenAI-compatible servers, which the chat path reaches without auth.
KEYLESS_UPSTREAM_API_KEY = "no-key"
# Anthropic-Messages-only fields that don't translate to OpenAI # Anthropic-Messages-only fields that don't translate to OpenAI
# Chat Completions. ``litellm.drop_params`` only filters *known* # Chat Completions. ``litellm.drop_params`` only filters *known*
# unsupported params; these newer/extension fields get passed through # unsupported params; these newer/extension fields get passed through
@@ -78,6 +84,34 @@ ALLOWED_MESSAGES_REQUEST_FIELDS: frozenset[str] = frozenset(
) )
def prune_blank_system_blocks(body: dict) -> None:
"""Drop whitespace-only ``system`` text.
Anthropic accepts a blank system prompt; OpenAI-compatible upstreams
reject it with ``text content blocks must contain non-whitespace text``.
"""
system = body.get("system")
if isinstance(system, str):
if not system.strip():
body.pop("system", None)
return
if not isinstance(system, list):
return
kept = [
block
for block in system
if not (
isinstance(block, dict)
and block.get("type") == "text"
and not str(block.get("text") or "").strip()
)
]
if kept:
body["system"] = kept
else:
body.pop("system", None)
def coerce_litellm_payload(payload: object) -> dict: def coerce_litellm_payload(payload: object) -> dict:
"""Convert a litellm event into a plain dict. """Convert a litellm event into a plain dict.
@@ -372,16 +406,9 @@ def annotate_event(event: dict, requested_model: str | None) -> AnnotatedEvent:
_coerce_float(root_cost_details.get("output_cost")), _coerce_float(root_cost_details.get("output_cost")),
) )
event_type = str(event.get("type") or "")
payload = json.dumps(event)
if event_type:
sse_bytes = f"event: {event_type}\ndata: {payload}\n\n".encode()
else:
sse_bytes = f"data: {payload}\n\n".encode()
return AnnotatedEvent( return AnnotatedEvent(
event, event,
sse_bytes, encode_sse(event),
in_tokens, in_tokens,
out_tokens, out_tokens,
cache_read_tokens, cache_read_tokens,
@@ -393,6 +420,14 @@ def annotate_event(event: dict, requested_model: str | None) -> AnnotatedEvent:
) )
def encode_sse(event: dict) -> bytes:
event_type = str(event.get("type") or "")
payload = json.dumps(event)
if event_type:
return f"event: {event_type}\ndata: {payload}\n\n".encode()
return f"data: {payload}\n\n".encode()
async def stream_annotated_events( async def stream_annotated_events(
iterator: AsyncIterator[Any], iterator: AsyncIterator[Any],
requested_model: str | None, requested_model: str | None,
@@ -450,6 +485,76 @@ def compute_refund(amount: int, unit: str, cost_msats: int) -> int:
raise ValueError(f"Invalid unit: {unit}") raise ValueError(f"Invalid unit: {unit}")
_MAX_UPSTREAM_MESSAGE_CHARS = 300
def collapse_litellm_message(message: str) -> str:
"""Keep the innermost provider message and cap its length."""
tail = message.rsplit("Original exception:", 1)[-1].strip()
while True:
stripped = tail.removeprefix("litellm.")
head, _, rest = stripped.partition(": ")
if rest and head.endswith(("Error", "Exception")):
stripped = rest.strip()
if stripped == tail:
break
tail = stripped
if len(tail) > _MAX_UPSTREAM_MESSAGE_CHARS:
tail = tail[: _MAX_UPSTREAM_MESSAGE_CHARS - 1].rstrip() + "…"
return tail
def is_provider_exception(exc: BaseException) -> bool:
"""Distinguish SDK failures from bugs in our stream handling."""
return type(exc).__module__.split(".", 1)[0] in {"litellm", "openai"}
def upstream_error_from_exception(
exc: Exception,
*,
log_message: str,
log_extra: dict[str, Any] | None = None,
) -> UpstreamError:
"""Redact and classify provider failures, including mid-stream errors."""
raw_message = getattr(exc, "message", None) or str(exc) or repr(exc)
# Redact provider account ids before the message reaches logs or the client.
exc_message = collapse_litellm_message(redact_org_ids(raw_message))
exc_status = getattr(exc, "status_code", None)
exc_response = getattr(exc, "response", None)
response_text = None
if exc_response is not None:
try:
response_text = redact_org_ids(
getattr(exc_response, "text", str(exc_response))
)
except Exception:
response_text = "<unreadable>"
status_for_classify = exc_status if isinstance(exc_status, int) else 502
rate_limit = classify_rate_limit(
status_for_classify, exc_message, getattr(exc, "headers", None)
)
logger.error(
log_message,
extra={
"error": exc_message,
"error_type": type(exc).__name__,
"status_code": exc_status,
"error_code": rate_limit.code if rate_limit else None,
"llm_provider": getattr(exc, "llm_provider", None),
"body": redact_org_ids(str(getattr(exc, "body", "") or "")) or None,
"response_text": response_text,
**(log_extra or {}),
},
)
return UpstreamError(
f"Upstream error via litellm: {exc_message}",
status_code=status_for_classify,
code=rate_limit.code if rate_limit else None,
details=rate_limit.as_details() if rate_limit else None,
from_upstream_response=True,
)
async def dispatch_anthropic_messages( async def dispatch_anthropic_messages(
*, *,
request_body: bytes | None, request_body: bytes | None,
@@ -458,6 +563,8 @@ async def dispatch_anthropic_messages(
api_key: str, api_key: str,
provider_prefix: str, provider_prefix: str,
transform_model_name: Callable[[str], str], transform_model_name: Callable[[str], str],
adapt_request: Callable[[dict], str] | None = None,
transform_stream: Callable[[AsyncIterator[Any]], AsyncIterator[Any]] | None = None,
log_extra: dict[str, Any] | None = None, log_extra: dict[str, Any] | None = None,
) -> tuple[bool, Any, str | None]: ) -> tuple[bool, Any, str | None]:
"""Call ``litellm.anthropic.messages.acreate`` and return """Call ``litellm.anthropic.messages.acreate`` and return
@@ -465,6 +572,15 @@ async def dispatch_anthropic_messages(
Shared by the bearer-key and x-cashu paths. Raises :class:`UpstreamError` Shared by the bearer-key and x-cashu paths. Raises :class:`UpstreamError`
on bad input or upstream failure. on bad input or upstream failure.
``adapt_request`` is the provider's last word on the allowlisted body: it
may rewrite it in place and returns a suffix for the upstream model name,
which is how a provider expresses a feature litellm would otherwise
translate into a parameter the upstream rejects.
``transform_stream`` rewrites the upstream event stream before it is
aggregated or handed to the client, so a provider can repair events
litellm translates faithfully but clients cannot use.
""" """
if not request_body: if not request_body:
raise UpstreamError("Missing request body for /v1/messages", status_code=400) raise UpstreamError("Missing request body for /v1/messages", status_code=400)
@@ -499,19 +615,49 @@ async def dispatch_anthropic_messages(
) )
body = {k: v for k, v in body.items() if k in ALLOWED_MESSAGES_REQUEST_FIELDS} body = {k: v for k, v in body.items() if k in ALLOWED_MESSAGES_REQUEST_FIELDS}
prune_blank_system_blocks(body)
model_suffix = adapt_request(body) if adapt_request else ""
# LiteLLM turns Anthropic's server-side web_search tool into the OpenAI
# `web_search_options` parameter. Generic OpenAI-compatible chat endpoints
# (including those serving Claude through a proxy) may reject that field.
# Only a provider with an explicit adaptation (e.g. Venice's model suffix)
# can preserve search semantics; do not silently remove the tool and return
# an answer that never searched. Native /v1/messages providers bypass this
# dispatcher and receive the original tool unchanged.
tools = body.get("tools")
if provider_prefix == "openai/" and isinstance(tools, list) and any(
isinstance(tool, dict)
and (
(
isinstance(tool.get("type"), str)
and tool["type"].startswith("web_search")
)
or tool.get("name") == "web_search"
)
for tool in tools
):
raise UpstreamError(
"This upstream does not support Anthropic web search through "
"OpenAI-compatible /v1/messages translation",
status_code=400,
code="UNSUPPORTED_WEB_SEARCH",
)
# Convention: `model.id` is the canonical upstream model name; # Convention: `model.id` is the canonical upstream model name;
# `forwarded_model_id` is the public alias the internal API exposes # `forwarded_model_id` is the public alias the internal API exposes
# and echoes back to the client. # and echoes back to the client.
requested_model = ( requested_model = (
(model_obj.forwarded_model_id or model_obj.id) if model_obj else None (model_obj.forwarded_model_id or model_obj.id) if model_obj else None
) )
upstream_model = transform_model_name(model_obj.id) upstream_model = f"{transform_model_name(model_obj.id)}{model_suffix}"
litellm_model = f"{provider_prefix}{upstream_model}" litellm_model = f"{provider_prefix}{upstream_model}"
kwargs: dict = { kwargs: dict = {
"model": litellm_model, "model": litellm_model,
"api_base": base_url, "api_base": base_url,
"api_key": api_key, "api_key": api_key or KEYLESS_UPSTREAM_API_KEY,
"stream": upstream_stream, "stream": upstream_stream,
**body, **body,
} }
@@ -530,45 +676,15 @@ async def dispatch_anthropic_messages(
try: try:
result = await litellm.anthropic.messages.acreate(**kwargs) result = await litellm.anthropic.messages.acreate(**kwargs)
except Exception as exc: except Exception as exc:
raw_message = getattr(exc, "message", None) or str(exc) or repr(exc) raise upstream_error_from_exception(
# Redact provider account identifiers before the message reaches logs exc,
# or the surfaced error. log_message="litellm dispatch failed",
exc_message = redact_org_ids(raw_message) log_extra={"model": litellm_model, "api_base": base_url},
exc_status = getattr(exc, "status_code", None)
exc_response = getattr(exc, "response", None)
response_text = None
if exc_response is not None:
try:
response_text = redact_org_ids(
getattr(exc_response, "text", str(exc_response))
)
except Exception:
response_text = "<unreadable>"
status_for_classify = exc_status if isinstance(exc_status, int) else 502
rate_limit = classify_rate_limit(
status_for_classify, exc_message, getattr(exc, "headers", None)
)
logger.error(
"litellm dispatch failed",
extra={
"error": exc_message,
"error_type": type(exc).__name__,
"status_code": exc_status,
"error_code": rate_limit.code if rate_limit else None,
"llm_provider": getattr(exc, "llm_provider", None),
"body": redact_org_ids(str(getattr(exc, "body", "") or "")) or None,
"response_text": response_text,
"model": litellm_model,
"api_base": base_url,
},
)
raise UpstreamError(
f"Upstream error via litellm: {exc_message}",
status_code=status_for_classify,
code=rate_limit.code if rate_limit else None,
details=rate_limit.as_details() if rate_limit else None,
) from exc ) from exc
if transform_stream is not None and hasattr(result, "__aiter__"):
result = transform_stream(cast(AsyncIterator[Any], result))
if not client_stream and hasattr(result, "__aiter__"): if not client_stream and hasattr(result, "__aiter__"):
# Client asked for a non-streaming response but we always stream # Client asked for a non-streaming response but we always stream
# from upstream — drain the events into a single Anthropic Message # from upstream — drain the events into a single Anthropic Message
@@ -581,6 +697,13 @@ async def dispatch_anthropic_messages(
cast(AsyncIterator[Any], result) cast(AsyncIterator[Any], result)
) )
except Exception as exc: except Exception as exc:
if is_provider_exception(exc):
# Upstream failed part-way through, not an aggregation bug.
raise upstream_error_from_exception(
exc,
log_message="Upstream stream failed mid-flight",
log_extra={"model": litellm_model, "api_base": base_url},
) from exc
logger.error( logger.error(
"Failed to aggregate streamed events into message", "Failed to aggregate streamed events into message",
extra={ extra={
+3
View File
@@ -17,6 +17,7 @@ provider they named.
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import functools
import ipaddress import ipaddress
import json import json
import random import random
@@ -102,6 +103,8 @@ class ProviderPathSnapshot:
preserve_model_ids: frozenset[str] = frozenset() preserve_model_ids: frozenset[str] = frozenset()
# Streaming paths stamp this onto every chunk; the configured base URL set is small.
@functools.lru_cache(maxsize=256)
def public_provider_url(base_url: str) -> str: def public_provider_url(base_url: str) -> str:
"""Mask private IP addresses and URLs with explicit ports.""" """Mask private IP addresses and URLs with explicit ports."""
parsed = urlsplit(base_url) parsed = urlsplit(base_url)
+39
View File
@@ -1,11 +1,23 @@
import json
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
from litellm.llms.openai.chat.o_series_transformation import OpenAIOSeriesConfig
from ..payment.models import Model, async_fetch_openrouter_models from ..payment.models import Model, async_fetch_openrouter_models
from .base import BaseUpstreamProvider from .base import BaseUpstreamProvider
if TYPE_CHECKING: if TYPE_CHECKING:
from ..core.db import UpstreamProviderRow from ..core.db import UpstreamProviderRow
_O_SERIES = OpenAIOSeriesConfig()
def _rejects_max_tokens(model: str) -> bool:
return OpenAIGPT5Config.is_model_gpt_5_model(
model
) or _O_SERIES.is_model_o_series_model(model)
class OpenAIUpstreamProvider(BaseUpstreamProvider): class OpenAIUpstreamProvider(BaseUpstreamProvider):
"""Upstream provider specifically configured for OpenAI API.""" """Upstream provider specifically configured for OpenAI API."""
@@ -42,6 +54,33 @@ class OpenAIUpstreamProvider(BaseUpstreamProvider):
"""Strip 'openai/' prefix for OpenAI API compatibility.""" """Strip 'openai/' prefix for OpenAI API compatibility."""
return model_id.removeprefix("openai/") return model_id.removeprefix("openai/")
def prepare_request_body(
self,
body: bytes | None,
model_obj: Model,
include_stream_usage: bool = False,
) -> bytes | None:
body = super().prepare_request_body(body, model_obj, include_stream_usage)
if not body:
return body
try:
data = json.loads(body)
except ValueError:
return body
# Reasoning models 400 on max_tokens; renaming up front saves the
# reject-and-retry round trip. Names litellm doesn't know yet still
# fall through to request_correction's reactive rename.
if (
isinstance(data, dict)
and "messages" in data
and "max_tokens" in data
and "max_completion_tokens" not in data
and _rejects_max_tokens(self.transform_model_name(model_obj.id))
):
data["max_completion_tokens"] = data.pop("max_tokens")
return json.dumps(data).encode()
return body
async def fetch_models(self) -> list[Model]: async def fetch_models(self) -> list[Model]:
"""Fetch OpenAI models from OpenRouter API filtered by openai source.""" """Fetch OpenAI models from OpenRouter API filtered by openai source."""
models_data = await async_fetch_openrouter_models(source_filter="openai") models_data = await async_fetch_openrouter_models(source_filter="openai")
+35 -5
View File
@@ -2,12 +2,27 @@ from typing import TYPE_CHECKING
import httpx import httpx
from ..core.logging import get_logger
from ..payment.models import Model, async_fetch_openrouter_models from ..payment.models import Model, async_fetch_openrouter_models
from .base import BaseUpstreamProvider from .base import BaseUpstreamProvider, _reported_provider
from .model_paths import public_provider_url
if TYPE_CHECKING: if TYPE_CHECKING:
from ..core.db import UpstreamProviderRow from ..core.db import UpstreamProviderRow
logger = get_logger(__name__)
_UNKNOWN_SUB_PROVIDER = "unknown"
def _carries_usage(payload: dict) -> bool:
"""Whether a payload holds usage, at top level or in the Anthropic
``message`` / Responses ``response`` envelope."""
return any(
isinstance(obj, dict) and isinstance(obj.get("usage"), dict)
for obj in (payload, payload.get("message"), payload.get("response"))
)
class OpenRouterUpstreamProvider(BaseUpstreamProvider): class OpenRouterUpstreamProvider(BaseUpstreamProvider):
"""Upstream provider specifically configured for OpenRouter API.""" """Upstream provider specifically configured for OpenRouter API."""
@@ -26,22 +41,37 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider):
- Real upstream sub-provider (e.g. ``"GMICloud"``) -> ``"openrouter:GMICloud"``. - Real upstream sub-provider (e.g. ``"GMICloud"``) -> ``"openrouter:GMICloud"``.
- Missing sub-provider, or one that merely echoes ``"openrouter"`` -> - Missing sub-provider, or one that merely echoes ``"openrouter"`` ->
``"unknown"``. ``"openrouter:unknown"``: the router is still known even when the
serving provider is not (e.g. the Responses API never reports it).
- Idempotent: re-stamping never produces ``"openrouter:openrouter:..."``; - Idempotent: re-stamping never produces ``"openrouter:openrouter:..."``;
the ``openrouter:`` prefix appears at most once. the ``openrouter:`` prefix appears at most once.
""" """
if not isinstance(response_json, dict): if not isinstance(response_json, dict):
return return
response_json["provider_url"] = public_provider_url(self.base_url)
provider_type = (self.provider_type or "").strip() provider_type = (self.provider_type or "").strip()
existing = response_json.get("provider") sub = _reported_provider(response_json) or ""
sub = existing.strip() if isinstance(existing, str) else ""
# Strip any already-applied "openrouter:" prefixes (idempotency). # Strip any already-applied "openrouter:" prefixes (idempotency).
prefix = f"{provider_type}:" prefix = f"{provider_type}:"
while sub.lower().startswith(prefix.lower()): while sub.lower().startswith(prefix.lower()):
sub = sub[len(prefix) :].strip() sub = sub[len(prefix) :].strip()
# Already stamped as unknown on an earlier pass; keep it without
# warning again.
if sub.lower() == _UNKNOWN_SUB_PROVIDER:
response_json["provider"] = f"{provider_type}:{_UNKNOWN_SUB_PROVIDER}"
return
# No real sub-provider, or it just echoes our own router name. # No real sub-provider, or it just echoes our own router name.
if not sub or sub.lower() == provider_type.lower(): if not sub or sub.lower() == provider_type.lower():
response_json["provider"] = "unknown" # Warn only on the billed payload, not on every stream chunk.
if _carries_usage(response_json):
logger.warning(
"OpenRouter did not report the serving provider",
extra={
"model": response_json.get("model"),
"response_id": response_json.get("id"),
},
)
response_json["provider"] = f"{provider_type}:{_UNKNOWN_SUB_PROVIDER}"
return return
response_json["provider"] = f"{provider_type}:{sub}" response_json["provider"] = f"{provider_type}:{sub}"
+63 -1
View File
@@ -43,6 +43,18 @@ _UNSUPPORTED_PARAM_RE = re.compile(
) )
# Matches upstream error text that rejects a param and names its replacement,
# e.g. OpenAI's "Unsupported parameter: 'max_tokens' is not supported with this
# model. Use 'max_completion_tokens' instead." Both names must be quoted so a
# free-form hint like "use gpt-4 instead" never reads as a rename.
_RENAMED_PARAM_RE = re.compile(
r"[`'\"](?P<param>[a-zA-Z_][a-zA-Z0-9_]*)[`'\"]\s+is\s+"
r"(?:deprecated|not\s+supported|unsupported|no\s+longer\s+supported)\b"
r".*?\buse\s+[`'\"](?P<replacement>[a-zA-Z_][a-zA-Z0-9_]*)[`'\"]\s+instead",
re.IGNORECASE | re.DOTALL,
)
# A corrector inspects the parsed request body and the upstream error message # A corrector inspects the parsed request body and the upstream error message
# and returns ``(new_body_dict, label)`` for a fix it can apply, or ``None`` to # and returns ``(new_body_dict, label)`` for a fix it can apply, or ``None`` to
# decline. ``label`` identifies the fix so it is applied at most once per request. # decline. ``label`` identifies the fix so it is applied at most once per request.
@@ -100,6 +112,51 @@ _SPEND_SHAPING_PARAMS = frozenset(
} }
) )
# Output caps are interchangeable spellings of the same limit, so moving the
# value from one to another keeps the priced bound intact.
_OUTPUT_CAP_PARAMS = frozenset(
{
"max_tokens",
"max_completion_tokens",
"max_output_tokens",
"max_tokens_to_sample",
}
)
def rename_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None:
"""Move a rejected top-level param to the name the upstream asked for.
Returns ``(new_body, label)`` with the value carried over unchanged, or
``None`` when the error names no replacement, the param is absent, or the
replacement is already set.
A spend-shaping field is only renamed to another output cap: that keeps the
reservation's bound, whereas renaming into or out of any other spend-shaping
field could uncap or fan out the retry.
"""
match = _RENAMED_PARAM_RE.search(error_message)
if not match:
return None
param, replacement = match.group("param"), match.group("replacement")
if param == replacement or param not in body or replacement in body:
return None
param_spend = param.lower() in _SPEND_SHAPING_PARAMS
replacement_spend = replacement.lower() in _SPEND_SHAPING_PARAMS
if (param_spend or replacement_spend) and not (
param.lower() in _OUTPUT_CAP_PARAMS
and replacement.lower() in _OUTPUT_CAP_PARAMS
):
logger.warning(
"Upstream asked to rename '%s' to '%s'; refusing because it would "
"change the request's spend bound — surfacing the error",
param,
replacement,
)
return None
new_body = {(replacement if k == param else k): v for k, v in body.items()}
return new_body, f"{param}->{replacement}"
def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None: def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None:
"""Drop a top-level param the upstream named as unsupported/deprecated. """Drop a top-level param the upstream named as unsupported/deprecated.
@@ -130,7 +187,12 @@ def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str]
# Ordered pipeline of correctors tried on each recoverable rejection. # Ordered pipeline of correctors tried on each recoverable rejection.
DEFAULT_CORRECTORS: tuple[Corrector, ...] = (strip_unsupported_param,) # Renaming runs first so a param with a named replacement keeps its value
# instead of being dropped.
DEFAULT_CORRECTORS: tuple[Corrector, ...] = (
rename_unsupported_param,
strip_unsupported_param,
)
def correct_request( def correct_request(
+44
View File
@@ -0,0 +1,44 @@
"""Incremental SSE event splitting that stays linear in stream size."""
class SSEEventSplitter:
"""Split upstream bytes into SSE events delimited by a blank line.
CRLF is normalized to LF. Each call only scans newly received bytes, so a
large event arriving over many network chunks (e.g. a Responses API
``response.completed`` carrying the full output) costs O(n) rather than
rescanning the buffered prefix on every chunk.
"""
def __init__(self) -> None:
self._buffer = bytearray()
# A trailing CR may be the first half of a CRLF split across chunks.
self._pending_cr = False
def feed(self, chunk: bytes) -> list[bytes]:
"""Add ``chunk`` and return the events it completed, without delimiters."""
if self._pending_cr:
chunk = b"\r" + chunk
self._pending_cr = chunk.endswith(b"\r")
if self._pending_cr:
chunk = chunk[:-1]
# The delimiter may straddle the old tail and the new chunk.
scan_from = max(len(self._buffer) - 1, 0)
self._buffer += chunk.replace(b"\r\n", b"\n")
events: list[bytes] = []
start = 0
while (end := self._buffer.find(b"\n\n", scan_from)) != -1:
events.append(bytes(self._buffer[start:end]))
start = scan_from = end + 2
if start:
del self._buffer[:start]
return events
def flush(self) -> bytes:
"""Return any trailing bytes that never saw a closing blank line."""
tail = bytes(self._buffer) + (b"\r" if self._pending_cr else b"")
self._buffer.clear()
self._pending_cr = False
return tail
+203
View File
@@ -0,0 +1,203 @@
from __future__ import annotations
import asyncio
import inspect
from collections.abc import AsyncIterator, Awaitable, Callable
from typing import Any, Self, cast
import httpx
from fastapi.responses import StreamingResponse
from starlette.types import Receive, Scope, Send
from ..core import get_logger
from ..core.settings import settings
logger = get_logger(__name__)
async def aclose_if_needed(resource: object | None) -> None:
if resource is None:
return
close = getattr(resource, "aclose", None)
if close is None:
return
result = close()
if inspect.isawaitable(result):
await result
async def shielded_aclose(resource: object | None) -> None:
await asyncio.shield(aclose_if_needed(resource))
class ResponseHandoff:
"""Close a response unless ownership is transferred to a stream."""
def __init__(self) -> None:
self._response: object | None = None
def acquire(self, response: object) -> None:
self._response = response
def handoff(self) -> None:
self._response = None
async def close(self, *, suppress_errors: bool = False) -> None:
response = self._response
self._response = None
if response is None:
return
try:
await shielded_aclose(response)
except BaseException:
if not suppress_errors:
raise
logger.exception("Failed to close upstream response before handoff")
async def finalize_and_close_stream(
finalize: Callable[[], Awaitable[None]] | None,
response: object | None,
) -> None:
"""Settle billing, then return the response connection to its pool."""
try:
if finalize is not None:
await finalize()
finally:
await aclose_if_needed(response)
class PersistentStreamFinalizer:
"""Run one stream finalizer to completion across cancellation boundaries."""
def __init__(self, finalize: Callable[[], Awaitable[None]]) -> None:
self._finalize = finalize
self._task: asyncio.Future[None] | None = None
self._lock = asyncio.Lock()
async def _bounded_finalize(self) -> None:
async with asyncio.timeout(settings.request_cleanup_timeout_seconds):
await self._finalize()
async def run(self) -> None:
async with self._lock:
if self._task is None:
self._task = asyncio.ensure_future(self._bounded_finalize())
task = self._task
await asyncio.shield(task)
class FinalizingAsyncIterator:
"""Tie iterator shutdown to a finalizer created before streaming starts."""
def __init__(
self,
iterator: AsyncIterator[bytes],
finalizer: PersistentStreamFinalizer,
) -> None:
self._iterator = iterator
self._finalizer = finalizer
def __aiter__(self) -> Self:
return self
async def __anext__(self) -> bytes:
try:
return await self._iterator.__anext__()
except BaseException:
await self._finalizer.run()
raise
async def aclose(self) -> None:
try:
await aclose_if_needed(self._iterator)
finally:
await self._finalizer.run()
class ClosingStreamingResponse(StreamingResponse):
"""Close the body iterator even when downstream ASGI sends fail."""
def __init__(
self,
content: AsyncIterator[bytes],
*,
finalizer: PersistentStreamFinalizer | None = None,
**kwargs: Any,
) -> None:
if finalizer is not None:
content = FinalizingAsyncIterator(content, finalizer)
super().__init__(content, **kwargs)
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
try:
await super().__call__(scope, receive, send)
finally:
await asyncio.shield(aclose_if_needed(self.body_iterator))
class OwnedUpstreamStream:
"""Keep a one-shot HTTP client alive for the lifetime of its response."""
def __init__(
self,
iterator: AsyncIterator[bytes],
response: httpx.Response,
client: httpx.AsyncClient,
) -> None:
self._iterator = iterator
self._response = response
self._client = client
self._cleanup_complete = False
self._cleanup_task: asyncio.Task[None] | None = None
self._close_lock = asyncio.Lock()
def __aiter__(self) -> Self:
return self
async def __anext__(self) -> bytes:
try:
return await self._iterator.__anext__()
except StopAsyncIteration:
await self.aclose()
raise
async def _cleanup(self) -> None:
try:
await aclose_if_needed(self._iterator)
finally:
try:
await self._response.aclose()
finally:
await self._client.aclose()
self._cleanup_complete = True
async def aclose(self) -> None:
async with self._close_lock:
if self._cleanup_complete:
return
if self._cleanup_task is None or self._cleanup_task.done():
self._cleanup_task = asyncio.create_task(self._cleanup())
cleanup_task = self._cleanup_task
await asyncio.shield(cleanup_task)
def attach_upstream_stream_owner(
result: StreamingResponse,
response: httpx.Response,
client: httpx.AsyncClient,
) -> StreamingResponse:
result.body_iterator = OwnedUpstreamStream(
cast(AsyncIterator[bytes], result.body_iterator), response, client
)
return result
async def close_upstream_exchange(
response: httpx.Response | None, client: httpx.AsyncClient
) -> None:
try:
if response is not None:
await response.aclose()
finally:
await client.aclose()
+117
View File
@@ -0,0 +1,117 @@
"""Timeout guards for upstream streaming responses."""
from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator, Callable
import httpx
from ..core import get_logger
from ..core.error_scope import UPSTREAM_ERROR_STATUS
from ..core.exceptions import UpstreamError
from ..core.settings import settings
from .sse_splitter import SSEEventSplitter
logger = get_logger(__name__)
class GuardedStream(AsyncIterator[bytes]):
def __init__(
self,
first: bytes | None,
chunks: AsyncIterator[bytes],
provider_type: str,
on_idle_timeout: Callable[[], None] | None,
) -> None:
self.timed_out = False
self._chunks = self._resume(first, chunks, provider_type, on_idle_timeout)
def __aiter__(self) -> GuardedStream:
return self
async def __anext__(self) -> bytes:
return await anext(self._chunks)
async def _resume(
self,
first: bytes | None,
chunks: AsyncIterator[bytes],
provider_type: str,
on_idle_timeout: Callable[[], None] | None,
) -> AsyncIterator[bytes]:
chunk = first
while chunk is not None:
yield chunk
try:
chunk = await _next_chunk(
chunks, settings.upstream_stream_idle_timeout_seconds
)
except TimeoutError:
self.timed_out = True
logger.warning(
"Upstream stream stalled; aborting and billing actual usage",
extra={
"provider": provider_type,
"idle_timeout_seconds": settings.upstream_stream_idle_timeout_seconds,
},
)
if on_idle_timeout is not None:
on_idle_timeout()
return
def _has_data(event: bytes) -> bool:
return any(
line.startswith(b"data:") and line[5:].strip() for line in event.split(b"\n")
)
async def _sse_events(chunks: AsyncIterator[bytes]) -> AsyncIterator[bytes]:
"""Yield only deliverable SSE data events; comments cannot reset deadlines."""
splitter = SSEEventSplitter()
async for chunk in chunks:
for event in splitter.feed(chunk):
if _has_data(event):
yield event + b"\n\n"
tail = splitter.flush()
if _has_data(tail):
# Keep an unterminated tail unterminated: the caller's final flush must
# not mistake truncated JSON for a complete SSE frame.
yield tail
async def open_guarded_stream(
response: httpx.Response,
provider_type: str,
*,
sse: bool = False,
on_idle_timeout: Callable[[], None] | None = None,
) -> GuardedStream:
"""Prefetch a deliverable event before handing a response to the client.
Once the first event is sent, a stall cannot fail over; the stream ends and
the caller's finalizer settles usage observed before the interruption.
"""
chunks = response.aiter_bytes().__aiter__()
guarded_chunks = _sse_events(chunks) if sse else chunks
timeout = settings.upstream_first_token_timeout_seconds
try:
first = await _next_chunk(guarded_chunks, timeout)
except TimeoutError:
await response.aclose()
raise UpstreamError(
f"Upstream {provider_type} sent no first chunk within {timeout}s",
status_code=UPSTREAM_ERROR_STATUS,
code="UPSTREAM_TIMEOUT",
) from None
return GuardedStream(first, guarded_chunks, provider_type, on_idle_timeout)
async def _next_chunk(chunks: AsyncIterator[bytes], timeout: float) -> bytes | None:
"""Next chunk, or ``None`` at end of stream. ``timeout <= 0`` disables it."""
step = anext(chunks)
try:
return await (asyncio.wait_for(step, timeout) if timeout > 0 else step)
except StopAsyncIteration:
return None
+31
View File
@@ -1,5 +1,6 @@
from __future__ import annotations from __future__ import annotations
import json
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING, Optional
import httpx import httpx
@@ -7,6 +8,12 @@ from fastapi import Request
from fastapi.responses import Response, StreamingResponse from fastapi.responses import Response, StreamingResponse
from pydantic.v1 import BaseModel from pydantic.v1 import BaseModel
from ..core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS,
UPSTREAM_UNAVAILABLE,
)
from ..core.exceptions import UpstreamError from ..core.exceptions import UpstreamError
from ..core.logging import get_logger from ..core.logging import get_logger
from ..payment.models import Architecture, Model, Pricing from ..payment.models import Architecture, Model, Pricing
@@ -138,6 +145,30 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider):
response_headers = dict(resp.headers) response_headers = dict(resp.headers)
response_headers.pop("content-encoding", None) response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None) response_headers.pop("content-length", None)
if resp.status_code >= 500:
logger.warning(
"Tinfoil attestation upstream returned %s",
resp.status_code,
extra={"status_code": resp.status_code},
)
return Response(
content=json.dumps(
{
"error": {
"type": "upstream_error",
"code": UPSTREAM_UNAVAILABLE,
"message": (
"Attestation upstream returned "
f"{resp.status_code}"
),
"upstream_status": resp.status_code,
}
}
),
status_code=UPSTREAM_ERROR_STATUS,
media_type="application/json",
headers={ERROR_SCOPE_HEADER: ERROR_SCOPE_UPSTREAM},
)
return Response( return Response(
content=resp.content, content=resp.content,
status_code=resp.status_code, status_code=resp.status_code,
+16 -1
View File
@@ -20,7 +20,7 @@ from urllib.parse import urlsplit
import h11 import h11
from ..core import get_logger from ..core import get_logger
from ..core.exceptions import EhbpTimeoutError from ..core.exceptions import EhbpConnectionError, EhbpTimeoutError
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -110,6 +110,21 @@ async def forward_with_trailer(
raise EhbpTimeoutError( raise EhbpTimeoutError(
f"EHBP upstream {host} timed out after {timeout_seconds:g}s connecting" f"EHBP upstream {host} timed out after {timeout_seconds:g}s connecting"
) from exc ) from exc
except ConnectionAbortedError as exc:
# CPython's ssl module aborts a TLS handshake that outlives its
# internal timer ("SSL handshake is taking longer than N seconds")
# with ConnectionAbortedError. That is a connect timeout on the
# provider hop, not a local node fault, so surface it as a timeout.
raise EhbpTimeoutError(
f"EHBP upstream {host} TLS handshake timed out while connecting"
) from exc
except (ssl.SSLError, ConnectionError, OSError) as exc:
# DNS failure, connection refused/reset, or a non-timeout TLS error:
# the provider could not be reached. Attribute it to the upstream hop
# rather than letting it become a node-scoped 500.
raise EhbpConnectionError(
f"Unable to connect to EHBP upstream {host}: {type(exc).__name__}"
) from exc
try: try:
# Build HTTP/1.1 request # Build HTTP/1.1 request
+420
View File
@@ -0,0 +1,420 @@
from __future__ import annotations
from collections.abc import AsyncGenerator, AsyncIterator
from typing import TYPE_CHECKING, Any, cast
import httpx
from ..core.exceptions import UpstreamError
from ..core.logging import get_logger
from ..payment.models import Architecture, Model, Pricing, TopProvider
from . import messages_dispatch
from .base import BaseUpstreamProvider
from .stream_ownership import aclose_if_needed
if TYPE_CHECKING:
from ..core.db import UpstreamProviderRow
logger = get_logger(__name__)
# ``GET /models`` defaults to ``type=text``, which is why a Venice account
# configured as a generic upstream never sees the rest of its catalog.
_MODELS_TYPE_PARAM = "all"
# Families this proxy can both route and price. Image, audio, music and video
# are billed per clip or per second and return no usage object to settle
# against, so exposing them would hand out unpriced inference.
_SUPPORTED_TYPES = frozenset({"text", "embedding"})
# Venice prices text in USD per million tokens; Routstr prices per token.
_USD_PER_MILLION = 1_000_000.0
_ARCHITECTURES: dict[str, tuple[str, list[str], list[str]]] = {
"text": ("text->text", ["text"], ["text"]),
"embedding": ("text->embedding", ["text"], ["embedding"]),
}
# Venice runs search itself and reports it back through ``venice_parameters``;
# it has no Anthropic-shaped server tool and rejects the ``web_search_options``
# that litellm's Anthropic adapter derives from one. ``auto`` matches Anthropic
# semantics, where declaring the tool leaves the decision to the model.
# Citations are asked for because litellm's Anthropic response translation
# carries no ``venice_parameters``, so the inline ``^n^`` markers Venice writes
# into the text are the only way a caller sees that sources were used.
_WEB_SEARCH_SUFFIX = ":enable_web_search=auto&enable_web_citations=true"
# Anthropic web-search constraints with no Venice equivalent. Honouring the
# request means enforcing them, so a request that sets one is refused rather
# than answered by a search that ignored it. ``max_uses`` is absent on purpose:
# ``auto`` runs at most one search per request, so any cap of 1 or more is
# already met, while domain filters and location would be silently ignored.
# Only ``max_uses: 0``, a request for no search at all, cannot be honoured.
_UNENFORCEABLE_WEB_SEARCH_KEYS = frozenset(
{"allowed_domains", "blocked_domains", "user_location"}
)
# Venice streams OpenAI reasoning models' encrypted reasoning as a trailing
# ``reasoning_content`` delta carrying this marker. litellm turns it into a
# plaintext ``thinking`` block after the answer, which clients render as
# gibberish and which makes Claude Code report an empty final result.
_ENCRYPTED_REASONING_MARKER = "__ENCRYPTED_REASONING__"
async def _drop_encrypted_reasoning(
upstream: AsyncIterator[Any],
) -> AsyncGenerator[bytes, None]:
"""A thinking block's start carries no text, so it is held until its first
delta shows whether it is the encrypted payload; later indices shift down
to close the gap."""
encode = messages_dispatch.encode_sse
sse_buffer = b""
dropped: set[int] = set()
held: list[dict] | None = None
held_index: int | None = None
def shift(event: dict) -> dict:
index = event.get("index")
if not isinstance(index, int):
return event
gap = sum(1 for d in dropped if d < index)
return {**event, "index": index - gap} if gap else event
try:
async for chunk in upstream:
events, sse_buffer = messages_dispatch.events_from_chunk(chunk, sse_buffer)
for event in events:
etype = event.get("type")
index = event.get("index")
if held is not None:
delta = event.get("delta") or {}
is_own_delta = (
index == held_index and etype == "content_block_delta"
)
thinking = str(delta.get("thinking") or "")
if is_own_delta and thinking.startswith(
_ENCRYPTED_REASONING_MARKER
):
dropped.add(cast(int, index))
held = None
continue
if is_own_delta and not thinking:
held.append(event)
continue
for pending in held:
yield encode(shift(pending))
held = None
if index in dropped:
continue
block = event.get("content_block") or {}
if (
etype == "content_block_start"
and block.get("type") == "thinking"
and not block.get("thinking")
):
held, held_index = [event], index
continue
yield encode(shift(event))
if held is not None:
for pending in held:
yield encode(shift(pending))
finally:
await aclose_if_needed(upstream)
def _is_web_search_tool(tool: Any) -> bool:
"""An Anthropic server-side web-search tool, by either of its markers.
Matches litellm's own detection (``litellm/llms/anthropic/
experimental_pass_through/adapters/transformation.py``), so every tool it
would turn into ``web_search_options`` is caught here first.
"""
if not isinstance(tool, dict):
return False
tool_type = tool.get("type")
return (
isinstance(tool_type, str) and tool_type.startswith("web_search")
) or tool.get("name") == "web_search"
def _merge_cache_marked_system(body: dict) -> None:
"""Venice rejects an OpenAI ``system`` message with two or more text parts
when any part carries ``cache_control`` (``400 system: text content blocks
must contain non-whitespace text``), even though every part is non-blank.
Claude Code always sends that shape. A single marked block is accepted and
still caches, so the prefix stays cacheable under the last marker.
"""
system = body.get("system")
if not isinstance(system, list) or len(system) < 2:
return
if not all(
isinstance(block, dict)
and block.get("type") == "text"
and isinstance(block.get("text"), str)
for block in system
):
return
markers = [block["cache_control"] for block in system if block.get("cache_control")]
if not markers:
return
body["system"] = [
{
"type": "text",
"text": "\n\n".join(block["text"] for block in system),
"cache_control": markers[-1],
}
]
def _usd(entry: Any) -> float | None:
"""Read the USD leg of a Venice ``{usd, diem}`` price pair."""
if isinstance(entry, dict):
value = entry.get("usd")
if isinstance(value, (int, float)) and not isinstance(value, bool):
return float(value)
return None
class VeniceUpstreamProvider(BaseUpstreamProvider):
"""Upstream provider for the Venice.ai API.
Venice publishes a complete price book on its own catalog, so models are
built from that rather than matched against OpenRouter, which has never
heard of most of Venice's catalog.
"""
provider_type = "venice"
default_base_url = "https://api.venice.ai/api/v1"
platform_url = "https://venice.ai/settings/api"
def __init__(self, api_key: str, provider_fee: float = 1.01):
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"
) -> "VeniceUpstreamProvider":
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": "Venice AI",
"default_base_url": cls.default_base_url,
"fixed_base_url": True,
"platform_url": cls.platform_url,
}
def transform_model_name(self, model_id: str) -> str:
return model_id.removeprefix("venice/")
def transform_messages_stream(
self, stream: AsyncIterator[Any]
) -> AsyncIterator[Any]:
return _drop_encrypted_reasoning(stream)
def adapt_messages_request(self, body: dict, model_obj: Model) -> str:
_merge_cache_marked_system(body)
return self._adapt_web_search(body)
def _adapt_web_search(self, body: dict) -> str:
"""Trade an Anthropic web-search tool for Venice's own search switch.
Left in the body, litellm's Anthropic adapter rewrites the tool into a
top-level ``web_search_options``, which Venice answers with a 400. The
tool is lifted out here and the same intent re-expressed as a model
feature suffix, the one form of ``venice_parameters`` that survives
that adapter.
"""
tools = body.get("tools")
if not isinstance(tools, list):
return ""
search_tools = [tool for tool in tools if _is_web_search_tool(tool)]
if not search_tools:
return ""
# A key carrying null or an empty list states no constraint, so it is
# read as absent rather than refused. ``auto`` runs at most one search,
# so only an integer ``max_uses`` of one or more is known to be met.
unenforceable = sorted(
{
key
for tool in search_tools
for key, value in tool.items()
if (
key in _UNENFORCEABLE_WEB_SEARCH_KEYS
and value is not None
and value != []
)
or (
key == "max_uses"
and value is not None
and not (
isinstance(value, int)
and not isinstance(value, bool)
and value >= 1
)
)
}
)
if unenforceable:
raise UpstreamError(
"Venice web search cannot honour these Anthropic web_search "
f"options: {', '.join(unenforceable)}",
status_code=400,
code="UNSUPPORTED_WEB_SEARCH_OPTION",
details={"unsupported_options": unenforceable},
)
tool_choice = body.get("tool_choice")
if isinstance(tool_choice, dict) and tool_choice.get("name") == "web_search":
raise UpstreamError(
"Venice web search cannot be forced through tool_choice; it is "
"decided by the model",
status_code=400,
code="UNSUPPORTED_WEB_SEARCH_OPTION",
details={"unsupported_options": ["tool_choice"]},
)
remaining = [tool for tool in tools if not _is_web_search_tool(tool)]
if remaining:
# A caller's ``tool_choice: any`` is kept and litellm maps it to
# OpenAI ``required``, so one of the remaining function tools must
# now be called where Anthropic would have let a search satisfy it.
# Deliberate: OpenRouter never rewrites tool_choice for web search
# either, and guessing an alternative would change caller intent.
body["tools"] = remaining
else:
body.pop("tools", None)
# tool_choice without tools is rejected by OpenAI-shaped upstreams.
body.pop("tool_choice", None)
return _WEB_SEARCH_SUFFIX
async def _fetch_provider_models(self) -> dict:
url = f"{self.base_url.rstrip('/')}/models"
headers = {"Authorization": f"Bearer {self.api_key}"} if self.api_key else None
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.get(
url, params={"type": _MODELS_TYPE_PARAM}, headers=headers
)
response.raise_for_status()
return response.json()
async def fetch_models(self) -> list[Model]:
try:
payload = await self._fetch_provider_models()
except Exception as e:
logger.error(
"Error fetching Venice models",
extra={"error": str(e), "error_type": type(e).__name__},
)
return []
models: list[Model] = []
skipped: list[str] = []
for entry in payload.get("data", []):
if not isinstance(entry, dict):
continue
try:
model = self._parse_model(entry)
except Exception as e:
logger.warning(
"Failed to parse Venice model",
extra={
"model_id": entry.get("id", "unknown"),
"error": str(e),
"error_type": type(e).__name__,
},
)
continue
if model is None:
skipped.append(str(entry.get("id", "unknown")))
continue
models.append(model)
if skipped:
logger.debug(
f"({len(skipped)}) Venice models skipped as unsupported or unpriced",
extra={"skipped_models": skipped},
)
return models
def _parse_model(self, entry: dict[str, Any]) -> Model | None:
model_type = entry.get("type")
model_id = entry.get("id")
spec = entry.get("model_spec")
if not model_id or model_type not in _SUPPORTED_TYPES:
return None
if not isinstance(spec, dict) or spec.get("offline"):
return None
pricing = self._parse_pricing(spec.get("pricing"), str(model_type))
if pricing is None:
return None
modality, input_modalities, output_modalities = _ARCHITECTURES[str(model_type)]
capabilities = spec.get("capabilities")
if (
model_type == "text"
and isinstance(capabilities, dict)
and capabilities.get("supportsVision")
):
input_modalities = [*input_modalities, "image"]
modality = "text+image->text"
context_length = spec.get("availableContextTokens")
max_completion_tokens = spec.get("maxCompletionTokens")
name = spec.get("name") or str(model_id)
return Model(
id=str(model_id),
name=str(name),
created=int(entry.get("created") or 0),
description=str(spec.get("description") or f"Venice {model_type} model"),
context_length=int(context_length) if context_length else 0,
architecture=Architecture(
modality=modality,
input_modalities=input_modalities,
output_modalities=output_modalities,
tokenizer="Unknown",
instruct_type=None,
),
pricing=pricing,
top_provider=TopProvider(
context_length=int(context_length) if context_length else None,
max_completion_tokens=int(max_completion_tokens)
if max_completion_tokens
else None,
),
)
def _parse_pricing(self, raw: Any, model_type: str) -> Pricing | None:
if not isinstance(raw, dict):
return None
# The ``extended`` tier some models charge past a context threshold is
# ignored: billing it would overcharge every request staying under it.
input_usd = _usd(raw.get("input"))
output_usd = _usd(raw.get("output"))
# Embeddings produce no completion tokens, so only they may omit an
# output price. Anywhere else a missing or all-zero price would serve
# completions free and a negative one would credit the caller, the
# same guards ``generic.py`` applies to this price book.
if output_usd is None and model_type == "embedding":
output_usd = 0.0
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
return Pricing(
prompt=input_usd / _USD_PER_MILLION,
completion=output_usd / _USD_PER_MILLION,
input_cache_read=(_usd(raw.get("cache_input")) or 0.0) / _USD_PER_MILLION,
input_cache_write=(_usd(raw.get("cache_write")) or 0.0) / _USD_PER_MILLION,
)
+294 -77
View File
@@ -15,12 +15,14 @@ import httpx
from cashu.core.base import MeltQuote, Proof, Token from cashu.core.base import MeltQuote, Proof, Token
from cashu.core.mint_info import MintInfo as _CashuMintInfo from cashu.core.mint_info import MintInfo as _CashuMintInfo
from cashu.wallet.crud import get_keysets as get_cashu_keysets from cashu.wallet.crud import get_keysets as get_cashu_keysets
from cashu.wallet.crud import get_proofs as get_cashu_proofs
from cashu.wallet.helpers import deserialize_token_from_string from cashu.wallet.helpers import deserialize_token_from_string
from cashu.wallet.wallet import Wallet as _CashuWallet from cashu.wallet.wallet import Wallet as _CashuWallet
from pydantic_core import PydanticUndefined from pydantic_core import PydanticUndefined
from sqlmodel import col, select, update from sqlmodel import col, select, update
from .cashu_compat import install_cashu_httpx_shim from .cashu_compat import install_cashu_httpx_shim
from .checkstate import filter_unspent_proofs
from .core import db, get_logger 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_with_retry as store_cashu_transaction
from .core.settings import settings from .core.settings import settings
@@ -35,7 +37,7 @@ from .mint import (
mint_cooldown_remaining, mint_cooldown_remaining,
run_mint_operation, run_mint_operation,
) )
from .payment.lnurl import raw_send_to_lnurl from .payment.lnurl import MeltUnpaidError, raw_send_to_lnurl
# cashu 0.20.x passes the `proxies` kwarg httpx removed in 0.28; see the module # cashu 0.20.x passes the `proxies` kwarg httpx removed in 0.28; see the module
# docstring. Installed at import so no mint call can run before the patch. # docstring. Installed at import so no mint call can run before the patch.
@@ -120,7 +122,7 @@ def _msats_to_sats_ceil(amount: int) -> int:
def _mints_to_inspect() -> list[str]: def _mints_to_inspect() -> list[str]:
"""Return configured mints plus the primary mint, without duplicates.""" """Return configured mints plus the primary mint, without duplicates."""
mint_urls = list(settings.cashu_mints) mint_urls = list(dict.fromkeys(settings.cashu_mints))
if settings.primary_mint and settings.primary_mint not in mint_urls: if settings.primary_mint and settings.primary_mint not in mint_urls:
mint_urls.append(settings.primary_mint) mint_urls.append(settings.primary_mint)
return mint_urls return mint_urls
@@ -143,6 +145,12 @@ class Wallet(_CashuWallet):
request=resp.request, request=resp.request,
response=resp, response=resp,
) )
if resp.status_code in {413, 500} and resp.request.url.path.endswith(
"/v1/checkstate"
):
# Preserve size/HTTP diagnostics even when a proxy or mint returns
# JSON with a detail field. Mutation error handling stays unchanged.
resp.raise_for_status()
try: try:
response_data = resp.json() response_data = resp.json()
except json.JSONDecodeError: except json.JSONDecodeError:
@@ -177,9 +185,12 @@ class Wallet(_CashuWallet):
pass pass
await self.load_mint_keysets(force_old_keysets) await self.load_mint_keysets(force_old_keysets)
await self.activate_keyset(keyset_id)
await self.load_mint_info(reload=True) await self.load_mint_info(reload=True)
# Arm on the fetch, not the activation: a unit the mint does not
# serve makes ``activate_keyset`` raise, and arming after it would
# refetch keysets on every call.
_mint_metadata_last_load[mint_url] = time.monotonic() _mint_metadata_last_load[mint_url] = time.monotonic()
await self.activate_keyset(keyset_id)
class MintConnectionError(Exception): class MintConnectionError(Exception):
@@ -694,18 +705,64 @@ class Bolt11PaymentPlan:
return maximum if self.unit == "sat" else (maximum + 999) // 1000 return maximum if self.unit == "sat" else (maximum + 999) // 1000
def _to_msats(amount: int, unit: str) -> int:
return _sats_to_msats(amount) if unit == "sat" else amount
async def _other_wallets_unreserved_msats(mint_url: str, unit: str) -> int:
"""Sum unreserved proofs of every other trusted wallet, in msats.
Every wallet shares one db, so two queries answer for all of them. Loading
a wallet per mint and unit instead refetched keysets from each mint on
every call and rate-limited them.
Read fresh, not from a wallet's snapshot: this total only ever raises the
payout ceiling, and a snapshot up to 30s stale could hide another
process's reservation.
"""
wallet = await get_wallet(mint_url, unit, load=False)
trusted = set(_mints_to_inspect())
origins: dict[str, tuple[str, str]] = {}
for keyset in await get_cashu_keysets(db=wallet.db):
keyset_unit = keyset.unit if isinstance(keyset.unit, str) else keyset.unit.name
origin = (keyset.mint_url, keyset_unit)
if origin == (mint_url, unit):
continue
if keyset.mint_url in trusted and keyset_unit in ("sat", "msat"):
origins[keyset.id] = origin
total = 0
for proof in await get_cashu_proofs(db=wallet.db):
proof_origin = origins.get(proof.id)
if proof_origin is None or proof.reserved:
continue
total += _to_msats(proof.amount, proof_origin[1])
return total
async def _owner_balance_for_mint_and_unit( async def _owner_balance_for_mint_and_unit(
mint_url: str, unit: str, proofs_balance: int mint_url: str, unit: str, proofs_balance: int
) -> int: ) -> int:
"""Return spendable node-owned funds without crossing user liabilities.""" """Return owner funds in one wallet, in that wallet's unit.
A key's refund mint is a preference, not funding provenance: a key topped
up from a second mint keeps its original refund mint. Hence two bounds —
the per-mint one keeps refunds serviceable from the mint they name, the
global one stops misattributed customer funds being paid out as profit.
"""
others_msats = await _other_wallets_unreserved_msats(mint_url, unit)
async with db.create_session() as session: async with db.create_session() as session:
# Refund mint is a preference, not funding provenance. Mirror payout's mint_liability = await db.user_liability_for_mint_and_unit(
# conservative rule and protect the full liability at every mint. session, mint_url, unit
user_liability = await db.total_user_liability(session) )
# API-key balances are stored in msats. Cashu ``sat`` proofs are not. total_liability = await db.total_user_liability(session)
if unit == "sat": proofs_msats = _to_msats(proofs_balance, unit)
user_liability = _msats_to_sats_ceil(user_liability) surplus_msats = min(
return max(0, proofs_balance - user_liability) proofs_msats - mint_liability,
proofs_msats + others_msats - total_liability,
)
# Cashu ``sat`` proofs are whole sats.
surplus = _msats_to_sats(surplus_msats) if unit == "sat" else surplus_msats
return max(0, surplus)
async def maximum_owner_cashu_balance_sats() -> int: async def maximum_owner_cashu_balance_sats() -> int:
@@ -780,9 +837,7 @@ async def _prepare_bolt11_payment(invoice: str) -> Bolt11PaymentPlan:
) )
if owner_balance < required: if owner_balance < required:
continue continue
owner_balance_msats = ( owner_balance_msats = _to_msats(owner_balance, unit)
owner_balance * 1000 if unit == "sat" else owner_balance
)
candidates.append( candidates.append(
(owner_balance_msats, wallet, proofs, quote, mint_url, unit) (owner_balance_msats, wallet, proofs, quote, mint_url, unit)
) )
@@ -1155,6 +1210,9 @@ _wallets: dict[str, Wallet] = {}
# Proofs require a shorter refresh interval than remote mint metadata. # Proofs require a shorter refresh interval than remote mint metadata.
_wallet_last_load: dict[str, float] = {} _wallet_last_load: dict[str, float] = {}
_wallet_last_mint_load: dict[str, float] = {} _wallet_last_mint_load: dict[str, float] = {}
# Metadata loads the mint answered but that left the wallet unusable, replayed
# for the reload interval so the failure costs one request, not one per call.
_wallet_mint_load_errors: dict[str, tuple[float, Exception]] = {}
_wallet_load_locks: dict[str, asyncio.Lock] = {} _wallet_load_locks: dict[str, asyncio.Lock] = {}
@@ -1165,6 +1223,7 @@ async def get_wallet(
retry_on_rate_limit: bool = True, retry_on_rate_limit: bool = True,
force_reload: bool = False, force_reload: bool = False,
load_proofs: bool = True, load_proofs: bool = True,
force_reload_proofs: bool = False,
) -> Wallet: ) -> Wallet:
global _wallets, _wallet_last_load, _wallet_last_mint_load, _wallet_load_locks global _wallets, _wallet_last_load, _wallet_last_mint_load, _wallet_load_locks
id = f"{mint_url}_{unit}" id = f"{mint_url}_{unit}"
@@ -1181,22 +1240,41 @@ async def get_wallet(
or last_mint_load is None or last_mint_load is None
or now - last_mint_load >= _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS or now - last_mint_load >= _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS
): ):
await run_mint_operation( cached_error = _wallet_mint_load_errors.get(id)
lambda: ( if (
_wallets[id].load_mint(force_refresh=True) not force_reload
if force_reload and cached_error is not None
else _wallets[id].load_mint() and now - cached_error[0] < _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS
), ):
op_name="load_mint", raise cached_error[1]
mint_url=mint_url, try:
retry_on_rate_limit=retry_on_rate_limit, await run_mint_operation(
) lambda: (
_wallets[id].load_mint(force_refresh=True)
if force_reload
else _wallets[id].load_mint()
),
op_name="load_mint",
mint_url=mint_url,
retry_on_rate_limit=retry_on_rate_limit,
)
except Exception as error:
# Transport failures and 429s stay retryable; the rate
# guard owns those. Anything else means the mint answered
# and still cannot serve this wallet.
if not (
is_mint_connection_error(error) or _is_mint_rate_limited(error)
):
_wallet_mint_load_errors[id] = (time.monotonic(), error)
raise
_wallet_mint_load_errors.pop(id, None)
_wallet_last_mint_load[id] = time.monotonic() _wallet_last_mint_load[id] = time.monotonic()
if load_proofs: if load_proofs:
last_proof_load = _wallet_last_load.get(id) last_proof_load = _wallet_last_load.get(id)
if ( if (
force_reload force_reload
or force_reload_proofs
or last_proof_load is None or last_proof_load is None
or now - last_proof_load or now - last_proof_load
>= _WALLET_PROOF_RELOAD_MIN_INTERVAL_SECONDS >= _WALLET_PROOF_RELOAD_MIN_INTERVAL_SECONDS
@@ -1231,29 +1309,9 @@ async def slow_filter_spend_proofs(
*, *,
retry_on_rate_limit: bool = True, retry_on_rate_limit: bool = True,
) -> list[Proof]: ) -> list[Proof]:
if not proofs: return await filter_unspent_proofs(
return [] proofs, wallet, retry_on_rate_limit=retry_on_rate_limit
_proofs = [] )
_spent_proofs = []
# Keep proof-state checks in large batches. Mint quotas count HTTP requests,
# so smaller batches make balance reads slower and more likely to hit 429s.
batch_size = 1000
for i in range(0, len(proofs), batch_size):
pb = proofs[i : i + batch_size]
proof_states = await run_mint_operation(
lambda: wallet.check_proof_state(pb),
op_name="check_proof_state",
mint_url=str(wallet.url),
retry_on_rate_limit=retry_on_rate_limit,
)
for proof, state in zip(pb, proof_states.states):
if str(state.state) != "spent":
_proofs.append(proof)
else:
_spent_proofs.append(proof)
if _spent_proofs:
await wallet.set_reserved_for_send(_spent_proofs, reserved=True)
return _proofs
class BalanceDetail(TypedDict, total=False): class BalanceDetail(TypedDict, total=False):
@@ -1540,17 +1598,117 @@ async def fetch_all_balances(
) )
PAYOUT_HISTORY_STALE_SECONDS = 600
async def _record_payout_history(
*,
quote_id: str,
bolt11: str,
amount_sats: int,
mint_url: str,
destination: str,
) -> None:
"""Best-effort history insert; a history failure must never block a payout."""
try:
async with db.create_session() as session:
await db.record_lightning_payout(
session,
quote_id=quote_id,
bolt11=bolt11,
amount_sats=amount_sats,
mint_url=mint_url,
destination=destination,
)
except Exception as e:
logger.error(
"Failed to record Lightning payout history",
extra={
"error": str(e),
"error_type": type(e).__name__,
"quote_id": quote_id,
"mint_url": mint_url,
},
)
async def _reconcile_stale_payout_history(mint_url: str, unit: str) -> None:
"""Resolve payout rows left pending by a crash or an ambiguous melt.
Runs under ``wallet_operation_guard``. Only writes what the mint asserts
(paid/unpaid); quotes still pending or unreachable are left for later.
"""
try:
cutoff = int(time.time()) - PAYOUT_HISTORY_STALE_SECONDS
async with db.create_session() as session:
stale = await db.list_unsettled_lightning_payouts(
session, mint_url, created_before=cutoff
)
for payout in stale:
quote_state = await _check_bolt11_payment_status_locked(
mint_url, unit, payout.payment_hash
)
if quote_state == "paid":
await _settle_payout_history(payout.payment_hash, status="paid")
elif quote_state == "unpaid":
await _settle_payout_history(payout.payment_hash, status="failed")
else:
continue
logger.info(
"Reconciled stale Lightning payout history",
extra={
"quote_id": payout.payment_hash,
"mint_url": mint_url,
"quote_state": quote_state,
},
)
except Exception as e:
logger.error(
"Failed to reconcile Lightning payout history",
extra={
"error": str(e),
"error_type": type(e).__name__,
"mint_url": mint_url,
},
)
async def _settle_payout_history(
quote_id: str, *, status: str, amount_sats: int | None = None
) -> None:
"""Best-effort history update after the external payment outcome is known."""
try:
async with db.create_session() as session:
await db.settle_lightning_payout(
session, quote_id, status=status, amount_sats=amount_sats
)
except Exception as e:
logger.error(
"Failed to update Lightning payout history",
extra={
"error": str(e),
"error_type": type(e).__name__,
"quote_id": quote_id,
"status": status,
},
)
async def _payout_mint_and_unit(mint_url: str, unit: str) -> None: async def _payout_mint_and_unit(mint_url: str, unit: str) -> None:
"""Send only conservatively proven owner funds for one wallet.""" """Send only conservatively proven owner funds for one wallet."""
try: try:
# Runs under wallet_operation_guard; a cached wallet may carry a proof # Runs under wallet_operation_guard; a cached wallet may carry a proof
# snapshot up to 30s stale from another process's reservation, so the # snapshot up to 30s stale from another process's reservation, so the
# cross-process lock is only safe with a fresh reload. # cross-process lock is only safe with fresh local proofs, not a
wallet = await get_wallet(mint_url, unit, force_reload=True) # forced network refresh of every keyset.
wallet = await get_wallet(mint_url, unit, force_reload_proofs=True)
proofs = get_proofs_per_mint_and_unit(wallet, mint_url, unit, not_reserved=True) proofs = get_proofs_per_mint_and_unit(wallet, mint_url, unit, not_reserved=True)
if not proofs: min_amount = (
# Nothing to pay out, so skip the settle delay rather than hold the settings.min_payout_sat
# cross-process guard (and block credits) for a wallet with no funds. if unit == "sat"
else _sats_to_msats(settings.min_payout_sat)
)
if sum(proof.amount for proof in proofs) <= min_amount:
return return
proofs = await slow_filter_spend_proofs(proofs, wallet) proofs = await slow_filter_spend_proofs(proofs, wallet)
await asyncio.sleep(5) await asyncio.sleep(5)
@@ -1561,15 +1719,13 @@ async def _payout_mint_and_unit(mint_url: str, unit: str) -> None:
) )
return return
# Fetch liability after the proofs snapshot and settle delay while the # Read liabilities and the other wallets' proofs after this wallet's proofs
# wallet operation guard excludes concurrent proof mutation and crediting. # snapshot and settle delay, while the wallet operation guard excludes
# concurrent proof mutation and crediting.
try: try:
async with db.create_session() as session: available_balance = await _owner_balance_for_mint_and_unit(
# ApiKey stores a refund preference, not funding provenance. Until mint_url, unit, sum(proof.amount for proof in proofs)
# liabilities have a durable per-credit ledger, subtract the total )
# liability from every wallet rather than risk calling customer
# funds owner profit on the wrong mint.
user_balance = await db.total_user_liability(session)
except Exception as e: except Exception as e:
logger.error( logger.error(
f"Error in periodic payout cycle: {type(e).__name__}", f"Error in periodic payout cycle: {type(e).__name__}",
@@ -1578,29 +1734,63 @@ async def _payout_mint_and_unit(mint_url: str, unit: str) -> None:
return return
try: try:
if unit == "sat": max_amount = (
user_balance = _msats_to_sats_ceil(user_balance) settings.max_payout_sat
proofs_balance = sum(proof.amount for proof in proofs)
available_balance = proofs_balance - user_balance
min_amount = (
settings.min_payout_sat
if unit == "sat" if unit == "sat"
else _sats_to_msats(settings.min_payout_sat) else _sats_to_msats(settings.max_payout_sat)
) )
if available_balance > min_amount: if available_balance > min_amount:
amount_received = await raw_send_to_lnurl( payout_amount = min(available_balance, max_amount)
wallet, payout_quote_id: str | None = None
proofs,
settings.receive_ln_address, async def record_payout(quote_id: str, bolt11: str) -> None:
unit, nonlocal payout_quote_id
amount=available_balance, payout_quote_id = quote_id
) await _record_payout_history(
quote_id=quote_id,
bolt11=bolt11,
amount_sats=(
payout_amount
if unit == "sat"
else _msats_to_sats(payout_amount)
),
mint_url=mint_url,
destination=settings.receive_ln_address,
)
try:
amount_received = await raw_send_to_lnurl(
wallet,
proofs,
settings.receive_ln_address,
unit,
amount=payout_amount,
on_melt_quote=record_payout,
)
except Exception as e:
if payout_quote_id is not None:
await _settle_payout_history(
payout_quote_id,
status=(
"failed"
if isinstance(e, MeltUnpaidError)
else "reconciliation_required"
),
)
raise
if payout_quote_id is not None:
await _settle_payout_history(
payout_quote_id,
status="paid",
amount_sats=_msats_to_sats(amount_received),
)
logger.info( logger.info(
"Payout sent successfully", "Payout sent successfully",
extra={ extra={
"mint_url": mint_url, "mint_url": mint_url,
"unit": unit, "unit": unit,
"balance": available_balance, "balance": available_balance,
"amount": payout_amount,
"amount_received": amount_received, "amount_received": amount_received,
}, },
) )
@@ -1640,6 +1830,7 @@ async def periodic_payout() -> None:
# Proof mutation, liability observation, and sending are one # Proof mutation, liability observation, and sending are one
# cross-process critical section. Credits take the same lock. # cross-process critical section. Credits take the same lock.
async with wallet_operation_guard(): async with wallet_operation_guard():
await _reconcile_stale_payout_history(mint_url, unit)
await _payout_mint_and_unit(mint_url, unit) await _payout_mint_and_unit(mint_url, unit)
except Exception as e: except Exception as e:
logger.error( logger.error(
@@ -1866,6 +2057,7 @@ async def periodic_routstr_fee_payout() -> None:
payout_unit, payout_unit,
) )
if completed: if completed:
await _settle_payout_history(payout_quote_id, status="paid")
logger.info( logger.info(
"Routstr fee payout reconciled as paid", "Routstr fee payout reconciled as paid",
extra={"payout_quote_id": payout_quote_id}, extra={"payout_quote_id": payout_quote_id},
@@ -1880,6 +2072,9 @@ async def periodic_routstr_fee_payout() -> None:
payout_unit, payout_unit,
) )
if restored: if restored:
await _settle_payout_history(
payout_quote_id, status="failed"
)
logger.warning( logger.warning(
"Routstr fee payout reconciled as unpaid and restored for retry", "Routstr fee payout reconciled as unpaid and restored for retry",
extra={"payout_quote_id": payout_quote_id}, extra={"payout_quote_id": payout_quote_id},
@@ -1910,7 +2105,7 @@ async def periodic_routstr_fee_payout() -> None:
attempt_quote_id: str | None = None attempt_quote_id: str | None = None
async def checkpoint_quote(quote_id: str) -> None: async def checkpoint_quote(quote_id: str, bolt11: str) -> None:
nonlocal attempt_quote_id nonlocal attempt_quote_id
async with db.create_session() as session: async with db.create_session() as session:
checkpointed = await db.reset_routstr_fee( checkpointed = await db.reset_routstr_fee(
@@ -1923,6 +2118,13 @@ async def periodic_routstr_fee_payout() -> None:
if not checkpointed: if not checkpointed:
raise _RoutstrFeePayoutAlreadyClaimed raise _RoutstrFeePayoutAlreadyClaimed
attempt_quote_id = quote_id attempt_quote_id = quote_id
await _record_payout_history(
quote_id=quote_id,
bolt11=bolt11,
amount_sats=accumulated_sats,
mint_url=settings.primary_mint,
destination=ROUTSTR_LN_ADDRESS,
)
try: try:
amount_received = await raw_send_to_lnurl( amount_received = await raw_send_to_lnurl(
@@ -1949,6 +2151,9 @@ async def periodic_routstr_fee_payout() -> None:
extra={"payout_in_progress_msats": paid_msats}, extra={"payout_in_progress_msats": paid_msats},
exc_info=isinstance(e, Exception), exc_info=isinstance(e, Exception),
) )
await _settle_payout_history(
attempt_quote_id, status="reconciliation_required"
)
if not isinstance(e, Exception): if not isinstance(e, Exception):
raise raise
continue continue
@@ -1969,6 +2174,9 @@ async def periodic_routstr_fee_payout() -> None:
extra={"payout_in_progress_msats": paid_msats}, extra={"payout_in_progress_msats": paid_msats},
exc_info=isinstance(e, Exception), exc_info=isinstance(e, Exception),
) )
await _settle_payout_history(
attempt_quote_id, status="reconciliation_required"
)
if not isinstance(e, Exception): if not isinstance(e, Exception):
raise raise
continue continue
@@ -1977,8 +2185,17 @@ async def periodic_routstr_fee_payout() -> None:
"Routstr fee payout sent but checkpoint was not completed; awaiting quote reconciliation", "Routstr fee payout sent but checkpoint was not completed; awaiting quote reconciliation",
extra={"payout_in_progress_msats": paid_msats}, extra={"payout_in_progress_msats": paid_msats},
) )
await _settle_payout_history(
attempt_quote_id, status="reconciliation_required"
)
continue continue
await _settle_payout_history(
attempt_quote_id,
status="paid",
amount_sats=_msats_to_sats(amount_received),
)
logger.info( logger.info(
"Routstr fee payout sent", "Routstr fee payout sent",
extra={ extra={
@@ -1995,8 +2212,8 @@ async def periodic_routstr_fee_payout() -> None:
def _quote_callback( def _quote_callback(
notify: Callable[[str, str], Awaitable[None]], mint: str notify: Callable[[str, str], Awaitable[None]], mint: str
) -> Callable[[str], Awaitable[None]]: ) -> Callable[[str, str], Awaitable[None]]:
async def callback(quote_id: str) -> None: async def callback(quote_id: str, _bolt11: str) -> None:
await notify(quote_id, mint) await notify(quote_id, mint)
return callback return callback
+10
View File
@@ -31,3 +31,13 @@ def _isolate_redemption_negative_cache() -> Iterator[None]:
redemption_negative_cache.clear() redemption_negative_cache.clear()
yield yield
redemption_negative_cache.clear() redemption_negative_cache.clear()
@pytest.fixture(autouse=True)
def _isolate_upstream_cooldowns() -> Iterator[None]:
"""Clear process-wide upstream cooldowns so one test's failures can't skip providers in the next."""
from routstr.upstream.cooldown import reset_cooldowns
reset_cooldowns()
yield
reset_cooldowns()
+5
View File
@@ -365,6 +365,10 @@ async def test_database_url(tmp_path: Any) -> str:
@pytest_asyncio.fixture @pytest_asyncio.fixture
async def integration_engine(test_database_url: str) -> AsyncGenerator[Any, None]: async def integration_engine(test_database_url: str) -> AsyncGenerator[Any, None]:
"""Create an async engine for integration tests""" """Create an async engine for integration tests"""
from routstr.core.settings import settings
# Match the production engine's busy timeout; sqlite3's 5s default makes
# concurrency tests flake with "database is locked" on slow CI runners.
engine = create_async_engine( engine = create_async_engine(
test_database_url, test_database_url,
echo=False, echo=False,
@@ -372,6 +376,7 @@ async def integration_engine(test_database_url: str) -> AsyncGenerator[Any, None
pool_pre_ping=True, pool_pre_ping=True,
pool_size=5, pool_size=5,
max_overflow=10, max_overflow=10,
connect_args={"timeout": settings.database_busy_timeout},
) )
# Initialize database schema # Initialize database schema
+126 -3
View File
@@ -27,6 +27,12 @@ EXPENSIVE_BASE_URL = "https://expensive.example.com/v1"
THIRD_BASE_URL = "https://third.example.com/v1" THIRD_BASE_URL = "https://third.example.com/v1"
@pytest.fixture(autouse=True)
def _no_upstream_5xx_retry_backoff(monkeypatch: pytest.MonkeyPatch) -> None:
"""Keep the same-upstream retry backoff out of the test runtime."""
monkeypatch.setattr("routstr.proxy._UPSTREAM_5XX_RETRY_BACKOFF_SECONDS", 0)
def _make_model( def _make_model(
model_id: str, model_id: str,
prompt_sats: float, prompt_sats: float,
@@ -191,14 +197,15 @@ async def test_failover_serve_billed_at_serving_providers_rate(
assert response.status_code == 200 assert response.status_code == 200
payload = response.json() payload = response.json()
# Both providers were attempted, cheapest first. # The winner is tried twice: its 502 is retried in place before failover.
assert [r.url.host for r in sent_requests] == [ assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"cheap.example.com", "cheap.example.com",
"expensive.example.com", "expensive.example.com",
] ]
# The fallback must be asked for ITS OWN model spelling, not the winner's. # The fallback must be asked for ITS OWN model spelling, not the winner's.
forwarded_body = json.loads(sent_requests[1].content) forwarded_body = json.loads(sent_requests[2].content)
assert forwarded_body["model"] == "provb/dual-model" assert forwarded_body["model"] == "provb/dual-model"
# The response echo names the model that actually served. # The response echo names the model that actually served.
@@ -319,7 +326,9 @@ async def test_same_id_failover_settles_at_serving_price(
) )
assert response.status_code == 200 assert response.status_code == 200
# The winner's 502 is retried in place before failover.
assert [r.url.host for r in sent_requests] == [ assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"cheap.example.com", "cheap.example.com",
"expensive.example.com", "expensive.example.com",
] ]
@@ -459,7 +468,9 @@ async def test_usd_cost_serve_carries_serving_providers_fee(
) )
assert response.status_code == 200 assert response.status_code == 200
# The winner's 502 is retried in place before failover.
assert [r.url.host for r in sent_requests] == [ assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"cheap.example.com", "cheap.example.com",
"expensive.example.com", "expensive.example.com",
] ]
@@ -530,7 +541,11 @@ async def test_failover_beyond_balance_envelope_is_rejected(
# The 20_000-sat envelope exceeds the key's 10_000-sat balance: the # The 20_000-sat envelope exceeds the key's 10_000-sat balance: the
# fallback must be rejected before its upstream is ever contacted. # fallback must be rejected before its upstream is ever contacted.
assert response.status_code == 402 assert response.status_code == 402
assert [r.url.host for r in sent_requests] == ["cheap.example.com"] # The winner is retried in place; the fallback is still never contacted.
assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"cheap.example.com",
]
@pytest.fixture @pytest.fixture
async def raised_envelope_provider_maps( async def raised_envelope_provider_maps(
patched_db_engine: None, patched_db_engine: None,
@@ -593,7 +608,9 @@ async def test_failover_reserves_serving_candidates_envelope(
) )
assert response.status_code == 200 assert response.status_code == 200
# The winner's 502 is retried in place before failover.
assert [r.url.host for r in sent_requests] == [ assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"cheap.example.com", "cheap.example.com",
"expensive.example.com", "expensive.example.com",
] ]
@@ -610,3 +627,109 @@ async def test_failover_reserves_serving_candidates_envelope(
charged = next(record for record in records if record.status == "charged") charged = next(record for record in records if record.status == "charged")
assert charged.reserved_msats > released.reserved_msats assert charged.reserved_msats > released.reserved_msats
assert all(record.status != "active" for record in records) assert all(record.status != "active" for record in records)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_transient_502_retries_same_upstream_before_failing_over(
authenticated_client: AsyncClient,
dual_provider_maps: tuple[_StaticProvider, _StaticProvider],
) -> None:
"""A transient 502 is retried on the same upstream, not failed over at once.
The winner answers the first attempt with a 502 and the second with a
completion, so the pricier fallback is never contacted and the request is
billed at the winner's rate (2_000 msats, not the fallback's 10_000).
"""
sent_requests: list[httpx.Request] = []
cheap_attempts = 0
async def fake_transport(
request: httpx.Request, *args: Any, **kwargs: Any
) -> httpx.Response:
nonlocal cheap_attempts
sent_requests.append(request)
if request.url.host == "cheap.example.com":
cheap_attempts += 1
if cheap_attempts == 1:
return httpx.Response(
502,
content=json.dumps({"error": {"message": "bad gateway"}}).encode(),
headers={"content-type": "application/json"},
)
return _successful_upstream_response()
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
# Retried in place: same host twice, fallback never consulted.
assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"cheap.example.com",
]
payload = response.json()
assert payload["model"] == "prova/dual-model"
# Billed at the winner's rate, not the fallback's 10_000.
assert payload["cost"]["total_msats"] == 2_000
@pytest.mark.integration
@pytest.mark.asyncio
async def test_transport_failure_is_not_retried_in_place(
authenticated_client: AsyncClient,
dual_provider_maps: tuple[_StaticProvider, _StaticProvider],
) -> None:
"""A transport failure fails over at once instead of retrying in place.
The proxy maps a connect/timeout error to a 502 of its own, so the upstream
may already have accepted and billed the request: re-sending it is not safe.
"""
sent_requests: list[httpx.Request] = []
async def fake_transport(
request: httpx.Request, *args: Any, **kwargs: Any
) -> httpx.Response:
sent_requests.append(request)
if request.url.host == "cheap.example.com":
raise httpx.ConnectError("connection refused", request=request)
return _successful_upstream_response()
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",
]
@@ -17,8 +17,10 @@ import pytest
from cashu.core.base import Proof from cashu.core.base import Proof
from sqlalchemy import inspect from sqlalchemy import inspect
from sqlalchemy.ext.asyncio import AsyncEngine from sqlalchemy.ext.asyncio import AsyncEngine
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core import db
from routstr.core.db import ApiKey, LightningInvoice from routstr.core.db import ApiKey, LightningInvoice
from routstr.lightning import _create_api_key_record from routstr.lightning import _create_api_key_record
@@ -81,6 +83,41 @@ async def test_invoice_persists_validity_date(
assert stored.validity_date == expiry assert stored.validity_date == expiry
@pytest.mark.asyncio
async def test_outgoing_payout_history_is_recorded_and_settled(
integration_session: AsyncSession,
) -> None:
await db.record_lightning_payout(
integration_session,
quote_id="payout-quote",
bolt11="lnbc1payout",
amount_sats=1_000,
mint_url="https://mint.test",
destination="owner@example.com",
)
result = await integration_session.exec(
select(LightningInvoice).where(LightningInvoice.payment_hash == "payout-quote")
)
payout = result.one()
assert payout.direction == "out"
assert payout.purpose == "payout"
assert payout.status == "pending"
assert payout.mint_url == "https://mint.test"
await db.settle_lightning_payout(
integration_session,
"payout-quote",
status="paid",
amount_sats=995,
)
await integration_session.refresh(payout)
assert payout.status == "paid"
assert payout.amount_sats == 995
assert payout.paid_at is not None
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Propagation to ApiKey # Propagation to ApiKey
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
+47 -9
View File
@@ -5,6 +5,7 @@ from unittest.mock import AsyncMock, Mock, patch
import pytest import pytest
from cashu.core.base import Proof from cashu.core.base import Proof
from fastapi import HTTPException
from sqlalchemy.ext.asyncio import AsyncEngine from sqlalchemy.ext.asyncio import AsyncEngine
from sqlmodel import col, update from sqlmodel import col, update
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
@@ -13,12 +14,15 @@ from routstr.core.db import ApiKey, LightningInvoice
from routstr.lightning import ( from routstr.lightning import (
INVOICE_EXPIRY_GRACE_SECONDS, INVOICE_EXPIRY_GRACE_SECONDS,
INVOICE_WATCH_BATCH_LIMIT, INVOICE_WATCH_BATCH_LIMIT,
InvoiceRecoverRequest,
_expire_invoice_if_authoritatively_unpaid, _expire_invoice_if_authoritatively_unpaid,
_expire_overdue_invoices, _expire_overdue_invoices,
_finalize_invoice_settlement, _finalize_invoice_settlement,
_InvoiceSettlement, _InvoiceSettlement,
_process_invoice_watch_batch, _process_invoice_watch_batch,
check_invoice_payment, check_invoice_payment,
get_invoice_status,
recover_invoice,
) )
@@ -58,7 +62,8 @@ async def test_invoice_read_transaction_closes_before_external_mint_io(
return wallet return wallet
with patch( with patch(
"routstr.lightning.get_wallet", side_effect=get_wallet_without_open_db_transaction "routstr.lightning.get_wallet",
side_effect=get_wallet_without_open_db_transaction,
): ):
await check_invoice_payment(stored, integration_session) await check_invoice_payment(stored, integration_session)
@@ -191,9 +196,7 @@ async def test_failed_final_commit_rolls_back_claim_and_credit_for_retry(
assert unchanged.balance == 100_000 assert unchanged.balance == 100_000
async with AsyncSession(integration_engine, expire_on_commit=False) as retry: async with AsyncSession(integration_engine, expire_on_commit=False) as retry:
settled, _ = await _finalize_invoice_settlement( settled, _ = await _finalize_invoice_settlement(snapshot, retry, 1_700_000_001)
snapshot, retry, 1_700_000_001
)
assert settled assert settled
async with AsyncSession(integration_engine, expire_on_commit=False) as verify: async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
@@ -302,9 +305,7 @@ async def test_expiry_cas_cannot_overwrite_concurrent_paid_invoice(
assert result.rowcount == 1 assert result.rowcount == 1
await paid.commit() await paid.commit()
expired = await _expire_invoice_if_authoritatively_unpaid( expired = await _expire_invoice_if_authoritatively_unpaid(stale, caller, True)
stale, caller, True
)
assert expired is False assert expired is False
assert stale.status == "paid" assert stale.status == "paid"
@@ -381,8 +382,9 @@ async def test_sweep_expires_only_overdue_pending_invoices(
overdue = _lightning_invoice(expires_at=now - 1) overdue = _lightning_invoice(expires_at=now - 1)
fresh = _lightning_invoice(expires_at=now + 3600) fresh = _lightning_invoice(expires_at=now + 3600)
settling = _lightning_invoice(expires_at=now - 1, status="settlement_pending") settling = _lightning_invoice(expires_at=now - 1, status="settlement_pending")
outgoing = _lightning_invoice(expires_at=now - 1, direction="out")
async with AsyncSession(integration_engine, expire_on_commit=False) as seed: async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
seed.add_all([overdue, fresh, settling]) seed.add_all([overdue, fresh, settling, outgoing])
await seed.commit() await seed.commit()
await _expire_overdue_invoices(now) await _expire_overdue_invoices(now)
@@ -392,6 +394,7 @@ async def test_sweep_expires_only_overdue_pending_invoices(
(overdue, "expired"), (overdue, "expired"),
(fresh, "pending"), (fresh, "pending"),
(settling, "settlement_pending"), (settling, "settlement_pending"),
(outgoing, "pending"),
): ):
stored = await verify.get(LightningInvoice, invoice.id) stored = await verify.get(LightningInvoice, invoice.id)
assert stored is not None assert stored is not None
@@ -409,8 +412,11 @@ async def test_watch_batch_expires_overdue_invoices_and_keeps_settling_rows(
settling = _lightning_invoice( settling = _lightning_invoice(
expires_at=now - 86_400, created_at=now - 86_400, status="settlement_pending" expires_at=now - 86_400, created_at=now - 86_400, status="settlement_pending"
) )
outgoing = _lightning_invoice(
expires_at=now + 3600, created_at=now, direction="out"
)
async with AsyncSession(integration_engine, expire_on_commit=False) as seed: async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
seed.add_all([overdue, fresh, settling]) seed.add_all([overdue, fresh, settling, outgoing])
await seed.commit() await seed.commit()
polled: list[str] = [] polled: list[str] = []
@@ -425,6 +431,7 @@ async def test_watch_batch_expires_overdue_invoices_and_keeps_settling_rows(
assert fresh.id in polled assert fresh.id in polled
assert settling.id in polled assert settling.id in polled
assert outgoing.id not in polled
async with AsyncSession(integration_engine, expire_on_commit=False) as verify: async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
stored = await verify.get(LightningInvoice, overdue.id) stored = await verify.get(LightningInvoice, overdue.id)
@@ -618,3 +625,34 @@ async def test_recovery_tail_cannot_starve_owed_or_live_invoices(
assert len(polled) == INVOICE_WATCH_BATCH_LIMIT assert len(polled) == INVOICE_WATCH_BATCH_LIMIT
assert {inv.id for inv in settling} <= set(polled) assert {inv.id for inv in settling} <= set(polled)
assert {inv.id for inv in fresh} <= set(polled) assert {inv.id for inv in fresh} <= set(polled)
@pytest.mark.asyncio
async def test_public_invoice_endpoints_ignore_payout_rows(
integration_engine: AsyncEngine,
patched_db_engine: None,
) -> None:
"""A payout's bolt11/id must not let /recover or /status touch the row."""
payout = _lightning_invoice(
direction="out",
purpose="payout",
expires_at=int(time.time()) - 1,
)
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
seed.add(payout)
await seed.commit()
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
with pytest.raises(HTTPException) as recover_error:
await recover_invoice(
InvoiceRecoverRequest(bolt11=payout.bolt11), session, False
)
with pytest.raises(HTTPException) as status_error:
await get_invoice_status(payout.id, session, False)
assert recover_error.value.status_code == 404
assert status_error.value.status_code == 404
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
stored = await verify.get(LightningInvoice, payout.id)
assert stored is not None
assert stored.status == "pending"
@@ -33,7 +33,6 @@ async def test_authenticated_proxy_releases_db_connection_before_upstream_header
request = MagicMock() request = MagicMock()
request.method = "POST" request.method = "POST"
request.headers = {"authorization": "Bearer test-key"} request.headers = {"authorization": "Bearer test-key"}
request.body = AsyncMock(return_value=json.dumps({"model": "test-model"}).encode())
request.url.path = "/v1/chat/completions" request.url.path = "/v1/chat/completions"
request.state.request_id = "pool-hold-regression" request.state.request_id = "pool-hold-regression"
@@ -59,7 +58,10 @@ async def test_authenticated_proxy_releases_db_connection_before_upstream_header
patch("routstr.proxy.get_bearer_token_key", AsyncMock(return_value=key)), patch("routstr.proxy.get_bearer_token_key", AsyncMock(return_value=key)),
): ):
response = await proxy_module._proxy( response = await proxy_module._proxy(
request, "v1/chat/completions", integration_session request,
"v1/chat/completions",
integration_session,
json.dumps({"model": "test-model"}).encode(),
) )
assert response.status_code == 200 assert response.status_code == 200
@@ -55,9 +55,13 @@ async def test_reserve_increases_reserved_balance(
cost = 100 cost = 100
key = await _persist(integration_session, _make_key(balance=500)) key = await _persist(integration_session, _make_key(balance=500))
await pay_for_request(key, cost, integration_session) reservation = await pay_for_request(key, cost, integration_session)
await integration_session.refresh(key) await integration_session.refresh(key)
assert reservation.key_hash == key.hashed_key
assert reservation.billing_key_hash == key.hashed_key
assert reservation.reserved_msats == cost
assert reservation.release_id
assert key.reserved_balance == cost assert key.reserved_balance == cost
assert key.balance == 500 # balance column is NOT decremented on reserve assert key.balance == 500 # balance column is NOT decremented on reserve
assert key.total_balance == 500 - cost # available = balance - reserved assert key.total_balance == 500 - cost # available = balance - reserved
@@ -75,11 +79,11 @@ async def test_revert_releases_reservation(
cost = 150 cost = 150
key = await _persist(integration_session, _make_key(balance=300)) key = await _persist(integration_session, _make_key(balance=300))
await pay_for_request(key, cost, integration_session) reservation = await pay_for_request(key, cost, integration_session)
await integration_session.refresh(key) await integration_session.refresh(key)
assert key.reserved_balance == cost assert key.reserved_balance == cost
await revert_pay_for_request(key, integration_session, cost) await revert_pay_for_request(key, integration_session, cost, reservation)
await integration_session.refresh(key) await integration_session.refresh(key)
assert key.reserved_balance == 0 assert key.reserved_balance == 0
@@ -470,7 +470,6 @@ async def test_startup_runs_bootstrap_before_settings_initialize(
return None return None
monkeypatch.setattr(main, "configure_litellm", lambda: 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, "run_migrations", lambda: None)
monkeypatch.setattr(main, "init_db", noop_init_db) monkeypatch.setattr(main, "init_db", noop_init_db)
monkeypatch.setattr(main, "create_session", fake_create_session) monkeypatch.setattr(main, "create_session", fake_create_session)
@@ -0,0 +1,131 @@
"""What Routstr actually puts on the wire for a Venice web-search request.
The unit tests stop at the kwargs handed to litellm. Everything that produced
the reported ``400 Unrecognized key(s) in object: 'web_search_options'``
happened *after* that point, inside litellm's Anthropic adapter, so this test
runs the whole dispatch against a loopback OpenAI-compatible server and reads
the bytes Venice would have received.
"""
from __future__ import annotations
import json
import threading
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import Any, Iterator
import pytest
from routstr.payment.models import Architecture, Model, Pricing
from routstr.upstream.litellm_routing import configure_litellm
from routstr.upstream.venice import VeniceUpstreamProvider
_CHUNKS = [
{
"id": "chatcmpl-1",
"object": "chat.completion.chunk",
"created": 0,
"model": "deepseek-v4-flash-0731",
"choices": [{"index": 0, "delta": {"role": "assistant", "content": "ok"}}],
},
{
"id": "chatcmpl-1",
"object": "chat.completion.chunk",
"created": 0,
"model": "deepseek-v4-flash-0731",
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7},
},
]
@pytest.fixture
def upstream() -> Iterator[tuple[str, dict[str, Any]]]:
"""A loopback stand-in for ``api.venice.ai`` that records one request."""
captured: dict[str, Any] = {}
class Handler(BaseHTTPRequestHandler):
def do_POST(self) -> None: # noqa: N802 - http.server's spelling
length = int(self.headers.get("Content-Length", 0))
captured["path"] = self.path
captured["body"] = json.loads(self.rfile.read(length))
self.send_response(200)
self.send_header("Content-Type", "text/event-stream")
self.end_headers()
for chunk in _CHUNKS:
self.wfile.write(f"data: {json.dumps(chunk)}\n\n".encode())
self.wfile.write(b"data: [DONE]\n\n")
def log_message(self, *args: Any) -> None:
return None
server = HTTPServer(("127.0.0.1", 0), Handler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield f"http://127.0.0.1:{server.server_address[1]}/v1", captured
finally:
server.shutdown()
thread.join(timeout=5)
def _model() -> Model:
return Model(
id="deepseek-v4-flash-0731",
name="deepseek-v4-flash-0731",
created=0,
description="",
context_length=8192,
architecture=Architecture(
modality="text->text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="Unknown",
instruct_type=None,
),
pricing=Pricing(prompt=0.0, completion=0.0),
)
@pytest.mark.asyncio
async def test_web_search_request_reaches_venice_in_its_own_shape(
upstream: tuple[str, dict[str, Any]],
) -> None:
base_url, captured = upstream
# The app applies this at startup; without it litellm posts the Anthropic
# body to /responses, which Venice serves only in alpha.
configure_litellm()
provider = VeniceUpstreamProvider(api_key="sk-test")
provider.base_url = base_url
await provider._dispatch_anthropic_messages(
request_body=json.dumps(
{
"model": "venice/deepseek-v4-flash-0731",
"messages": [{"role": "user", "content": "what shipped today?"}],
"max_tokens": 64,
"stream": True,
"tools": [
{"type": "web_search_20250305", "name": "web_search"},
{
"name": "lookup",
"description": "Look something up",
"input_schema": {"type": "object", "properties": {}},
},
],
}
).encode(),
model_obj=_model(),
)
body = captured["body"]
assert captured["path"] == "/v1/chat/completions"
# The reported 400, at the only place it could be observed.
assert "web_search_options" not in body
assert body["model"] == (
"deepseek-v4-flash-0731:enable_web_search=auto&enable_web_citations=true"
)
# The function tool still travels, in OpenAI's shape.
assert [tool["function"]["name"] for tool in body["tools"]] == ["lookup"]
+27
View File
@@ -0,0 +1,27 @@
"""Helpers for driving ``routstr.proxy.proxy`` with mocked request and session."""
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from typing import Any
from unittest.mock import MagicMock, patch
from routstr import proxy as proxy_module
def mock_request_stream(request: MagicMock, body: bytes) -> None:
"""Give a mocked request a readable body stream (the proxy reads the stream)."""
async def stream() -> AsyncIterator[bytes]:
yield body
request.stream = stream
def patch_proxy_session(session: Any) -> Any:
"""Make the proxy route use ``session`` instead of opening its own."""
@asynccontextmanager
async def factory() -> AsyncIterator[Any]:
yield session
return patch.object(proxy_module, "create_session", factory)
+140
View File
@@ -0,0 +1,140 @@
"""Bounded request-body read: size cap, read timeout, and late DB session."""
import asyncio
from collections.abc import AsyncIterator
from typing import Any
from unittest.mock import ANY, AsyncMock, MagicMock, patch
import pytest
from fastapi.responses import Response
from starlette.requests import Request
from routstr import proxy as proxy_module
from routstr.core.settings import settings
def _make_request(headers: dict[str, str], chunks: list[bytes]) -> MagicMock:
request = MagicMock()
request.method = "POST"
request.headers = headers
request.state.request_id = "req-bounded-body"
request.consumed = []
async def stream() -> AsyncIterator[bytes]:
for chunk in chunks:
request.consumed.append(chunk)
yield chunk
request.stream = stream
return request
def _slow_request(delay: float) -> MagicMock:
request = MagicMock()
request.method = "POST"
request.headers = {}
request.state.request_id = "req-slow-body"
async def stream() -> AsyncIterator[bytes]:
yield b"{"
await asyncio.sleep(delay)
yield b"}"
request.stream = stream
return request
async def _run(request: MagicMock) -> tuple[Any, MagicMock, AsyncMock]:
"""Run the proxy route with the session factory and _proxy stubbed out."""
session_factory = MagicMock()
inner = AsyncMock(return_value=Response(status_code=200))
with (
patch.object(proxy_module, "create_session", session_factory),
patch.object(proxy_module, "_proxy", inner),
):
response = await proxy_module.proxy(request, "v1/chat/completions")
return response, session_factory, inner
@pytest.mark.asyncio
async def test_oversize_content_length_rejected_without_reading() -> None:
request = _make_request({"content-length": "999999999"}, [b"x" * 16])
response, session_factory, inner = await _run(request)
assert response.status_code == 413
assert request.consumed == []
inner.assert_not_awaited()
session_factory.assert_not_called()
@pytest.mark.asyncio
async def test_oversize_chunked_body_rejected_mid_stream() -> None:
with patch.object(settings, "max_request_body_bytes", 8):
request = _make_request({}, [b"1234", b"5678", b"9012", b"3456"])
response, session_factory, inner = await _run(request)
assert response.status_code == 413
# Reading stops as soon as the cap is exceeded; the last chunk is never read.
assert request.consumed == [b"1234", b"5678", b"9012"]
inner.assert_not_awaited()
session_factory.assert_not_called()
@pytest.mark.asyncio
async def test_slow_body_times_out() -> None:
with patch.object(settings, "request_body_timeout_seconds", 0.05):
request = _slow_request(delay=5)
response, session_factory, inner = await _run(request)
assert response.status_code == 408
inner.assert_not_awaited()
session_factory.assert_not_called()
@pytest.mark.asyncio
async def test_normal_request_reaches_proxy_with_body() -> None:
body = b'{"model": "test-model"}'
request = _make_request({"content-length": str(len(body))}, [body])
response, session_factory, inner = await _run(request)
assert response.status_code == 200
session_factory.assert_called_once()
inner.assert_awaited_once_with(request, "v1/chat/completions", ANY, body)
def _starlette_request(body: bytes) -> Request:
messages: list[dict[str, Any]] = [
{"type": "http.request", "body": body, "more_body": False}
]
async def receive() -> dict[str, Any]:
return messages.pop(0) if messages else {"type": "http.disconnect"}
return Request(
{
"type": "http",
"method": "POST",
"headers": [(b"content-length", str(len(body)).encode())],
"path": "/v1/chat/completions",
"query_string": b"",
"state": {},
},
receive,
)
@pytest.mark.asyncio
async def test_body_stays_readable_after_bounded_read() -> None:
"""EHBP forwarding and upstream passthrough re-read the same request."""
body = b'{"model": "test-model"}'
request = _starlette_request(body)
assert await proxy_module._read_bounded_body(request) == body
assert await request.body() == body
streamed = bytearray()
async for chunk in request.stream():
streamed += chunk
assert bytes(streamed) == body
-9
View File
@@ -31,15 +31,6 @@ from routstr.payment.models import (
backfill_cache_pricing, backfill_cache_pricing,
) )
from routstr.upstream import GenericUpstreamProvider from routstr.upstream import GenericUpstreamProvider
from routstr.upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing
@pytest.fixture(autouse=True)
def _deepseek_v4_pricing() -> None:
# litellm's bundled cost map lacks the DeepSeek V4 entries (they only
# appear when its remote map is reachable); production injects them at
# startup via this same shim.
register_deepseek_v4_pricing()
def _make_model(model_id: str, pricing: Pricing) -> Model: def _make_model(model_id: str, pricing: Pricing) -> Model:
+314
View File
@@ -0,0 +1,314 @@
import asyncio
from collections.abc import Callable, Iterator
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock, patch
import httpx
import pytest
from cashu.core.base import Proof, ProofSpentState
from routstr import checkstate
from routstr.checkstate import _learned_sizes, filter_unspent_proofs
from routstr.mint import MintRateGuard, fail_fast_mint_operations
@pytest.fixture(autouse=True)
def isolate() -> Iterator[None]:
_learned_sizes.clear()
MintRateGuard._guards.clear()
yield
_learned_sizes.clear()
MintRateGuard._guards.clear()
def proofs(count: int) -> list[Proof]:
return [Mock(Y=str(i)) for i in range(count)]
def response(batch: list[Proof]) -> SimpleNamespace:
return SimpleNamespace(
states=[SimpleNamespace(Y=p.Y, state=ProofSpentState.unspent) for p in batch]
)
def rejection(status: int) -> httpx.HTTPStatusError:
request = httpx.Request("POST", "https://mint.test/v1/checkstate")
return httpx.HTTPStatusError(
"rejected",
request=request,
response=httpx.Response(
status,
request=request,
headers={"content-type": "text/html", "retry-after": "120"},
),
)
def wallet(check: Callable[[list[Proof]], object]) -> Mock:
return Mock(
url="https://mint.test",
check_proof_state=AsyncMock(side_effect=check),
set_reserved_for_send=AsyncMock(),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("status", [413, 500])
async def test_adapts_and_reuses_size_without_skipping_proofs(status: int) -> None:
checked: list[Proof] = []
async def check(batch: list[Proof]) -> SimpleNamespace:
if len(batch) > 120:
raise rejection(status)
checked.extend(batch)
return response(batch)
w = wallet(check)
ps = proofs(1001)
assert await filter_unspent_proofs(ps, w) == ps
assert checked == ps
sizes = [len(c.args[0]) for c in w.check_proof_state.await_args_list]
assert sizes[:5] == [1000, 500, 250, 125, 62]
w.check_proof_state.reset_mock()
assert await filter_unspent_proofs(ps, w) == ps
assert max(len(c.args[0]) for c in w.check_proof_state.await_args_list) == 62
@pytest.mark.asyncio
async def test_size_fallback_works_inside_cooldown_probe_under_wallet_guard() -> None:
async def check(batch: list[Proof]) -> SimpleNamespace:
if len(batch) > 2:
raise rejection(500)
return response(batch)
w = wallet(check)
MintRateGuard.get(w.url).apply_cooldown(0, reason="transport")
async with fail_fast_mint_operations():
ps = proofs(8)
assert await filter_unspent_proofs(ps, w) == ps
assert MintRateGuard.get(w.url).cooldown_remaining() == 0
@pytest.mark.asyncio
async def test_429_is_not_a_size_signal() -> None:
w = wallet(Mock(side_effect=rejection(429)))
with pytest.raises(httpx.HTTPStatusError):
await filter_unspent_proofs(proofs(1000), w, retry_on_rate_limit=False)
assert w.check_proof_state.await_count == 1
assert not _learned_sizes
assert MintRateGuard.get(w.url).cooldown_remaining() > 100
@pytest.mark.asyncio
@pytest.mark.parametrize("status", [400, 401, 422, 503])
async def test_other_http_errors_are_not_split(status: int) -> None:
w = wallet(Mock(side_effect=rejection(status)))
with pytest.raises(httpx.HTTPStatusError):
await filter_unspent_proofs(proofs(10), w)
assert w.check_proof_state.await_count == 1
@pytest.mark.asyncio
async def test_singleton_failure_is_bounded_and_does_not_poison_cache() -> None:
w = wallet(Mock(side_effect=rejection(500)))
with pytest.raises(httpx.HTTPStatusError):
await filter_unspent_proofs(proofs(1000), w)
assert [len(c.args[0]) for c in w.check_proof_state.await_args_list] == [
1000,
500,
250,
125,
62,
31,
15,
7,
3,
1,
]
assert not _learned_sizes
w.set_reserved_for_send.assert_not_awaited()
@pytest.mark.asyncio
async def test_request_budget_counts_successes_and_failures() -> None:
w = wallet(response)
with (
patch.object(checkstate, "_DEFAULT_BATCH_SIZE", 1),
patch.object(checkstate, "_MAX_REQUESTS", 2),
pytest.raises(ValueError, match="budget"),
):
await filter_unspent_proofs(proofs(3), w)
assert w.check_proof_state.await_count == 2
w.set_reserved_for_send.assert_not_awaited()
@pytest.mark.asyncio
async def test_total_deadline_cancels_slow_check() -> None:
async def check(batch: list[Proof]) -> None:
await asyncio.Event().wait()
w = wallet(check)
with (
patch.object(checkstate, "_SCAN_TIMEOUT_SECONDS", 0.01),
pytest.raises(TimeoutError),
):
await filter_unspent_proofs(proofs(1), w)
assert w.check_proof_state.await_count == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("malformation", ["missing", "reordered", "unknown"])
async def test_invalid_response_fails_closed(malformation: str) -> None:
def check(batch: list[Proof]) -> SimpleNamespace:
result = response(batch)
if malformation == "missing":
result.states.pop()
elif malformation == "reordered":
result.states.reverse()
else:
result.states[0].state = "UNKNOWN"
return result
w = wallet(check)
with pytest.raises(ValueError, match="Invalid proof-state"):
await filter_unspent_proofs(proofs(3), w)
w.set_reserved_for_send.assert_not_awaited()
@pytest.mark.asyncio
async def test_only_unspent_proofs_are_spendable() -> None:
ps = proofs(3)
states = [ProofSpentState.unspent, ProofSpentState.pending, ProofSpentState.spent]
w = wallet(
lambda batch: SimpleNamespace(
states=[SimpleNamespace(Y=p.Y, state=s) for p, s in zip(batch, states)]
)
)
assert await filter_unspent_proofs(ps, w) == ps[:1]
w.set_reserved_for_send.assert_awaited_once_with(ps[2:], reserved=True)
@pytest.mark.asyncio
async def test_learned_size_is_per_mint_and_expires() -> None:
w = wallet(response)
ps = proofs(5)
_learned_sizes[w.url] = (1, 0)
other = wallet(response)
other.url = "https://other.test"
_learned_sizes[other.url] = (1, float("inf"))
with patch.object(checkstate, "_DEFAULT_BATCH_SIZE", 2):
assert await filter_unspent_proofs(ps, w) == ps
assert [len(c.args[0]) for c in w.check_proof_state.await_args_list] == [
2,
2,
1,
]
assert await filter_unspent_proofs(ps, other) == ps
assert [len(c.args[0]) for c in other.check_proof_state.await_args_list] == [
1
] * 5
@pytest.mark.asyncio
async def test_smaller_later_batch_failure_does_not_skip_or_return_partial() -> None:
ps = proofs(9)
checked: list[Proof] = []
def check(batch: list[Proof]) -> SimpleNamespace:
if batch[0] is not ps[0] and len(batch) > 1:
raise rejection(500)
checked.extend(batch)
return response(batch)
w = wallet(check)
with patch.object(checkstate, "_DEFAULT_BATCH_SIZE", 4):
assert await filter_unspent_proofs(ps, w) == ps
assert checked == ps
@pytest.mark.parametrize("status", [413, 500])
@pytest.mark.parametrize("body", [{"detail": "too big"}, "<html>error</html>"])
def test_wallet_adapter_preserves_checkstate_http_status(
status: int, body: dict[str, str] | str
) -> None:
from routstr.wallet import Wallet
request = httpx.Request("POST", "https://mint.test/v1/checkstate")
reply = (
httpx.Response(status, request=request, json=body)
if isinstance(body, dict)
else httpx.Response(status, request=request, text=body)
)
with pytest.raises(httpx.HTTPStatusError) as error:
Wallet.raise_on_error_request(reply)
assert error.value.response is reply
@pytest.mark.asyncio
async def test_default_batch_fits_real_sdk_model() -> None:
from cashu.core.models import PostCheckStateRequest
limit = PostCheckStateRequest.model_json_schema()["properties"]["Ys"]["maxItems"]
ps = [
Proof(id="00", amount=1, secret=f"sdk-{i}", C="02" + "00" * 32)
for i in range(limit + 1)
]
sizes: list[int] = []
def check(batch: list[Proof]) -> SimpleNamespace:
payload = PostCheckStateRequest(Ys=[p.Y for p in batch])
sizes.append(len(payload.Ys))
return response(batch)
w = wallet(check)
assert await filter_unspent_proofs(ps, w) == ps
assert sizes == [limit, 1]
@pytest.mark.asyncio
async def test_scan_deadline_opens_cooldown_for_next_guarded_scan() -> None:
from routstr.mint import MintCooldownError
async def check(batch: list[Proof]) -> None:
await asyncio.Event().wait()
w = wallet(check)
with patch.object(checkstate, "_SCAN_TIMEOUT_SECONDS", 0.01):
async with fail_fast_mint_operations():
with pytest.raises(TimeoutError):
await filter_unspent_proofs(proofs(1), w)
with pytest.raises(MintCooldownError):
await filter_unspent_proofs(proofs(1), w)
assert w.check_proof_state.await_count == 1
assert MintRateGuard.get(w.url).cooldown_reason() == "transport"
@pytest.mark.asyncio
async def test_external_cancellation_does_not_open_cooldown() -> None:
started = asyncio.Event()
async def check(batch: list[Proof]) -> None:
started.set()
await asyncio.Event().wait()
w = wallet(check)
task = asyncio.create_task(filter_unspent_proofs(proofs(1), w))
await started.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert MintRateGuard.get(w.url).cooldown_remaining() == 0
@pytest.mark.asyncio
async def test_scan_deadline_preserves_longer_rate_limit_cooldown() -> None:
w = wallet(response)
guard = MintRateGuard.get(w.url)
guard.apply_rate_limit_cooldown(120)
until = guard._cooldown_until
with patch.object(checkstate, "_SCAN_TIMEOUT_SECONDS", 0.01):
with pytest.raises(TimeoutError):
await filter_unspent_proofs(proofs(1), w)
assert guard._cooldown_until == until
assert guard.cooldown_reason() == "rate_limited"
w.check_proof_state.assert_not_awaited()
+69 -9
View File
@@ -1,14 +1,24 @@
from __future__ import annotations from __future__ import annotations
import json
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from routstr.core.exceptions import EhbpTimeoutError, UpstreamError from routstr.core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS,
)
from routstr.core.exceptions import (
EhbpConnectionError,
EhbpTimeoutError,
UpstreamError,
)
from routstr.upstream import ehbp as ehbp_module from routstr.upstream import ehbp as ehbp_module
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# forward_ehbp_x_cashu_request — timeout fails closed with a refund + 504 # forward_ehbp_x_cashu_request — timeout fails closed with a refund + 424
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -48,7 +58,7 @@ def _ehbp_upstream_mocks() -> tuple[MagicMock, MagicMock]:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_x_cashu_timeout_refunds_and_returns_504( async def test_x_cashu_timeout_refunds_and_returns_424(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
monkeypatch.setattr( monkeypatch.setattr(
@@ -78,27 +88,32 @@ async def test_x_cashu_timeout_refunds_and_returns_504(
upstream=upstream, upstream=upstream,
) )
assert response.status_code == 504 assert response.status_code == UPSTREAM_ERROR_STATUS
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
assert response.headers["X-Cashu"] == "refund-token" assert response.headers["X-Cashu"] == "refund-token"
body = json.loads(bytes(response.body))
assert body["error"]["type"] == "upstream_timeout"
assert body["error"]["code"] == "UPSTREAM_TIMEOUT"
send_cashu_refund_mock.assert_awaited_once_with(1000, "msat", None, "req-123") send_cashu_refund_mock.assert_awaited_once_with(1000, "msat", None, "req-123")
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# forward_ehbp_request — the bearer path must let the timeout through, so # forward_ehbp_request — the bearer path must let the timeout through, so
# proxy.py can answer 504 instead of flattening it to a generic 500 # proxy.py can answer 424 instead of flattening it to a generic 500
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_bearer_timeout_propagates_504( async def test_bearer_timeout_propagates_424(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
"""A timed-out bearer request must not be rewritten to a 500. """A timed-out bearer request must not be rewritten to a 500.
``forward_ehbp_request`` ends in a bare ``except Exception`` that turns any ``forward_ehbp_request`` ends in a bare ``except Exception`` that turns any
error into ``UpstreamError(..., status_code=500)``. The ``except error into ``UpstreamError(..., status_code=500)``. The ``except
UpstreamError: raise`` above it is the only thing preserving the 504 that UpstreamError: raise`` above it is the only thing preserving the upstream
``proxy.py`` returns to the client, so this test pins that handler. timeout status that ``proxy.py`` returns to the client, so this test pins
that handler.
""" """
monkeypatch.setattr( monkeypatch.setattr(
ehbp_module, ehbp_module,
@@ -126,6 +141,51 @@ async def test_bearer_timeout_propagates_504(
model_obj=model_obj, model_obj=model_obj,
) )
assert exc_info.value.status_code == 504 assert exc_info.value.status_code == UPSTREAM_ERROR_STATUS
assert exc_info.value.code == "UPSTREAM_TIMEOUT" assert exc_info.value.code == "UPSTREAM_TIMEOUT"
assert exc_info.value.scope == ERROR_SCOPE_UPSTREAM
assert isinstance(exc_info.value, UpstreamError)
@pytest.mark.asyncio
async def test_bearer_connection_error_propagates_upstream_scope(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A connect failure must stay upstream-scoped instead of becoming a 500.
``forward_with_trailer`` classifies TLS/connection failures as
:class:`EhbpConnectionError`; ``forward_ehbp_request``'s ``except
UpstreamError: raise`` must let it through so ``proxy.py`` answers 424 with
the upstream scope header rather than a node-scoped 500.
"""
monkeypatch.setattr(
ehbp_module,
"forward_with_trailer",
AsyncMock(
side_effect=EhbpConnectionError(
"Unable to connect to EHBP upstream inference.tinfoil.sh: "
"ConnectionAbortedError"
)
),
)
upstream, model_obj = _ehbp_upstream_mocks()
key = MagicMock()
key.hashed_key = "abcdef1234567890"
with pytest.raises(EhbpConnectionError) as exc_info:
await ehbp_module.forward_ehbp_request(
request=await _request(),
path="v1/chat/completions",
headers={},
request_body=b"opaque",
upstream=upstream,
key=key,
max_cost_for_model=5000,
session=MagicMock(),
model_obj=model_obj,
)
assert exc_info.value.status_code == UPSTREAM_ERROR_STATUS
assert exc_info.value.code == "UPSTREAM_UNAVAILABLE"
assert exc_info.value.scope == ERROR_SCOPE_UPSTREAM
assert isinstance(exc_info.value, UpstreamError) assert isinstance(exc_info.value, UpstreamError)
+37 -10
View File
@@ -1,5 +1,5 @@
import asyncio import asyncio
from collections.abc import AsyncGenerator from collections.abc import AsyncGenerator, Generator
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock, patch from unittest.mock import AsyncMock, Mock, patch
@@ -29,6 +29,21 @@ def _session_context(session: Mock) -> _SessionContext:
return _SessionContext(session) return _SessionContext(session)
@pytest.fixture(autouse=True)
def _mock_lightning_payout_history() -> Generator[
tuple[AsyncMock, AsyncMock], None, None
]:
with (
patch(
"routstr.wallet.db.record_lightning_payout", new_callable=AsyncMock
) as record,
patch(
"routstr.wallet.db.settle_lightning_payout", new_callable=AsyncMock
) as settle,
):
yield record, settle
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_fee_payout_checkpoint_is_atomic_and_durable() -> None: async def test_fee_payout_checkpoint_is_atomic_and_durable() -> None:
engine = create_async_engine("sqlite+aiosqlite://") engine = create_async_engine("sqlite+aiosqlite://")
@@ -178,7 +193,7 @@ async def test_fee_payout_prepares_wallet_then_checkpoints_before_sending() -> N
async def send(*_args: object, **kwargs: object) -> int: async def send(*_args: object, **kwargs: object) -> int:
checkpoint_quote = kwargs["on_melt_quote"] checkpoint_quote = kwargs["on_melt_quote"]
await checkpoint_quote("quote-1") # type: ignore[operator] await checkpoint_quote("quote-1", "lnbc1payout") # type: ignore[operator]
events.append("send") events.append("send")
return 5 return 5
@@ -256,7 +271,7 @@ async def test_fee_payout_lost_checkpoint_race_does_not_send() -> None:
async def send(*_args: object, **kwargs: object) -> int: async def send(*_args: object, **kwargs: object) -> int:
checkpoint_quote = kwargs["on_melt_quote"] checkpoint_quote = kwargs["on_melt_quote"]
await checkpoint_quote("quote-1") # type: ignore[operator] await checkpoint_quote("quote-1", "lnbc1payout") # type: ignore[operator]
await dispatched() await dispatched()
return 5 return 5
@@ -289,7 +304,9 @@ async def test_fee_payout_lost_checkpoint_race_does_not_send() -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_fee_payout_finalizes_a_paid_unresolved_quote_without_resending() -> None: async def test_fee_payout_finalizes_a_paid_unresolved_quote_without_resending(
_mock_lightning_payout_history: tuple[AsyncMock, AsyncMock],
) -> None:
session = Mock() session = Mock()
fee = SimpleNamespace( fee = SimpleNamespace(
accumulated_msats=10_000, accumulated_msats=10_000,
@@ -331,10 +348,14 @@ async def test_fee_payout_finalizes_a_paid_unresolved_quote_without_resending()
) )
restore.assert_not_awaited() restore.assert_not_awaited()
send.assert_not_awaited() send.assert_not_awaited()
_, settle = _mock_lightning_payout_history
settle.assert_awaited_once_with(session, "quote-1", status="paid", amount_sats=None)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_fee_payout_restores_only_an_unpaid_quote_and_retries() -> None: async def test_fee_payout_restores_only_an_unpaid_quote_and_retries(
_mock_lightning_payout_history: tuple[AsyncMock, AsyncMock],
) -> None:
session = Mock() session = Mock()
unresolved_fee = SimpleNamespace( unresolved_fee = SimpleNamespace(
accumulated_msats=10_000, accumulated_msats=10_000,
@@ -351,7 +372,9 @@ async def test_fee_payout_restores_only_an_unpaid_quote_and_retries() -> None:
) )
async def send(*_args: object, **kwargs: object) -> int: async def send(*_args: object, **kwargs: object) -> int:
await kwargs["on_melt_quote"]("quote-2") # type: ignore[index,operator] await kwargs["on_melt_quote"]( # type: ignore[index,operator]
"quote-2", "lnbc1payout"
)
return 15 return 15
with ( with (
@@ -398,6 +421,8 @@ async def test_fee_payout_restores_only_an_unpaid_quote_and_retries() -> None:
session, 15_000, "quote-2", wallet.settings.primary_mint, "sat" session, 15_000, "quote-2", wallet.settings.primary_mint, "sat"
) )
raw_send.assert_awaited_once() raw_send.assert_awaited_once()
_, settle = _mock_lightning_payout_history
settle.assert_any_await(session, "quote-1", status="failed", amount_sats=None)
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -579,7 +604,7 @@ async def test_fee_payout_keeps_checkpoint_when_send_outcome_is_unknown() -> Non
async def send(*_args: object, **kwargs: object) -> int: async def send(*_args: object, **kwargs: object) -> int:
checkpoint_quote = kwargs["on_melt_quote"] checkpoint_quote = kwargs["on_melt_quote"]
await checkpoint_quote("quote-1") # type: ignore[operator] await checkpoint_quote("quote-1", "lnbc1payout") # type: ignore[operator]
raise TimeoutError("unknown outcome") raise TimeoutError("unknown outcome")
with ( with (
@@ -620,7 +645,7 @@ async def test_fee_payout_cancellation_during_send_alerts_and_propagates() -> No
async def cancel_send(*_args: object, **kwargs: object) -> int: async def cancel_send(*_args: object, **kwargs: object) -> int:
checkpoint_quote = kwargs["on_melt_quote"] checkpoint_quote = kwargs["on_melt_quote"]
await checkpoint_quote("quote-1") # type: ignore[operator] await checkpoint_quote("quote-1", "lnbc1payout") # type: ignore[operator]
raise asyncio.CancelledError raise asyncio.CancelledError
with ( with (
@@ -666,6 +691,8 @@ async def test_fee_payout_completion_failures_use_sent_checkpoint_alert(
side_effect=[ side_effect=[
_session_context(session), _session_context(session),
_session_context(session), _session_context(session),
_session_context(session),
RuntimeError("pool unavailable"),
RuntimeError("pool unavailable"), RuntimeError("pool unavailable"),
] ]
) )
@@ -675,7 +702,7 @@ async def test_fee_payout_completion_failures_use_sent_checkpoint_alert(
async def send(*_args: object, **kwargs: object) -> int: async def send(*_args: object, **kwargs: object) -> int:
checkpoint_quote = kwargs["on_melt_quote"] checkpoint_quote = kwargs["on_melt_quote"]
await checkpoint_quote("quote-1") # type: ignore[operator] await checkpoint_quote("quote-1", "lnbc1payout") # type: ignore[operator]
return 5 return 5
with ( with (
@@ -726,7 +753,7 @@ async def test_fee_payout_releases_db_connection_during_send(tmp_path: object) -
async def send(*_args: object, **kwargs: object) -> int: async def send(*_args: object, **kwargs: object) -> int:
assert engine.pool.checkedout() == 0 # type: ignore[attr-defined] assert engine.pool.checkedout() == 0 # type: ignore[attr-defined]
checkpoint_quote = kwargs["on_melt_quote"] checkpoint_quote = kwargs["on_melt_quote"]
await checkpoint_quote("quote-1") # type: ignore[operator] await checkpoint_quote("quote-1", "lnbc1payout") # type: ignore[operator]
assert engine.pool.checkedout() == 0 # type: ignore[attr-defined] assert engine.pool.checkedout() == 0 # type: ignore[attr-defined]
return 5 return 5
+1
View File
@@ -37,6 +37,7 @@ def _invoice(**overrides: object) -> SimpleNamespace:
"payment_hash": "quote-1", "payment_hash": "quote-1",
"amount_sats": 100, "amount_sats": 100,
"purpose": "create", "purpose": "create",
"direction": "in",
"status": "pending", "status": "pending",
"paid_at": None, "paid_at": None,
"api_key_hash": None, "api_key_hash": None,
+21
View File
@@ -91,3 +91,24 @@ def test_detect_litellm_prefix_custom_default() -> None:
assert detect_litellm_prefix("https://example.com", default="anthropic/") == ( assert detect_litellm_prefix("https://example.com", default="anthropic/") == (
"anthropic/" "anthropic/"
) )
@pytest.mark.parametrize("model", ["gpt-6", "gpt-6-luna", "gpt-5.5"])
def test_litellm_sends_max_completion_tokens_for_gpt_5_and_later(model: str) -> None:
"""OpenAI rejects ``max_tokens`` on these models; litellm <1.101 only
rewrote it for names containing ``gpt-5``, so gpt-6 got a 400."""
import litellm
from litellm.utils import ProviderConfigManager
config = ProviderConfigManager.get_provider_chat_config(
model=model, provider=litellm.LlmProviders.OPENAI
)
assert config is not None
mapped = config.map_openai_params(
non_default_params={"max_tokens": 10},
optional_params={},
model=model,
drop_params=True,
)
assert mapped == {"max_completion_tokens": 10}
@@ -192,7 +192,7 @@ async def test_raw_send_to_lnurl_requotes_for_exact_input_fees_without_recursion
assert paid == 485_000 assert paid == 485_000
assert wallet.melt_quote.await_count == 2 assert wallet.melt_quote.await_count == 2
checkpoint.assert_awaited_once_with("q2") checkpoint.assert_awaited_once_with("q2", "lnbc1...")
wallet.select_to_send.assert_not_called() wallet.select_to_send.assert_not_called()
selected = wallet.melt.await_args.kwargs["proofs"] selected = wallet.melt.await_args.kwargs["proofs"]
assert sum(proof.amount for proof in selected) == 500 assert sum(proof.amount for proof in selected) == 500
@@ -358,9 +358,7 @@ def _patch_getaddrinfo(ip: str) -> Any:
loop = MagicMock() loop = MagicMock()
loop.getaddrinfo = fake_getaddrinfo loop.getaddrinfo = fake_getaddrinfo
return patch.object( return patch.object(lnurl_module.asyncio, "get_running_loop", return_value=loop)
lnurl_module.asyncio, "get_running_loop", return_value=loop
)
@pytest.mark.asyncio @pytest.mark.asyncio
+181
View File
@@ -0,0 +1,181 @@
from collections.abc import AsyncIterator, Iterator
from contextlib import asynccontextmanager
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock, patch
import pytest
from cashu.core.base import BlindedMessage, BlindedSignature, Proof, Unit
from cashu.core.crypto import b_dhke
from cashu.core.models import PostMeltQuoteResponse
from cashu.wallet.v1_api import LedgerAPI
from cashu.wallet.wallet import Wallet as CashuWallet
from routstr.core.settings import settings
from routstr.mint import MintRateGuard
from routstr.wallet import _payout_mint_and_unit
@pytest.fixture(autouse=True)
def empty_cross_wallet_proofs() -> Iterator[None]:
"""No other wallet holds proofs, so only this wallet's own bound applies."""
with (
patch("routstr.wallet.get_cashu_keysets", AsyncMock(return_value=[])),
patch("routstr.wallet.get_cashu_proofs", AsyncMock(return_value=[])),
):
yield
@pytest.mark.asyncio
@pytest.mark.parametrize("unit,scale", [("sat", 1), ("msat", 1000)])
@pytest.mark.parametrize(
"liability,input_fee,reserve,actual_fee",
[(0, 0, 0, 0), (0, 7, 10, 3), (300000, 7, 10, 3)],
)
async def test_capped_payout_recovers_all_change_with_real_cashu_sdk(
unit: str,
scale: int,
liability: int,
input_fee: int,
reserve: int,
actual_fee: int,
) -> None:
MintRateGuard._guards.clear()
private_key = b_dhke.PrivateKey()
proof = Proof(
id="00",
amount=524288 * scale,
secret="input-proof",
C=private_key.public_key.format().hex(),
)
w = CashuWallet.__new__(CashuWallet)
w.url = "https://mint.test"
w.unit = Unit[unit]
w.keyset_id = "00"
w.keysets = {
"00": SimpleNamespace(
public_keys={2**i: private_key.public_key for i in range(40)}
)
}
w.proofs = [proof]
w.db = Mock()
w.get_fees_for_proofs = Mock(return_value=input_fee * scale)
w.set_reserved_for_send = AsyncMock()
w.set_reserved_for_melt = AsyncMock()
w.sign_proofs_inplace_melt = Mock(side_effect=lambda ps, outputs, quote: ps)
w._store_proofs = AsyncMock()
async def invalidate(ps: list[Proof]) -> None:
w.proofs = [p for p in w.proofs if p not in ps]
w.invalidate = AsyncMock(side_effect=invalidate)
w.generate_n_secrets = AsyncMock(
side_effect=lambda n: (
[f"change-{i}" for i in range(n)],
[],
[f"path-{i}" for i in range(n)],
)
)
quotes: dict[str, PostMeltQuoteResponse] = {}
async def quote(invoice: str) -> PostMeltQuoteResponse:
amount_msat = int(invoice)
amount = amount_msat // 1000 if unit == "sat" else amount_msat
q = PostMeltQuoteResponse(
quote=str(amount),
amount=amount,
unit=unit,
request=invoice,
fee_reserve=reserve * scale,
state="UNPAID",
expiry=None,
)
quotes[q.quote] = q
return q
w.melt_quote = AsyncMock(side_effect=quote)
selected_total = 0
returned_change = 0
blank_count = 0
paid_amount = 0
async def mint_melt(
quote_id: str, inputs: list[Proof], outputs: list[BlindedMessage]
) -> PostMeltQuoteResponse:
nonlocal selected_total, returned_change, blank_count, paid_amount
q = quotes[quote_id]
selected_total = sum(p.amount for p in inputs)
paid_amount = q.amount
blank_count = len(outputs)
assert q.fee_reserve == reserve * scale
change = selected_total - q.amount - (input_fee + actual_fee) * scale
amounts = [2**i for i in range(change.bit_length()) if change & (2**i)]
signatures = []
for amount, output in zip(amounts, outputs):
blinded, _, _ = b_dhke.step2_bob(
b_dhke.PublicKey(bytes.fromhex(output.B_)), private_key
)
signatures.append(
BlindedSignature(id="00", amount=amount, C_=blinded.format().hex())
)
returned_change = sum(s.amount for s in signatures)
assert returned_change == change
return q.model_copy(update={"state": "PAID", "change": signatures})
@asynccontextmanager
async def session() -> AsyncIterator[Mock]:
yield Mock()
with (
patch.object(settings, "max_payout_sat", 250000),
patch.object(settings, "min_payout_sat", 210),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=w)),
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[proof]),
patch(
"routstr.wallet.slow_filter_spend_proofs", AsyncMock(return_value=[proof])
),
patch("routstr.wallet.asyncio.sleep", AsyncMock()),
patch("routstr.wallet.db.create_session", session),
patch(
"routstr.wallet.db.total_user_liability",
AsyncMock(return_value=liability * 1000),
),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=liability * 1000),
),
patch(
"routstr.payment.lnurl.get_lnurl_data",
AsyncMock(
return_value={
"callback_url": "https://ln.test/cb",
"min_sendable": 1000,
"max_sendable": 10**12,
}
),
),
patch(
"routstr.payment.lnurl.get_lnurl_invoice",
AsyncMock(side_effect=lambda callback, amount: (str(amount), {})),
),
patch.object(LedgerAPI, "melt", AsyncMock(side_effect=mint_melt)) as transport,
patch("cashu.wallet.wallet.update_bolt11_melt_quote", AsyncMock()),
):
await _payout_mint_and_unit(w.url, unit)
transport.assert_awaited_once()
assert selected_total == 524288 * scale
assert blank_count > 0
assert sum(p.amount for p in w.proofs) == returned_change
assert all(
b_dhke.verify(private_key, b_dhke.PublicKey(bytes.fromhex(p.C)), p.secret)
for p in w.proofs
)
net_debit = selected_total - returned_change
assert net_debit == paid_amount + (input_fee + actual_fee) * scale
assert net_debit <= min(250000, 524288 - liability) * scale
assert returned_change >= liability * scale
if liability == input_fee == reserve == actual_fee == 0:
assert returned_change == 274288 * scale
assert net_debit == 250000 * scale
w._store_proofs.assert_awaited_once()
MintRateGuard._guards.clear()
+29 -1
View File
@@ -65,6 +65,33 @@ def _lnurl_patches() -> tuple[Any, Any]:
) )
@pytest.mark.asyncio
@pytest.mark.parametrize("outcome", ["timeout", "pending"])
async def test_oversized_proof_change_budget_preserves_ambiguous_melt(
outcome: str,
) -> None:
wallet, proofs = _wallet()
proofs[0].amount = 524288
if outcome == "timeout":
wallet.melt.side_effect = httpx.ReadTimeout("response lost")
else:
wallet.melt.return_value = MagicMock(state=MeltQuoteState.pending)
wallet.get_melt_quote = AsyncMock(
return_value=MagicMock(state=MeltQuoteState.pending)
)
data_patch, invoice_patch = _lnurl_patches()
with data_patch, invoice_patch, pytest.raises(MeltOutcomeAmbiguousError):
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
wallet.melt.assert_awaited_once()
assert wallet.melt.await_args.kwargs["fee_reserve_sat"] == 524288 - QUOTE_AMOUNT_SAT
wallet.set_reserved_for_send.assert_awaited_once_with(proofs, reserved=True)
if outcome == "timeout":
wallet.set_reserved_for_melt.assert_awaited_once_with(
proofs, reserved=True, quote_id="q"
)
wallet.get_melt_quote.assert_awaited_once_with("q")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_raw_send_to_lnurl_direct_unpaid_is_retry_safe() -> None: async def test_raw_send_to_lnurl_direct_unpaid_is_retry_safe() -> None:
wallet, proofs = _wallet() wallet, proofs = _wallet()
@@ -319,8 +346,9 @@ async def test_raw_send_to_lnurl_checkpoints_quote_before_melt_dispatch() -> Non
wallet, proofs = _wallet() wallet, proofs = _wallet()
events: list[str] = [] events: list[str] = []
async def checkpoint(quote_id: str) -> None: async def checkpoint(quote_id: str, bolt11: str) -> None:
assert quote_id == "q" assert quote_id == "q"
assert bolt11 == "lnbc1..."
events.append("checkpoint") events.append("checkpoint")
async def melt(**_kwargs: object) -> MagicMock: async def melt(**_kwargs: object) -> MagicMock:
@@ -0,0 +1,383 @@
"""Model and provider attribution on the request completion/failure log lines."""
import json
import logging
from collections.abc import Iterator
from contextlib import contextmanager
from pathlib import Path
from typing import Any
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
from fastapi import FastAPI, HTTPException, Request
from fastapi.responses import Response
from fastapi.testclient import TestClient
from httpx import ASGITransport, AsyncClient
from pythonjsonlogger import jsonlogger
from routstr import proxy as proxy_module
from routstr.core.db import get_session
from routstr.core.exceptions import UpstreamError
from routstr.core.logging import (
DailyRotatingFileHandler,
RequestIdFilter,
SecurityFilter,
VersionFilter,
)
from routstr.core.middleware import LoggingMiddleware, _attribution
@pytest.fixture
def handler(tmp_path: Path) -> Iterator[DailyRotatingFileHandler]:
log_dir = tmp_path / "logs"
log_dir.mkdir()
h = DailyRotatingFileHandler(
str(log_dir / "app.log"), when="midnight", interval=1, backupCount=30
)
h.setLevel(logging.DEBUG)
h.setFormatter(
jsonlogger.JsonFormatter(
"%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s "
"%(lineno)d %(version)s %(request_id)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
)
for f in (VersionFilter(), RequestIdFilter(), SecurityFilter()):
h.addFilter(f)
try:
yield h
finally:
h.close()
@contextmanager
def _middleware_logs_to(handler: DailyRotatingFileHandler) -> Iterator[None]:
middleware_logger = logging.getLogger("routstr.core.middleware")
saved = (
middleware_logger.handlers,
middleware_logger.level,
middleware_logger.propagate,
)
middleware_logger.handlers = [handler]
middleware_logger.setLevel(logging.INFO)
middleware_logger.propagate = False
try:
yield
finally:
(
middleware_logger.handlers,
middleware_logger.level,
middleware_logger.propagate,
) = saved
def _record(handler: DailyRotatingFileHandler, message: str) -> dict[str, Any]:
handler.flush()
lines = Path(handler.baseFilename).read_text().strip().splitlines()
matches = [r for r in map(json.loads, lines) if r.get("message") == message]
assert matches, f"no {message!r} record was written"
return matches[-1]
# --------------------------------------------------------------------------- #
# Middleware: fields land on the log lines.
# --------------------------------------------------------------------------- #
def _middleware_app() -> FastAPI:
app = FastAPI()
app.add_middleware(LoggingMiddleware)
@app.post("/v1/chat/completions")
async def completions(request: Request) -> dict:
request.state.model = "z-ai/glm-5.3-flash"
request.state.provider = "openrouter"
return {"ok": True}
@app.post("/v1/models")
async def models() -> dict:
return {"ok": True}
@app.post("/v1/broken")
async def broken(request: Request) -> dict:
request.state.model = "deepseek/deepseek-v4.1-flash"
request.state.provider = "venice"
raise RuntimeError("upstream exploded")
return app
def test_completion_log_carries_model_and_provider(
handler: DailyRotatingFileHandler,
) -> None:
with _middleware_logs_to(handler):
with TestClient(_middleware_app(), raise_server_exceptions=False) as client:
response = client.post(
"/v1/chat/completions", json={"model": "glm-5.3-flash"}
)
assert response.status_code == 200
rec = _record(handler, "Request completed")
assert rec["model"] == "z-ai/glm-5.3-flash"
assert rec["provider"] == "openrouter"
assert rec["status_code"] == 200
def test_completion_log_omits_attribution_when_route_sets_none(
handler: DailyRotatingFileHandler,
) -> None:
with _middleware_logs_to(handler):
with TestClient(_middleware_app(), raise_server_exceptions=False) as client:
response = client.post("/v1/models", json={})
assert response.status_code == 200
rec = _record(handler, "Request completed")
assert "model" not in rec
assert "provider" not in rec
def test_failed_request_log_carries_attribution(
handler: DailyRotatingFileHandler,
) -> None:
with _middleware_logs_to(handler):
with TestClient(_middleware_app(), raise_server_exceptions=False) as client:
response = client.post("/v1/broken", json={})
assert response.status_code == 500
rec = _record(handler, "Request failed")
assert rec["model"] == "deepseek/deepseek-v4.1-flash"
assert rec["provider"] == "venice"
assert rec["error_type"] == "RuntimeError"
def test_middleware_logger_state_is_restored(
handler: DailyRotatingFileHandler,
) -> None:
middleware_logger = logging.getLogger("routstr.core.middleware")
before = (
list(middleware_logger.handlers),
middleware_logger.level,
middleware_logger.propagate,
)
with _middleware_logs_to(handler):
pass
after = (
list(middleware_logger.handlers),
middleware_logger.level,
middleware_logger.propagate,
)
assert after == before
# --------------------------------------------------------------------------- #
# Proxy: which model/provider each routing path attributes the request to.
# --------------------------------------------------------------------------- #
def _model(model_id: str) -> MagicMock:
return MagicMock(id=model_id)
def _upstream(provider_type: str) -> MagicMock:
upstream = MagicMock()
upstream.provider_type = provider_type
upstream.prepare_headers = MagicMock(return_value={})
upstream.on_upstream_error_redirect = AsyncMock()
return upstream
@pytest.fixture
def captured() -> dict[str, object]:
return {}
@pytest.fixture
def proxy_app(captured: dict[str, object]) -> FastAPI:
app = FastAPI()
app.include_router(proxy_module.proxy_router)
app.dependency_overrides[get_session] = lambda: AsyncMock()
@app.middleware("http")
async def capture(request: Request, call_next: Any) -> Response:
try:
return await call_next(request)
finally:
captured.update(_attribution(request))
return app
@pytest.fixture
def routing(monkeypatch: pytest.MonkeyPatch) -> dict[str, Any]:
"""Stub pricing/reservation so only candidate routing drives the test."""
max_costs: dict[str, int] = {}
async def max_cost(
model: str, session: object, model_obj: MagicMock | None = None
) -> int:
return max_costs.get(getattr(model_obj, "id", ""), 100)
async def discounted(cost: int, body: object, model_obj: object = None) -> int:
return cost
state: dict[str, Any] = {
"candidates": [],
"max_costs": max_costs,
"pay": AsyncMock(return_value=MagicMock()),
}
monkeypatch.setattr(
proxy_module, "get_candidates", lambda _model_id: state["candidates"]
)
monkeypatch.setattr(proxy_module, "get_max_cost_for_model", max_cost)
monkeypatch.setattr(proxy_module, "calculate_discounted_max_cost", discounted)
monkeypatch.setattr(proxy_module, "check_token_balance", lambda *_a: None)
monkeypatch.setattr(
proxy_module,
"get_bearer_token_key",
AsyncMock(return_value=MagicMock(hashed_key="abcdef123456", balance=0)),
)
monkeypatch.setattr(proxy_module, "pay_for_request", state["pay"])
monkeypatch.setattr(proxy_module, "revert_pay_for_request", AsyncMock())
monkeypatch.setattr(proxy_module, "_finish_read_transaction", AsyncMock())
return state
async def _send(app: FastAPI, method: str, path: str, **kwargs: Any) -> httpx.Response:
async with AsyncClient(
transport=ASGITransport(app=app), # type: ignore[arg-type]
base_url="http://test",
) as client:
return await client.request(method, path, **kwargs)
@pytest.mark.asyncio
async def test_unauthenticated_request_is_attributed_to_the_requested_model(
proxy_app: FastAPI, routing: dict[str, Any], captured: dict[str, object]
) -> None:
routing["candidates"] = [(_model("prov/model-a"), _upstream("prov"))]
response = await _send(
proxy_app, "POST", "/v1/chat/completions", json={"model": "model-a"}
)
assert response.status_code == 401
assert captured == {"model": "model-a"}
@pytest.mark.parametrize("body", [{}, {"model": "unknown"}, {"model": 123}])
@pytest.mark.asyncio
async def test_request_without_a_model_is_not_attributed(
proxy_app: FastAPI,
routing: dict[str, Any],
captured: dict[str, object],
body: dict[str, object],
) -> None:
response = await _send(proxy_app, "POST", "/v1/chat/completions", json=body)
assert response.status_code == 400
assert captured == {}
@pytest.mark.asyncio
async def test_paid_fallback_is_attributed_to_the_serving_candidate(
proxy_app: FastAPI, routing: dict[str, Any], captured: dict[str, object]
) -> None:
primary, fallback = _upstream("prov-a"), _upstream("prov-b")
primary.forward_request = AsyncMock(
side_effect=UpstreamError("down", status_code=502)
)
fallback.forward_request = AsyncMock(return_value=Response(status_code=200))
routing["candidates"] = [
(_model("prov-a/model-a"), primary),
(_model("prov-b/model-a-v2"), fallback),
]
response = await _send(
proxy_app,
"POST",
"/v1/chat/completions",
json={"model": "model-a"},
headers={"authorization": "Bearer sk-test"},
)
assert response.status_code == 200
assert captured == {"model": "prov-b/model-a-v2", "provider": "prov-b"}
@pytest.mark.asyncio
async def test_fallback_rejected_at_reservation_keeps_last_attempted_attribution(
proxy_app: FastAPI, routing: dict[str, Any], captured: dict[str, object]
) -> None:
"""A pricier fallback the key cannot reserve is never tried, so the line
stays with the upstream that actually handled (and failed) the request."""
primary, fallback = _upstream("prov-a"), _upstream("prov-b")
primary.forward_request = AsyncMock(
side_effect=UpstreamError("down", status_code=502)
)
fallback.forward_request = AsyncMock()
routing["candidates"] = [
(_model("prov-a/model-a"), primary),
(_model("prov-b/model-a"), fallback),
]
routing["max_costs"]["prov-b/model-a"] = 200
routing["pay"].side_effect = [
MagicMock(),
HTTPException(status_code=402, detail="Insufficient balance"),
]
response = await _send(
proxy_app,
"POST",
"/v1/chat/completions",
json={"model": "model-a"},
headers={"authorization": "Bearer sk-test"},
)
assert response.status_code == 402
fallback.forward_request.assert_not_awaited()
assert captured == {"model": "prov-a/model-a", "provider": "prov-a"}
@pytest.mark.asyncio
async def test_x_cashu_fallback_is_attributed_to_the_serving_candidate(
proxy_app: FastAPI, routing: dict[str, Any], captured: dict[str, object]
) -> None:
primary, fallback = _upstream("prov-a"), _upstream("prov-b")
primary.handle_x_cashu = AsyncMock(
side_effect=UpstreamError("down", status_code=502)
)
fallback.handle_x_cashu = AsyncMock(return_value=Response(status_code=200))
routing["candidates"] = [
(_model("prov-a/model-a"), primary),
(_model("prov-b/model-a"), fallback),
]
response = await _send(
proxy_app,
"POST",
"/v1/chat/completions",
json={"model": "model-a"},
headers={"x-cashu": "cashuAtoken"},
)
assert response.status_code == 200
assert captured == {"model": "prov-b/model-a", "provider": "prov-b"}
@pytest.mark.asyncio
async def test_unauthenticated_get_fallback_is_attributed_to_the_serving_upstream(
proxy_app: FastAPI, routing: dict[str, Any], captured: dict[str, object]
) -> None:
primary, fallback = _upstream("prov-a"), _upstream("prov-b")
primary.forward_get_request = AsyncMock(return_value=Response(status_code=502))
fallback.forward_get_request = AsyncMock(return_value=Response(status_code=200))
routing["candidates"] = [
(_model("prov-a/model-a"), primary),
(_model("prov-b/model-a"), fallback),
]
response = await _send(proxy_app, "GET", "/v1/models")
assert response.status_code == 200
assert captured["provider"] == "prov-b"
-1
View File
@@ -42,7 +42,6 @@ def log_dir(tmp_path: Path) -> Path:
@pytest.fixture @pytest.fixture
def handler(log_dir: Path) -> Iterator[DailyRotatingFileHandler]: def handler(log_dir: Path) -> Iterator[DailyRotatingFileHandler]:
"""A file handler configured exactly like the production ``file`` handler."""
handler = DailyRotatingFileHandler( handler = DailyRotatingFileHandler(
str(log_dir / "app.log"), str(log_dir / "app.log"),
when="midnight", when="midnight",
+68 -2
View File
@@ -23,6 +23,9 @@ from routstr.core.db import ApiKey # noqa: E402
from routstr.payment.cost_calculation import CostData # noqa: E402 from routstr.payment.cost_calculation import CostData # noqa: E402
from routstr.payment.models import Architecture, Model, Pricing # noqa: E402 from routstr.payment.models import Architecture, Model, Pricing # noqa: E402
from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 from routstr.upstream.base import BaseUpstreamProvider # noqa: E402
from routstr.upstream.messages_dispatch import ( # noqa: E402
prune_blank_system_blocks,
)
from routstr.wallet import MintConnectionError, TokenConsumedError # noqa: E402 from routstr.wallet import MintConnectionError, TokenConsumedError # noqa: E402
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -113,6 +116,39 @@ def _make_request(request_id: str | None = "req-test") -> Any:
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def test_prune_blank_system_blocks_drops_blank_blocks() -> None:
body = {
"system": [
{"type": "text", "text": " \n"},
{"type": "text", "text": "real prompt"},
]
}
prune_blank_system_blocks(body)
assert body["system"] == [{"type": "text", "text": "real prompt"}]
def test_prune_blank_system_blocks_drops_key_when_all_blank() -> None:
body = {"system": [{"type": "text", "text": "\n"}], "max_tokens": 8}
prune_blank_system_blocks(body)
assert body == {"max_tokens": 8}
def test_prune_blank_system_blocks_handles_string_system() -> None:
blank = {"system": " "}
prune_blank_system_blocks(blank)
assert blank == {}
kept = {"system": "be brief"}
prune_blank_system_blocks(kept)
assert kept == {"system": "be brief"}
def test_prune_blank_system_blocks_keeps_non_text_blocks() -> None:
body = {"system": [{"type": "image", "source": {}}]}
prune_blank_system_blocks(body)
assert body["system"] == [{"type": "image", "source": {}}]
def test_coerce_litellm_payload_handles_dict() -> None: def test_coerce_litellm_payload_handles_dict() -> None:
out = BaseUpstreamProvider._coerce_litellm_payload({"a": 1}) out = BaseUpstreamProvider._coerce_litellm_payload({"a": 1})
assert out == {"a": 1} assert out == {"a": 1}
@@ -1597,7 +1633,7 @@ async def test_x_cashu_transport_error_after_redemption_is_not_retryable(
handler_name: str, forward_attr: str handler_name: str, forward_attr: str
) -> None: ) -> None:
"""A transport failure while forwarding (after the token is spent) maps to """A transport failure while forwarding (after the token is spent) maps to
502 upstream_error, never a retryable cashu_mint_unreachable.""" 424 + UPSTREAM_UNAVAILABLE, never a retryable cashu_mint_unreachable."""
provider = _make_provider() provider = _make_provider()
model = _make_model() model = _make_model()
request = _make_request() request = _make_request()
@@ -1623,9 +1659,11 @@ async def test_x_cashu_transport_error_after_redemption_is_not_retryable(
model_obj=model, model_obj=model,
) )
assert response.status_code == 502 assert response.status_code == 424
assert response.headers["X-Routstr-Error-Scope"] == "upstream"
body = json.loads(bytes(response.body)) body = json.loads(bytes(response.body))
assert body["error"]["type"] == "upstream_error" assert body["error"]["type"] == "upstream_error"
assert body["error"]["code"] == "UPSTREAM_UNAVAILABLE"
assert body["error"]["code"] != "cashu_mint_unreachable" assert body["error"]["code"] != "cashu_mint_unreachable"
@@ -1753,3 +1791,31 @@ async def test_x_cashu_zero_value_rejected_not_forwarded(
assert body["error"]["code"] == "cashu_token_zero_value" assert body["error"]["code"] == "cashu_token_zero_value"
# Spent-to-zero token must not be echoed back for retry. # Spent-to-zero token must not be echoed back for retry.
assert "X-Cashu" not in response.headers assert "X-Cashu" not in response.headers
@pytest.mark.asyncio
async def test_dispatch_passes_placeholder_key_for_keyless_upstream() -> None:
"""A blank upstream key must not reach litellm, which would fall back to
OPENAI_API_KEY and fail with an AuthenticationError."""
provider = BaseUpstreamProvider(base_url="http://localhost:8000/v1", api_key="")
captured_kwargs: dict[str, Any] = {}
async def fake_acreate(**kwargs: Any) -> AsyncIterator[dict]:
captured_kwargs.update(kwargs)
async def no_events() -> AsyncIterator[dict]:
return
yield
return no_events()
with patch(
"litellm.anthropic.messages.acreate",
new=AsyncMock(side_effect=fake_acreate),
):
await provider._dispatch_anthropic_messages(
request_body=_anthropic_request_body(stream=True),
model_obj=_make_model(),
)
assert captured_kwargs["api_key"] == "no-key"
+152
View File
@@ -0,0 +1,152 @@
import os
from typing import Any, AsyncIterator
from unittest.mock import AsyncMock, patch
import litellm
import pytest
from litellm.exceptions import MidStreamFallbackError
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
from routstr.core.exceptions import UpstreamError # noqa: E402
from routstr.payment.models import Architecture, Model, Pricing # noqa: E402
from routstr.upstream.base import BaseUpstreamProvider # noqa: E402
from routstr.upstream.messages_dispatch import ( # noqa: E402
collapse_litellm_message,
)
_MIDSTREAM_FAILURE = MidStreamFallbackError(
message="No credits.",
model="x",
llm_provider="openai",
original_exception=litellm.APIError(
status_code=500, message="No credits.", llm_provider="openai", model="x"
),
)
@pytest.mark.parametrize(
("message", "expected"),
[
("You have no credits remaining.", "You have no credits remaining."),
# upstream_error_from_exception reads `.message`, which omits the
# "Original exception:" chain that only `str()` appends.
(_MIDSTREAM_FAILURE.message, "No credits."),
(str(_MIDSTREAM_FAILURE), "No credits."),
("x" * 301, "x" * 299 + "…"),
],
)
def test_collapse_litellm_message(message: str, expected: str) -> None:
assert collapse_litellm_message(message) == expected
_RATE_LIMIT = litellm.RateLimitError(
message=(
"Rate limit reached for gpt-4o on tokens per min (TPM): Limit 30000, "
"Used 29000, Requested 2000. Please try again in 1.2s."
),
llm_provider="openai",
model="gpt-4o",
)
_BAD_REQUEST = litellm.BadRequestError(
message="context length exceeded", model="gpt-4o", llm_provider="openai"
)
_MID_STREAM_CASES = [
pytest.param(_RATE_LIMIT, 429, "UPSTREAM_RATE_LIMIT", id="rate-limit"),
pytest.param(_BAD_REQUEST, 400, None, id="bad-request"),
pytest.param(_MIDSTREAM_FAILURE, 500, None, id="midstream-fallback"),
]
def _make_model() -> Model:
return Model(
id="gpt-4o",
name="gpt-4o",
created=0,
description="",
context_length=8192,
architecture=Architecture(
modality="text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="x",
instruct_type=None,
),
pricing=Pricing(
prompt=0.0,
completion=0.0,
request=0.0,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
max_cost=0.0,
),
)
def _failing_stream(exc: Exception) -> AsyncIterator[dict]:
async def gen() -> AsyncIterator[dict]:
yield {
"type": "message_start",
"message": {"id": "msg_1", "model": "gpt-4o", "usage": {}},
}
raise exc
return gen()
def _assert_upstream_error(
err: UpstreamError, status_code: int, code: str | None
) -> None:
assert err.status_code == status_code
assert err.code == code
assert err.from_upstream_response is True
assert "litellm." not in str(err)
@pytest.mark.asyncio
@pytest.mark.parametrize(("exc", "status_code", "code"), _MID_STREAM_CASES)
async def test_non_streaming_aggregation_surfaces_mid_stream_failure(
exc: Exception, status_code: int, code: str | None
) -> None:
async def fake_acreate(**kwargs: Any) -> AsyncIterator[dict]:
return _failing_stream(exc)
with (
patch(
"litellm.anthropic.messages.acreate",
new=AsyncMock(side_effect=fake_acreate),
),
pytest.raises(UpstreamError) as exc_info,
):
await BaseUpstreamProvider(
base_url="http://test", api_key="k"
)._dispatch_anthropic_messages(
request_body=b'{"messages": [], "max_tokens": 8, "stream": false}',
model_obj=_make_model(),
)
_assert_upstream_error(exc_info.value, status_code, code)
@pytest.mark.asyncio
@pytest.mark.parametrize(("exc", "status_code", "code"), _MID_STREAM_CASES)
async def test_x_cashu_buffered_stream_surfaces_mid_stream_failure(
exc: Exception, status_code: int, code: str | None
) -> None:
provider = BaseUpstreamProvider(base_url="http://test", api_key="k")
with pytest.raises(UpstreamError) as exc_info:
await provider._stream_x_cashu_litellm_messages(
_failing_stream(exc),
amount=5_000,
unit="sat",
max_cost_for_model=10_000,
requested_model="gpt-4o",
mint=None,
request_id="req-test",
)
_assert_upstream_error(exc_info.value, status_code, code)
+173 -12
View File
@@ -9,8 +9,16 @@ import pytest
from routstr import proxy as proxy_module from routstr import proxy as proxy_module
from routstr.auth import ReservationSnapshot from routstr.auth import ReservationSnapshot
from routstr.core.db import ApiKey from routstr.core.db import ApiKey
from routstr.core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_NODE,
ERROR_SCOPE_UPSTREAM,
UPSTREAM_UNAVAILABLE,
)
from routstr.upstream.model_paths import decode_model_path, encode_model_path from routstr.upstream.model_paths import decode_model_path, encode_model_path
from .proxy_test_utils import mock_request_stream, patch_proxy_session
MODEL_ID = "test-model" MODEL_ID = "test-model"
@@ -32,7 +40,7 @@ def _make_request(headers: dict[str, str], body: bytes) -> MagicMock:
request = MagicMock() request = MagicMock()
request.method = "POST" request.method = "POST"
request.headers = headers request.headers = headers
request.body = AsyncMock(return_value=body) mock_request_stream(request, body)
request.state = MagicMock() request.state = MagicMock()
request.state.request_id = "req-model-path" request.state.request_id = "req-model-path"
return request return request
@@ -62,15 +70,13 @@ async def _run_proxy(
), ),
patch.object(proxy_module, "check_token_balance", MagicMock()), patch.object(proxy_module, "check_token_balance", MagicMock()),
patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)), patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)),
patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)),
patch.object( patch.object(
proxy_module, proxy_module, "pay_for_request", AsyncMock(return_value=reservation)
"get_reservation_snapshot",
AsyncMock(return_value=reservation),
), ),
patch.object(proxy_module, "revert_pay_for_request", AsyncMock()), patch.object(proxy_module, "revert_pay_for_request", AsyncMock()),
patch_proxy_session(MagicMock()),
): ):
return await proxy_module.proxy(request, path, session=MagicMock()) return await proxy_module.proxy(request, path)
def test_decode_model_path_round_trips_encode() -> None: def test_decode_model_path_round_trips_encode() -> None:
@@ -391,8 +397,13 @@ def test_model_path_header_is_not_forwarded() -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
@pytest.mark.parametrize("path", ["v1/chat/completions", "v1/responses"]) @pytest.mark.parametrize("path", ["v1/chat/completions", "v1/responses"])
@pytest.mark.parametrize("status_code", [200, 429, 502]) @pytest.mark.parametrize(
async def test_cashu_pin_reaches_http_transport(path: str, status_code: int) -> None: "status_code,client_status",
[(200, 200), (429, 429), (502, 424)],
)
async def test_cashu_pin_reaches_http_transport(
path: str, status_code: int, client_status: int
) -> None:
import httpx import httpx
from fastapi.responses import Response from fastapi.responses import Response
@@ -443,7 +454,11 @@ async def test_cashu_pin_reaches_http_transport(path: str, status_code: int) ->
response = await _run_proxy( response = await _run_proxy(
request, [(model, upstream), (model, fallback)], path request, [(model, upstream), (model, fallback)], path
) )
assert response.status_code == status_code assert response.status_code == client_status
if client_status == 424:
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
body = json.loads(bytes(response.body))
assert body["error"]["code"] == UPSTREAM_UNAVAILABLE
redeem.assert_awaited_once() redeem.assert_awaited_once()
assert len(sent) == 1 assert len(sent) == 1
assert sent[0].url.host == "openrouter.ai" assert sent[0].url.host == "openrouter.ai"
@@ -486,7 +501,10 @@ async def test_pinned_exception_does_not_fall_back() -> None:
response = await _run_proxy( response = await _run_proxy(
request, [(MagicMock(), first), (MagicMock(), fallback)] request, [(MagicMock(), first), (MagicMock(), fallback)]
) )
assert response.status_code == 503 # Pinned: no fallback.
assert response.status_code == 424
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
assert json.loads(bytes(response.body))["error"]["code"] == UPSTREAM_UNAVAILABLE
first.forward_request.assert_awaited_once() first.forward_request.assert_awaited_once()
fallback.forward_request.assert_not_awaited() fallback.forward_request.assert_not_awaited()
@@ -514,8 +532,9 @@ async def test_unsupported_endpoint_pins_fail_before_payment(
patch.object( patch.object(
proxy_module, "get_candidates", return_value=[(MagicMock(), upstream)] proxy_module, "get_candidates", return_value=[(MagicMock(), upstream)]
), ),
patch_proxy_session(MagicMock()),
): ):
response = await proxy_module.proxy(request, path, MagicMock()) response = await proxy_module.proxy(request, path)
assert response.status_code == 400 assert response.status_code == 400
assert json.loads(response.body)["error"]["type"] == "unsupported_request" assert json.loads(response.body)["error"]["type"] == "unsupported_request"
payment.assert_not_called() payment.assert_not_called()
@@ -545,7 +564,8 @@ async def test_ehbp_pin_does_not_fall_back(cashu: bool) -> None:
response = await _run_proxy( response = await _run_proxy(
request, [(MagicMock(), selected), (MagicMock(), fallback)] request, [(MagicMock(), selected), (MagicMock(), fallback)]
) )
assert response.status_code == 503 assert response.status_code == 424
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
forward.assert_awaited_once() forward.assert_awaited_once()
assert forward.await_args is not None assert forward.await_args is not None
assert forward.await_args.kwargs["upstream"] is selected assert forward.await_args.kwargs["upstream"] is selected
@@ -684,3 +704,144 @@ async def test_pinned_recovery_preserves_routing_fields(
assert response.status_code == 400 assert response.status_code == 400
selected.forward_request.assert_awaited_once() selected.forward_request.assert_awaited_once()
fallback.forward_request.assert_not_awaited() fallback.forward_request.assert_not_awaited()
_OPENAI_MAX_TOKENS_ERROR = json.dumps(
{
"error": {
"message": "Unsupported parameter: 'max_tokens' is not supported "
"with this model. Use 'max_completion_tokens' instead.",
"type": "invalid_request_error",
"param": "max_tokens",
"code": "unsupported_parameter",
}
}
).encode()
@pytest.mark.asyncio
@pytest.mark.parametrize("pinned", [False, True])
async def test_rejected_max_tokens_is_renamed_and_retried_on_same_upstream(
pinned: bool,
) -> None:
selected, fallback = _make_upstream(1), _make_upstream(2)
selected.forward_request = AsyncMock(
side_effect=[
MagicMock(status_code=400, body=_OPENAI_MAX_TOKENS_ERROR),
MagicMock(status_code=200, body=b"{}"),
]
)
headers = {"authorization": "Bearer key"}
if pinned:
headers["x-routstr-model-path"] = encode_model_path(selected.base_url, MODEL_ID)
request = _make_request(
headers,
json.dumps(
{"model": MODEL_ID, "max_tokens": 300, "messages": [], "stream": True}
).encode(),
)
response = await _run_proxy(
request, [(MagicMock(), selected), (MagicMock(), fallback)]
)
assert response.status_code == 200
assert selected.forward_request.await_count == 2
before, after = [
json.loads(call.args[3]) for call in selected.forward_request.await_args_list
]
assert before["max_tokens"] == 300 and "max_completion_tokens" not in before
assert after["max_completion_tokens"] == 300 and "max_tokens" not in after
assert {k: v for k, v in after.items() if k != "max_completion_tokens"} == {
k: v for k, v in before.items() if k != "max_tokens"
}
fallback.forward_request.assert_not_awaited()
@pytest.mark.asyncio
async def test_rename_that_changes_spend_bound_is_not_retried() -> None:
selected = _make_upstream(1, 400)
selected.forward_request.return_value.body = json.dumps(
{"error": {"message": "'max_tokens' is not supported. Use 'n' instead."}}
).encode()
request = _make_request(
{"authorization": "Bearer key"},
json.dumps({"model": MODEL_ID, "max_tokens": 300}).encode(),
)
response = await _run_proxy(request, [(MagicMock(), selected)])
assert response.status_code == 400
selected.forward_request.assert_awaited_once()
# --------------------------------------------------------------------------- #
# Upstream 5xx -> 424 + UPSTREAM_UNAVAILABLE + scope header; node faults stay 500.
# --------------------------------------------------------------------------- #
@pytest.mark.asyncio
async def test_upstream_424_fails_over_to_a_healthy_provider() -> None:
"""An upstream-attributed 424 is still retryable: the caller only ever
sees the healthy provider's 200."""
from routstr.core.exceptions import UpstreamError
first, healthy = _make_upstream(1), _make_upstream(2)
first.forward_request.side_effect = UpstreamError("bad gateway", status_code=502)
request = _make_request(
{"authorization": "Bearer key"}, json.dumps({"model": MODEL_ID}).encode()
)
response = await _run_proxy(request, [(MagicMock(), first), (MagicMock(), healthy)])
assert response.status_code == 200
first.forward_request.assert_awaited_once()
healthy.forward_request.assert_awaited_once()
# The caller never sees the upstream error body or any scope header.
assert ERROR_SCOPE_HEADER not in response.headers
@pytest.mark.asyncio
async def test_last_candidate_upstream_failure_reports_424() -> None:
"""Every candidate failed on the provider hop: 424 + upstream scope, with
the provider's own status preserved for operators."""
from routstr.core.exceptions import UpstreamError
only = _make_upstream(1)
only.forward_request.side_effect = UpstreamError("bad gateway", status_code=502)
request = _make_request(
{"authorization": "Bearer key"}, json.dumps({"model": MODEL_ID}).encode()
)
response = await _run_proxy(request, [(MagicMock(), only)])
assert response.status_code == 424
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
body = json.loads(bytes(response.body))
assert body["error"]["type"] == "upstream_error"
assert body["error"]["code"] == UPSTREAM_UNAVAILABLE
assert body["error"]["details"]["upstream_status"] == 502
@pytest.mark.asyncio
async def test_node_fault_stays_500_without_scope_header() -> None:
"""A genuine node fault keeps its 500 and carries no scope header, so a
client can still tell this node is the broken one."""
from routstr.core.exceptions import UpstreamError
only = _make_upstream(1)
only.forward_request.side_effect = UpstreamError(
"An unexpected server error occurred",
status_code=500,
scope=ERROR_SCOPE_NODE,
)
request = _make_request(
{"authorization": "Bearer key"}, json.dumps({"model": MODEL_ID}).encode()
)
response = await _run_proxy(request, [(MagicMock(), only)])
assert response.status_code == 500
assert ERROR_SCOPE_HEADER not in response.headers
body = json.loads(bytes(response.body))
assert body["error"]["code"] != UPSTREAM_UNAVAILABLE
+90
View File
@@ -0,0 +1,90 @@
"""OpenAI reasoning models get ``max_completion_tokens`` before the request is sent."""
from __future__ import annotations
import json
import os
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to")
import pytest
from routstr.payment.models import Architecture, Model, Pricing
from routstr.upstream import GenericUpstreamProvider
from routstr.upstream.openai import OpenAIUpstreamProvider
def _model(model_id: str) -> Model:
return Model(
id=model_id,
name="test",
created=0,
description="",
context_length=128000,
architecture=Architecture(
modality="text->text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="x",
instruct_type=None,
),
pricing=Pricing(prompt=0.0, completion=0.0),
)
def _chat(model_id: str, **fields: object) -> bytes:
return json.dumps(
{"model": model_id, "messages": [{"role": "user", "content": "hi"}], **fields}
).encode()
def _prepare(provider: object, model_id: str, body: bytes) -> dict:
out = provider.prepare_request_body(body, _model(model_id)) # type: ignore[attr-defined]
assert out is not None
return json.loads(out)
@pytest.mark.parametrize(
"model_id", ["gpt-5.6-sol", "openai/gpt-6-sol", "openai/gpt-5", "o3", "o4-mini"]
)
def test_reasoning_model_max_tokens_is_renamed(model_id: str) -> None:
provider = OpenAIUpstreamProvider(api_key="k")
data = _prepare(provider, model_id, _chat(model_id, max_tokens=300))
assert data["max_completion_tokens"] == 300
assert "max_tokens" not in data
@pytest.mark.parametrize("model_id", ["gpt-4o", "openai/gpt-4.1"])
def test_non_reasoning_model_keeps_max_tokens(model_id: str) -> None:
provider = OpenAIUpstreamProvider(api_key="k")
data = _prepare(provider, model_id, _chat(model_id, max_tokens=300))
assert data["max_tokens"] == 300
assert "max_completion_tokens" not in data
def test_both_caps_set_is_left_for_upstream() -> None:
provider = OpenAIUpstreamProvider(api_key="k")
data = _prepare(
provider,
"gpt-5.6-sol",
_chat("gpt-5.6-sol", max_tokens=300, max_completion_tokens=200),
)
assert data["max_tokens"] == 300
assert data["max_completion_tokens"] == 200
def test_non_chat_body_is_untouched() -> None:
provider = OpenAIUpstreamProvider(api_key="k")
body = json.dumps({"model": "gpt-5.6-sol", "input": "hi", "max_tokens": 5}).encode()
data = _prepare(provider, "gpt-5.6-sol", body)
assert data["max_tokens"] == 5
assert "max_completion_tokens" not in data
def test_other_upstreams_keep_max_tokens() -> None:
provider = GenericUpstreamProvider(base_url="http://test", api_key="k")
data = _prepare(provider, "gpt-5.6-sol", _chat("gpt-5.6-sol", max_tokens=300))
assert data["max_tokens"] == 300
assert "max_completion_tokens" not in data
@@ -0,0 +1,186 @@
"""Retry behaviour for the OpenRouter catalogue fetch."""
from __future__ import annotations
import json
from typing import Any, Callable
import httpx
import pytest
from routstr.payment import models as models_module
from routstr.payment.models import async_fetch_openrouter_models
MODELS_URL = "https://openrouter.ai/api/v1/models"
EMBEDDINGS_URL = "https://openrouter.ai/api/v1/embeddings/models"
def _model(model_id: str) -> dict[str, Any]:
return {
"id": model_id,
"name": model_id,
"pricing": {"prompt": "0.000001", "completion": "0.000002"},
}
def _ok_response(url: str, payload: dict[str, Any]) -> httpx.Response:
return httpx.Response(
200,
request=httpx.Request("GET", url),
content=json.dumps(payload).encode(),
headers={"content-type": "application/json"},
)
def _error_response(url: str, status: int) -> httpx.Response:
return httpx.Response(status, request=httpx.Request("GET", url), content=b"nope")
def _truncated_response(url: str) -> httpx.Response:
"""A body cut mid-JSON — the shape OpenRouter actually sent the node."""
return httpx.Response(
200,
request=httpx.Request("GET", url),
content=b'{"data": [{"id": "vendor/model-a", "name": "Model A", "pric',
headers={"content-type": "application/json"},
)
def _payload_for(url: str) -> dict[str, Any]:
if url.endswith("/embeddings/models"):
return {"data": [_model("vendor/embed-1")]}
return {"data": [_model("vendor/model-a")]}
@pytest.fixture(autouse=True)
def _no_retry_backoff(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(models_module, "OPENROUTER_MODELS_RETRY_BACKOFF_SECONDS", 0)
def _install_get(
monkeypatch: pytest.MonkeyPatch,
handler: Callable[[str, int], httpx.Response],
) -> dict[str, int]:
"""Patch ``httpx.AsyncClient.get`` and count calls per endpoint."""
counts: dict[str, int] = {}
async def fake_get(
self: httpx.AsyncClient, url: Any, **kwargs: Any
) -> httpx.Response:
key = str(url)
counts[key] = counts.get(key, 0) + 1
return handler(key, counts[key])
monkeypatch.setattr(httpx.AsyncClient, "get", fake_get)
return counts
@pytest.mark.asyncio
async def test_truncated_body_is_retried_and_recovers(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A truncated body on the first attempt must not empty the catalogue."""
def handler(url: str, call: int) -> httpx.Response:
if call == 1:
return _truncated_response(url)
return _ok_response(url, _payload_for(url))
counts = _install_get(monkeypatch, handler)
result = await async_fetch_openrouter_models()
assert [model["id"] for model in result] == ["vendor/model-a", "vendor/embed-1"]
assert counts[MODELS_URL] == 2
assert counts[EMBEDDINGS_URL] == 2
@pytest.mark.asyncio
async def test_gives_up_after_max_attempts_and_logs_the_error(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""After every attempt fails: log once and return an empty catalogue."""
errors: list[str] = []
monkeypatch.setattr(models_module.logger, "error", lambda msg: errors.append(msg))
counts = _install_get(monkeypatch, lambda url, call: _truncated_response(url))
result = await async_fetch_openrouter_models()
assert result == []
assert counts[MODELS_URL] == models_module.OPENROUTER_MODELS_MAX_ATTEMPTS
assert len(errors) == 1
assert "after 3 attempt(s)" in errors[0]
@pytest.mark.asyncio
async def test_embeddings_outage_still_yields_the_main_catalogue(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""The secondary endpoint is best-effort: it cannot empty the catalogue."""
def handler(url: str, call: int) -> httpx.Response:
if url == EMBEDDINGS_URL:
return _error_response(url, 503)
return _ok_response(url, _payload_for(url))
counts = _install_get(monkeypatch, handler)
result = await async_fetch_openrouter_models()
assert [model["id"] for model in result] == ["vendor/model-a"]
assert counts[MODELS_URL] == 1
assert counts[EMBEDDINGS_URL] == 1
@pytest.mark.asyncio
async def test_main_catalogue_server_error_is_retried(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A 5xx on /models is transient, so the attempt is retried."""
def handler(url: str, call: int) -> httpx.Response:
if url == MODELS_URL and call == 1:
return _error_response(url, 503)
return _ok_response(url, _payload_for(url))
counts = _install_get(monkeypatch, handler)
result = await async_fetch_openrouter_models()
assert [model["id"] for model in result] == ["vendor/model-a", "vendor/embed-1"]
assert counts[MODELS_URL] == 2
@pytest.mark.asyncio
async def test_client_error_is_not_retried(monkeypatch: pytest.MonkeyPatch) -> None:
"""Retrying a 4xx only adds load to an upstream that already said no."""
counts = _install_get(monkeypatch, lambda url, call: _error_response(url, 401))
result = await async_fetch_openrouter_models()
assert result == []
assert counts[MODELS_URL] == 1
@pytest.mark.asyncio
async def test_source_filter_and_free_models_are_still_applied(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""The moved filter loop still strips prefixes and drops free tiers."""
payload = {
"data": [
_model("openai/gpt-x"),
_model("openai/gpt-x:free"),
_model("other/y"),
]
}
def handler(url: str, call: int) -> httpx.Response:
return _ok_response(url, payload if url == MODELS_URL else {"data": []})
_install_get(monkeypatch, handler)
result = await async_fetch_openrouter_models(source_filter="openai")
assert [model["id"] for model in result] == ["gpt-x"]
@@ -0,0 +1,55 @@
from typing import Any
from unittest.mock import Mock
import pytest
from sqlmodel.ext.asyncio.session import AsyncSession
import routstr.auth as auth_module
from routstr.core.db import ApiKey
@pytest.mark.asyncio
async def test_payment_settlement_logs_its_duration(
monkeypatch: pytest.MonkeyPatch,
) -> None:
async def settle(*_args: Any, **_kwargs: Any) -> dict[str, int]:
return {"total_cost": 1}
log_info = Mock()
monkeypatch.setattr(auth_module, "_adjust_payment_for_tokens", settle)
monkeypatch.setattr(auth_module.logger, "info", log_info)
key = ApiKey(hashed_key="abcdefgh1234", balance=0)
session = AsyncSession()
result = await auth_module.adjust_payment_for_tokens(key, {}, session, 10)
await session.close()
assert result == {"total_cost": 1}
log_info.assert_called_once()
(message,) = log_info.call_args.args
extra = log_info.call_args.kwargs["extra"]
assert message == "Payment settlement finished"
assert extra["settlement_duration_ms"] >= 0
assert extra["settlement_succeeded"] is True
@pytest.mark.asyncio
async def test_payment_settlement_logs_failure_without_swallowing_it(
monkeypatch: pytest.MonkeyPatch,
) -> None:
async def fail(*_args: Any, **_kwargs: Any) -> dict:
raise RuntimeError("database locked")
log_info = Mock()
monkeypatch.setattr(auth_module, "_adjust_payment_for_tokens", fail)
monkeypatch.setattr(auth_module.logger, "info", log_info)
key = ApiKey(hashed_key="abcdefgh1234", balance=0)
session = AsyncSession()
with pytest.raises(RuntimeError, match="database locked"):
await auth_module.adjust_payment_for_tokens(key, {}, session, 10)
await session.close()
extra = log_info.call_args.kwargs["extra"]
assert extra["settlement_duration_ms"] >= 0
assert extra["settlement_succeeded"] is False
+195
View File
@@ -0,0 +1,195 @@
"""Owner payout keeps each wallet's declared liability and the global total.
Regression for multi-mint payout starvation: subtracting the *total* user
liability from every wallet hid the owner surplus on any mint holding less
than the whole liability, so only the largest wallet could ever pay out.
"""
from collections.abc import AsyncIterator, Iterator
from contextlib import asynccontextmanager, contextmanager
from unittest.mock import AsyncMock, Mock, patch
import pytest
from routstr.core.settings import settings
from routstr.wallet import _owner_balance_for_mint_and_unit, _payout_mint_and_unit
MINT_A = "https://a.test"
MINT_B = "https://b.test"
@asynccontextmanager
async def _session() -> AsyncIterator[Mock]:
yield Mock()
def _keyset_id(mint_url: str, unit: str) -> str:
return f"{mint_url}|{unit}"
@contextmanager
def _wallet_db(
sat_proofs: dict[str, int], reserved: frozenset[str] = frozenset()
) -> Iterator[AsyncMock]:
"""One sat keyset per mint, one proof behind it. Yields the get_wallet mock."""
keysets = [
Mock(id=_keyset_id(mint_url, "sat"), mint_url=mint_url, unit="sat")
for mint_url in sat_proofs
]
proofs = [
Mock(
id=_keyset_id(mint_url, "sat"),
amount=amount,
reserved=mint_url in reserved,
)
for mint_url, amount in sat_proofs.items()
]
get_wallet = AsyncMock(return_value=Mock(url=MINT_B, db=Mock()))
with (
patch("routstr.wallet.get_wallet", get_wallet),
patch("routstr.wallet.get_cashu_keysets", AsyncMock(return_value=keysets)),
patch("routstr.wallet.get_cashu_proofs", AsyncMock(return_value=proofs)),
):
yield get_wallet
@contextmanager
def _env() -> Iterator[None]:
with (
patch("routstr.wallet.db.create_session", _session),
patch.object(settings, "cashu_mints", [MINT_A, MINT_B]),
patch.object(settings, "primary_mint", MINT_A),
):
yield
@contextmanager
def _liabilities(per_mint_sats: dict[str, int], total_sats: int) -> Iterator[None]:
async def per_mint(_session: object, mint_url: str, unit: str) -> int:
return per_mint_sats.get(mint_url, 0) * 1000
with (
_env(),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(side_effect=per_mint),
),
patch(
"routstr.wallet.db.total_user_liability",
AsyncMock(return_value=total_sats * 1000),
),
):
yield
@pytest.mark.asyncio
async def test_owner_balance_keeps_only_the_wallets_own_liability() -> None:
with (
_liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250),
_wallet_db({MINT_A: 400, MINT_B: 270}),
):
assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 236
assert await _owner_balance_for_mint_and_unit(MINT_A, "sat", 400) == 184
@pytest.mark.asyncio
async def test_owner_balance_never_exceeds_global_surplus() -> None:
"""Liability nobody declared against a mint is still covered in aggregate."""
with (
_liabilities({}, total_sats=250),
_wallet_db({MINT_A: 100, MINT_B: 270}),
):
assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 120
@pytest.mark.asyncio
async def test_proofs_of_an_untrusted_mint_do_not_raise_the_bound() -> None:
"""Only configured mints back the global surplus."""
with (
_liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250),
patch.object(settings, "cashu_mints", [MINT_B]),
patch.object(settings, "primary_mint", MINT_B),
_wallet_db({MINT_A: 400, MINT_B: 270}),
):
assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 20
@pytest.mark.asyncio
async def test_reserved_proofs_do_not_raise_the_bound() -> None:
"""Another process may already be spending them."""
with (
_liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250),
_wallet_db({MINT_A: 400, MINT_B: 270}, reserved=frozenset({MINT_A})),
):
assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 20
@pytest.mark.asyncio
async def test_duplicate_configured_mint_is_counted_once() -> None:
"""A mint listed twice in CASHU_MINTS would otherwise raise the global bound."""
with (
_liabilities({}, total_sats=600),
patch.object(settings, "cashu_mints", [MINT_A, MINT_A, MINT_B]),
_wallet_db({MINT_A: 400, MINT_B: 270}),
):
assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 70
@pytest.mark.asyncio
async def test_cross_wallet_bound_asks_no_mint_for_metadata() -> None:
"""The sum is local. Loading a wallet per mint and unit rate-limited mints."""
with (
_liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250),
_wallet_db({MINT_A: 400, MINT_B: 270}) as get_wallet,
):
await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270)
assert get_wallet.await_args_list
assert all(c.kwargs.get("load") is False for c in get_wallet.await_args_list)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"mint_liability,total_liability,expected",
[(1_500, 2_200, 1_800), (2_700, 1_500, 1_300)],
)
async def test_msat_wallet_surplus_is_not_rounded(
mint_liability: int, total_liability: int, expected: int
) -> None:
"""Either bound can bind, and neither is rounded to whole sats."""
with (
_env(),
_wallet_db({}),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=mint_liability),
),
patch(
"routstr.wallet.db.total_user_liability",
AsyncMock(return_value=total_liability),
),
):
assert await _owner_balance_for_mint_and_unit(MINT_B, "msat", 4_000) == expected
@pytest.mark.asyncio
async def test_payout_sends_the_smaller_wallets_surplus() -> None:
send = AsyncMock(return_value=236_000)
with (
_liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250),
patch.object(settings, "min_payout_sat", 50),
patch.object(settings, "max_payout_sat", 250_000),
_wallet_db({MINT_A: 400, MINT_B: 270}),
patch(
"routstr.wallet.get_proofs_per_mint_and_unit",
Mock(return_value=[Mock(amount=270)]),
),
patch(
"routstr.wallet.slow_filter_spend_proofs",
AsyncMock(side_effect=lambda proofs, wallet: proofs),
),
patch("routstr.wallet.asyncio.sleep", AsyncMock()),
patch("routstr.wallet.raw_send_to_lnurl", send),
):
await _payout_mint_and_unit(MINT_B, "sat")
assert send.await_args is not None
assert send.await_args.kwargs["amount"] == 236
+80
View File
@@ -0,0 +1,80 @@
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from unittest.mock import AsyncMock, Mock, call, patch
import pytest
from routstr.core.settings import settings
from routstr.wallet import _payout_mint_and_unit
@asynccontextmanager
async def session() -> AsyncIterator[Mock]:
yield Mock()
@pytest.mark.asyncio
@pytest.mark.parametrize("unit,scale", [("sat", 1), ("msat", 1000)])
@pytest.mark.parametrize(
"balance,liability,expected",
[(1000, 0, 100), (80, 30000, 50), (20, 20000, None), (0, 0, None), (10, 0, None)],
)
async def test_payout_limits_and_proof_refresh(
unit: str, scale: int, balance: int, liability: int, expected: int | None
) -> None:
send = AsyncMock()
get_wallet = AsyncMock()
check = AsyncMock(side_effect=lambda ps, w: ps)
sleep = AsyncMock()
with (
patch.object(settings, "min_payout_sat", 10),
patch.object(settings, "max_payout_sat", 100),
patch("routstr.wallet.get_wallet", get_wallet),
patch(
"routstr.wallet.get_proofs_per_mint_and_unit",
return_value=[Mock(amount=balance * scale)],
),
patch("routstr.wallet.slow_filter_spend_proofs", check),
patch("routstr.wallet.db.create_session", session),
patch(
"routstr.wallet.db.total_user_liability", AsyncMock(return_value=liability)
),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=liability),
),
patch("routstr.wallet.asyncio.sleep", sleep),
patch("routstr.wallet.raw_send_to_lnurl", send),
):
await _payout_mint_and_unit("https://mint.test", unit)
# Later awaits belong to the other-wallet scan, which forces a reload too.
assert get_wallet.await_args_list[0] == call(
"https://mint.test", unit, force_reload_proofs=True
)
if expected is None:
send.assert_not_awaited()
else:
assert send.await_args is not None
assert send.await_args.kwargs["amount"] == expected * scale
if balance <= 10:
check.assert_not_awaited()
sleep.assert_not_awaited()
@pytest.mark.asyncio
async def test_failed_proof_check_never_pays_partial_balance() -> None:
send = AsyncMock()
with (
patch("routstr.wallet.get_wallet", AsyncMock()),
patch(
"routstr.wallet.get_proofs_per_mint_and_unit",
return_value=[Mock(amount=1_000_000)],
),
patch(
"routstr.wallet.slow_filter_spend_proofs",
AsyncMock(side_effect=ValueError("Invalid proof-state response")),
),
patch("routstr.wallet.raw_send_to_lnurl", send),
):
await _payout_mint_and_unit("https://mint.test", "sat")
send.assert_not_awaited()
+359 -16
View File
@@ -10,14 +10,38 @@ Covers two regressions from the auto-payout / primary-mint audit
mint/units in the same cycle (the try/except is now per mint/unit). mint/units in the same cycle (the try/except is now per mint/unit).
""" """
from collections.abc import Callable, Coroutine from collections.abc import Callable, Coroutine, Iterator
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from pathlib import Path
from typing import Any from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import ANY, AsyncMock, MagicMock, patch
import pytest import pytest
from routstr.wallet import _payout_units, periodic_payout from routstr.payment.lnurl import MeltOutcomeAmbiguousError, MeltUnpaidError
from routstr.wallet import (
_payout_units,
_reconcile_stale_payout_history,
periodic_payout,
)
@pytest.fixture(autouse=True)
def empty_cross_wallet_proofs() -> Iterator[None]:
"""No other wallet holds proofs, so only this wallet's own bound applies."""
with (
patch("routstr.wallet.get_cashu_keysets", AsyncMock(return_value=[])),
patch("routstr.wallet.get_cashu_proofs", AsyncMock(return_value=[])),
):
yield
@pytest.fixture(autouse=True)
def isolate_wallet_lock(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(
"routstr.wallet._WALLET_OPERATION_LOCK", tmp_path / "wallet.lock"
)
# Sentinel interval used to break the otherwise-infinite payout loop after # Sentinel interval used to break the otherwise-infinite payout loop after
# exactly one full cycle. # exactly one full cycle.
@@ -53,11 +77,20 @@ def _one_cycle_sleep() -> Callable[[float], Coroutine[Any, Any, None]]:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> None: async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> None:
"""primary_mint absent from cashu_mints is still paid out.""" """primary_mint absent from cashu_mints is paid out and recorded."""
from routstr.core.settings import settings from routstr.core.settings import settings
get_wallet = AsyncMock(return_value=MagicMock()) get_wallet = AsyncMock(return_value=MagicMock())
raw_send = AsyncMock(return_value=1000) record_payout = AsyncMock()
settle_payout = AsyncMock()
async def send(*args: object, **kwargs: object) -> int:
await kwargs["on_melt_quote"]( # type: ignore[index,operator]
"quote-1", "lnbc1payout"
)
return 1_000_000
raw_send = AsyncMock(side_effect=send)
with ( with (
patch.object(settings, "cashu_mints", []), patch.object(settings, "cashu_mints", []),
@@ -84,6 +117,12 @@ async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> Non
"routstr.wallet.db.total_user_liability", "routstr.wallet.db.total_user_liability",
AsyncMock(return_value=0), AsyncMock(return_value=0),
), ),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch("routstr.wallet.db.record_lightning_payout", record_payout),
patch("routstr.wallet.db.settle_lightning_payout", settle_payout),
patch("routstr.wallet.raw_send_to_lnurl", raw_send), patch("routstr.wallet.raw_send_to_lnurl", raw_send),
): ):
with pytest.raises(_LoopBreak): with pytest.raises(_LoopBreak):
@@ -92,6 +131,20 @@ async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> Non
processed = {call.args[0] for call in get_wallet.await_args_list} processed = {call.args[0] for call in get_wallet.await_args_list}
assert processed == {"http://primary:3338"} assert processed == {"http://primary:3338"}
assert raw_send.await_count >= 1 assert raw_send.await_count >= 1
record_payout.assert_awaited_once_with(
ANY,
quote_id="quote-1",
bolt11="lnbc1payout",
amount_sats=100_000,
mint_url="http://primary:3338",
destination="owner@ln.tld",
)
settle_payout.assert_awaited_once_with(
ANY,
"quote-1",
status="paid",
amount_sats=1_000,
)
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -142,6 +195,10 @@ async def test_periodic_payout_releases_session_before_slow_mint_send() -> None:
"routstr.wallet.db.total_user_liability", "routstr.wallet.db.total_user_liability",
AsyncMock(return_value=0), AsyncMock(return_value=0),
), ),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(side_effect=raw_send)), patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(side_effect=raw_send)),
): ):
with pytest.raises(_LoopBreak): with pytest.raises(_LoopBreak):
@@ -155,9 +212,7 @@ async def test_periodic_payout_isolates_failing_mint() -> None:
"""A failing mint does not prevent payout for the other mints.""" """A failing mint does not prevent payout for the other mints."""
from routstr.core.settings import settings from routstr.core.settings import settings
async def _get_wallet( async def _get_wallet(mint_url: str, unit: str, **_: object) -> MagicMock:
mint_url: str, unit: str, force_reload: bool = False
) -> MagicMock:
if mint_url == "http://bad:3338": if mint_url == "http://bad:3338":
raise RuntimeError("mint unreachable") raise RuntimeError("mint unreachable")
return MagicMock() return MagicMock()
@@ -190,6 +245,10 @@ async def test_periodic_payout_isolates_failing_mint() -> None:
"routstr.wallet.db.total_user_liability", "routstr.wallet.db.total_user_liability",
AsyncMock(return_value=0), AsyncMock(return_value=0),
), ),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch("routstr.wallet.raw_send_to_lnurl", raw_send), patch("routstr.wallet.raw_send_to_lnurl", raw_send),
): ):
with pytest.raises(_LoopBreak): with pytest.raises(_LoopBreak):
@@ -197,10 +256,13 @@ async def test_periodic_payout_isolates_failing_mint() -> None:
# The bad mint raised on get_wallet for both units, yet the good mint was # 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. # still reached and paid out for both units — failures are isolated.
good_calls = [ good_reloads = [
c for c in get_wallet.await_args_list if c.args[0] == "http://good:3338" c
for c in get_wallet.await_args_list
if c.args[0] == "http://good:3338" and c.kwargs.get("force_reload_proofs")
] ]
assert len(good_calls) == 2 # sat + msat # One proof read per unit, and no extra mint load for the bound.
assert len(good_reloads) == 2
assert raw_send.await_count == 2 # good mint paid for both units assert raw_send.await_count == 2 # good mint paid for both units
@@ -219,6 +281,7 @@ async def test_periodic_payout_handles_session_creation_failure() -> None:
patch.object(settings, "payout_interval_seconds", _INTERVAL), patch.object(settings, "payout_interval_seconds", _INTERVAL),
patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()), patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()),
patch("routstr.wallet.db.create_session", create_session), patch("routstr.wallet.db.create_session", create_session),
patch.object(settings, "min_payout_sat", 10),
patch( patch(
"routstr.wallet._get_supported_mint_units", "routstr.wallet._get_supported_mint_units",
AsyncMock(return_value=["sat", "msat"]), AsyncMock(return_value=["sat", "msat"]),
@@ -237,11 +300,12 @@ async def test_periodic_payout_handles_session_creation_failure() -> None:
with pytest.raises(_LoopBreak): with pytest.raises(_LoopBreak):
await periodic_payout() await periodic_payout()
# The liability session is opened per mint/unit (sat + msat), and each # Per mint/unit (sat + msat) a session is opened twice: once by the stale
# DB failure retains the cycle-specific alert wording while remaining # payout-history sweep and once for the liability read. Each DB failure is
# isolated to its own iteration. # logged and isolated to its own step; the liability error keeps the
assert create_session.call_count == 2 # cycle-specific alert wording.
assert logger.error.call_count == 2 assert create_session.call_count == 4
assert logger.error.call_count == 4
message = logger.error.call_args.args[0] message = logger.error.call_args.args[0]
extra = logger.error.call_args.kwargs["extra"] extra = logger.error.call_args.kwargs["extra"]
assert message == "Error in periodic payout cycle: RuntimeError" assert message == "Error in periodic payout cycle: RuntimeError"
@@ -255,3 +319,282 @@ async def test_payout_units_excludes_units_the_sender_cannot_pay() -> None:
AsyncMock(return_value=["usd", "sat", "eur", "msat"]), AsyncMock(return_value=["usd", "sat", "eur", "msat"]),
): ):
assert await _payout_units("http://mint:3338") == ["sat", "msat"] assert await _payout_units("http://mint:3338") == ["sat", "msat"]
@pytest.mark.asyncio
async def test_periodic_payout_caps_amount_at_max_payout_sat() -> None:
"""Available balance above max_payout_sat is capped for a single payout."""
from routstr.core.settings import settings
raw_send = AsyncMock(return_value=1000)
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.object(settings, "min_payout_sat", 10),
patch.object(settings, "max_payout_sat", 250_000),
patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()),
patch("routstr.wallet.db.create_session", _fake_session),
patch(
"routstr.wallet._get_supported_mint_units",
AsyncMock(return_value=["sat"]),
),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
patch(
"routstr.wallet.get_proofs_per_mint_and_unit",
MagicMock(return_value=[MagicMock(amount=1_000_000)]),
),
patch(
"routstr.wallet.slow_filter_spend_proofs",
AsyncMock(side_effect=lambda proofs, wallet: proofs),
),
patch(
"routstr.wallet.db.total_user_liability",
AsyncMock(return_value=0),
),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch("routstr.wallet.raw_send_to_lnurl", raw_send),
):
with pytest.raises(_LoopBreak):
await periodic_payout()
assert raw_send.await_count >= 1
assert raw_send.await_args_list[0].kwargs["amount"] == 250_000
@pytest.mark.asyncio
async def test_payout_history_records_the_capped_amount() -> None:
"""History stores what is actually sent, not the uncapped balance."""
from routstr.core.settings import settings
record_payout = AsyncMock()
settle_payout = AsyncMock()
async def send(*args: object, **kwargs: object) -> int:
await kwargs["on_melt_quote"]( # type: ignore[index,operator]
"quote-capped", "lnbc1capped"
)
return 250_000_000
raw_send = AsyncMock(side_effect=send)
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.object(settings, "min_payout_sat", 10),
patch.object(settings, "max_payout_sat", 250_000),
patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()),
patch("routstr.wallet.db.create_session", _fake_session),
patch(
"routstr.wallet._get_supported_mint_units",
AsyncMock(return_value=["sat"]),
),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
patch(
"routstr.wallet.get_proofs_per_mint_and_unit",
MagicMock(return_value=[MagicMock(amount=1_000_000)]),
),
patch(
"routstr.wallet.slow_filter_spend_proofs",
AsyncMock(side_effect=lambda proofs, wallet: proofs),
),
patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=0)),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch(
"routstr.wallet.db.list_unsettled_lightning_payouts",
AsyncMock(return_value=[]),
),
patch("routstr.wallet.db.record_lightning_payout", record_payout),
patch("routstr.wallet.db.settle_lightning_payout", settle_payout),
patch("routstr.wallet.raw_send_to_lnurl", raw_send),
):
with pytest.raises(_LoopBreak):
await periodic_payout()
record_payout.assert_awaited_once_with(
ANY,
quote_id="quote-capped",
bolt11="lnbc1capped",
amount_sats=250_000,
mint_url="http://mint:3338",
destination="owner@ln.tld",
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("error", "expected_status"),
[
(MeltUnpaidError("mint confirmed unpaid"), "failed"),
(MeltOutcomeAmbiguousError("outcome unknown"), "reconciliation_required"),
(RuntimeError("HTTP 500 after dispatch"), "reconciliation_required"),
],
)
async def test_payout_history_marks_failed_only_on_proven_non_payment(
error: Exception, expected_status: str
) -> None:
"""Only a mint-confirmed unpaid melt is recorded as failed."""
from routstr.core.settings import settings
settle_payout = AsyncMock()
async def send(*args: object, **kwargs: object) -> int:
await kwargs["on_melt_quote"]( # type: ignore[index,operator]
"quote-err", "lnbc1err"
)
raise error
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.object(settings, "min_payout_sat", 10),
patch.object(settings, "max_payout_sat", 250_000),
patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()),
patch("routstr.wallet.db.create_session", _fake_session),
patch(
"routstr.wallet._get_supported_mint_units",
AsyncMock(return_value=["sat"]),
),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
patch(
"routstr.wallet.get_proofs_per_mint_and_unit",
MagicMock(return_value=[MagicMock(amount=1_000_000)]),
),
patch(
"routstr.wallet.slow_filter_spend_proofs",
AsyncMock(side_effect=lambda proofs, wallet: proofs),
),
patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=0)),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch(
"routstr.wallet.db.list_unsettled_lightning_payouts",
AsyncMock(return_value=[]),
),
patch("routstr.wallet.db.record_lightning_payout", AsyncMock()),
patch("routstr.wallet.db.settle_lightning_payout", settle_payout),
patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(side_effect=send)),
):
with pytest.raises(_LoopBreak):
await periodic_payout()
settle_payout.assert_awaited_once_with(
ANY, "quote-err", status=expected_status, amount_sats=None
)
@pytest.mark.asyncio
async def test_payout_history_write_failure_does_not_block_payout() -> None:
"""A failing history insert is logged; the melt and settlement still run."""
from routstr.core.settings import settings
get_wallet = AsyncMock(return_value=MagicMock())
record_payout = AsyncMock(side_effect=RuntimeError("database is locked"))
settle_payout = AsyncMock()
logger = MagicMock()
async def send(*args: object, **kwargs: object) -> int:
await kwargs["on_melt_quote"]( # type: ignore[index,operator]
"quote-1", "lnbc1payout"
)
return 1_000_000
raw_send = AsyncMock(side_effect=send)
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.object(settings, "max_payout_sat", 250_000),
patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()),
patch("routstr.wallet.db.create_session", _fake_session),
patch(
"routstr.wallet._get_supported_mint_units",
AsyncMock(return_value=["sat"]),
),
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.total_user_liability", AsyncMock(return_value=0)),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch(
"routstr.wallet.db.list_unsettled_lightning_payouts",
AsyncMock(return_value=[]),
),
patch("routstr.wallet.db.record_lightning_payout", record_payout),
patch("routstr.wallet.db.settle_lightning_payout", settle_payout),
patch("routstr.wallet.raw_send_to_lnurl", raw_send),
patch("routstr.wallet.logger", logger),
):
with pytest.raises(_LoopBreak):
await periodic_payout()
record_payout.assert_awaited_once()
assert raw_send.await_count == 1
settle_payout.assert_awaited_once_with(
ANY, "quote-1", status="paid", amount_sats=1_000
)
messages = [call.args[0] for call in logger.error.call_args_list]
assert "Failed to record Lightning payout history" in messages
@pytest.mark.asyncio
async def test_stale_payout_history_is_reconciled_from_mint_state() -> None:
"""Stale out-rows follow the mint's verdict; pending/unknown are left alone."""
stale = [
MagicMock(payment_hash="q-paid"),
MagicMock(payment_hash="q-unpaid"),
MagicMock(payment_hash="q-pending"),
MagicMock(payment_hash="q-unknown"),
]
states = {
"q-paid": "paid",
"q-unpaid": "unpaid",
"q-pending": "pending",
"q-unknown": "unknown",
}
settle_payout = AsyncMock()
async def _state(_mint: str, _unit: str, quote_id: str) -> str:
return states[quote_id]
with (
patch("routstr.wallet.db.create_session", _fake_session),
patch(
"routstr.wallet.db.list_unsettled_lightning_payouts",
AsyncMock(return_value=stale),
),
patch("routstr.wallet._check_bolt11_payment_status_locked", _state),
patch("routstr.wallet.db.settle_lightning_payout", settle_payout),
):
await _reconcile_stale_payout_history("http://mint:3338", "sat")
assert settle_payout.await_args_list == [
((ANY, "q-paid"), {"status": "paid", "amount_sats": None}),
((ANY, "q-unpaid"), {"status": "failed", "amount_sats": None}),
]
@@ -0,0 +1,168 @@
import asyncio
from collections.abc import AsyncGenerator, AsyncIterator
from typing import cast
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi.responses import StreamingResponse
from routstr.upstream.base import BaseUpstreamProvider
async def _chunks() -> AsyncIterator[bytes]:
yield b"chunk"
def _forwarding_case() -> tuple[
BaseUpstreamProvider,
MagicMock,
MagicMock,
MagicMock,
MagicMock,
MagicMock,
MagicMock,
]:
provider = BaseUpstreamProvider("https://api.example.com", "test-key")
request = MagicMock()
request.method = "POST"
request.query_params = {}
key = MagicMock()
key.hashed_key = "key-hash"
session = MagicMock()
model = MagicMock()
model.forwarded_model_id = None
model.id = "model"
response = MagicMock(spec=httpx.Response)
response.status_code = 200
response.headers = {"content-type": "application/octet-stream"}
response.aclose = AsyncMock()
response.aiter_bytes = MagicMock(side_effect=_chunks)
client = MagicMock()
client.build_request.return_value = MagicMock()
client.send = AsyncMock(return_value=response)
return provider, request, key, session, model, response, client
async def _forward(
method_name: str,
*,
reservation_snapshot: object | None,
) -> tuple[StreamingResponse, MagicMock, BaseUpstreamProvider]:
provider, request, key, session, model, response, client = _forwarding_case()
prepare_method = (
"prepare_request_body"
if method_name == "forward_request"
else "prepare_responses_request_body"
)
with (
patch(
"routstr.upstream.base.acquire_upstream_http_client", return_value=client
),
patch.object(provider, "normalize_request_path", return_value="audio/speech"),
patch.object(
provider,
"build_request_url",
return_value="https://api.example.com/audio/speech",
),
patch.object(provider, prepare_method, return_value=b"{}"),
patch.object(provider, "prepare_params", return_value={}),
):
result = await getattr(provider, method_name)(
request=request,
path="audio/speech",
headers={},
request_body=b"{}",
key=key,
max_cost_for_model=1_000,
session=session,
model_obj=model,
reservation_snapshot=reservation_snapshot,
)
assert isinstance(result, StreamingResponse)
return result, response, provider
@pytest.mark.asyncio
@pytest.mark.parametrize(
"method_name", ["forward_request", "forward_responses_request"]
)
async def test_cancellation_before_stream_handoff_closes_response_once(
method_name: str,
) -> None:
provider, request, key, session, model, response, client = _forwarding_case()
prepare_method = (
"prepare_request_body"
if method_name == "forward_request"
else "prepare_responses_request_body"
)
lookup_started = asyncio.Event()
async def wait_for_reservation(*_: object) -> None:
lookup_started.set()
await asyncio.Future()
with (
patch(
"routstr.upstream.base.acquire_upstream_http_client", return_value=client
),
patch.object(provider, "normalize_request_path", return_value="audio/speech"),
patch.object(
provider,
"build_request_url",
return_value="https://api.example.com/audio/speech",
),
patch.object(provider, prepare_method, return_value=b"{}"),
patch.object(provider, "prepare_params", return_value={}),
patch(
"routstr.upstream.base.get_reservation_snapshot",
side_effect=wait_for_reservation,
),
):
task = asyncio.create_task(
getattr(provider, method_name)(
request=request,
path="audio/speech",
headers={},
request_body=b"{}",
key=key,
max_cost_for_model=1_000,
session=session,
model_obj=model,
)
)
await lookup_started.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
response.aclose.assert_awaited_once_with()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"method_name", ["forward_request", "forward_responses_request"]
)
async def test_successful_stream_handoff_does_not_close_response_early(
method_name: str,
) -> None:
result, response, provider = await _forward(
method_name,
reservation_snapshot=MagicMock(),
)
response.aclose.assert_not_awaited()
iterator = cast(AsyncGenerator[bytes, None], result.body_iterator)
with patch.object(
provider,
"_finalize_generic_streaming_payment",
new=AsyncMock(),
):
assert await anext(iterator) == b"chunk"
await iterator.aclose()
response.aclose.assert_awaited_once_with()
+114 -6
View File
@@ -1,5 +1,8 @@
from unittest.mock import patch
from routstr.upstream.anthropic import AnthropicUpstreamProvider from routstr.upstream.anthropic import AnthropicUpstreamProvider
from routstr.upstream.base import BaseUpstreamProvider from routstr.upstream.base import BaseUpstreamProvider
from routstr.upstream.generic import GenericUpstreamProvider
from routstr.upstream.openrouter import OpenRouterUpstreamProvider from routstr.upstream.openrouter import OpenRouterUpstreamProvider
@@ -32,12 +35,12 @@ def test_apply_provider_field_openrouter_passthrough() -> None:
def test_apply_provider_field_openrouter_no_upstream_provider() -> None: def test_apply_provider_field_openrouter_no_upstream_provider() -> None:
"""If OpenRouter omits the provider field, the real serving provider is """If OpenRouter omits the provider field, the serving provider is
unknown — a bare ``openrouter`` value carries no information.""" unknown but the router is not."""
p = _make_provider(OpenRouterUpstreamProvider, "openrouter") p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"id": "gen-abc"} data: dict = {"id": "gen-abc"}
p._apply_provider_field(data) p._apply_provider_field(data)
assert data["provider"] == "unknown" assert data["provider"] == "openrouter:unknown"
def test_apply_provider_field_openrouter_echoes_router_name() -> None: def test_apply_provider_field_openrouter_echoes_router_name() -> None:
@@ -45,7 +48,35 @@ def test_apply_provider_field_openrouter_echoes_router_name() -> None:
p = _make_provider(OpenRouterUpstreamProvider, "openrouter") p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"provider": "openrouter"} data: dict = {"provider": "openrouter"}
p._apply_provider_field(data) p._apply_provider_field(data)
assert data["provider"] == "unknown" assert data["provider"] == "openrouter:unknown"
def test_apply_provider_field_openrouter_unknown_is_idempotent() -> None:
"""Re-stamping an unknown payload (e.g. in inject_cost_metadata) keeps
``openrouter:unknown`` instead of reading ``unknown`` as a sub-provider."""
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"id": "gen-abc"}
p._apply_provider_field(data)
p._apply_provider_field(data)
assert data["provider"] == "openrouter:unknown"
def test_apply_provider_field_openrouter_warns_once_on_billed_payload() -> None:
"""A missing provider is logged on the payload carrying usage, not on
every stream chunk or on a re-stamp."""
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
chunk: dict = {"type": "response.output_text.delta", "delta": "hi"}
completed: dict = {
"type": "response.completed",
"response": {"id": "gen-abc", "usage": {"input_tokens": 1}},
}
with patch("routstr.upstream.openrouter.logger.warning") as warning:
p._apply_provider_field(chunk)
p._apply_provider_field(completed)
p._apply_provider_field(completed)
warning.assert_called_once()
assert chunk["provider"] == completed["provider"] == "openrouter:unknown"
def test_apply_provider_field_openrouter_idempotent_no_double_prefix() -> None: def test_apply_provider_field_openrouter_idempotent_no_double_prefix() -> None:
@@ -78,14 +109,41 @@ def test_apply_provider_field_blank_upstream_treated_as_missing() -> None:
p = _make_provider(OpenRouterUpstreamProvider, "openrouter") p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"provider": " "} data: dict = {"provider": " "}
p._apply_provider_field(data) p._apply_provider_field(data)
assert data["provider"] == "unknown" assert data["provider"] == "openrouter:unknown"
def test_apply_provider_field_non_string_upstream_treated_as_missing() -> None: def test_apply_provider_field_non_string_upstream_treated_as_missing() -> None:
p = _make_provider(OpenRouterUpstreamProvider, "openrouter") p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"provider": 42} data: dict = {"provider": 42}
p._apply_provider_field(data) p._apply_provider_field(data)
assert data["provider"] == "unknown" assert data["provider"] == "openrouter:unknown"
def test_apply_provider_field_openrouter_reads_nested_envelopes() -> None:
"""Anthropic ``message`` and Responses ``response`` envelopes nest the
upstream provider; it must not be reported as unknown."""
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
message_start: dict = {
"type": "message_start",
"message": {"provider": "Anthropic"},
}
p._apply_provider_field(message_start)
assert message_start["provider"] == "openrouter:Anthropic"
created: dict = {"type": "response.created", "response": {"provider": "OpenAI"}}
p._apply_provider_field(created)
assert created["provider"] == "openrouter:OpenAI"
def test_stamp_streamed_provider_carries_earlier_provider() -> None:
"""Events without their own provider inherit the one reported earlier in
the stream instead of becoming ``unknown``."""
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
first: dict = {"provider": "Fireworks"}
carried = p._stamp_streamed_provider(first, None)
delta: dict = {"type": "content_block_delta"}
assert p._stamp_streamed_provider(delta, carried) == "Fireworks"
assert first["provider"] == delta["provider"] == "openrouter:Fireworks"
def test_apply_provider_field_idempotent_for_direct_upstream() -> None: def test_apply_provider_field_idempotent_for_direct_upstream() -> None:
@@ -127,3 +185,53 @@ def test_inject_cost_metadata_sets_provider() -> None:
p.inject_cost_metadata(response_json, cost_data, key) p.inject_cost_metadata(response_json, cost_data, key)
assert response_json["provider"] == "openrouter:Anthropic" assert response_json["provider"] == "openrouter:Anthropic"
def test_apply_provider_field_generic_uses_upstream_host() -> None:
"""A generic upstream has no router-reported provider; the serving host
identifies it, mirroring ``openrouter:<sub-provider>``."""
p = GenericUpstreamProvider(base_url="https://api.deepseek.com/v1", api_key="k")
data: dict = {"id": "chatcmpl-1", "model": "deepseek-chat"}
p._apply_provider_field(data)
assert data["provider"] == "generic:api.deepseek.com"
def test_apply_provider_field_generic_keeps_upstream_reported_provider() -> None:
p = GenericUpstreamProvider(base_url="https://api.deepseek.com/v1", api_key="k")
data: dict = {"provider": "Fireworks"}
p._apply_provider_field(data)
assert data["provider"] == "generic:Fireworks"
def test_apply_provider_field_generic_idempotent() -> None:
p = GenericUpstreamProvider(base_url="https://api.deepseek.com/v1", api_key="k")
data: dict = {}
p._apply_provider_field(data)
p._apply_provider_field(data)
assert data["provider"] == "generic:api.deepseek.com"
def test_apply_provider_field_sets_provider_url() -> None:
"""Every provider exposes the upstream base URL it served from."""
generic = GenericUpstreamProvider(
base_url="https://api.deepseek.com/v1", api_key="k"
)
data: dict = {}
generic._apply_provider_field(data)
assert data["provider_url"] == "https://api.deepseek.com/v1"
openrouter = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data = {"provider": "Anthropic"}
openrouter._apply_provider_field(data)
assert data["provider_url"] == "https://openrouter.ai/api/v1"
def test_apply_provider_field_masks_private_upstream() -> None:
"""Private or port-bearing upstream URLs are masked the same way model
paths mask them, so neither ``provider`` nor ``provider_url`` leaks a
local address."""
p = GenericUpstreamProvider(base_url="http://10.0.0.5:11434/v1", api_key="k")
data: dict = {}
p._apply_provider_field(data)
assert data["provider"] == "generic:localhost"
assert data["provider_url"] == "http://localhost"
+21
View File
@@ -51,6 +51,8 @@ def test_ambiguous_paths_are_rejected(path: str) -> None:
"v1/chat/completions", "v1/chat/completions",
"chat/completions", "chat/completions",
"v1/responses", "v1/responses",
"v1/messages",
"v1/messages/count_tokens",
"v1/embeddings", "v1/embeddings",
"models", "models",
"v1/models/gpt-4", "v1/models/gpt-4",
@@ -138,6 +140,7 @@ def test_known_prefix_does_not_carry_an_unknown_endpoint(path: str) -> None:
("completions", "POST"), ("completions", "POST"),
("v1/responses", "POST"), ("v1/responses", "POST"),
("v1/messages", "POST"), ("v1/messages", "POST"),
("v1/messages/count_tokens", "POST"),
("v1/embeddings", "POST"), ("v1/embeddings", "POST"),
("models", "GET"), ("models", "GET"),
("attestation", "GET"), ("attestation", "GET"),
@@ -163,6 +166,24 @@ def test_method_must_match_the_endpoint(path: str, method: str) -> None:
assert _forwarding_allowed(path, method) is False assert _forwarding_allowed(path, method) is False
@pytest.mark.parametrize(
"path",
[
"messages/count_tokens",
"v1/messages/count_tokens",
"v1/messages/count_tokens/",
],
)
def test_count_tokens_endpoint_stays_allowed(path: str) -> None:
# Regression guard: /v1/messages/count_tokens is supported end-to-end
# (local handler when the upstream lacks native Anthropic support, plain
# forward otherwise), but the exact-match allowlist once omitted it, so
# Claude Code and the Anthropic SDKs were 404'd on every request. It must
# always be reachable, on POST only.
assert _forwarding_allowed(path, "POST") is True
assert _forwarding_allowed(path, "GET") is False
def test_operator_additions_are_parsed_per_endpoint() -> None: def test_operator_additions_are_parsed_per_endpoint() -> None:
parsed = _parse_extra_allowed_endpoints("POST:v1/rerank, GET:batches ,post:audio/x") parsed = _parse_extra_allowed_endpoints("POST:v1/rerank, GET:batches ,post:audio/x")
assert parsed == { assert parsed == {
+12 -5
View File
@@ -6,6 +6,8 @@ from fastapi.responses import StreamingResponse
from routstr import proxy as proxy_module from routstr import proxy as proxy_module
from .proxy_test_utils import mock_request_stream, patch_proxy_session
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_proxy_closes_request_session_before_returning_response() -> None: async def test_proxy_closes_request_session_before_returning_response() -> None:
@@ -15,9 +17,11 @@ async def test_proxy_closes_request_session_before_returning_response() -> None:
request.headers = {"accept": "application/json"} request.headers = {"accept": "application/json"}
request.url.path = "/not-an-api-route" request.url.path = "/not-an-api-route"
request.state.request_id = "test-request" request.state.request_id = "test-request"
mock_request_stream(request, b"")
session = AsyncMock() session = AsyncMock()
response = await proxy_module.proxy(request, "not-an-api-route", session=session) with patch_proxy_session(session):
response = await proxy_module.proxy(request, "not-an-api-route")
assert response.status_code == 404 assert response.status_code == 404
session.close.assert_awaited_once() session.close.assert_awaited_once()
@@ -26,6 +30,8 @@ async def test_proxy_closes_request_session_before_returning_response() -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_proxy_session_is_closed_before_first_stream_chunk() -> None: async def test_proxy_session_is_closed_before_first_stream_chunk() -> None:
request = MagicMock() request = MagicMock()
request.headers = {}
mock_request_stream(request, b"")
session = AsyncMock() session = AsyncMock()
async def stream() -> AsyncIterator[bytes]: async def stream() -> AsyncIterator[bytes]:
@@ -33,10 +39,11 @@ async def test_proxy_session_is_closed_before_first_stream_chunk() -> None:
yield b"chunk" yield b"chunk"
upstream_response = StreamingResponse(stream()) upstream_response = StreamingResponse(stream())
with patch("routstr.proxy._proxy", AsyncMock(return_value=upstream_response)): with (
response = await proxy_module.proxy( patch("routstr.proxy._proxy", AsyncMock(return_value=upstream_response)),
request, "v1/chat/completions", session=session patch_proxy_session(session),
) ):
response = await proxy_module.proxy(request, "v1/chat/completions")
assert isinstance(response, StreamingResponse) assert isinstance(response, StreamingResponse)
chunks = [chunk async for chunk in response.body_iterator] chunks = [chunk async for chunk in response.body_iterator]
@@ -1,13 +1,21 @@
from __future__ import annotations from __future__ import annotations
import json
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest import pytest
from fastapi import FastAPI from fastapi import FastAPI
from fastapi.responses import Response from fastapi.responses import Response
from httpx import ASGITransport, AsyncClient from httpx import ASGITransport, AsyncClient
from routstr import proxy as proxy_module from routstr import proxy as proxy_module
from routstr.core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS,
UPSTREAM_UNAVAILABLE,
)
@pytest.fixture @pytest.fixture
@@ -150,3 +158,141 @@ def test_attestation_upstream_selection_is_tinfoil_only() -> None:
assert proxy_module._select_unauthenticated_get_upstreams( assert proxy_module._select_unauthenticated_get_upstreams(
"attestationjunk", [non_tinfoil, tinfoil] "attestationjunk", [non_tinfoil, tinfoil]
) == [non_tinfoil, tinfoil] ) == [non_tinfoil, tinfoil]
# --------------------------------------------------------------------------- #
# Unauthenticated GET: upstream 5xx -> 424 + scope header, still retryable.
# --------------------------------------------------------------------------- #
def _attributed_424() -> Response:
"""The response a provider hands back for an upstream-attributed 5xx."""
import json as _json
return Response(
content=_json.dumps(
{
"error": {
"type": "upstream_error",
"code": UPSTREAM_UNAVAILABLE,
"message": "Attestation upstream returned 503",
"upstream_status": 503,
}
}
).encode(),
status_code=UPSTREAM_ERROR_STATUS,
media_type="application/json",
headers={ERROR_SCOPE_HEADER: ERROR_SCOPE_UPSTREAM},
)
def _attestation_provider(forward: AsyncMock) -> MagicMock:
provider = MagicMock()
provider.provider_type = "tinfoil"
provider.prepare_headers = MagicMock(return_value={})
provider.forward_get_request = forward
return provider
@pytest.mark.asyncio
async def test_unauthenticated_get_returns_attributed_424_when_all_fail(
monkeypatch: pytest.MonkeyPatch, proxy_app: FastAPI
) -> None:
tinfoil = _attestation_provider(AsyncMock(return_value=_attributed_424()))
monkeypatch.setattr(proxy_module, "_upstreams", [tinfoil])
async with AsyncClient(
transport=ASGITransport(app=proxy_app), # type: ignore[arg-type]
base_url="http://test",
) as client:
response = await client.get("/attestation")
assert response.status_code == UPSTREAM_ERROR_STATUS
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
payload = json.loads(response.content)
assert payload["error"]["code"] == UPSTREAM_UNAVAILABLE
assert payload["error"]["upstream_status"] == 503
@pytest.mark.asyncio
async def test_unauthenticated_get_fails_over_past_an_attributed_424(
monkeypatch: pytest.MonkeyPatch, proxy_app: FastAPI
) -> None:
"""An upstream-attributed 424 stays retryable: the caller sees the healthy
provider's response and never the upstream error."""
failing = _attestation_provider(AsyncMock(return_value=_attributed_424()))
healthy = _attestation_provider(
AsyncMock(return_value=Response(status_code=200, content=b'{"ok":true}'))
)
monkeypatch.setattr(proxy_module, "_upstreams", [failing, healthy])
async with AsyncClient(
transport=ASGITransport(app=proxy_app), # type: ignore[arg-type]
base_url="http://test",
) as client:
response = await client.get("/attestation")
assert response.status_code == 200
assert response.content == b'{"ok":true}'
failing.forward_get_request.assert_awaited_once()
healthy.forward_get_request.assert_awaited_once()
assert ERROR_SCOPE_HEADER not in response.headers
@pytest.mark.asyncio
async def test_attestation_host_5xx_is_attributed_to_the_upstream(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""The Tinfoil attestation hop itself maps its 5xx to 424 + upstream scope."""
from routstr.upstream.tinfoil import TinfoilUpstreamProvider
class _FakeClient:
async def __aenter__(self) -> "_FakeClient":
return self
async def __aexit__(self, *_exc: object) -> bool:
return False
async def get(self, _url: str, headers: dict | None = None) -> httpx.Response:
return httpx.Response(status_code=503, content=b"atc down")
monkeypatch.setattr(
"routstr.upstream.tinfoil.httpx.AsyncClient", lambda **_kw: _FakeClient()
)
provider = TinfoilUpstreamProvider(api_key="k")
response = await provider._proxy_attestation({})
assert response.status_code == UPSTREAM_ERROR_STATUS
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
payload = json.loads(bytes(response.body))
assert payload["error"]["type"] == "upstream_error"
assert payload["error"]["code"] == UPSTREAM_UNAVAILABLE
assert payload["error"]["upstream_status"] == 503
@pytest.mark.asyncio
async def test_attestation_host_4xx_passes_through(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from routstr.upstream.tinfoil import TinfoilUpstreamProvider
class _FakeClient:
async def __aenter__(self) -> "_FakeClient":
return self
async def __aexit__(self, *_exc: object) -> bool:
return False
async def get(self, _url: str, headers: dict | None = None) -> httpx.Response:
return httpx.Response(status_code=404, content=b"missing")
monkeypatch.setattr(
"routstr.upstream.tinfoil.httpx.AsyncClient", lambda **_kw: _FakeClient()
)
provider = TinfoilUpstreamProvider(api_key="k")
response = await provider._proxy_attestation({})
assert response.status_code == 404
assert bytes(response.body) == b"missing"
+312
View File
@@ -0,0 +1,312 @@
import logging
import subprocess
import sys
import textwrap
import threading
from pathlib import Path
import pytest
import routstr.core.logging as routstr_logging
from routstr.core.logging import QueuedDailyRotatingFileHandler
def _log_text(tmp_path: Path) -> str:
return "".join(path.read_text() for path in sorted(tmp_path.glob("app_*.log")))
def _make_handler(
tmp_path: Path, name: str
) -> tuple[logging.Logger, QueuedDailyRotatingFileHandler]:
handler = QueuedDailyRotatingFileHandler(
str(tmp_path / "app.log"), when="midnight", backupCount=1
)
handler.setFormatter(logging.Formatter("%(message)s"))
logger = logging.Logger(name)
logger.addHandler(handler)
return logger, handler
def test_queued_file_handler_flushes_records_on_close(tmp_path: Path) -> None:
logger, handler = _make_handler(tmp_path, "queued-file-test")
try:
logger.info("written from listener")
handler.flush()
assert "written from listener" in _log_text(tmp_path)
finally:
handler.close()
def test_queued_file_handler_loses_no_records_on_close(tmp_path: Path) -> None:
logger, handler = _make_handler(tmp_path, "queued-file-drain-test")
try:
for index in range(400):
logger.info("Payment processed successfully %d", index)
finally:
handler.close()
written = _log_text(tmp_path)
assert written.count("Payment processed successfully") == 400
def test_queued_file_handler_keeps_logging_after_close(tmp_path: Path) -> None:
"""dictConfig closes live handlers; uvicorn runs one after app import."""
logger, handler = _make_handler(tmp_path, "queued-file-reopen-test")
logger.info("before close")
handler.close()
logger.info("after close")
handler.close()
assert "after close" in _log_text(tmp_path)
handler_list = getattr(logging, "_handlerList")
handler_list[:] = [
reference for reference in handler_list if reference() is not handler
]
handler.close()
logger.info("after reopen")
assert any(reference() is handler for reference in handler_list)
logging.shutdown(
handlerList=[reference for reference in handler_list if reference() is handler]
)
assert "after reopen" in _log_text(tmp_path)
def test_queued_file_handler_contains_reopen_failures(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
logger, handler = _make_handler(tmp_path, "queued-file-failure-test")
handler.close()
attempts = 0
def fail_to_open(*args: object, **kwargs: object) -> None:
nonlocal attempts
attempts += 1
raise OSError("disk unavailable")
errors: list[logging.LogRecord] = []
monkeypatch.setattr(routstr_logging, "DailyRotatingFileHandler", fail_to_open)
monkeypatch.setattr(type(handler), "handleError", lambda _self, r: errors.append(r))
for _ in range(50):
logger.info("must not reach billing")
assert attempts == 1
assert len(errors) == 1
handler.close()
def test_queued_file_handler_reports_records_dropped_during_backoff(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str]
) -> None:
logger, handler = _make_handler(tmp_path, "queued-file-drop-report-test")
handler.close()
def fail_to_open(*args: object, **kwargs: object) -> None:
raise OSError("disk unavailable")
monkeypatch.setattr(routstr_logging, "DailyRotatingFileHandler", fail_to_open)
monkeypatch.setattr(type(handler), "handleError", lambda _self, _r: None)
for _ in range(50):
logger.info("must not vanish without a trace")
stderr = capsys.readouterr().err
assert stderr.count("dropping records") == 1
assert "is unavailable" in stderr
handler.close()
def test_queued_file_handler_emit_does_not_raise_into_caller(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
logger, handler = _make_handler(tmp_path, "queued-file-emit-failure-test")
handled: list[logging.LogRecord] = []
class BrokenQueue:
def put_nowait(self, _record: logging.LogRecord) -> None:
raise OSError("queue is gone")
monkeypatch.setattr(handler, "_queue", BrokenQueue())
monkeypatch.setattr(type(handler), "handleError", lambda _s, r: handled.append(r))
logger.info("settlement line")
assert len(handled) == 1
handler.close()
def test_queued_file_handler_recovers_after_close_timeout(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
logger, handler = _make_handler(tmp_path, "queued-file-timeout-test")
listener_blocked = threading.Event()
allow_listener = threading.Event()
old_target = handler._target
original_handle = old_target.handle
original_close = old_target.close
target_closed = threading.Event()
close_count = 0
def blocked_handle(record: logging.LogRecord) -> bool:
listener_blocked.set()
assert allow_listener.wait(timeout=10)
return original_handle(record)
def track_close() -> None:
nonlocal close_count
close_count += 1
original_close()
target_closed.set()
monkeypatch.setattr(old_target, "handle", blocked_handle)
monkeypatch.setattr(old_target, "close", track_close)
handler._drain_timeout_seconds = 0.01
logger.info("blocked record")
assert listener_blocked.wait(timeout=10)
handler.close()
logger.info("record after timeout")
allow_listener.set()
handler.close()
assert target_closed.wait(timeout=10)
assert close_count == 1
assert "record after timeout" in _log_text(tmp_path)
def test_queued_file_handler_reopens_when_close_wins_emit_race(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
logger, handler = _make_handler(tmp_path, "queued-file-atomic-race-test")
emitter_waiting = threading.Event()
allow_emitter = threading.Event()
original_acquire = handler.acquire
emitter_thread: threading.Thread | None = None
gated = True
def gated_acquire() -> None:
nonlocal gated
if gated and threading.current_thread() is emitter_thread:
gated = False
emitter_waiting.set()
assert allow_emitter.wait(timeout=10)
original_acquire()
monkeypatch.setattr(handler, "acquire", gated_acquire)
emitter_thread = threading.Thread(target=logger.info, args=("racing record",))
try:
emitter_thread.start()
assert emitter_waiting.wait(timeout=10)
handler.close()
allow_emitter.set()
emitter_thread.join(timeout=10)
assert not emitter_thread.is_alive()
handler.close()
assert "racing record" in _log_text(tmp_path)
finally:
allow_emitter.set()
handler.close()
def test_queued_file_handler_survives_close_racing_with_emit(tmp_path: Path) -> None:
logger, handler = _make_handler(tmp_path, "queued-file-race-test")
done = threading.Event()
def spam() -> None:
while not done.is_set():
logger.info("racing record")
def churn() -> None:
for _ in range(50):
handler.close()
emitter = threading.Thread(target=spam, daemon=True)
closer = threading.Thread(target=churn, daemon=True)
try:
emitter.start()
closer.start()
closer.join(timeout=10)
done.set()
emitter.join(timeout=10)
assert not closer.is_alive(), "close() deadlocked against a concurrent emit()"
assert not emitter.is_alive(), "emit() deadlocked against a concurrent close()"
logger.info("final record")
handler.flush()
assert "final record" in _log_text(tmp_path)
finally:
done.set()
handler.close()
def test_queued_file_handler_does_not_deadlock_against_dictconfig(
tmp_path: Path,
) -> None:
script = textwrap.dedent(
"""
import logging
import logging.config
import sys
import threading
import time
from pathlib import Path
from routstr.core.logging import QueuedDailyRotatingFileHandler
log_dir = Path(sys.argv[1])
handler = QueuedDailyRotatingFileHandler(
str(log_dir / "app.log"), when="midnight", backupCount=1
)
handler.setFormatter(logging.Formatter("%(message)s"))
logger = logging.Logger("queued-file-dictconfig-test")
logger.addHandler(handler)
emitted = threading.Event()
def spam():
for _ in range(100):
logger.info("racing record")
emitted.set()
handler.close()
time.sleep(0.001)
def reconfigure():
assert emitted.wait(timeout=10)
for _ in range(10):
logging.config.dictConfig(
{
"version": 1,
"disable_existing_loggers": False,
"handlers": {},
"loggers": {},
"root": {"level": "INFO"},
}
)
time.sleep(0.001)
emitter = threading.Thread(target=spam)
configurer = threading.Thread(target=reconfigure)
emitter.start()
configurer.start()
emitter.join(timeout=20)
configurer.join(timeout=20)
assert not emitter.is_alive(), "logging deadlocked against dictConfig"
assert not configurer.is_alive(), "dictConfig deadlocked against logging"
handler.close()
"""
)
result = subprocess.run(
[sys.executable, "-c", script, str(tmp_path)],
capture_output=True,
text=True,
timeout=30,
)
assert result.returncode == 0, result.stderr
assert "racing record" in _log_text(tmp_path)
+190
View File
@@ -16,9 +16,15 @@ from routstr.upstream.request_correction import (
Correction, Correction,
correct_request, correct_request,
extract_error_message, extract_error_message,
rename_unsupported_param,
strip_unsupported_param, strip_unsupported_param,
) )
OPENAI_MAX_TOKENS_ERROR = (
"Unsupported parameter: 'max_tokens' is not supported with this model. "
"Use 'max_completion_tokens' instead."
)
def _body(**kwargs: object) -> bytes: def _body(**kwargs: object) -> bytes:
return json.dumps(kwargs).encode() return json.dumps(kwargs).encode()
@@ -139,6 +145,190 @@ class TestStripUnsupportedParam:
assert strip_unsupported_param(body, "`Max_Tokens` is deprecated") is None assert strip_unsupported_param(body, "`Max_Tokens` is deprecated") is None
class TestRenameUnsupportedParam:
def test_renames_max_tokens_for_openai_reasoning_models(self) -> None:
body = {"model": "gpt-5.6-sol", "max_tokens": 256, "messages": []}
result = rename_unsupported_param(body, OPENAI_MAX_TOKENS_ERROR)
assert result is not None
new_body, label = result
assert label == "max_tokens->max_completion_tokens"
assert new_body == {
"model": "gpt-5.6-sol",
"max_completion_tokens": 256,
"messages": [],
}
def test_preserves_key_order(self) -> None:
body = {"model": "m", "max_tokens": 1, "stream": True}
result = rename_unsupported_param(body, OPENAI_MAX_TOKENS_ERROR)
assert result is not None
assert list(result[0]) == ["model", "max_completion_tokens", "stream"]
def test_does_not_mutate_input(self) -> None:
body = {"model": "m", "max_tokens": 8}
assert rename_unsupported_param(body, OPENAI_MAX_TOKENS_ERROR) is not None
assert body == {"model": "m", "max_tokens": 8}
def test_renames_between_any_output_caps(self) -> None:
caps = (
"max_tokens",
"max_completion_tokens",
"max_output_tokens",
"max_tokens_to_sample",
)
for param in caps:
for replacement in caps:
if param == replacement:
continue
message = f"`{param}` is deprecated. Use `{replacement}` instead."
result = rename_unsupported_param({param: 7}, message)
assert result == ({replacement: 7}, f"{param}->{replacement}"), (
param,
replacement,
)
def test_renames_non_spend_param(self) -> None:
message = "'functions' is deprecated. Use 'tools' instead."
result = rename_unsupported_param({"functions": [{"name": "f"}]}, message)
assert result == ({"tools": [{"name": "f"}]}, "functions->tools")
def test_matches_across_quote_styles_case_and_newlines(self) -> None:
for message in (
'Unsupported parameter: "max_tokens" is not supported.\nUse '
'"max_completion_tokens" instead.',
"`max_tokens` IS UNSUPPORTED here; please USE `max_completion_tokens`"
" INSTEAD",
"'max_tokens' is no longer supported, use 'max_completion_tokens' instead",
):
result = rename_unsupported_param({"max_tokens": 3}, message)
assert result is not None, message
assert result[0] == {"max_completion_tokens": 3}
def test_refuses_renames_that_change_the_spend_bound(self) -> None:
for param, replacement in (
("max_tokens", "n"),
("n", "best_of"),
("best_of", "n"),
("temperature", "max_tokens"),
("max_tokens", "temperature"),
("n", "max_tokens"),
):
message = f"'{param}' is not supported. Use '{replacement}' instead."
assert rename_unsupported_param({param: 2}, message) is None, (
param,
replacement,
)
def test_spend_guard_is_case_insensitive(self) -> None:
ok = "'Max_Tokens' is not supported. Use 'MAX_COMPLETION_TOKENS' instead."
assert rename_unsupported_param({"Max_Tokens": 4}, ok) == (
{"MAX_COMPLETION_TOKENS": 4},
"Max_Tokens->MAX_COMPLETION_TOKENS",
)
bad = "'Max_Tokens' is not supported. Use 'N' instead."
assert rename_unsupported_param({"Max_Tokens": 4}, bad) is None
def test_declines_when_replacement_already_present(self) -> None:
body = {"max_tokens": 4, "max_completion_tokens": 8}
assert rename_unsupported_param(body, OPENAI_MAX_TOKENS_ERROR) is None
def test_declines_when_param_absent(self) -> None:
assert rename_unsupported_param({"model": "m"}, OPENAI_MAX_TOKENS_ERROR) is None
def test_declines_self_rename(self) -> None:
message = "'max_tokens' is deprecated. Use 'max_tokens' instead."
assert rename_unsupported_param({"max_tokens": 1}, message) is None
def test_declines_unquoted_or_missing_replacement(self) -> None:
for message in (
"`gpt-3` is deprecated, use gpt-4 instead",
"'max_tokens' is not supported, use max_completion_tokens instead",
"'max_tokens' is not supported with this model.",
"Use 'max_completion_tokens' instead.",
):
assert rename_unsupported_param({"max_tokens": 1}, message) is None, message
def test_declines_nested_only_param(self) -> None:
body = {"reasoning": {"max_tokens": 5}}
assert rename_unsupported_param(body, OPENAI_MAX_TOKENS_ERROR) is None
class TestCorrectRequestRename:
def test_openai_max_tokens_error_is_renamed_not_refused(self) -> None:
body = _body(model="gpt-5.6-sol", max_tokens=512, messages=[])
result = correct_request(body, OPENAI_MAX_TOKENS_ERROR, set())
assert isinstance(result, Correction)
assert result.label == "max_tokens->max_completion_tokens"
decoded = json.loads(result.body)
assert "max_tokens" not in decoded
assert decoded["max_completion_tokens"] == 512
def test_rename_wins_over_strip_for_non_spend_param(self) -> None:
body = _body(model="m", functions=[1])
result = correct_request(
body, "'functions' is deprecated. Use 'tools' instead.", set()
)
assert result is not None
assert json.loads(result.body) == {"model": "m", "tools": [1]}
def test_unsafe_rename_of_cap_still_surfaces_error(self) -> None:
body = _body(model="m", max_tokens=5)
assert (
correct_request(
body, "'max_tokens' is not supported. Use 'n' instead.", set()
)
is None
)
def test_applied_rename_does_not_repeat_or_strip_cap(self) -> None:
body = _body(model="m", max_tokens=5)
applied = {"max_tokens->max_completion_tokens"}
assert correct_request(body, OPENAI_MAX_TOKENS_ERROR, applied) is None
def test_rename_ping_pong_terminates(self) -> None:
"""An upstream that flip-flops between names cannot loop forever."""
forward = OPENAI_MAX_TOKENS_ERROR
backward = "'max_completion_tokens' is not supported. Use 'max_tokens' instead."
body = _body(model="m", max_tokens=5)
applied: set[str] = set()
for attempt in range(10):
message = forward if attempt % 2 == 0 else backward
result = correct_request(body, message, applied)
if result is None:
break
body, applied = result.body, applied | {result.label}
else:
raise AssertionError("correction loop did not terminate")
assert applied == {
"max_tokens->max_completion_tokens",
"max_completion_tokens->max_tokens",
}
assert json.loads(body) == {"model": "m", "max_tokens": 5}
def test_buffered_openai_error_response_is_renamed(self) -> None:
resp = Response(
content=json.dumps(
{
"error": {
"message": OPENAI_MAX_TOKENS_ERROR,
"type": "invalid_request_error",
"param": "max_tokens",
"code": "unsupported_parameter",
}
}
).encode(),
status_code=400,
)
body = _body(model="gpt-5.6-sol", max_tokens=64, stream=True)
result = correct_request(body, extract_error_message(resp), set())
assert result is not None
assert json.loads(result.body) == {
"model": "gpt-5.6-sol",
"max_completion_tokens": 64,
"stream": True,
}
class TestExtractErrorMessage: class TestExtractErrorMessage:
def test_extracts_nested_error_message(self) -> None: def test_extracts_nested_error_message(self) -> None:
resp = Response( resp = Response(
+285
View File
@@ -0,0 +1,285 @@
import asyncio
from collections.abc import AsyncGenerator
from contextlib import asynccontextmanager
from pathlib import Path
from unittest.mock import patch
import pytest
from sqlalchemy.ext.asyncio import create_async_engine
from sqlalchemy.pool import NullPool
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.responses import PlainTextResponse
from starlette.routing import Route
from starlette.types import Message, Receive, Scope, Send
import routstr.core.db as db_module
from routstr.auth import (
ReservationSnapshot,
_claim_reservation_for_charge,
_stop_reservation_heartbeat,
pay_for_request,
)
from routstr.core.db import ApiKey, ReservationRelease
from routstr.core.lifecycle import RequestLifecycleMiddleware
from routstr.core.middleware import LoggingMiddleware
from routstr.core.settings import settings
@pytest.mark.asyncio
@pytest.mark.parametrize("reason", ["disconnect", "deadline", "send"])
async def test_lifecycle_stops_live_work(reason: str) -> None:
closed = asyncio.Event()
receive_queue: asyncio.Queue[Message] = asyncio.Queue()
await receive_queue.put({"type": "http.request", "body": b"", "more_body": False})
sent: list[Message] = []
async def app(scope: Scope, receive: Receive, send: Send) -> None:
try:
assert (await receive())["type"] == "http.request"
await send({"type": "http.response.start", "status": 200, "headers": []})
while True:
await send(
{"type": "http.response.body", "body": b"x", "more_body": True}
)
await asyncio.sleep(0.01)
finally:
closed.set()
async def send(message: Message) -> None:
sent.append(message)
if reason == "send" and message["type"] == "http.response.body":
await asyncio.sleep(100)
async def disconnect() -> None:
await asyncio.sleep(0.02)
await receive_queue.put({"type": "http.disconnect"})
task = asyncio.create_task(disconnect()) if reason == "disconnect" else None
with (
patch.object(settings, "max_request_lifetime_seconds", 0.08),
patch.object(settings, "downstream_send_timeout_seconds", 0.03),
patch.object(settings, "request_cleanup_timeout_seconds", 0.1),
):
try:
await asyncio.wait_for(
RequestLifecycleMiddleware(app)(
{"type": "http"}, receive_queue.get, send
),
1,
)
except TimeoutError:
assert reason == "send"
if task:
await task
assert closed.is_set()
assert sent
@pytest.mark.asyncio
async def test_unrelated_oserror_still_propagates() -> None:
receive_queue: asyncio.Queue[Message] = asyncio.Queue()
await receive_queue.put({"type": "http.request", "body": b"", "more_body": False})
async def app(scope: Scope, receive: Receive, send: Send) -> None:
await receive()
raise OSError("Connection reset by peer")
async def send(message: Message) -> None:
pass
with pytest.raises(OSError, match="Connection reset by peer"):
await asyncio.wait_for(
RequestLifecycleMiddleware(app)({"type": "http"}, receive_queue.get, send),
1,
)
@pytest.mark.asyncio
async def test_disconnect_before_headers_preserves_wallet_work() -> None:
receive_queue: asyncio.Queue[Message] = asyncio.Queue()
await receive_queue.put({"type": "http.request", "body": b"", "more_body": False})
entered = asyncio.Event()
finish_wallet = asyncio.Event()
wallet_credited = asyncio.Event()
async def app(scope: Scope, receive: Receive, send: Send) -> None:
await receive()
entered.set()
await finish_wallet.wait() # The mint accepted the token; credit is still pending.
wallet_credited.set()
await send({"type": "http.response.start", "status": 200, "headers": []})
run = asyncio.create_task(
RequestLifecycleMiddleware(app)(
{"type": "http"}, receive_queue.get, lambda message: asyncio.sleep(0)
)
)
await asyncio.wait_for(entered.wait(), 1)
await receive_queue.put({"type": "http.disconnect"})
await asyncio.sleep(0.02)
assert not run.done()
finish_wallet.set()
await asyncio.wait_for(run, 1) # No propagated exception for an expected disconnect.
assert wallet_credited.is_set()
@pytest.mark.asyncio
async def test_disconnect_before_headers_with_logging_middleware() -> None:
receive_queue: asyncio.Queue[Message] = asyncio.Queue()
await receive_queue.put({"type": "http.request", "body": b"", "more_body": False})
entered = asyncio.Event()
finish_wallet = asyncio.Event()
wallet_credited = asyncio.Event()
async def wallet(request: Request) -> PlainTextResponse:
await request.body()
entered.set()
await finish_wallet.wait()
wallet_credited.set()
return PlainTextResponse("settled")
app = RequestLifecycleMiddleware(
LoggingMiddleware(Starlette(routes=[Route("/wallet", wallet, methods=["POST"])]))
)
scope: Scope = {
"type": "http",
"asgi": {"version": "3.0", "spec_version": "2.4"},
"http_version": "1.1",
"method": "POST",
"scheme": "http",
"path": "/wallet",
"raw_path": b"/wallet",
"root_path": "",
"query_string": b"",
"headers": [],
"client": ("test", 1234),
"server": ("test", 80),
}
async def send(message: Message) -> None:
pass
run = asyncio.create_task(app(scope, receive_queue.get, send))
try:
await asyncio.wait_for(entered.wait(), 1)
await receive_queue.put({"type": "http.disconnect"})
await asyncio.sleep(0.02)
assert not run.done()
finish_wallet.set()
await asyncio.wait_for(run, 1) # No propagated exception for an expected disconnect.
assert wallet_credited.is_set()
finally:
finish_wallet.set()
if not run.done():
run.cancel()
await asyncio.gather(run, return_exceptions=True)
@pytest.mark.asyncio
async def test_deadline_cancels_app_before_sending_504() -> None:
receive_queue: asyncio.Queue[Message] = asyncio.Queue()
await receive_queue.put({"type": "http.request", "body": b"", "more_body": False})
sent: list[Message] = []
app_stopped = asyncio.Event()
async def app(scope: Scope, receive: Receive, send: Send) -> None:
await receive()
try:
await asyncio.sleep(100)
finally:
with pytest.raises(OSError, match="Downstream request terminated"):
await send({"type": "http.response.start", "status": 200, "headers": []})
app_stopped.set()
async def send(message: Message) -> None:
assert app_stopped.is_set()
sent.append(message)
with (
patch.object(settings, "max_request_lifetime_seconds", 0.02),
patch.object(settings, "request_cleanup_timeout_seconds", 0.1),
):
await asyncio.wait_for(
RequestLifecycleMiddleware(app)({"type": "http"}, receive_queue.get, send),
1,
)
assert [message["type"] for message in sent] == [
"http.response.start",
"http.response.body",
]
assert sent[0]["status"] == 504
@pytest.mark.asyncio
async def test_disconnect_does_not_release_before_stream_settles(tmp_path: Path) -> None:
engine = create_async_engine(
f"sqlite+aiosqlite:///{tmp_path / 'reservations.db'}", poolclass=NullPool
)
async with engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.create_all)
@asynccontextmanager
async def session() -> AsyncGenerator[AsyncSession, None]:
async with AsyncSession(engine, expire_on_commit=False) as db:
yield db
with patch.object(db_module, "create_session", session):
async with session() as db:
db.add(ApiKey(hashed_key="stream-key", balance=10_000))
await db.commit()
started = asyncio.Event()
finalizer_started = asyncio.Event()
settle = asyncio.Event()
result: asyncio.Future[bool] = asyncio.get_running_loop().create_future()
snapshot: ReservationSnapshot | None = None
receive_queue: asyncio.Queue[Message] = asyncio.Queue()
await receive_queue.put({"type": "http.request", "body": b"", "more_body": False})
async def app(scope: Scope, receive: Receive, send: Send) -> None:
nonlocal snapshot
async with session() as db:
key = await db.get(ApiKey, "stream-key")
assert key is not None
snapshot = await pay_for_request(key, 1000, db)
await receive()
await send({"type": "http.response.start", "status": 200, "headers": []})
started.set()
try:
await asyncio.sleep(100)
finally:
async def finalize() -> None:
assert snapshot is not None
finalizer_started.set()
await settle.wait()
async with session() as db:
claimed = await _claim_reservation_for_charge(snapshot, db)
await db.commit()
await _stop_reservation_heartbeat(snapshot.release_id)
result.set_result(claimed)
asyncio.create_task(finalize())
run = asyncio.create_task(
RequestLifecycleMiddleware(app)(
{"type": "http"}, receive_queue.get, lambda message: asyncio.sleep(0)
)
)
try:
await asyncio.wait_for(started.wait(), 1)
await receive_queue.put({"type": "http.disconnect"})
await asyncio.wait_for(finalizer_started.wait(), 1)
await asyncio.wait_for(run, 1)
settle.set()
assert await asyncio.wait_for(result, 1)
assert snapshot is not None
async with session() as db:
row = await db.get(ReservationRelease, snapshot.release_id)
assert row is not None and row.status == "charged"
finally:
settle.set()
if snapshot is not None:
await _stop_reservation_heartbeat(snapshot.release_id)
await engine.dispose()
+191
View File
@@ -0,0 +1,191 @@
"""Tests for stage timings, the duration header and skipped-path error logging."""
import asyncio
import logging
from collections.abc import AsyncIterator, Iterator
import pytest
from fastapi import FastAPI, HTTPException, Request
from fastapi.responses import StreamingResponse
from fastapi.testclient import TestClient
from routstr.core.middleware import LoggingMiddleware, mark
from routstr.core.settings import settings
class _RecordingHandler(logging.Handler):
def __init__(self) -> None:
super().__init__(level=logging.DEBUG)
self.records: list[logging.LogRecord] = []
def emit(self, record: logging.LogRecord) -> None:
self.records.append(record)
def completions(self) -> list[logging.LogRecord]:
return [r for r in self.records if r.getMessage() == "Request completed"]
@pytest.fixture
def records() -> Iterator[_RecordingHandler]:
handler = _RecordingHandler()
middleware_logger = logging.getLogger("routstr.core.middleware")
middleware_logger.setLevel(logging.INFO)
original_propagate = middleware_logger.propagate
original_handlers = middleware_logger.handlers
middleware_logger.propagate = False
middleware_logger.handlers = [handler]
try:
yield handler
finally:
middleware_logger.handlers = original_handlers
middleware_logger.propagate = original_propagate
@pytest.fixture
def client() -> Iterator[TestClient]:
app = FastAPI()
app.add_middleware(LoggingMiddleware)
# /v1/wallet/info is in _SKIP_LOG_EXACT, so it exercises the suppression path.
@app.get("/v1/wallet/info")
async def wallet_info(fail: bool = False) -> dict[str, str]:
if fail:
raise HTTPException(status_code=400, detail="spent token")
return {"status": "ok"}
@app.post("/v1/chat/completions")
async def completions(request: Request) -> dict[str, str]:
mark(request, "body_read")
mark(request, "auth")
return {"status": "ok"}
@app.post("/v1/chat/completions/stream")
async def streamed(request: Request) -> StreamingResponse:
mark(request, "body_read")
async def body() -> AsyncIterator[bytes]:
yield b"data: one\n\n"
await asyncio.sleep(0.05)
yield b"data: [DONE]\n\n"
return StreamingResponse(body(), media_type="text/event-stream")
@app.get("/admin/api/boom")
async def boom() -> dict[str, str]:
raise HTTPException(status_code=500, detail="boom")
with TestClient(app, raise_server_exceptions=False) as test_client:
yield test_client
def test_skipped_path_logs_4xx(client: TestClient, records: _RecordingHandler) -> None:
assert client.get("/v1/wallet/info", params={"fail": True}).status_code == 400
completions = records.completions()
assert len(completions) == 1
record = completions[0]
assert record.status_code == 400 # type: ignore[attr-defined]
assert record.path == "/v1/wallet/info" # type: ignore[attr-defined]
assert record.method == "GET" # type: ignore[attr-defined]
assert record.duration_ms >= 0 # type: ignore[attr-defined]
def test_skipped_path_does_not_log_2xx(
client: TestClient, records: _RecordingHandler
) -> None:
assert client.get("/v1/wallet/info").status_code == 200
assert records.completions() == []
def test_duration_header_present(client: TestClient) -> None:
response = client.get("/v1/wallet/info")
assert float(response.headers["x-routstr-duration-ms"]) >= 0
def test_stage_fields_on_completion_log(
client: TestClient, records: _RecordingHandler
) -> None:
response = client.post("/v1/chat/completions", json={"model": "m"})
assert response.status_code == 200
completions = records.completions()
assert len(completions) == 1
record = completions[0]
assert record.body_read_ms >= 0 # type: ignore[attr-defined]
assert record.auth_ms >= record.body_read_ms # type: ignore[attr-defined]
assert record.content_length == int( # type: ignore[attr-defined]
response.request.headers["content-length"]
)
def test_bogus_content_length_is_dropped(
client: TestClient, records: _RecordingHandler
) -> None:
assert (
client.get(
"/v1/wallet/info",
params={"fail": True},
headers={"content-length": "not-a-number"},
).status_code
== 400
)
assert records.completions()[0].content_length is None # type: ignore[attr-defined]
def test_streamed_duration_covers_the_body(
client: TestClient, records: _RecordingHandler
) -> None:
response = client.post("/v1/chat/completions/stream", json={"model": "m"})
assert response.status_code == 200
assert response.text.endswith("data: [DONE]\n\n")
completions = records.completions()
assert len(completions) == 1
record = completions[0]
# The body sleeps 50ms, so a duration that stopped at the headers would be
# well under it.
assert record.duration_ms >= 50 # type: ignore[attr-defined]
assert record.time_to_headers_ms < record.duration_ms # type: ignore[attr-defined]
logged_request_id = record.request_id # type: ignore[attr-defined]
assert logged_request_id == response.headers["x-routstr-request-id"]
assert record.body_read_ms >= 0 # type: ignore[attr-defined]
def test_slow_streamed_request_logs_warning(
client: TestClient, records: _RecordingHandler, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(settings, "slow_request_warn_seconds", 0.02)
assert (
client.post("/v1/chat/completions/stream", json={"model": "m"}).status_code
== 200
)
assert records.completions()[0].levelno == logging.WARNING
def test_prefix_skipped_path_still_hides_client_errors(
client: TestClient, records: _RecordingHandler
) -> None:
# /admin/api/* is polled on a timer, so an expired session must not turn
# into one log line per poll; a 500 on the same prefix must still be logged.
assert client.get("/admin/api/balances").status_code == 404
assert records.completions() == []
assert client.get("/admin/api/boom").status_code == 500
assert len(records.completions()) == 1
def test_slow_request_logs_warning(
client: TestClient, records: _RecordingHandler, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(settings, "slow_request_warn_seconds", 0.0)
assert client.post("/v1/chat/completions", json={"model": "m"}).status_code == 200
completions = records.completions()
assert len(completions) == 1
assert completions[0].levelno == logging.WARNING
+20 -5
View File
@@ -1,5 +1,6 @@
import json import json
import os import os
from pathlib import Path
import pytest import pytest
from pydantic.v1 import ValidationError from pydantic.v1 import ValidationError
@@ -7,7 +8,7 @@ from sqlalchemy.ext.asyncio import create_async_engine
from sqlmodel import text from sqlmodel import text
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.settings import Settings, SettingsService, settings from routstr.core.settings import ENV_ONLY_FIELDS, Settings, SettingsService, settings
NSEC_HEX = "1" * 64 NSEC_HEX = "1" * 64
@@ -72,6 +73,20 @@ def test_database_pool_defaults_provide_concurrency_headroom() -> None:
assert s.database_pool_hold_warn_seconds == 10.0 assert s.database_pool_hold_warn_seconds == 10.0
def test_env_only_settings_are_documented() -> None:
env_example = Path(__file__).parents[2] / ".env.example"
documented = {
line.lstrip("# ").split("=", 1)[0]
for line in env_example.read_text().splitlines()
if "=" in line
}
aliases = {
Settings.__fields__[field].field_info.extra["env"] for field in ENV_ONLY_FIELDS
}
assert aliases <= documented
@pytest.mark.parametrize( @pytest.mark.parametrize(
("field", "bad_value"), ("field", "bad_value"),
[ [
@@ -207,9 +222,7 @@ async def test_settings_initialize_discards_unknown_keys() -> None:
# Simulate older persisted key name and an unknown key. # Simulate older persisted key name and an unknown key.
await session.exec( # type: ignore await session.exec( # type: ignore
text( text("UPDATE settings SET data = :data WHERE id = 1").bindparams(
"UPDATE settings SET data = :data WHERE id = 1"
).bindparams(
data='{"name":"LegacyNode","nostr_analytics_enabled":false,"unknown_key":123}' data='{"name":"LegacyNode","nostr_analytics_enabled":false,"unknown_key":123}'
) )
) )
@@ -279,7 +292,9 @@ async def test_upstream_api_key_survives_persistence(
await SettingsService.initialize(session) await SettingsService.initialize(session)
await session.exec( # type: ignore await session.exec( # type: ignore
text("UPDATE settings SET data = :d WHERE id = 1").bindparams( text("UPDATE settings SET data = :d WHERE id = 1").bindparams(
d=json.dumps({"name": "LegacyNode", "upstream_api_key": "sk-only-in-db"}) d=json.dumps(
{"name": "LegacyNode", "upstream_api_key": "sk-only-in-db"}
)
) )
) )
await session.commit() await session.commit()
+74
View File
@@ -0,0 +1,74 @@
import random
import pytest
from routstr.upstream.sse_splitter import SSEEventSplitter
def _reference_split(chunks: list[bytes]) -> tuple[list[bytes], bytes]:
"""The original rescanning implementation the splitter replaces."""
events: list[bytes] = []
buffer = b""
for chunk in chunks:
buffer = (buffer + chunk).replace(b"\r\n", b"\n")
while b"\n\n" in buffer:
raw_event, buffer = buffer.split(b"\n\n", 1)
events.append(raw_event)
return events, buffer
def _split(chunks: list[bytes]) -> tuple[list[bytes], bytes]:
splitter = SSEEventSplitter()
events = [event for chunk in chunks for event in splitter.feed(chunk)]
return events, splitter.flush()
STREAMS = [
b'data: {"a":1}\n\ndata: {"b":2}\n\ndata: [DONE]\n\n',
b'data: {"a":1}\r\n\r\ndata: {"b":2}\r\n\r\ndata: [DONE]\r\n\r\n',
b': OPENROUTER PROCESSING\n\ndata: {"a":1}\n\n: keepalive\n\ndata: [DONE]\n\n',
b'event: response.created\ndata: {"type":"x"}\n\nevent: done\ndata: {"t":1}\n\n',
b'data: {"part":\ndata: "two"}\n\n\n\ndata: {"trailing":true}',
b'data: {"a":1}\r\n\r\ndata: {"tail":1}\r',
b"\n\n\n\n",
b"",
]
@pytest.mark.parametrize("stream", STREAMS)
def test_matches_reference_at_every_two_way_split(stream: bytes) -> None:
for cut in range(len(stream) + 1):
chunks = [stream[:cut], stream[cut:]]
assert _split(chunks) == _reference_split(chunks)
@pytest.mark.parametrize("stream", STREAMS)
def test_matches_reference_on_random_chunkings(stream: bytes) -> None:
rng = random.Random(0)
for _ in range(200):
cuts = sorted(rng.sample(range(len(stream) + 1), min(len(stream), 6)))
bounds = [0, *cuts, len(stream)]
chunks = [stream[a:b] for a, b in zip(bounds, bounds[1:])]
assert _split(chunks) == _reference_split(chunks)
def test_byte_at_a_time_crlf_stream() -> None:
stream = b'data: {"a":1}\r\n\r\ndata: {"b":2}\r\n\r\n'
events, tail = _split([bytes([b]) for b in stream])
assert events == [b'data: {"a":1}', b'data: {"b":2}']
assert tail == b""
def test_flush_returns_held_back_carriage_return() -> None:
splitter = SSEEventSplitter()
assert splitter.feed(b"data: x\r") == []
assert splitter.flush() == b"data: x\r"
assert splitter.flush() == b""
def test_large_event_over_many_chunks() -> None:
payload = b"data: " + b"x" * 200_000 + b"\n\n"
chunks = [payload[i : i + 64] for i in range(0, len(payload), 64)]
events, tail = _split(chunks)
assert events == [payload[:-2]]
assert tail == b""
+155 -7
View File
@@ -9,6 +9,7 @@ Covers:
""" """
import asyncio import asyncio
import math
import time import time
from typing import AsyncGenerator from typing import AsyncGenerator
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
@@ -16,17 +17,21 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlalchemy.pool import StaticPool from sqlalchemy.pool import StaticPool
from sqlmodel import SQLModel from sqlmodel import SQLModel, select
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
import routstr.auth as auth_module
from routstr.auth import pay_for_request from routstr.auth import pay_for_request
from routstr.balance import refund_wallet_endpoint from routstr.balance import refund_wallet_endpoint
from routstr.core.db import ( from routstr.core.db import (
ApiKey, ApiKey,
ReservationRelease,
release_stale_reservations, release_stale_reservations,
reset_all_reserved_balances, reset_all_reserved_balances,
) )
from .proxy_test_utils import mock_request_stream, patch_proxy_session
def _make_engine() -> AsyncEngine: def _make_engine() -> AsyncEngine:
return create_async_engine( return create_async_engine(
@@ -55,11 +60,16 @@ async def session() -> "AsyncGenerator[AsyncSession, None]":
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_pay_for_request_sets_reserved_at(session: AsyncSession) -> None: async def test_pay_for_request_sets_reserved_at(
session: AsyncSession, monkeypatch: pytest.MonkeyPatch
) -> None:
key = ApiKey(hashed_key="paykey", balance=10_000) key = ApiKey(hashed_key="paykey", balance=10_000)
session.add(key) session.add(key)
await session.commit() await session.commit()
logger_info = MagicMock()
payments_info = MagicMock()
monkeypatch.setattr(auth_module.logger, "info", logger_info)
monkeypatch.setattr(auth_module.payments_logger, "info", payments_info)
before = int(time.time()) before = int(time.time())
await pay_for_request(key, 1_000, session) await pay_for_request(key, 1_000, session)
@@ -67,6 +77,85 @@ async def test_pay_for_request_sets_reserved_at(session: AsyncSession) -> None:
assert key.reserved_balance == 1_000 assert key.reserved_balance == 1_000
assert key.reserved_at is not None assert key.reserved_at is not None
assert key.reserved_at >= before assert key.reserved_at >= before
success_logs = [
call
for call in logger_info.call_args_list
if call.args == ("Payment processed successfully",)
]
assert len(success_logs) == 1
payments_info.assert_called_once()
assert payments_info.call_args.args == ("RESERVE",)
@pytest.mark.asyncio
async def test_pay_for_request_expires_at_has_floor_margin(
session: AsyncSession, monkeypatch: pytest.MonkeyPatch
) -> None:
"""reserved_at_now floors to the second; expires_at must add 1s so a
finalizer finishing exactly at the nominal deadline isn't fenced out."""
key = ApiKey(hashed_key="floorkey", balance=10_000)
session.add(key)
await session.commit()
fixed_time = 1_700_000_000.9 # fractional second, floors when int()'d
monkeypatch.setattr(auth_module.time, "time", lambda: fixed_time)
snapshot = await pay_for_request(key, 1_000, session)
row = await session.get(ReservationRelease, snapshot.release_id)
assert row is not None
expected = (
int(fixed_time)
+ math.ceil(
auth_module.settings.max_request_lifetime_seconds
+ auth_module.settings.request_cleanup_timeout_seconds
)
+ 1
)
assert row.expires_at == expected
@pytest.mark.asyncio
@pytest.mark.asyncio
async def test_pay_for_request_releases_reservation_when_validation_fails(
session: AsyncSession, monkeypatch: pytest.MonkeyPatch
) -> None:
key = ApiKey(hashed_key="invalid-reservation", balance=10_000)
session.add(key)
await session.commit()
async def reject_reservation(*_args: object, **_kwargs: object) -> None:
raise RuntimeError("reservation identity changed")
logger_info = MagicMock()
payments_info = MagicMock()
monkeypatch.setattr(
auth_module, "_validate_reservation_snapshot", reject_reservation
)
monkeypatch.setattr(auth_module.logger, "info", logger_info)
monkeypatch.setattr(auth_module.payments_logger, "info", payments_info)
with pytest.raises(RuntimeError, match="identity changed"):
await pay_for_request(key, 1_000, session)
assert not any(
call.args == ("Payment processed successfully",)
for call in logger_info.call_args_list
)
payments_info.assert_not_called()
await session.refresh(key)
release = (
await session.exec(
select(ReservationRelease).where(
ReservationRelease.key_hash == key.hashed_key
)
)
).one()
assert key.reserved_balance == 0
assert key.total_requests == 0
assert release.status == "released"
assert release.id not in auth_module._reservation_heartbeats
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -328,7 +417,7 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None:
request = MagicMock() request = MagicMock()
request.method = "POST" request.method = "POST"
request.headers = {"authorization": "Bearer sk-cancelkey"} request.headers = {"authorization": "Bearer sk-cancelkey"}
request.body = AsyncMock(return_value=b'{"model": "test-model"}') mock_request_stream(request, b'{"model": "test-model"}')
upstream = MagicMock() upstream = MagicMock()
upstream.provider_type = "test" upstream.provider_type = "test"
@@ -355,15 +444,74 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None:
), ),
patch.object(proxy_module, "check_token_balance", MagicMock()), patch.object(proxy_module, "check_token_balance", MagicMock()),
patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)), patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)),
patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)),
patch.object( patch.object(
proxy_module, proxy_module,
"get_reservation_snapshot", "pay_for_request",
AsyncMock(return_value=reservation_snapshot), AsyncMock(return_value=reservation_snapshot),
), ),
patch.object(proxy_module, "revert_pay_for_request", revert_mock), patch.object(proxy_module, "revert_pay_for_request", revert_mock),
patch_proxy_session(session),
): ):
with pytest.raises(asyncio.CancelledError): with pytest.raises(asyncio.CancelledError):
await proxy_module.proxy(request, "v1/chat/completions", session=session) await proxy_module.proxy(request, "v1/chat/completions")
revert_mock.assert_awaited_once_with(key, session, 1000, reservation_snapshot) revert_mock.assert_awaited_once_with(key, session, 1000, reservation_snapshot)
@pytest.mark.asyncio
async def test_absolute_expiry_releases_fresh_lease(session: AsyncSession) -> None:
now = int(time.time())
key = ApiKey(
hashed_key="expired-deadline",
balance=5000,
reserved_balance=1000,
reserved_at=now,
)
session.add(key)
session.add(
ReservationRelease(
id="expired",
key_hash=key.hashed_key,
billing_key_hash=key.hashed_key,
reserved_msats=1000,
created_at=now,
started_at=now - 100,
expires_at=now - 1,
)
)
await session.commit()
assert await release_stale_reservations(session, 300) == 1
await session.refresh(key)
assert key.reserved_balance == 0
assert key.balance == 5000
@pytest.mark.asyncio
async def test_expired_reservation_cannot_renew_or_claim_charge(
session: AsyncSession,
) -> None:
from routstr.auth import (
ReservationSnapshot,
_claim_reservation_for_charge,
renew_reservation,
)
snapshot = ReservationSnapshot(
release_id="fenced",
key_hash="fenced-key",
billing_key_hash="fenced-key",
reserved_msats=1000,
)
session.add(ApiKey(hashed_key="fenced-key", balance=5000, reserved_balance=1000))
session.add(
ReservationRelease(
id="fenced",
key_hash="fenced-key",
billing_key_hash="fenced-key",
reserved_msats=1000,
expires_at=int(time.time()) - 1,
)
)
await session.commit()
assert not await renew_reservation(snapshot, session)
assert not await _claim_reservation_for_charge(snapshot, session)
-3
View File
@@ -42,8 +42,6 @@ async def test_stream_with_id_injection() -> None:
key.hashed_key = "test_hash" key.hashed_key = "test_hash"
key.balance = 1000 key.balance = 1000
background_tasks = MagicMock()
# We need to mock adjust_payment_for_tokens since it's called at the end # We need to mock adjust_payment_for_tokens since it's called at the end
with MagicMock(): with MagicMock():
from routstr.upstream import base from routstr.upstream import base
@@ -66,7 +64,6 @@ async def test_stream_with_id_injection() -> None:
response=mock_response, response=mock_response,
key=key, key=key,
max_cost_for_model=100, max_cost_for_model=100,
background_tasks=background_tasks,
requested_model="test-model", requested_model="test-model",
reservation_snapshot=ReservationSnapshot( reservation_snapshot=ReservationSnapshot(
release_id="test-release", release_id="test-release",
+379 -33
View File
@@ -6,13 +6,13 @@ from unittest.mock import AsyncMock, MagicMock, patch
import httpx import httpx
import pytest import pytest
from fastapi import BackgroundTasks
from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlmodel import SQLModel from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
import routstr.auth as auth_module import routstr.auth as auth_module
import routstr.upstream.gemini_messages as gemini_messages
from routstr.auth import ( from routstr.auth import (
ReservationSnapshot, ReservationSnapshot,
adjust_payment_for_tokens, adjust_payment_for_tokens,
@@ -160,7 +160,7 @@ async def test_post_commit_failure_cannot_release_charged_reservation() -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_generic_background_settlement_uses_explicit_reservation() -> None: async def test_generic_stream_settlement_uses_explicit_reservation() -> None:
engine = await _engine() engine = await _engine()
provider = BaseUpstreamProvider( provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key", provider_fee=1.0 base_url="https://api.example.com", api_key="test-key", provider_fee=1.0
@@ -212,8 +212,285 @@ async def test_generic_background_settlement_uses_explicit_reservation() -> None
await engine.dispose() await engine.dispose()
def _opaque_stream_response(*chunks: bytes) -> MagicMock:
async def aiter_bytes() -> AsyncGenerator[bytes, None]:
for chunk in chunks:
yield chunk
response = MagicMock(spec=httpx.Response)
response.aiter_bytes = aiter_bytes
response.aclose = AsyncMock()
return response
class _CountingAsyncByteStream(httpx.AsyncByteStream):
def __init__(self, *chunks: bytes) -> None:
self._chunks = chunks
self.close_count = 0
async def __aiter__(self) -> AsyncGenerator[bytes, None]:
for chunk in self._chunks:
yield chunk
async def aclose(self) -> None:
self.close_count += 1
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_streaming_release_is_terminal_and_suppresses_background_charge() -> None: async def test_generic_stream_completion_settles_and_closes_once() -> None:
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key"
)
finalize = AsyncMock()
provider._finalize_generic_streaming_payment = finalize # type: ignore[method-assign]
response = _opaque_stream_response(b"first", b"second")
reservation = MagicMock(spec=ReservationSnapshot)
stream = provider._stream_generic_with_settlement(
response,
"key-hash",
500,
"audio/speech",
None,
provider.provider_fee,
reservation,
)
assert [chunk async for chunk in stream] == [b"first", b"second"]
await stream.aclose()
finalize.assert_awaited_once_with(
"key-hash",
500,
"audio/speech",
None,
provider.provider_fee,
reservation,
)
response.aclose.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_generic_stream_abort_settles_and_closes_once() -> None:
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key"
)
finalize = AsyncMock()
provider._finalize_generic_streaming_payment = finalize # type: ignore[method-assign]
response = _opaque_stream_response(b"first", b"second")
reservation = MagicMock(spec=ReservationSnapshot)
stream = provider._stream_generic_with_settlement(
response,
"key-hash",
500,
"audio/speech",
None,
provider.provider_fee,
reservation,
)
assert await anext(stream) == b"first"
await stream.aclose()
await stream.aclose()
finalize.assert_awaited_once_with(
"key-hash",
500,
"audio/speech",
None,
provider.provider_fee,
reservation,
)
response.aclose.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_streaming_response_closes_iterator_when_downstream_send_is_cancelled() -> (
None
):
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key"
)
finalize = AsyncMock()
provider._finalize_generic_streaming_payment = finalize # type: ignore[method-assign]
upstream_response = _opaque_stream_response(b"first", b"second")
reservation = MagicMock(spec=ReservationSnapshot)
upstream_response.status_code = 201
upstream_response.headers = {"x-upstream": "preserved"}
response = await provider._generic_streaming_response(
upstream_response,
"key-hash",
500,
"audio/speech",
None,
provider.provider_fee,
reservation,
)
sent: list[dict[str, object]] = []
async def receive() -> dict[str, str]:
return {"type": "http.disconnect"}
async def send(message: dict[str, object]) -> None:
sent.append(message)
if message["type"] == "http.response.body" and message.get("body"):
raise asyncio.CancelledError
scope = {
"type": "http",
"asgi": {"version": "3.0", "spec_version": "2.4"},
"method": "GET",
"path": "/v1/audio/speech",
"raw_path": b"/v1/audio/speech",
"query_string": b"",
"headers": [],
"client": ("127.0.0.1", 1),
"server": ("testserver", 80),
"scheme": "http",
}
with pytest.raises(asyncio.CancelledError):
await response(scope, receive, send) # type: ignore[arg-type]
assert sent[0]["type"] == "http.response.start"
assert sent[0]["status"] == 201
headers = cast(list[tuple[bytes, bytes]], sent[0]["headers"])
assert (b"x-upstream", b"preserved") in headers
finalize.assert_awaited_once_with(
"key-hash",
500,
"audio/speech",
None,
provider.provider_fee,
reservation,
)
upstream_response.aclose.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_generic_stream_settles_when_response_start_fails() -> None:
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key"
)
finalize = AsyncMock()
provider._finalize_generic_streaming_payment = finalize # type: ignore[method-assign]
upstream_response = _opaque_stream_response(b"never-read")
upstream_response.status_code = 201
upstream_response.headers = {"x-upstream": "preserved"}
reservation = MagicMock(spec=ReservationSnapshot)
response = await provider._generic_streaming_response(
upstream_response,
"key-hash",
500,
"audio/speech",
None,
provider.provider_fee,
reservation,
)
async def receive() -> dict[str, str]:
return {"type": "http.disconnect"}
async def send(message: dict[str, object]) -> None:
assert message["type"] == "http.response.start"
raise RuntimeError("response start failed")
scope = {
"type": "http",
"asgi": {"version": "3.0", "spec_version": "2.4"},
"method": "GET",
"path": "/v1/audio/speech",
"raw_path": b"/v1/audio/speech",
"query_string": b"",
"headers": [],
"client": ("127.0.0.1", 1),
"server": ("testserver", 80),
"scheme": "http",
}
with pytest.raises(RuntimeError, match="response start failed"):
await response(scope, receive, send) # type: ignore[arg-type]
finalize.assert_awaited_once_with(
"key-hash",
500,
"audio/speech",
None,
provider.provider_fee,
reservation,
)
upstream_response.aclose.assert_awaited_once_with()
@pytest.mark.asyncio
@pytest.mark.parametrize("api", ["chat", "responses", "messages"])
async def test_parsed_stream_finalizes_when_response_start_fails(api: str) -> None:
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key"
)
upstream_response = _opaque_stream_response(b"never-read")
upstream_response.status_code = 200
upstream_response.headers = {"content-type": "text/event-stream"}
key = MagicMock(spec=ApiKey)
key.hashed_key = f"{api}-start-failure"
key.balance = 10_000
snapshot = ReservationSnapshot(
release_id=f"{api}-start-failure-release",
key_hash=key.hashed_key,
billing_key_hash=key.hashed_key,
reserved_msats=500,
)
session = MagicMock()
session.get = AsyncMock(return_value=key)
session_context = MagicMock()
session_context.__aenter__ = AsyncMock(return_value=session)
session_context.__aexit__ = AsyncMock(return_value=None)
adjust = AsyncMock(return_value={"input_tokens": 0, "output_tokens": 0})
with (
patch("routstr.upstream.base.adjust_payment_for_tokens", adjust),
patch("routstr.upstream.base.create_session", return_value=session_context),
):
if api == "chat":
response = await provider.handle_streaming_chat_completion(
upstream_response, key, 500, reservation_snapshot=snapshot
)
elif api == "responses":
response = await provider.handle_streaming_responses_completion(
upstream_response, key, 500, reservation_snapshot=snapshot
)
else:
response = await provider.handle_streaming_messages_completion(
upstream_response, key, 500, reservation_snapshot=snapshot
)
async def receive() -> dict[str, str]:
return {"type": "http.disconnect"}
async def send(message: dict[str, object]) -> None:
assert message["type"] == "http.response.start"
raise RuntimeError("response start failed")
scope = {
"type": "http",
"asgi": {"version": "3.0", "spec_version": "2.4"},
"method": "GET",
"path": f"/v1/{api}",
"raw_path": f"/v1/{api}".encode(),
"query_string": b"",
"headers": [],
"client": ("127.0.0.1", 1),
"server": ("testserver", 80),
"scheme": "http",
}
with pytest.raises(RuntimeError, match="response start failed"):
await response(scope, receive, send) # type: ignore[arg-type]
adjust.assert_awaited_once()
upstream_response.aclose.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_streaming_release_is_terminal_before_error_propagates() -> None:
provider = BaseUpstreamProvider( provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key" base_url="https://api.example.com", api_key="test-key"
) )
@@ -237,7 +514,6 @@ async def test_streaming_release_is_terminal_and_suppresses_background_charge()
release = AsyncMock(return_value=True) release = AsyncMock(return_value=True)
reservation_snapshot = MagicMock() reservation_snapshot = MagicMock()
reservation_snapshot.reserved_msats = 500 reservation_snapshot.reserved_msats = 500
background_tasks = MagicMock()
with ( with (
patch( patch(
@@ -255,7 +531,6 @@ async def test_streaming_release_is_terminal_and_suppresses_background_charge()
response=upstream_response, response=upstream_response,
key=key, key=key,
max_cost_for_model=500, max_cost_for_model=500,
background_tasks=background_tasks,
) )
with pytest.raises(SQLAlchemyError, match="database unavailable"): with pytest.raises(SQLAlchemyError, match="database unavailable"):
@@ -264,7 +539,6 @@ async def test_streaming_release_is_terminal_and_suppresses_background_charge()
session.rollback.assert_awaited_once() session.rollback.assert_awaited_once()
release.assert_awaited_once_with(reservation_snapshot, session, 500) release.assert_awaited_once_with(reservation_snapshot, session, 500)
background_tasks.add_task.assert_not_called()
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -357,8 +631,6 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once(
) )
upstream_response.aiter_bytes = aiter_bytes upstream_response.aiter_bytes = aiter_bytes
upstream_response.aclose = AsyncMock() upstream_response.aclose = AsyncMock()
client = MagicMock()
client.aclose = AsyncMock()
key = MagicMock(spec=ApiKey) key = MagicMock(spec=ApiKey)
key.hashed_key = f"{api}-partial" key.hashed_key = f"{api}-partial"
key.balance = 10_000 key.balance = 10_000
@@ -391,9 +663,7 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once(
response=upstream_response, response=upstream_response,
key=key, key=key,
max_cost_for_model=500, max_cost_for_model=500,
background_tasks=BackgroundTasks(),
reservation_snapshot=snapshot, reservation_snapshot=snapshot,
client=client,
) )
else: else:
response = await provider.handle_streaming_responses_completion( response = await provider.handle_streaming_responses_completion(
@@ -401,14 +671,10 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once(
key=key, key=key,
max_cost_for_model=500, max_cost_for_model=500,
reservation_snapshot=snapshot, reservation_snapshot=snapshot,
client=client,
) )
emitted = bytearray() emitted = bytearray()
with pytest.raises(httpx.RemoteProtocolError): async for chunk in response.body_iterator:
async for chunk in response.body_iterator: emitted.extend(chunk.encode() if isinstance(chunk, str) else bytes(chunk))
emitted.extend(
chunk.encode() if isinstance(chunk, str) else bytes(chunk)
)
adjust.assert_awaited_once() adjust.assert_awaited_once()
if finalization_fails: if finalization_fails:
@@ -417,13 +683,12 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once(
else: else:
release.assert_not_awaited() release.assert_not_awaited()
upstream_response.aclose.assert_awaited_once() upstream_response.aclose.assert_awaited_once()
client.aclose.assert_awaited_once()
assert b"[DONE]" not in emitted assert b"[DONE]" not in emitted
@pytest.mark.asyncio @pytest.mark.asyncio
@pytest.mark.parametrize("api", ["chat", "responses"]) @pytest.mark.parametrize("api", ["chat", "responses"])
async def test_partial_stream_preserves_transport_error_when_billing_db_is_down( async def test_partial_stream_closes_when_billing_db_is_down(
api: str, api: str,
) -> None: ) -> None:
provider = BaseUpstreamProvider( provider = BaseUpstreamProvider(
@@ -439,8 +704,6 @@ async def test_partial_stream_preserves_transport_error_when_billing_db_is_down(
) )
upstream_response.aiter_bytes = aiter_bytes upstream_response.aiter_bytes = aiter_bytes
upstream_response.aclose = AsyncMock() upstream_response.aclose = AsyncMock()
client = MagicMock()
client.aclose = AsyncMock()
key = MagicMock(spec=ApiKey) key = MagicMock(spec=ApiKey)
key.hashed_key = f"{api}-database-down" key.hashed_key = f"{api}-database-down"
key.balance = 10_000 key.balance = 10_000
@@ -464,9 +727,7 @@ async def test_partial_stream_preserves_transport_error_when_billing_db_is_down(
response=upstream_response, response=upstream_response,
key=key, key=key,
max_cost_for_model=500, max_cost_for_model=500,
background_tasks=BackgroundTasks(),
reservation_snapshot=snapshot, reservation_snapshot=snapshot,
client=client,
) )
else: else:
response = await provider.handle_streaming_responses_completion( response = await provider.handle_streaming_responses_completion(
@@ -474,14 +735,11 @@ async def test_partial_stream_preserves_transport_error_when_billing_db_is_down(
key=key, key=key,
max_cost_for_model=500, max_cost_for_model=500,
reservation_snapshot=snapshot, reservation_snapshot=snapshot,
client=client,
) )
with pytest.raises(httpx.RemoteProtocolError, match="incomplete chunked read"): async for _ in response.body_iterator:
async for _ in response.body_iterator: pass
pass
upstream_response.aclose.assert_awaited_once() upstream_response.aclose.assert_awaited_once()
client.aclose.assert_awaited_once()
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -637,6 +895,100 @@ async def test_messages_streaming_releases_and_raises_on_billing_failure(
release.assert_awaited_once_with(snapshot, session, 500) release.assert_awaited_once_with(snapshot, session, 500)
@pytest.mark.asyncio
async def test_gemini_messages_finalizes_when_response_start_fails() -> None:
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key"
)
key = MagicMock(spec=ApiKey)
key.hashed_key = "gemini-start-failure"
key.balance = 10_000
snapshot = ReservationSnapshot(
release_id="gemini-start-failure-release",
key_hash=key.hashed_key,
billing_key_hash=key.hashed_key,
reserved_msats=500,
)
model = MagicMock(spec=Model)
model.id = "gemini-test"
model.forwarded_model_id = None
upstream_stream = _CountingAsyncByteStream(
b'data: {"choices":[{"delta":{"content":"unused"}}]}\n\n'
)
upstream_response = httpx.Response(
200,
request=httpx.Request("POST", "https://gemini.example/chat/completions"),
stream=upstream_stream,
)
session = MagicMock()
session.get = AsyncMock(return_value=key)
session_context = MagicMock()
session_context.__aenter__ = AsyncMock(return_value=session)
session_context.__aexit__ = AsyncMock(return_value=None)
adjust = AsyncMock(return_value={"input_tokens": 0, "output_tokens": 0})
post_and_stream = AsyncMock(return_value=upstream_response)
with (
patch("routstr.upstream.base.adjust_payment_for_tokens", adjust),
patch("routstr.upstream.base.create_session", return_value=session_context),
patch.object(
gemini_messages,
"_translate_anthropic_to_openai",
return_value={"messages": []},
),
patch.object(gemini_messages, "_post_and_stream", post_and_stream),
):
(
client_stream,
iterator,
requested_model,
) = await gemini_messages.dispatch_gemini_messages(
request_body=json.dumps(
{"model": model.id, "messages": [], "stream": True}
).encode(),
model_obj=model,
base_url="https://gemini.example",
api_key="test-key",
transform_model_name=lambda name: name,
)
assert client_stream is True
assert requested_model == model.id
response = provider._stream_litellm_messages(
iterator=iterator,
key=key,
max_cost_for_model=500,
requested_model=requested_model,
reservation_snapshot=snapshot,
)
async def receive() -> dict[str, str]:
return {"type": "http.disconnect"}
async def send(message: dict[str, object]) -> None:
assert message["type"] == "http.response.start"
raise RuntimeError("response start failed")
scope = {
"type": "http",
"asgi": {"version": "3.0", "spec_version": "2.4"},
"method": "GET",
"path": "/v1/messages",
"raw_path": b"/v1/messages",
"query_string": b"",
"headers": [],
"client": ("127.0.0.1", 1),
"server": ("testserver", 80),
"scheme": "http",
}
with pytest.raises(RuntimeError, match="response start failed"):
await response(scope, receive, send) # type: ignore[arg-type]
post_and_stream.assert_awaited_once()
adjust.assert_awaited_once()
assert upstream_response.is_closed
assert upstream_stream.close_count == 1
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() -> None: async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() -> None:
engine = await _engine() engine = await _engine()
@@ -721,7 +1073,6 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat()
{"model": model.id, "messages": [{"role": "user", "content": "hi"}]} {"model": model.id, "messages": [{"role": "user", "content": "hi"}]}
).encode() ).encode()
background_tasks = BackgroundTasks()
try: try:
with ( with (
patch( patch(
@@ -746,7 +1097,6 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat()
response=upstream_response, response=upstream_response,
key=key, key=key,
max_cost_for_model=500, max_cost_for_model=500,
background_tasks=background_tasks,
model_obj=model, model_obj=model,
reservation_snapshot=snapshot, reservation_snapshot=snapshot,
request_body=request_body, request_body=request_body,
@@ -754,10 +1104,6 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat()
iterator = cast(AsyncGenerator[bytes, None], response.body_iterator) iterator = cast(AsyncGenerator[bytes, None], response.body_iterator)
await iterator.__anext__() # first chunk reaches the client await iterator.__anext__() # first chunk reaches the client
await iterator.aclose() # client aborts the socket here await iterator.aclose() # client aborts the socket here
# Starlette runs the response's background tasks after the abort.
for task in background_tasks.tasks:
await task()
finally: finally:
await auth_module._stop_reservation_heartbeat(snapshot.release_id) await auth_module._stop_reservation_heartbeat(snapshot.release_id)
+3 -2
View File
@@ -42,7 +42,9 @@ def _make_response(chunks: list[bytes]) -> MagicMock:
return mock_response return mock_response
async def _drive(chunks: list[bytes], requested_model: str | None = None) -> list[bytes]: async def _drive(
chunks: list[bytes], requested_model: str | None = None
) -> list[bytes]:
"""Run the real streaming generator over ``chunks`` and collect output bytes.""" """Run the real streaming generator over ``chunks`` and collect output bytes."""
provider = BaseUpstreamProvider( provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test_key" base_url="https://api.example.com", api_key="test_key"
@@ -66,7 +68,6 @@ async def _drive(chunks: list[bytes], requested_model: str | None = None) -> lis
response=_make_response(chunks), response=_make_response(chunks),
key=key, key=key,
max_cost_for_model=100, max_cost_for_model=100,
background_tasks=MagicMock(),
requested_model=requested_model, requested_model=requested_model,
reservation_snapshot=ReservationSnapshot( reservation_snapshot=ReservationSnapshot(
release_id="test-release", release_id="test-release",
+5 -5
View File
@@ -31,6 +31,8 @@ from routstr.upstream.tinfoil import (
) )
from routstr.upstream.tinfoil_trailer import TrailerResponse from routstr.upstream.tinfoil_trailer import TrailerResponse
from .proxy_test_utils import patch_proxy_session
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# parse_tinfoil_usage_metrics # parse_tinfoil_usage_metrics
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -1293,10 +1295,9 @@ async def test_bearer_key_config_422_releases_reservation_and_passes_through() -
), ),
patch.object(proxy_module, "check_token_balance", MagicMock()), patch.object(proxy_module, "check_token_balance", MagicMock()),
patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)), patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)),
patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)),
patch.object( patch.object(
proxy_module, proxy_module,
"get_reservation_snapshot", "pay_for_request",
AsyncMock(return_value=reservation_snapshot), AsyncMock(return_value=reservation_snapshot),
), ),
patch.object(proxy_module, "revert_pay_for_request", revert_mock), patch.object(proxy_module, "revert_pay_for_request", revert_mock),
@@ -1304,10 +1305,9 @@ async def test_bearer_key_config_422_releases_reservation_and_passes_through() -
"routstr.upstream.ehbp.forward_with_trailer", "routstr.upstream.ehbp.forward_with_trailer",
AsyncMock(return_value=upstream_resp), AsyncMock(return_value=upstream_resp),
), ),
patch_proxy_session(session),
): ):
response = await proxy_module.proxy( response = await proxy_module.proxy(request, "v1/chat/completions")
request, "v1/chat/completions", session=session
)
# The reservation was released despite the early passthrough return. # The reservation was released despite the early passthrough return.
revert_mock.assert_awaited_once_with(key, session, 1_000, reservation_snapshot) revert_mock.assert_awaited_once_with(key, session, 1_000, reservation_snapshot)
+83 -3
View File
@@ -1,11 +1,13 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import ssl
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from routstr.core.exceptions import EhbpTimeoutError, UpstreamError from routstr.core.error_scope import ERROR_SCOPE_UPSTREAM, UPSTREAM_ERROR_STATUS
from routstr.core.exceptions import EhbpConnectionError, EhbpTimeoutError, UpstreamError
from routstr.upstream.tinfoil_trailer import forward_with_trailer from routstr.upstream.tinfoil_trailer import forward_with_trailer
@@ -168,6 +170,68 @@ async def test_forward_with_trailer_connect_timeout_raises_ehbp_timeout(
) )
@pytest.mark.asyncio
async def test_forward_with_trailer_tls_handshake_timeout_raises_ehbp_timeout(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""The stdlib TLS handshake timer surfaces as ConnectionAbortedError.
CPython aborts a slow handshake with ``ConnectionAbortedError`` rather
than ``asyncio.TimeoutError``, so the connect handler must classify it as
an upstream timeout — otherwise it escapes to the node-scoped 500 in
``forward_ehbp_request``.
"""
async def _handshake_timeout(*_args: object, **_kwargs: object) -> object:
raise ConnectionAbortedError(
"SSL handshake is taking longer than 60.0 seconds: aborting the connection"
)
monkeypatch.setattr(
"routstr.upstream.tinfoil_trailer.asyncio.open_connection",
_handshake_timeout,
)
with pytest.raises(EhbpTimeoutError, match="TLS handshake timed out"):
await forward_with_trailer(
method="POST",
url="https://inference.tinfoil.sh/v1/chat/completions",
headers={},
body=b"opaque",
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"exc",
[
ConnectionRefusedError("connection refused"),
ConnectionResetError("connection reset"),
ssl.SSLError("certificate verify failed"),
OSError("name resolution failed"),
],
)
async def test_forward_with_trailer_connection_failure_raises_ehbp_connection(
monkeypatch: pytest.MonkeyPatch, exc: Exception
) -> None:
"""Non-timeout connect failures must be upstream-scoped, not node 500s."""
async def _fail_connect(*_args: object, **_kwargs: object) -> object:
raise exc
monkeypatch.setattr(
"routstr.upstream.tinfoil_trailer.asyncio.open_connection", _fail_connect
)
with pytest.raises(EhbpConnectionError, match="Unable to connect"):
await forward_with_trailer(
method="POST",
url="https://inference.tinfoil.sh/v1/chat/completions",
headers={},
body=b"opaque",
)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_forward_with_trailer_read_timeout_raises_ehbp_timeout( async def test_forward_with_trailer_read_timeout_raises_ehbp_timeout(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
@@ -193,9 +257,10 @@ async def test_forward_with_trailer_read_timeout_raises_ehbp_timeout(
def test_ehbp_timeout_error_metadata() -> None: def test_ehbp_timeout_error_metadata() -> None:
exc = EhbpTimeoutError("boom") exc = EhbpTimeoutError("boom")
assert exc.status_code == 504 assert exc.status_code == UPSTREAM_ERROR_STATUS
assert exc.code == "UPSTREAM_TIMEOUT" assert exc.code == "UPSTREAM_TIMEOUT"
assert exc.details is None assert exc.details is None
assert exc.scope == ERROR_SCOPE_UPSTREAM
assert isinstance(exc, UpstreamError) assert isinstance(exc, UpstreamError)
@@ -203,5 +268,20 @@ def test_ehbp_timeout_error_forwards_details() -> None:
"""``details`` must survive so the response builder can forward it.""" """``details`` must survive so the response builder can forward it."""
exc = EhbpTimeoutError("boom", details={"phase": "connect"}) exc = EhbpTimeoutError("boom", details={"phase": "connect"})
assert exc.details == {"phase": "connect"} assert exc.details == {"phase": "connect"}
assert exc.status_code == 504 assert exc.status_code == UPSTREAM_ERROR_STATUS
assert exc.code == "UPSTREAM_TIMEOUT" assert exc.code == "UPSTREAM_TIMEOUT"
def test_ehbp_connection_error_metadata() -> None:
exc = EhbpConnectionError("boom")
assert exc.status_code == UPSTREAM_ERROR_STATUS
assert exc.code == "UPSTREAM_UNAVAILABLE"
assert exc.details is None
assert exc.scope == ERROR_SCOPE_UPSTREAM
assert isinstance(exc, UpstreamError)
def test_ehbp_connection_error_forwards_details() -> None:
exc = EhbpConnectionError("boom", details={"provider": "tinfoil"})
assert exc.details == {"provider": "tinfoil"}
assert exc.code == "UPSTREAM_UNAVAILABLE"
+297
View File
@@ -0,0 +1,297 @@
"""Unit tests for ``DeepSeekUpstreamProvider``.
DeepSeek is priced from the provider's own peak-rate table, never from litellm
or OpenRouter: litellm's ``deepseek-v4-flash`` entry is stale and OpenRouter
resells below DeepSeek's peak rate, so either would bill under cost. These
tests pin the table prices (including the cache-hit rate), that a model the
table misses imports disabled without consulting the fallback chain, and that
``reasoning_content`` in history reaches DeepSeek untouched — thinking mode
with ``tools`` answers 400 when it is stripped.
"""
from __future__ import annotations
import json
import threading
from collections.abc import Iterator
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import Any
from unittest.mock import AsyncMock, Mock, patch
import litellm
import pytest
from routstr.upstream import upstream_provider_classes
from routstr.upstream.deepseek import DeepSeekUpstreamProvider
class _FakeResponse:
def __init__(self, payload: dict[str, Any]) -> None:
self._payload = payload
def raise_for_status(self) -> None:
return None
def json(self) -> dict[str, Any]:
return self._payload
class _FakeAsyncClient:
def __init__(self, payload: dict[str, Any], calls: list[dict[str, Any]]) -> None:
self._payload = payload
self._calls = calls
async def __aenter__(self) -> "_FakeAsyncClient":
return self
async def __aexit__(self, *exc: object) -> bool:
return False
async def get(
self, url: str, headers: dict[str, str] | None = None
) -> _FakeResponse:
self._calls.append({"url": url, "headers": headers})
return _FakeResponse(self._payload)
# Shape of DeepSeek's ``GET /models``: bare ids, no pricing.
CATALOG: dict[str, Any] = {
"object": "list",
"data": [
{"id": "deepseek-flash", "object": "model", "owned_by": "deepseek"},
{"id": "deepseek-v4-pro", "object": "model", "owned_by": "deepseek"},
{"id": "deepseek-v4-flash", "object": "model", "owned_by": "deepseek"},
{"id": "deepseek-chat", "object": "model", "owned_by": "deepseek"},
],
}
async def _fetch(
catalog: dict[str, Any] = CATALOG,
) -> tuple[dict[str, Any], list[dict[str, Any]], AsyncMock]:
calls: list[dict[str, Any]] = []
fallback = AsyncMock(return_value=None)
provider = DeepSeekUpstreamProvider(api_key="sk-test")
with (
patch(
"routstr.upstream.generic.httpx.AsyncClient",
lambda *args, **kwargs: _FakeAsyncClient(catalog, calls),
),
patch("routstr.upstream.generic.FallbackPricingResolver.resolve", fallback),
):
models = await provider.fetch_models()
return {m.id: m for m in models}, calls, fallback
def test_metadata_and_registration() -> None:
assert DeepSeekUpstreamProvider in upstream_provider_classes
assert DeepSeekUpstreamProvider.get_provider_metadata() == {
"id": "deepseek",
"name": "DeepSeek",
"default_base_url": "https://api.deepseek.com",
"fixed_base_url": True,
"platform_url": "https://platform.deepseek.com/api_keys",
}
def test_build_from_row_ignores_row_base_url() -> None:
row = Mock(
api_key="sk-row", provider_fee=1.05, base_url="https://elsewhere.example"
)
provider = DeepSeekUpstreamProvider._build_from_row(row)
assert provider.api_key == "sk-row"
assert provider.provider_fee == 1.05
assert provider.base_url == "https://api.deepseek.com"
def test_litellm_prefix_is_deepseek() -> None:
provider = DeepSeekUpstreamProvider(api_key="sk-test")
assert provider.get_litellm_provider_prefix() == "deepseek/"
@pytest.mark.parametrize(
"model_id,expected",
[
("deepseek/deepseek-v4-flash", "deepseek-v4-flash"),
("deepseek-v4-flash", "deepseek-v4-flash"),
("deepseek/deepseek-flash", "deepseek-flash"),
],
)
def test_transform_model_name(model_id: str, expected: str) -> None:
provider = DeepSeekUpstreamProvider(api_key="sk-test")
assert provider.transform_model_name(model_id) == expected
def test_provider_field_names_deepseek_not_host() -> None:
provider = DeepSeekUpstreamProvider(api_key="sk-test")
payload: dict[str, Any] = {"id": "chatcmpl-1"}
provider._apply_provider_field(payload)
assert payload["provider"] == "deepseek"
@pytest.mark.asyncio
async def test_fetch_models_calls_deepseek_models_endpoint_with_key() -> None:
_, calls, _ = await _fetch()
assert calls == [
{
"url": "https://api.deepseek.com/models",
"headers": {"Authorization": "Bearer sk-test"},
}
]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"model_id,prompt,completion,cache_read",
[
("deepseek-flash", 0.30, 1.20, 0.006),
# Retired alias DeepSeek serves and bills as deepseek-flash.
("deepseek-v4-flash", 0.30, 1.20, 0.006),
("deepseek-v4-pro", 1.32, 3.96, 0.044),
],
)
async def test_table_models_priced_at_peak_rate(
model_id: str, prompt: float, completion: float, cache_read: float
) -> None:
models, _, _ = await _fetch()
model = models[model_id]
assert model.enabled is True
assert model.pricing.prompt == pytest.approx(prompt / 1_000_000)
assert model.pricing.completion == pytest.approx(completion / 1_000_000)
assert model.pricing.input_cache_read == pytest.approx(cache_read / 1_000_000)
assert model.context_length == 1_000_000
@pytest.mark.asyncio
async def test_vision_follows_the_model() -> None:
models, _, _ = await _fetch()
assert "image" in models["deepseek-flash"].architecture.input_modalities
assert models["deepseek-v4-pro"].architecture.input_modalities == ["text"]
@pytest.mark.asyncio
async def test_unlisted_model_imports_disabled_without_fallback() -> None:
"""litellm prices ``deepseek-chat``; the provider must not take that price."""
models, _, fallback = await _fetch()
model = models["deepseek-chat"]
assert model.enabled is False
assert model.pricing.prompt == 0.0
assert model.pricing.completion == 0.0
fallback.assert_not_awaited()
@pytest.mark.asyncio
async def test_cache_rate_survives_fee_and_is_not_replaced_by_litellm() -> None:
"""litellm's stale ``deepseek-v4-flash`` cache rate (1.4e-08 in the bundled
map) must not replace the table's; backfill only fills an absent rate. The
fee applies to the cache rate like every other component.
The litellm entry is pinned here because the remote cost map already
carries the table's rate, which would let an overwrite go unnoticed."""
models, _, _ = await _fetch()
provider = DeepSeekUpstreamProvider(api_key="sk-test", provider_fee=1.05)
stale = {"cache_read_input_token_cost": 1.4e-08}
with patch("routstr.payment.models.litellm_cost_entry", return_value=stale):
priced = provider._apply_provider_fee_to_model(models["deepseek-v4-flash"])
assert priced.pricing.input_cache_read == pytest.approx(0.006e-6 * 1.05)
assert priced.pricing.prompt == pytest.approx(0.30e-6 * 1.05)
# A cache hit costs 2% of a miss, not the full input rate.
assert priced.pricing.input_cache_read / priced.pricing.prompt == pytest.approx(
0.02
)
@pytest.mark.asyncio
async def test_reasoning_content_in_history_reaches_upstream() -> None:
models, _, _ = await _fetch()
provider = DeepSeekUpstreamProvider(api_key="sk-test")
messages = [
{"role": "user", "content": "weather in Paris?"},
{
"role": "assistant",
"content": "",
"reasoning_content": "Need the weather tool.",
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "get_weather", "arguments": "{}"},
}
],
},
{"role": "tool", "tool_call_id": "call_1", "content": "18C"},
]
body = json.dumps(
{
"model": "deepseek/deepseek-flash",
"messages": messages,
"tools": [{"type": "function", "function": {"name": "get_weather"}}],
}
).encode()
out = provider.prepare_request_body(body, models["deepseek-flash"])
assert out is not None
sent = json.loads(out)
assert sent["model"] == "deepseek-flash"
assert sent["messages"] == messages
_ANTHROPIC_SSE = (
b"event: message_start\n"
b'data: {"type":"message_start","message":{"id":"msg_1","type":"message",'
b'"role":"assistant","model":"deepseek-flash","content":[],'
b'"stop_reason":null,"usage":{"input_tokens":3,"output_tokens":0}}}\n\n'
b"event: message_stop\n"
b'data: {"type":"message_stop"}\n\n'
)
@pytest.fixture
def anthropic_stub() -> Iterator[tuple[str, list[tuple[str, dict[str, Any]]]]]:
"""Loopback stand-in for DeepSeek's Anthropic-format endpoint."""
seen: list[tuple[str, dict[str, Any]]] = []
class Handler(BaseHTTPRequestHandler):
def do_POST(self) -> None:
length = int(self.headers["Content-Length"])
seen.append((self.path, json.loads(self.rfile.read(length))))
self.send_response(200)
self.send_header("Content-Type", "text/event-stream")
self.send_header("Content-Length", str(len(_ANTHROPIC_SSE)))
self.end_headers()
self.wfile.write(_ANTHROPIC_SSE)
def log_message(self, *args: Any) -> None:
return None
server = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield f"http://127.0.0.1:{server.server_address[1]}", seen
finally:
server.shutdown()
server.server_close()
@pytest.mark.asyncio
async def test_messages_stream_reaches_deepseek_anthropic_endpoint(
anthropic_stub: tuple[str, list[tuple[str, dict[str, Any]]]],
) -> None:
# litellm sends deepseek/ Messages calls to DeepSeek's /anthropic endpoint;
# its stream iterator imports litellm.proxy, which needs ``backoff``.
api_base, seen = anthropic_stub
stream = await litellm.anthropic.messages.acreate(
model=DeepSeekUpstreamProvider.litellm_provider_prefix + "deepseek-flash",
messages=[{"role": "user", "content": "hi"}],
max_tokens=8,
stream=True,
api_key="sk-test",
api_base=api_base,
)
chunks = [chunk async for chunk in stream] # type: ignore[union-attr]
assert b"message_stop" in b"".join(chunks)
assert len(seen) == 1
assert seen[0][0] == "/anthropic/v1/messages"
assert seen[0][1]["model"] == "deepseek-flash"
+228 -2
View File
@@ -14,7 +14,19 @@ from unittest.mock import Mock
import httpx import httpx
import pytest import pytest
from routstr.core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_NODE,
ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS,
UPSTREAM_UNAVAILABLE,
client_code_for_upstream_error,
client_status_for_upstream_error,
)
from routstr.core.exceptions import UpstreamError
from routstr.payment.helpers import create_upstream_error_response
from routstr.upstream.base import BaseUpstreamProvider, _is_json_content_type from routstr.upstream.base import BaseUpstreamProvider, _is_json_content_type
from routstr.upstream.rate_limit import UPSTREAM_RATE_LIMIT
def _make_request(request_id: str = "req-123") -> Mock: def _make_request(request_id: str = "req-123") -> Mock:
@@ -105,10 +117,13 @@ async def test_plain_text_error_is_normalized(
_make_request(), "v1/messages", upstream _make_request(), "v1/messages", upstream
) )
assert response.status_code == 503 assert response.status_code == UPSTREAM_ERROR_STATUS
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
assert response.media_type == "application/json" assert response.media_type == "application/json"
payload = json.loads(bytes(response.body)) payload = json.loads(bytes(response.body))
assert payload["error"]["message"] == "Service Unavailable" assert payload["error"]["message"] == "Service Unavailable"
assert payload["error"]["code"] == UPSTREAM_UNAVAILABLE
assert payload["error"]["upstream_status"] == 503
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -123,10 +138,13 @@ async def test_empty_body_with_non_json_content_type_normalizes(
_make_request(), "v1/messages", upstream _make_request(), "v1/messages", upstream
) )
assert response.status_code == 502 assert response.status_code == UPSTREAM_ERROR_STATUS
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
assert response.media_type == "application/json" assert response.media_type == "application/json"
payload = json.loads(bytes(response.body)) payload = json.loads(bytes(response.body))
assert payload["error"]["type"] == "upstream_error" assert payload["error"]["type"] == "upstream_error"
assert payload["error"]["code"] == UPSTREAM_UNAVAILABLE
assert payload["error"]["upstream_status"] == 502
assert payload["error"]["upstream_body_preview"] is None assert payload["error"]["upstream_body_preview"] is None
@@ -148,3 +166,211 @@ async def test_json_error_body_is_passed_through_unchanged(
assert response.status_code == 400 assert response.status_code == 400
assert bytes(response.body) == json_body assert bytes(response.body) == json_body
assert response.media_type == "application/json" assert response.media_type == "application/json"
# --------------------------------------------------------------------------- #
# Upstream 5xx -> 424 + UPSTREAM_UNAVAILABLE + scope header; node faults stay
# 500 without it; rate limits keep 429.
# --------------------------------------------------------------------------- #
@pytest.mark.asyncio
@pytest.mark.parametrize("path", ["v1/chat/completions", "v1/messages", "v1/responses"])
@pytest.mark.parametrize("upstream_status", [500, 502, 503, 504])
async def test_upstream_5xx_is_attributed_to_the_upstream(
provider: BaseUpstreamProvider, path: str, upstream_status: int
) -> None:
body = json.dumps(
{"error": {"message": "provider exploded", "type": "server_error"}}
).encode()
upstream = _make_upstream_response(
body=body, status_code=upstream_status, content_type="application/json"
)
response = await provider.forward_upstream_error_response(
_make_request(), path, upstream
)
assert response.status_code == UPSTREAM_ERROR_STATUS
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
payload: dict[str, Any] = json.loads(bytes(response.body))
assert payload["error"]["code"] == UPSTREAM_UNAVAILABLE
assert payload["error"]["upstream_status"] == upstream_status
@pytest.mark.asyncio
async def test_upstream_5xx_non_json_body_keeps_scope_and_status(
provider: BaseUpstreamProvider,
) -> None:
"""The envelope for a non-JSON 5xx carries the same attribution."""
upstream = _make_upstream_response(
body=b"<html>bad gateway</html>", status_code=502, content_type="text/html"
)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/chat/completions", upstream
)
assert response.status_code == UPSTREAM_ERROR_STATUS
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
payload: dict[str, Any] = json.loads(bytes(response.body))
assert payload["error"]["type"] == "upstream_error"
assert payload["error"]["code"] == UPSTREAM_UNAVAILABLE
assert payload["error"]["upstream_status"] == 502
@pytest.mark.asyncio
@pytest.mark.parametrize("upstream_status", [400, 401, 403, 404, 422])
async def test_provider_4xx_passes_through_unchanged(
provider: BaseUpstreamProvider, upstream_status: int
) -> None:
"""A provider 4xx is its verdict on the request, not a node-health signal."""
body = json.dumps(
{"error": {"message": "bad request", "type": "invalid_request_error"}}
).encode()
upstream = _make_upstream_response(
body=body, status_code=upstream_status, content_type="application/json"
)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/chat/completions", upstream
)
assert response.status_code == upstream_status
@pytest.mark.asyncio
async def test_upstream_rate_limit_keeps_429(
provider: BaseUpstreamProvider,
) -> None:
"""429 + UPSTREAM_RATE_LIMIT is unchanged by the 424 mapping: the retry
hint is worth more than the status class."""
body = json.dumps(
{"error": {"message": "Rate limit reached, please try again"}}
).encode()
upstream = _make_upstream_response(body=body, status_code=429)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/chat/completions", upstream
)
assert response.status_code == 429
payload: dict[str, Any] = json.loads(bytes(response.body))
assert payload["error"]["code"] == UPSTREAM_RATE_LIMIT
def test_generic_upstream_error_response_reports_424() -> None:
"""``create_upstream_error_response`` maps a plain upstream failure to 424."""
err = UpstreamError("connection refused", status_code=502)
response = create_upstream_error_response(err, _make_request())
assert response.status_code == UPSTREAM_ERROR_STATUS
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
payload: dict[str, Any] = json.loads(bytes(response.body))
assert payload["error"]["type"] == "upstream_error"
assert payload["error"]["code"] == UPSTREAM_UNAVAILABLE
assert payload["error"]["details"]["upstream_status"] == 502
def test_rate_limit_error_response_keeps_429_and_code() -> None:
err = UpstreamError(
"slow down", status_code=429, code=UPSTREAM_RATE_LIMIT, details={"a": 1}
)
response = create_upstream_error_response(err, _make_request())
assert response.status_code == 429
payload: dict[str, Any] = json.loads(bytes(response.body))
assert payload["error"]["code"] == UPSTREAM_RATE_LIMIT
assert payload["error"]["details"] == {"a": 1}
def test_5xx_wrapped_rate_limit_error_response_keeps_429() -> None:
"""A rate limit wrapped in a provider 5xx still answers 429."""
err = UpstreamError("slow down", status_code=500, code=UPSTREAM_RATE_LIMIT)
response = create_upstream_error_response(err, _make_request())
assert response.status_code == 429
payload: dict[str, Any] = json.loads(bytes(response.body))
assert payload["error"]["code"] == UPSTREAM_RATE_LIMIT
def test_node_scoped_failure_stays_500_without_scope_header() -> None:
"""A genuine node fault must never be disguised as an upstream one."""
err = UpstreamError("mint unreachable", status_code=500, scope=ERROR_SCOPE_NODE)
response = create_upstream_error_response(err, _make_request())
assert response.status_code == 500
assert ERROR_SCOPE_HEADER not in response.headers
payload: dict[str, Any] = json.loads(bytes(response.body))
assert payload["error"]["code"] != UPSTREAM_UNAVAILABLE
def test_upstream_error_defaults_to_upstream_scope() -> None:
assert UpstreamError("boom").scope == ERROR_SCOPE_UPSTREAM
@pytest.mark.asyncio
async def test_json_body_without_error_mapping_gets_classification(
provider: BaseUpstreamProvider,
) -> None:
"""A rewritten status is never served without a matching ``error.code``."""
body = json.dumps({"detail": "internal failure"}).encode()
upstream = _make_upstream_response(
body=body, status_code=503, content_type="application/json"
)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/chat/completions", upstream
)
assert response.status_code == UPSTREAM_ERROR_STATUS
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
payload: dict[str, Any] = json.loads(bytes(response.body))
assert payload["detail"] == "internal failure"
assert payload["error"]["code"] == UPSTREAM_UNAVAILABLE
assert payload["error"]["upstream_status"] == 503
@pytest.mark.asyncio
async def test_json_body_with_non_mapping_error_is_left_alone(
provider: BaseUpstreamProvider,
) -> None:
"""A provider's own ``error`` value is never clobbered by the mapping."""
body = json.dumps({"error": "boom"}).encode()
upstream = _make_upstream_response(
body=body, status_code=503, content_type="application/json"
)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/chat/completions", upstream
)
assert response.status_code == UPSTREAM_ERROR_STATUS
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM
assert json.loads(bytes(response.body)) == {"error": "boom"}
@pytest.mark.parametrize("upstream_status", [429, 500, 502, 503, 529])
def test_rate_limit_status_and_code_never_disagree(upstream_status: int) -> None:
"""429 and ``UPSTREAM_RATE_LIMIT`` are one classification, not two: a caller
must never see ``424`` carrying the rate-limit code."""
assert client_status_for_upstream_error(upstream_status, UPSTREAM_RATE_LIMIT) == 429
assert (
client_code_for_upstream_error(upstream_status, UPSTREAM_RATE_LIMIT)
== UPSTREAM_RATE_LIMIT
)
@pytest.mark.parametrize("upstream_status", [400, 401, 403, 404, 422])
def test_provider_4xx_keeps_its_numeric_code(upstream_status: int) -> None:
"""The x-cashu envelopes pass the status as the code; a 4xx must keep the
legacy numeric ``error.code`` rather than degrade to ``null``."""
assert client_status_for_upstream_error(upstream_status) == upstream_status
assert (
client_code_for_upstream_error(upstream_status, upstream_status)
== upstream_status
)

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