mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
Merge branch 'main' into feat/upstream-certification-harness
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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/)**.
|
||||||
|
|||||||
@@ -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
@@ -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',
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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.
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -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
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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()
|
||||||
@@ -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
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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),
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -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,
|
||||||
|
),
|
||||||
|
)
|
||||||
@@ -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()
|
||||||
@@ -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={
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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}"
|
||||||
|
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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()
|
||||||
@@ -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
|
||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
@@ -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"]
|
||||||
@@ -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)
|
||||||
@@ -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
|
||||||
@@ -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:
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -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"
|
||||||
@@ -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",
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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()
|
||||||
@@ -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()
|
||||||
@@ -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"
|
||||||
|
|||||||
@@ -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 == {
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -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(
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -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
|
||||||
@@ -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()
|
||||||
|
|||||||
@@ -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""
|
||||||
@@ -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)
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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"
|
||||||
@@ -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
Reference in New Issue
Block a user