diff --git a/.env.example b/.env.example index 265ca2c4..a063da56 100644 --- a/.env.example +++ b/.env.example @@ -45,6 +45,7 @@ ROUTSTR_SECRET_KEY= # 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" # 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" # MINT_OPERATION_CONCURRENCY=4 # MINT_OPERATION_TIMEOUT_SECONDS=30 @@ -64,6 +65,24 @@ ROUTSTR_SECRET_KEY= # Network Configuration # CORS_ORIGINS=* # 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 # LOG_LEVEL=INFO diff --git a/README.md b/README.md index 0636c299..3a14bcb7 100644 --- a/README.md +++ b/README.md @@ -51,9 +51,25 @@ curl https://api.routstr.com/v1/chat/completions \ ## 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 # 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 @@ -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())" ``` -2. **Start the services**: +3. **Start the services**: ```bash 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 `/admin` URL. Read it from the logs: ```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`.) -4. **Configure**: +5. **Configure**: Open [http://localhost:8000/admin/](http://localhost:8000/admin/) to connect your AI providers and set pricing. For full instructions, see the **[Provider Quick Start Guide](https://docs.routstr.com/provider/quickstart/)**. diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index 5ead333e..d8b610ca 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -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 the known versioned ids so they can be requested directly. - 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:** diff --git a/docs/api/errors.md b/docs/api/errors.md index 82bb5be1..5f151bb6 100644 --- a/docs/api/errors.md +++ b/docs/api/errors.md @@ -51,11 +51,50 @@ legacy status behavior. | 403 | Forbidden | Access denied to resource | | 404 | Not Found | Endpoint or resource doesn't exist | | 422 | Unprocessable Entity | Validation errors | -| 429 | Too Many Requests | Rate limit exceeded | -| 500 | Internal Server Error | Server-side error | -| 502 | Bad Gateway | Upstream API error | +| 424 | Failed Dependency | An upstream inference provider failed. This node is healthy — see [Upstream attribution](#upstream-attribution-424-failed-dependency) | +| 429 | Too Many Requests | Rate limit exceeded (this node or an upstream provider) | +| 500 | Internal Server Error | Server-side error on this node | +| 502 | Bad Gateway | Gateway-level failure | | 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 ### Authentication Errors @@ -329,14 +368,18 @@ Retry-After: 45 ### 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 { "error": { "type": "upstream_error", "message": "Model is currently overloaded", - "code": "model_overloaded", + "code": "UPSTREAM_UNAVAILABLE", + "upstream_status": 503, "details": { "model": "gpt-4", "retry_after": 5 @@ -345,8 +388,11 @@ Retry-After: 45 } ``` -**Status:** 503 -**Resolution:** Retry request after delay +**Status:** 424 +**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 @@ -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 ### Content Policy @@ -416,7 +463,7 @@ def retry_with_backoff( # Check if error is retryable 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 delay = min( base_delay * (2 ** attempt) + random.uniform(0, 1), @@ -441,6 +488,9 @@ Group errors for handling: class ErrorHandler: # Errors that should be retried 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', 'upstream_timeout', 'model_overloaded', diff --git a/docs/api/overview.md b/docs/api/overview.md index c8069953..8a4a445d 100644 --- a/docs/api/overview.md +++ b/docs/api/overview.md @@ -91,7 +91,7 @@ All errors follow a consistent format: | `not_found` | 404 | Resource not found | | `rate_limit_exceeded` | 429 | Too many requests | | `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 @@ -268,9 +268,10 @@ X-Webhook-Signature: sha256=... | 402 | Payment required | | 403 | Forbidden | | 404 | Not found | +| 424 | Upstream provider failed (`X-Routstr-Error-Scope: upstream`) | | 429 | Rate limited | -| 500 | Server error | -| 502 | Upstream error | +| 500 | Server error (no scope header) | +| 502 | Gateway failure | | 503 | Service unavailable | ## CORS Support diff --git a/docs/provider/configuration.md b/docs/provider/configuration.md index 68433cc1..20668e60 100644 --- a/docs/provider/configuration.md +++ b/docs/provider/configuration.md @@ -48,6 +48,29 @@ Connect to your AI provider(s): | **Upstream URL** | API endpoint (e.g., `https://api.openai.com/v1`) | | **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 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 | +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 Automatic profit withdrawal: @@ -197,13 +229,14 @@ Use environment variables for: | `NPUB` | Nostr public key (bech32) | — | | `NSEC` | Legacy seed for the Nostr private key (otherwise set from the admin UI) | — | | `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_TIMEOUT_SECONDS` | Per-attempt timeout for mint network calls | `30` | | `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` | | `RECEIVE_LN_ADDRESS` | Lightning address for withdrawals | — | | `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` | | `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` | | `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 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 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. diff --git a/docs/provider/dashboard.md b/docs/provider/dashboard.md index 66852845..7e3db598 100644 --- a/docs/provider/dashboard.md +++ b/docs/provider/dashboard.md @@ -132,7 +132,9 @@ Connect to your AI provider: ### 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 - **Remove Mint** — Stop accepting from a mint diff --git a/docs/provider/deployment.md b/docs/provider/deployment.md index 784a764d..5cf8039b 100644 --- a/docs/provider/deployment.md +++ b/docs/provider/deployment.md @@ -2,179 +2,97 @@ 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. - -### Quick Start +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 +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 -docker run -d \ - --name routstr \ - -p 8000:8000 \ - -v routstr-data:/app/data \ - -e DATABASE_URL="sqlite:////app/data/routstr.db" \ - 9qeklajc/routstr:latest -``` +git clone https://github.com/Routstr/routstr-core.git +cd routstr-core -Access your node: -- **API & Admin Dashboard**: http://localhost:8000 +# Check out a release (v0.4.7 is current — see the releases page for the newest tag) +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 (Recommended) - -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. +Then open your node: +- **API & Admin Dashboard**: +- **Admin login**: the password is generated and logged once on first start ```bash -docker compose up -d +docker compose logs routstr | grep -i admin ``` -This will: -1. **Build the UI**: Compiles the frontend and copies it to a shared volume. -2. **Start Routstr**: Runs the Python node, mounting the built UI. -3. **Start Tor**: Provides anonymous access via a `.onion` address. +!!! note "The first start takes a few minutes" + `docker compose up` builds both images locally, and the Next.js dashboard + build is the slow part. Later starts reuse the built images. + +!!! 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 -services: - routstr: - image: ghcr.io/routstr/proxy:latest - container_name: routstr - 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. +1. **ui** — builds the Next.js admin dashboard and copies the result into the + shared `./ui_out` volume. +2. **routstr** — the Python node, serving the API and the dashboard built above. +3. **tor** — serves the node as a `.onion` hidden service, so no port forwarding + is needed. See [Tor Support](tor.md) for how to read your `.onion` address. --- ## Pre-Configuration (Optional) -While everything can be configured via the dashboard, you can pre-configure settings with environment variables for automated deployments. - -### 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`: +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: ```bash +# Upstream (optional — can also be set from the dashboard) UPSTREAM_BASE_URL=https://api.openai.com/v1 UPSTREAM_API_KEY=sk-proj-... -# Keep the database (and the key file generated beside it) on the mounted volume. -DATABASE_URL=sqlite:////app/data/routstr.db + # Encrypts node secrets at rest. Optional — if unset, a key is generated next to # your database (on the same volume) and its file is named once for backup. Set # it explicitly to manage the key yourself. ROUTSTR_SECRET_KEY= + +# Node identity NAME=My Provider Node +DESCRIPTION=Fast GPT-4 access via Lightning + +# Lightning withdrawals 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" If you leave `ROUTSTR_SECRET_KEY` unset, the node generates one and stores it - as `routstr_secret.key` **next to your database**, so it persists on the same - volume as your data — just include that volume in your backups. For stronger - isolation (keeping the key off the data volume), set `ROUTSTR_SECRET_KEY` from - a secrets manager instead. + as `routstr_secret.key` **next to your database**, so it persists alongside + your data — just include that in your backups. For stronger isolation + (keeping the key off the data volume), set `ROUTSTR_SECRET_KEY` from a + secrets manager instead. See [Configuration](configuration.md) for all available options. @@ -182,17 +100,21 @@ See [Configuration](configuration.md) for all available options. ## Persistence -Point `DATABASE_URL` inside `/app/data` (as the examples above do) so everything -Routstr persists lands on the mounted volume: +With the default `compose.yml` the repository directory is mounted into the +container, so everything Routstr persists stays in the directory you cloned: | 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 | | `.wallet/` | Cashu wallet data (your Bitcoin!) | +| `logs/` | Node logs | !!! 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 -Pull the latest image and restart: +Check out the new release and rebuild: ```bash -docker compose pull -docker compose up -d +git fetch --tags +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 -The easiest way to build everything from source: +`docker compose up -d` already builds from source. To build the images +explicitly without starting them: ```bash docker compose build ``` -### Individual Components -If you prefer building the node only (requires manual UI build first): +To build only the node image (the dashboard must already be built into +`./ui_out`): ```bash docker build -t routstr-node . -``` +``` \ No newline at end of file diff --git a/migrations/versions/a73d19b6c204_reservation_deadlines.py b/migrations/versions/a73d19b6c204_reservation_deadlines.py new file mode 100644 index 00000000..74702d61 --- /dev/null +++ b/migrations/versions/a73d19b6c204_reservation_deadlines.py @@ -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") diff --git a/migrations/versions/e4c7a1b9d520_add_direction_to_lightning_invoices.py b/migrations/versions/e4c7a1b9d520_add_direction_to_lightning_invoices.py new file mode 100644 index 00000000..a0a0964d --- /dev/null +++ b/migrations/versions/e4c7a1b9d520_add_direction_to_lightning_invoices.py @@ -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") diff --git a/pyproject.toml b/pyproject.toml index bb6f456f..9868e774 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -21,7 +21,9 @@ dependencies = [ "mdurl==0.1.2", "pillow>=10", "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] @@ -72,10 +74,12 @@ build-backend = "setuptools.build_meta" [tool.setuptools] packages = ["routstr"] +[tool.ruff] +extend-exclude = ["examples"] + [tool.ruff.lint] select = ["E", "F", "I"] ignore = ["E501"] -exclude = ["examples"] [tool.mypy] python_version = "3.11" @@ -111,6 +115,9 @@ override-dependencies = [ # Transitive deps whose dependents allow the patched version but don't require # it. Constraints raise the floor without bypassing any upstream pin. constraint-dependencies = [ + "anyio>=4.14.2", + "pyjwt>=2.15.0", + "urllib3>=2.8.0", "starlette>=1.3.1", "httpcore>=1.0.9", # 1.0.8 caps h11<0.15 # 1.76 is the first grpcio-tools release with CPython 3.14 wheels. diff --git a/routstr/auth.py b/routstr/auth.py index bf29e86e..99c1d18e 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -49,9 +49,7 @@ payments_logger = get_logger("routstr.payments") # Routstr platform fee constants ROUTSTR_FEE_PERCENT: float = 2.1 -ROUTSTR_LN_ADDRESS: str = ( - "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash" -) +ROUTSTR_LN_ADDRESS: str = "routstr-fees@rizful.com" ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS: int = 900 ROUTSTR_FEE_DEFAULT_PAYOUT: int = 200 @@ -552,8 +550,8 @@ async def _validate_bearer_key_locked( async def pay_for_request( key: ApiKey, cost_per_request: int, session: AsyncSession -) -> int: - """Process payment for a request.""" +) -> ReservationSnapshot: + """Reserve funds and return the durable identity for this request.""" # Ensure cost_per_request is at least the minimum allowed request cost 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 + 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()) stmt = ( update(ApiKey) @@ -686,6 +692,13 @@ async def pay_for_request( billing_key_hash=reservation.billing_key_hash, reserved_msats=reservation.reserved_msats, 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 @@ -738,6 +751,53 @@ async def pay_for_request( 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( "Payment processed successfully", extra={ @@ -762,7 +822,7 @@ async def pay_for_request( }, ) - return cost_per_request + return reservation async def revert_pay_for_request( @@ -828,6 +888,10 @@ async def renew_reservation( update(ReservationRelease) .where(col(ReservationRelease.id) == snapshot.release_id) .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())) ) await session.commit() @@ -855,12 +919,21 @@ def _start_reservation_heartbeat(snapshot: ReservationSnapshot) -> None: """ interval = max(1, settings.stale_reservation_timeout_seconds // 3) 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: try: while True: 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 # sweeper can release the reservation if no terminal # transition ever ran. @@ -1044,6 +1117,10 @@ async def _claim_reservation_for_charge( update(ReservationRelease) .where(col(ReservationRelease.id) == snapshot.release_id) .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.billing_key_hash) == snapshot.billing_key_hash) .where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats) @@ -1104,7 +1181,7 @@ async def _charge_reservation_rows( return True -async def adjust_payment_for_tokens( +async def _adjust_payment_for_tokens( key: ApiKey, response_data: dict, session: AsyncSession, @@ -1540,6 +1617,45 @@ async def adjust_payment_for_tokens( 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: """Periodically prune dead API keys. Interval <= 0 disables it. diff --git a/routstr/checkstate.py b/routstr/checkstate.py new file mode 100644 index 00000000..396498fc --- /dev/null +++ b/routstr/checkstate.py @@ -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 diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 31140ff9..2027cc16 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -2557,6 +2557,7 @@ async def get_transactions_api( async def get_lightning_invoices_api( status: str | None = None, purpose: str | None = None, + direction: str | None = None, search: str | None = None, limit: int = 50, offset: int = 0, @@ -2569,6 +2570,8 @@ async def get_lightning_invoices_api( base = base.where(LightningInvoice.status == status) if purpose: base = base.where(LightningInvoice.purpose == purpose) + if direction: + base = base.where(LightningInvoice.direction == direction) if search: pattern = f"%{search}%" base = base.where( diff --git a/routstr/core/db.py b/routstr/core/db.py index eb4ba2fe..84307a7e 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -174,7 +174,10 @@ async def _transition_stale_reservation( update(ReservationRelease) .where(col(ReservationRelease.id) == reservation_id) .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") ) return bool(transition.rowcount == 1) @@ -221,7 +224,10 @@ async def release_stale_reservations( query = ( select(ReservationRelease) .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: query = query.where( @@ -509,14 +515,17 @@ class LightningInvoice(SQLModel, table=True): # type: ignore status: str = Field( default="pending", description=( - "pending, settlement_pending, paid, expired, cancelled, " + "pending, settlement_pending, paid, failed, expired, cancelled, " "reconciliation_required" ), ) api_key_hash: str | None = Field( 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( default=None, 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) billing_key_hash: str = Field(index=True) reserved_msats: int + started_at: int | None = Field(default=None) + expires_at: int | None = Field(default=None, index=True) status: str = Field(default="active") created_at: int = Field(default_factory=lambda: int(time.time())) @@ -1005,6 +1016,80 @@ async def complete_routstr_fee_payout( 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: """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) +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( db_session: AsyncSession, mint_url: str, unit: str ) -> int: diff --git a/routstr/core/error_scope.py b/routstr/core/error_scope.py new file mode 100644 index 00000000..9cc8306a --- /dev/null +++ b/routstr/core/error_scope.py @@ -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 diff --git a/routstr/core/exceptions.py b/routstr/core/exceptions.py index b8dcdc38..ce6ac46e 100644 --- a/routstr/core/exceptions.py +++ b/routstr/core/exceptions.py @@ -5,6 +5,11 @@ from fastapi.encoders import jsonable_encoder from fastapi.exceptions import RequestValidationError from fastapi.responses import JSONResponse +from .error_scope import ( + ERROR_SCOPE_UPSTREAM, + UPSTREAM_ERROR_STATUS, + UPSTREAM_UNAVAILABLE, +) from .logging import get_logger logger = get_logger(__name__) @@ -18,6 +23,19 @@ class UpstreamError(Exception): string-matching the message. ``details`` holds optional structured, redaction-safe context. Both default to ``None`` for backwards 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__( @@ -26,11 +44,15 @@ class UpstreamError(Exception): status_code: int = 502, code: str | None = None, details: dict[str, object] | None = None, + from_upstream_response: bool = False, + scope: str = ERROR_SCOPE_UPSTREAM, ): self.message = message self.status_code = status_code self.code = code self.details = details + self.from_upstream_response = from_upstream_response + self.scope = scope super().__init__(message) @@ -38,8 +60,9 @@ class EhbpTimeoutError(UpstreamError): """Raised when an EHBP upstream times out waiting for a response. Distinct from a generic :class:`UpstreamError` so callers can map the - failure to a ``504 Gateway Timeout`` with a stable ``UPSTREAM_TIMEOUT`` - code instead of a misleading ``500`` internal server error. + failure to a stable ``UPSTREAM_TIMEOUT`` code instead of a misleading + ``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 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): super().__init__( message, - status_code=504, + status_code=UPSTREAM_ERROR_STATUS, code="UPSTREAM_TIMEOUT", 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: """Extract a message from an HTTPException ``detail``, capped at 200 chars.""" if isinstance(detail, dict): diff --git a/routstr/core/lifecycle.py b/routstr/core/lifecycle.py new file mode 100644 index 00000000..039ac046 --- /dev/null +++ b/routstr/core/lifecycle.py @@ -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"} + ) diff --git a/routstr/core/logging.py b/routstr/core/logging.py index 0fba407b..8ad91490 100644 --- a/routstr/core/logging.py +++ b/routstr/core/logging.py @@ -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 - 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 - 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 """ +import copy import logging.config import logging.handlers import os +import queue import re import sys +import threading +import time import tomllib from datetime import datetime from pathlib import Path @@ -127,6 +136,218 @@ class DailyRotatingFileHandler(logging.handlers.TimedRotatingFileHandler): 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: """Read the package version from pyproject.toml.""" try: @@ -369,7 +590,7 @@ def setup_logging() -> None: "handlers": { "console": console_handler, "file": { - "()": DailyRotatingFileHandler, + "()": QueuedDailyRotatingFileHandler, "level": log_level, "formatter": "json", "filename": "logs/app.log", diff --git a/routstr/core/main.py b/routstr/core/main.py index 5ce22d22..368202cd 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -34,7 +34,7 @@ from ..payment.price import update_prices_periodically from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically from ..refund import periodic_refund_reconcile from ..upstream.auto_topup import periodic_auto_topup -from ..upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing +from ..upstream.http_client import close_upstream_http_client from ..upstream.litellm_routing import configure_litellm from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout from .admin import admin_router @@ -44,6 +44,7 @@ from .exceptions import ( http_exception_handler, validation_exception_handler, ) +from .lifecycle import RequestLifecycleMiddleware from .logging import get_logger, setup_logging from .middleware import LoggingMiddleware 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. 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_migrations() @@ -260,6 +256,14 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: "Error stopping background tasks", 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): @@ -289,6 +293,7 @@ app.add_middleware( expose_headers=[ "x-routstr-request-id", "x-cashu", + "x-routstr-error-scope", "x-routstr-cost-msats", "x-routstr-cost-usd", "x-routstr-input-cost-msats", @@ -305,6 +310,10 @@ app.add_middleware( # Add logging middleware app.add_middleware(LoggingMiddleware) +# Outermost: observe the actual downstream connection, not middleware streams. + +app.add_middleware(RequestLifecycleMiddleware) + # Add exception handlers app.add_exception_handler(HTTPException, http_exception_handler) # type: ignore app.add_exception_handler(RequestValidationError, validation_exception_handler) diff --git a/routstr/core/middleware.py b/routstr/core/middleware.py index cbcab96b..e158b590 100644 --- a/routstr/core/middleware.py +++ b/routstr/core/middleware.py @@ -1,7 +1,7 @@ import time import uuid from contextvars import ContextVar -from typing import Callable +from typing import AsyncIterator, Callable from urllib.parse import urlsplit from fastapi import Request, Response @@ -9,6 +9,7 @@ from starlette.datastructures import Headers from starlette.middleware.base import BaseHTTPMiddleware from .logging import get_logger +from .settings import settings 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: 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: - 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) +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): """Middleware to log proxy interactions and page navigation. 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: # Generate request ID request_id = str(uuid.uuid4()) @@ -108,15 +229,17 @@ class LoggingMiddleware(BaseHTTPMiddleware): # Set request ID in context for logging token = request_id_context.set(request_id) - client_app_token = client_app_context.set( - client_app_from_headers(request.headers) - ) + client_app = client_app_from_headers(request.headers) + client_app_token = client_app_context.set(client_app) path = request.url.path should_log = _should_log(request.method, path) - # Start timing - start_time = time.time() + # Start timing. Monotonic throughout: a wall-clock step would otherwise + # 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: logger.info( @@ -135,33 +258,52 @@ class LoggingMiddleware(BaseHTTPMiddleware): try: response = await call_next(request) - if should_log: - 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, - ) + headers_duration = time.monotonic() - stage_start + if hasattr(response, "headers"): 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 except Exception as e: # 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( "Request failed", extra={ @@ -171,6 +313,7 @@ class LoggingMiddleware(BaseHTTPMiddleware): "duration_ms": round(duration * 1000, 2), "error": str(e), "error_type": type(e).__name__, + **_attribution(request), }, exc_info=True, ) @@ -185,5 +328,6 @@ __all__ = [ "LoggingMiddleware", "UNKNOWN_CLIENT_APP", "client_app_context", + "mark", "request_id_context", ] diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 2104aeb3..d77d6fc1 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -11,6 +11,14 @@ from typing import Any from pydantic.v1 import BaseModel, BaseSettings, Field 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 Config: @@ -28,6 +36,26 @@ class Settings(BaseSettings): # Core upstream_base_url: str = Field(default="", env="UPSTREAM_BASE_URL") 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 name: str = Field(default="ARoutstrNode", env="NAME") @@ -37,7 +65,12 @@ class Settings(BaseSettings): onion_url: str = Field(default="", env="ONION_URL") # 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") primary_mint: str = Field(default="", env="PRIMARY_MINT_URL") 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 # Lightning 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. payout_interval_seconds: int = Field( default=900, gt=0, env="PAYOUT_INTERVAL_SECONDS" @@ -98,6 +133,16 @@ class Settings(BaseSettings): 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 cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS") # 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 # and prefixes are not supported. 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") providers_refresh_interval_seconds: int = Field( default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS" @@ -158,9 +211,21 @@ class Settings(BaseSettings): 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 log_level: str = Field(default="INFO", env="LOG_LEVEL") 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 chat_completions_api_version: str = Field( @@ -219,6 +284,10 @@ ENV_ONLY_FIELDS = frozenset( "database_pool_pre_ping", "database_pool_hold_warn_seconds", "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: - 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: diff --git a/routstr/lightning.py b/routstr/lightning.py index b3ece68f..732f1847 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -443,7 +443,8 @@ async def get_invoice_status( structured_errors: bool = Depends(_uses_v2_errors), ) -> InvoiceStatusResponse: 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( 404, "Invoice not found", @@ -486,7 +487,9 @@ async def recover_invoice( structured_errors: bool = Depends(_uses_v2_errors), ) -> InvoiceStatusResponse: 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() @@ -942,6 +945,7 @@ async def _expire_overdue_invoices(now: int) -> int: expired = await expiry_session.exec( # type: ignore[call-overload] update(LightningInvoice) .where( + col(LightningInvoice.direction) == "in", col(LightningInvoice.status) == "pending", 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}) settling = await session.exec( select(LightningInvoice) - .where(col(LightningInvoice.status) == "settlement_pending") + .where( + col(LightningInvoice.direction) == "in", + col(LightningInvoice.status) == "settlement_pending", + ) .order_by(col(LightningInvoice.created_at)) .limit(INVOICE_WATCH_BATCH_LIMIT // 2) ) unpaid = await session.exec( select(LightningInvoice) .where( + col(LightningInvoice.direction) == "in", col(LightningInvoice.status) == "pending", col(LightningInvoice.expires_at) >= now, ) @@ -975,6 +983,7 @@ async def _process_invoice_watch_batch(session: AsyncSession, prev_now: int) -> recoverable = await session.exec( select(LightningInvoice) .where( + col(LightningInvoice.direction) == "in", col(LightningInvoice.status) == "expired", col(LightningInvoice.expires_at) > now - INVOICE_EXPIRY_GRACE_SECONDS, ) diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index d088151f..4582031b 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -15,6 +15,14 @@ from PIL import Image from sqlmodel.ext.asyncio.session import AsyncSession 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.redaction import redact_org_ids from ..core.settings import settings @@ -654,13 +662,15 @@ def create_error_response( token: str | None = None, code: str | int | None = None, details: dict[str, object] | None = None, + error_scope: str | None = None, ) -> Response: """Create a standardized error response. ``code`` is a stable, machine-readable classification (e.g. ``UPSTREAM_RATE_LIMIT``); when omitted it defaults to the HTTP status code 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] = { "message": redact_org_ids(message), @@ -669,6 +679,11 @@ def create_error_response( } if details is not None: 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( content=json.dumps( { @@ -678,7 +693,7 @@ def create_error_response( ), status_code=status_code, media_type="application/json", - headers={"X-Cashu": token} if token else {}, + headers=headers, ) @@ -687,13 +702,29 @@ def create_upstream_error_response( request: Request, fallback_status: int = 502, ) -> Response: - """Build an error response from an :class:`UpstreamError`, preserving its - structured ``code``, ``details``, and original ``status_code``.""" + """Build an error response from an :class:`UpstreamError`. + + 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( "upstream_error", str(error), - error.status_code or fallback_status, + client_status_for_upstream_error(status_code, code), request=request, - code=getattr(error, "code", None), - details=getattr(error, "details", None), + code=client_code_for_upstream_error(status_code, code), + details=upstream_status_details(details, status_code), + error_scope=ERROR_SCOPE_UPSTREAM, ) diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py index d5fb7e89..808de593 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -325,7 +325,7 @@ async def raw_send_to_lnurl( unit: str, amount: int | None = None, *, - on_melt_quote: Callable[[str], Awaitable[None]] | None = None, + on_melt_quote: Callable[[str, str], Awaitable[None]] | None = None, ) -> int: """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") 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 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) try: @@ -424,7 +426,7 @@ async def raw_send_to_lnurl( lambda: wallet.melt( proofs=proofs, invoice=bolt11_invoice, - fee_reserve_sat=melt_quote_resp.fee_reserve, + fee_reserve_sat=change_budget, quote_id=melt_quote_resp.quote, ), op_name="lnurl_melt", diff --git a/routstr/payment/models.py b/routstr/payment/models.py index bdd8fa73..8d7788a1 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -241,61 +241,98 @@ def _has_valid_pricing(model: dict) -> bool: return True -async def async_fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: - """Asynchronously fetch model information from OpenRouter API.""" +# OpenRouter occasionally answers /models with a truncated body, emptying the +# 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" + timeout = OPENROUTER_MODELS_TIMEOUT_SECONDS - try: - async with httpx.AsyncClient() as client: - models_response, embeddings_response = await asyncio.gather( - client.get(f"{base_url}/models", timeout=30), - client.get(f"{base_url}/embeddings/models", timeout=30), - return_exceptions=True, - ) + async with httpx.AsyncClient() as client: + models_response, embeddings_response = await asyncio.gather( + client.get(f"{base_url}/models", timeout=timeout), + client.get(f"{base_url}/embeddings/models", timeout=timeout), + return_exceptions=True, + ) - def process_models_response( - response: httpx.Response | BaseException, - ) -> list[dict]: - if not isinstance(response, BaseException): - response.raise_for_status() - data = response.json() - return [ - model - for model in data.get("data", []) - if ":free" not in model.get("id", "").lower() - ] + # Losing /models is what empties the node, so it fails the attempt and + # the caller retries. A missing embeddings half must not do the same. + models_data = _parse_models_response(models_response) + try: + models_data.extend(_parse_models_response(embeddings_response)) + except Exception as e: + logger.warning(f"Skipping OpenRouter embeddings models: {e}") + + # 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 + + +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 [] + 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] = [] - 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 [] + return [] def _build_model_from_row( diff --git a/routstr/proxy.py b/routstr/proxy.py index 2ada5d3a..a2cf9492 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -1,16 +1,16 @@ import asyncio import inspect import json +import re from typing import Any -from fastapi import APIRouter, Depends, HTTPException, Request +from fastapi import APIRouter, HTTPException, Request from fastapi.responses import Response, StreamingResponse from sqlmodel import select from .algorithm import create_model_mappings from .auth import ( ReservationSnapshot, - get_reservation_snapshot, pay_for_request, revert_pay_for_request, validate_bearer_key, @@ -22,9 +22,15 @@ from .core.db import ( ModelRow, UpstreamProviderRow, 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.middleware import mark from .core.not_found import build_not_found_response from .core.settings import settings from .payment.helpers import ( @@ -36,6 +42,12 @@ from .payment.helpers import ( ) from .payment.models import Model 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.helpers import init_upstreams from .upstream.model_paths import ( @@ -112,8 +124,6 @@ def get_candidates( if candidates := _provider_map.get(model_id_lower): return candidates - import re - base_model_id = re.sub(r"-\d{8}$", "", model_id_lower) if base_model_id != model_id_lower: if candidates := _provider_map.get(base_model_id): @@ -267,6 +277,10 @@ _ALLOWED_ENDPOINTS: dict[str, frozenset[str]] = { "completions": frozenset({"POST"}), "responses": 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"}), # TypeSafe System One decision endpoint: POST {state, model, questions} # -> {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)) -@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) -async def proxy( - request: Request, path: str, session: AsyncSession = Depends(get_session) -) -> Response | StreamingResponse: - """Run proxy setup in a short request session, never across response streaming.""" +# Gateway conditions a retry usually clears. 500 is excluded: as likely to be a +# deterministic rejection that fails identically on the next attempt. +_RETRYABLE_UPSTREAM_5XX = frozenset({502, 503, 504}) +_UPSTREAM_5XX_RETRY_BACKOFF_SECONDS = 0.5 + + +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: - return await _proxy(request, path, session) - finally: - # FastAPI yield dependencies normally close after the response body is - # sent. Close explicitly so a long stream cannot retain DB resources. - close_result = session.close() - if inspect.isawaitable(close_result): - await close_result + body = await asyncio.wait_for(read(), timeout) + except _BodyLimitExceeded: + error_type, message, status = ( + "invalid_request", + f"Request body exceeds the {max_bytes} byte limit", + 413, + ) + 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( - request: Request, path: str, session: AsyncSession + request: Request, path: str, session: AsyncSession, request_body: bytes ) -> Response | StreamingResponse: # Screen the path before any routing decision: reject ambiguous spellings, # 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) 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 # header and a binary HPKE-sealed body. The proxy cannot parse the body to @@ -450,6 +549,12 @@ async def _proxy( else: 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 # without model/cost/auth lookups. Do not prefix-match here: paths such as # /attestationjunk must continue through normal authentication. @@ -472,11 +577,12 @@ async def _proxy( last_error_response = None for i, upstream in enumerate(selected_upstreams): + request.state.provider = upstream.provider_type try: headers = upstream.prepare_headers(dict(request.headers)) response = await upstream.forward_get_request(request, path, headers) if ( - response.status_code in [502, 429] + response.status_code in [424, 502, 503, 429] and i < len(selected_upstreams) - 1 ): logger.warning( @@ -498,7 +604,12 @@ async def _proxy( last_error_response = create_upstream_error_response(e, request) continue return last_error_response or create_error_response( - "upstream_error", "All upstreams failed", 502, request=request + "upstream_error", + "All upstreams failed", + UPSTREAM_ERROR_STATUS, + request=request, + code=UPSTREAM_UNAVAILABLE, + error_scope=ERROR_SCOPE_UPSTREAM, ) selector: ModelPathSelector | None = None @@ -606,6 +717,20 @@ async def _proxy( 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 # below rebinds (model_obj, upstream) per candidate so forwarding and # 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): last_error = None for i, (model_obj, upstream) in enumerate(candidates): + _attribute_request(request, model_obj, upstream) try: if is_ehbp: if not upstream.supports_ehbp: @@ -632,7 +758,7 @@ async def _proxy( model_id, ) continue - return await forward_ehbp_x_cashu_request( + response = await forward_ehbp_x_cashu_request( request=request, x_cashu_token=x_cashu, path=path, @@ -641,7 +767,7 @@ async def _proxy( upstream=upstream, ) elif is_responses_api: - return await upstream.handle_x_cashu_responses( + response = await upstream.handle_x_cashu_responses( request, x_cashu, path, @@ -650,7 +776,7 @@ async def _proxy( request_body=request_body, ) else: - return await upstream.handle_x_cashu( + response = await upstream.handle_x_cashu( request, x_cashu, path, @@ -658,6 +784,12 @@ async def _proxy( model_obj, 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: logger.warning( "Upstream %s failed (x-cashu) for model=%s: %s", @@ -670,6 +802,13 @@ async def _proxy( "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: last_error = e continue @@ -677,13 +816,19 @@ async def _proxy( if last_error is not None: return create_upstream_error_response(last_error, request) 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): key = await get_bearer_token_key( headers, path, session, auth, max_cost_for_model, model_id ) + mark(request, "auth") else: if request.method not in ["GET"]: @@ -697,12 +842,16 @@ async def _proxy( logger.debug("Processing unauthenticated GET request", extra={"path": path}) 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: headers = upstream.prepare_headers(dict(request.headers)) response = await upstream.forward_get_request(request, path, headers) - if response.status_code in [502, 429] and i < len(candidates) - 1: + if ( + response.status_code in [424, 502, 503, 429] + and i < len(candidates) - 1 + ): error_message = "" try: if hasattr(response, "body"): @@ -736,14 +885,18 @@ async def _proxy( last_error_response = create_upstream_error_response(e, request) continue return last_error_response or create_error_response( - "upstream_error", "All upstreams failed", 502, request=request + "upstream_error", + "All upstreams failed", + UPSTREAM_ERROR_STATUS, + request=request, + code=UPSTREAM_UNAVAILABLE, + error_scope=ERROR_SCOPE_UPSTREAM, ) reservation_snapshot: ReservationSnapshot | None = None if is_ehbp or request_body_dict: - await pay_for_request(key, max_cost_for_model, session) - reservation_snapshot = await get_reservation_snapshot(key, session) - # Snapshot validation performs SELECTs after pay_for_request commits. + reservation_snapshot = await pay_for_request(key, max_cost_for_model, session) + # pay_for_request refreshes the key after committing the reservation. # End that read transaction before waiting on upstream response headers. await _finish_read_transaction(session) @@ -770,18 +923,25 @@ async def _proxy( key, session, max_cost_for_model, reservation_snapshot ) try: - await pay_for_request(key, candidate_max, session) + reservation_snapshot = await pay_for_request( + key, candidate_max, session + ) except HTTPException: if i == len(candidates) - 1: raise - await pay_for_request(key, max_cost_for_model, session) - reservation_snapshot = await get_reservation_snapshot(key, session) + reservation_snapshot = await pay_for_request( + key, max_cost_for_model, session + ) await _finish_read_transaction(session) continue - reservation_snapshot = await get_reservation_snapshot(key, session) await _finish_read_transaction(session) 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)) try: @@ -834,8 +994,39 @@ async def _proxy( model_obj, reservation_snapshot, ) - except UpstreamError: - # Let the outer UpstreamError handler manage retry/revert + except UpstreamError as e: + # 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 except Exception as e: # Unexpected error (not an upstream failure) — revert and propagate @@ -873,7 +1064,7 @@ async def _proxy( already_stripped.add(bad_param) logger.warning( "Upstream %s rejected param '%s' for model=%s; " - "stripping and retrying same upstream", + "correcting and retrying same upstream", upstream.provider_type, bad_param, model_id, @@ -888,8 +1079,23 @@ async def _proxy( break if response.status_code != 200: - # Check if we should retry (502 Upstream Error or 429 Rate Limit) - should_retry = response.status_code in [502, 429, 400, 401, 403, 404] + if _upstream_response_failure(response): + 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: error_message = "" try: @@ -965,6 +1171,13 @@ async def _proxy( raise 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( "Upstream %s failed for model=%s: %s", upstream.provider_type, @@ -990,7 +1203,12 @@ async def _proxy( # Should not be reached given logic above 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, ) diff --git a/routstr/upstream/__init__.py b/routstr/upstream/__init__.py index edac0020..c57e0c09 100644 --- a/routstr/upstream/__init__.py +++ b/routstr/upstream/__init__.py @@ -1,6 +1,7 @@ from .anthropic import AnthropicUpstreamProvider from .azure import AzureUpstreamProvider from .base import BaseUpstreamProvider +from .deepseek import DeepSeekUpstreamProvider from .fireworks import FireworksUpstreamProvider from .gemini import GeminiUpstreamProvider from .generic import GenericUpstreamProvider @@ -13,11 +14,13 @@ from .ppqai import PPQAIUpstreamProvider from .routstr import RoutstrUpstreamProvider from .tinfoil import TinfoilUpstreamProvider from .typesafe import TypeSafeUpstreamProvider +from .venice import VeniceUpstreamProvider from .xai import XAIUpstreamProvider upstream_provider_classes: list[type[BaseUpstreamProvider]] = [ AnthropicUpstreamProvider, AzureUpstreamProvider, + DeepSeekUpstreamProvider, FireworksUpstreamProvider, GeminiUpstreamProvider, GenericUpstreamProvider, @@ -30,6 +33,7 @@ upstream_provider_classes: list[type[BaseUpstreamProvider]] = [ RoutstrUpstreamProvider, TinfoilUpstreamProvider, TypeSafeUpstreamProvider, + VeniceUpstreamProvider, XAIUpstreamProvider, ] """List of all upstream classes""" diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 0b6be7f1..4c4b68e8 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -1,17 +1,16 @@ from __future__ import annotations import asyncio -import inspect import json import math import traceback import typing import uuid -from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterator +from collections.abc import AsyncGenerator, AsyncIterator, Iterator from typing import Any, Mapping, Self, cast import httpx -from fastapi import BackgroundTasks, HTTPException, Request +from fastapi import HTTPException, Request from fastapi.responses import Response, StreamingResponse from pydantic.v1 import BaseModel @@ -31,6 +30,15 @@ from ..core.db import ( from ..core.db import ( 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, + upstream_status_details, +) from ..core.exceptions import UpstreamError from ..core.redaction import redact_org_ids from ..payment.cost_calculation import ( @@ -56,15 +64,30 @@ from ..wallet import ( send_token, token_mint_url, ) -from . import messages_dispatch +from . import json_codec, messages_dispatch from .cache_breakpoints import ( inject_anthropic_cache_breakpoints, is_explicit_cache_model, ) +from .cooldown import model_identity, provider_identity, record_failure from .count_tokens import MissingUsageEstimator, count_tokens_locally +from .http_client import acquire_upstream_http_client, build_x_cashu_client from .litellm_routing import detect_litellm_prefix +from .model_paths import public_provider_url from .rate_limit import UPSTREAM_RATE_LIMIT, classify_rate_limit from .reasoning_effort import apply_reasoning_effort +from .sse_splitter import SSEEventSplitter +from .stream_ownership import ( + ClosingStreamingResponse, + OwnedUpstreamStream, + PersistentStreamFinalizer, + ResponseHandoff, + aclose_if_needed, + attach_upstream_stream_owner, + close_upstream_exchange, + finalize_and_close_stream, +) +from .stream_timeout import GuardedStream, open_guarded_stream if typing.TYPE_CHECKING: from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget @@ -72,32 +95,6 @@ if typing.TYPE_CHECKING: 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 _finalize_and_close_stream( - finalize: Callable[[], Awaitable[None]] | None, - response: object | None, - client: httpx.AsyncClient | None, -) -> None: - try: - if finalize is not None: - await finalize() - finally: - try: - await _aclose_if_needed(response) - finally: - await _aclose_if_needed(client) - - CostMetadata = CostData | MaxCostData | dict[str, Any] @@ -220,6 +217,20 @@ def _responses_usage_payload(data_json: dict) -> dict: return nested if isinstance(nested, dict) else data_json +def _reported_provider(payload: dict) -> str | None: + """Provider named by an upstream payload, if any. + + Checked at top level first, then inside the Anthropic ``message`` and + Responses ``response`` envelopes, which is where those dialects nest it. + """ + for obj in (payload, payload.get("message"), payload.get("response")): + if isinstance(obj, dict): + value = obj.get("provider") + if isinstance(value, str) and value.strip(): + return value.strip() + return None + + def _render_sse_event(field_lines: list[str], data: str) -> str: """Re-frame one parsed event, re-prefixing every line of a multi-line data.""" body = "".join(f"{line}\n" for line in field_lines) @@ -483,12 +494,15 @@ class BaseUpstreamProvider: Idempotent: re-stamping an already-stamped payload must not nest the prefix repeatedly (e.g. never ``"anthropic:anthropic"``). This matters because streaming paths can apply the field more than once per chunk. + + Also stamps ``provider_url`` with the upstream base URL that served + the request. """ if not isinstance(response_json, dict): return + response_json["provider_url"] = public_provider_url(self.base_url) provider_type = (self.provider_type or "").strip() - existing = response_json.get("provider") - existing_str = existing.strip() if isinstance(existing, str) else "" + existing_str = _reported_provider(response_json) or "" if not existing_str: response_json["provider"] = provider_type return @@ -500,6 +514,17 @@ class BaseUpstreamProvider: return response_json["provider"] = f"{provider_type}:{existing_str}" + def _stamp_streamed_provider( + self, payload: dict, carried: str | None + ) -> str | None: + """Stamp a streamed payload, falling back to a provider an earlier event + reported. Returns the provider to carry forward to later payloads.""" + reported = _reported_provider(payload) + if reported is None and carried is not None: + payload["provider"] = carried + self._apply_provider_field(payload) + return reported or carried + def _log_full_refund( self, *, @@ -938,6 +963,10 @@ class BaseUpstreamProvider: error_code = UPSTREAM_RATE_LIMIT error_details = rate_limit.as_details() + client_status = client_status_for_upstream_error(status_code, error_code) + client_code = client_code_for_upstream_error(status_code, error_code) + headers[ERROR_SCOPE_HEADER] = ERROR_SCOPE_UPSTREAM + logger.warning( "Upstream %s returned %s for model=%s path=%s: %s", self.provider_type, @@ -1007,23 +1036,42 @@ class BaseUpstreamProvider: # ``org-*`` regex preserves the surrounding JSON structure. redacted_text = redact_org_ids(body_bytes.decode("utf-8", errors="ignore")) redacted_body = redacted_text.encode() - # Surface the stable rate-limit classification on the forwarded - # body so callers can switch on ``error.code`` without parsing the - # provider-specific message. Fall back to the redacted bytes if the - # body is not a JSON object with an ``error`` mapping. - if rate_limit is not None: + # Surface the stable classification on the forwarded body so callers + # can switch on ``error.code`` without parsing the provider-specific + # message. Fall back to the redacted bytes if the body is not a JSON + # object with an ``error`` mapping. + if rate_limit is not None or client_status != status_code: try: parsed = json.loads(redacted_text) err = parsed.get("error") if isinstance(parsed, dict) else None if isinstance(err, dict): - err["code"] = UPSTREAM_RATE_LIMIT - err["details"] = error_details + if rate_limit is not None: + err["code"] = UPSTREAM_RATE_LIMIT + err["details"] = error_details + if client_status != status_code: + err["code"] = client_code + err["upstream_status"] = status_code + redacted_body = json.dumps(parsed).encode() + elif ( + client_status != status_code + and isinstance(parsed, dict) + and "error" not in parsed + ): + # JSON body without an ``error`` mapping (e.g. FastAPI's + # ``{"detail": ...}``). Add one so a rewritten status is + # never served without its classification. + parsed["error"] = { + "message": message or "Upstream returned an error response", + "type": "upstream_error", + "code": client_code, + "upstream_status": status_code, + } redacted_body = json.dumps(parsed).encode() except (ValueError, AttributeError): pass return Response( content=redacted_body, - status_code=status_code, + status_code=client_status, headers=headers, media_type=media_type, ) @@ -1036,7 +1084,7 @@ class BaseUpstreamProvider: error_obj: dict[str, object] = { "message": message or "Upstream returned a non-JSON error response", "type": "upstream_error", - "code": error_code, + "code": client_code, "upstream_status": status_code, "upstream_content_type": content_type or None, "upstream_body_preview": body_preview or None, @@ -1050,7 +1098,7 @@ class BaseUpstreamProvider: return Response( content=json.dumps(envelope).encode(), - status_code=status_code, + status_code=client_status, headers=headers, media_type="application/json", ) @@ -1105,16 +1153,25 @@ class BaseUpstreamProvider: ) return True + async def _guard_stream( + self, response: httpx.Response, model_obj: Model | None, *, sse: bool + ) -> GuardedStream: + def on_idle() -> None: + if model_obj is not None and model_obj.id: + record_failure(provider_identity(self), model_identity(model_obj.id)) + + return await open_guarded_stream( + response, self.provider_type, sse=sse, on_idle_timeout=on_idle + ) + async def handle_streaming_chat_completion( self, response: httpx.Response, key: ApiKey, max_cost_for_model: int, - background_tasks: BackgroundTasks, requested_model: str | None = None, model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, - client: httpx.AsyncClient | None = None, request_body: bytes | None = None, legacy_completion: bool = False, ) -> StreamingResponse: @@ -1128,6 +1185,8 @@ class BaseUpstreamProvider: Returns: StreamingResponse with cost data injected at the end """ + guarded_chunks = await self._guard_stream(response, model_obj, sse=True) + if reservation_snapshot is None: async with create_session() as snapshot_session: snapshot_key = await snapshot_session.get(key.__class__, key.hashed_key) @@ -1148,51 +1207,61 @@ class BaseUpstreamProvider: }, ) + usage_finalized = False + last_model_seen: str | None = None + provider_seen: str | None = None + + async def finalize_db_only() -> None: + nonlocal usage_finalized + if usage_finalized: + return + try: + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + return + try: + await adjust_payment_for_tokens( + fresh_key, + usage_estimator.response_data(last_model_seen), + new_session, + max_cost_for_model, + model_obj, + self.provider_fee, + reservation_snapshot, + ) + usage_finalized = True + except Exception: + logger.exception( + "Fallback stream billing finalization failed; releasing reservation", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, new_session, reservation_snapshot + ) + ) + except Exception: + logger.exception( + "Fallback stream billing recovery could not access the database", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) + + stream_finalizer = PersistentStreamFinalizer( + lambda: finalize_and_close_stream( + None if usage_finalized else finalize_db_only, + response, + ) + ) + async def stream_with_cost( max_cost_for_model: int, ) -> AsyncGenerator[bytes, None]: - usage_finalized: bool = False - last_model_seen: str | None = None + nonlocal usage_finalized, last_model_seen usage_chunk_data: dict | None = None done_seen: bool = False stream_id: str | None = None - async def finalize_db_only() -> None: - nonlocal usage_finalized - if usage_finalized: - return - try: - async with create_session() as new_session: - fresh_key = await new_session.get(key.__class__, key.hashed_key) - if not fresh_key: - return - try: - await adjust_payment_for_tokens( - fresh_key, - usage_estimator.response_data(last_model_seen), - new_session, - max_cost_for_model, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - usage_finalized = True - except Exception: - logger.exception( - "Fallback stream billing finalization failed; releasing reservation", - extra={"key_hash": key.hashed_key[:8] + "..."}, - ) - usage_finalized = ( - await self._release_failed_streaming_reservation( - fresh_key, new_session, reservation_snapshot - ) - ) - except Exception: - logger.exception( - "Fallback stream billing recovery could not access the database", - extra={"key_hash": key.hashed_key[:8] + "..."}, - ) - def _process_event( raw_event: bytes, final: bool = False ) -> Iterator[bytes]: @@ -1215,6 +1284,7 @@ class BaseUpstreamProvider: end of stream. """ nonlocal last_model_seen, usage_chunk_data, done_seen, stream_id + nonlocal provider_seen event = raw_event.strip(b"\r\n") if not event: @@ -1250,14 +1320,11 @@ class BaseUpstreamProvider: done_seen = True return - try: - obj = json.loads(data) - except Exception: - obj = None + obj = json_codec.loads(data) if isinstance(obj, dict): usage_estimator.observe(obj) - self._apply_provider_field(obj) + provider_seen = self._stamp_streamed_provider(obj, provider_seen) if obj.get("model"): last_model_seen = str(obj.get("model")) if requested_model: @@ -1295,15 +1362,12 @@ class BaseUpstreamProvider: # usage is reported exactly once (in the trailer). forward = {k: v for k, v in obj.items() if k != "usage"} yield ( - prefix - + b"data: " - + json.dumps(forward).encode() - + b"\n\n" + prefix + b"data: " + json_codec.dumps(forward) + b"\n\n" ) return usage_chunk_data = obj return - yield prefix + b"data: " + json.dumps(obj).encode() + b"\n\n" + yield prefix + b"data: " + json_codec.dumps(obj) + b"\n\n" else: if final: # Final flush of a truncated tail: the upstream closed @@ -1325,21 +1389,14 @@ class BaseUpstreamProvider: # byte boundaries, so a single event's JSON can span chunks and # multiple events can arrive together; buffering makes parsing # boundary-independent for every provider. - buffer = b"" - async for chunk in response.aiter_bytes(): - # Normalize the *joined* buffer, not each chunk in - # isolation: a CRLF event delimiter can straddle two - # ``aiter_bytes`` chunks (``...\r`` then ``\n...``). A - # per-chunk replace would leave a stray ``\r`` and the - # ``\n\n`` split would miss the delimiter, merging two - # events into one frame and breaking SSE clients. - buffer = (buffer + chunk).replace(b"\r\n", b"\n") - while b"\n\n" in buffer: - raw_event, buffer = buffer.split(b"\n\n", 1) + splitter = SSEEventSplitter() + async for chunk in guarded_chunks: + for raw_event in splitter.feed(chunk): for out in _process_event(raw_event): yield out # Flush any trailing event that lacked a final blank line. + buffer = splitter.flush() if buffer.strip(): for out in _process_event(buffer, final=True): yield out @@ -1393,6 +1450,7 @@ class BaseUpstreamProvider: if legacy_completion else "chat.completion.chunk", "model": last_model_seen or "unknown", + "provider": provider_seen, "choices": [], "usage": { "prompt_tokens": cost_data.get("input_tokens", 0), @@ -1418,9 +1476,19 @@ class BaseUpstreamProvider: yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() - if done_seen: + if guarded_chunks.timed_out: + yield b'data: {"error":{"code":"UPSTREAM_TIMEOUT","message":"Upstream stream stalled"}}\n\n' + elif done_seen: yield b"data: [DONE]\n\n" + except httpx.RemoteProtocolError as stream_error: + logger.warning( + "Upstream stream ended before the response was complete", + extra={ + "error": str(stream_error), + "key_hash": key.hashed_key[:8] + "...", + }, + ) except Exception as stream_error: logger.warning( "Streaming interrupted; finalizing before closing upstream", @@ -1432,23 +1500,16 @@ class BaseUpstreamProvider: ) raise finally: - # Shielded so a client disconnect cannot cancel billing - # finalization or leak the upstream connection. - await asyncio.shield( - _finalize_and_close_stream( - None if usage_finalized else finalize_db_only, - response, - client, - ) - ) + await stream_finalizer.run() # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) response_headers.pop("content-encoding", None) response_headers.pop("content-length", None) - return StreamingResponse( + return ClosingStreamingResponse( stream_with_cost(max_cost_for_model), + finalizer=stream_finalizer, status_code=response.status_code, headers=response_headers, ) @@ -1611,7 +1672,6 @@ class BaseUpstreamProvider: requested_model: str | None = None, model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, - client: httpx.AsyncClient | None = None, request_body: bytes | None = None, ) -> StreamingResponse: """Handle streaming Responses API responses with token usage tracking and cost adjustment. @@ -1624,6 +1684,8 @@ class BaseUpstreamProvider: Returns: StreamingResponse with cost data injected at the end """ + guarded_chunks = await self._guard_stream(response, model_obj, sse=True) + usage_estimator = MissingUsageEstimator(request_body, model_obj) logger.debug( @@ -1635,51 +1697,61 @@ class BaseUpstreamProvider: }, ) + usage_finalized = False + last_model_seen: str | None = None + provider_seen: str | None = None + + async def finalize_db_only() -> None: + nonlocal usage_finalized + if usage_finalized: + return + try: + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + return + try: + await adjust_payment_for_tokens( + fresh_key, + usage_estimator.response_data(last_model_seen), + new_session, + max_cost_for_model, + model_obj, + self.provider_fee, + reservation_snapshot, + ) + usage_finalized = True + except Exception: + logger.exception( + "Fallback Responses billing finalization failed; releasing reservation", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, new_session, reservation_snapshot + ) + ) + except Exception: + logger.exception( + "Fallback Responses billing recovery could not access the database", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) + + stream_finalizer = PersistentStreamFinalizer( + lambda: finalize_and_close_stream( + None if usage_finalized else finalize_db_only, + response, + ) + ) + async def stream_with_responses_cost( max_cost_for_model: int, ) -> AsyncGenerator[bytes, None]: - usage_finalized: bool = False - last_model_seen: str | None = None + nonlocal usage_finalized, last_model_seen reasoning_tokens: int = 0 usage_chunk_data: dict | None = None done_seen: bool = False - async def finalize_db_only() -> None: - nonlocal usage_finalized - if usage_finalized: - return - try: - async with create_session() as new_session: - fresh_key = await new_session.get(key.__class__, key.hashed_key) - if not fresh_key: - return - try: - await adjust_payment_for_tokens( - fresh_key, - usage_estimator.response_data(last_model_seen), - new_session, - max_cost_for_model, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - usage_finalized = True - except Exception: - logger.exception( - "Fallback Responses billing finalization failed; releasing reservation", - extra={"key_hash": key.hashed_key[:8] + "..."}, - ) - usage_finalized = ( - await self._release_failed_streaming_reservation( - fresh_key, new_session, reservation_snapshot - ) - ) - except Exception: - logger.exception( - "Fallback Responses billing recovery could not access the database", - extra={"key_hash": key.hashed_key[:8] + "..."}, - ) - def _process_event( raw_event: bytes, final: bool = False ) -> Iterator[bytes]: @@ -1691,7 +1763,7 @@ class BaseUpstreamProvider: and preserves ``event:``/``id:`` fields attached to their data line so Responses API event framing stays intact. """ - nonlocal last_model_seen, usage_chunk_data, done_seen + nonlocal last_model_seen, usage_chunk_data, done_seen, provider_seen nonlocal reasoning_tokens event = raw_event.strip(b"\r\n") @@ -1724,13 +1796,10 @@ class BaseUpstreamProvider: done_seen = True return - try: - obj = json.loads(data) - except json.JSONDecodeError: - obj = None + obj = json_codec.loads(data) if isinstance(obj, dict): - self._apply_provider_field(obj) + provider_seen = self._stamp_streamed_provider(obj, provider_seen) if obj.get("model"): last_model_seen = str(obj.get("model")) if requested_model: @@ -1753,7 +1822,7 @@ class BaseUpstreamProvider: return usage_estimator.observe(obj) - yield prefix + b"data: " + json.dumps(obj).encode() + b"\n\n" + yield prefix + b"data: " + json_codec.dumps(obj) + b"\n\n" else: if final: # Final flush of a truncated tail: upstream closed @@ -1768,20 +1837,13 @@ class BaseUpstreamProvider: try: # Buffer across network chunks; dispatch only on the SSE event # delimiter so parsing is independent of byte boundaries. - buffer = b"" - async for chunk in response.aiter_bytes(): - # Normalize the *joined* buffer, not each chunk in - # isolation: a CRLF event delimiter can straddle two - # ``aiter_bytes`` chunks (``...\r`` then ``\n...``). A - # per-chunk replace would leave a stray ``\r`` and the - # ``\n\n`` split would miss the delimiter, merging two - # events into one frame and breaking SSE clients. - buffer = (buffer + chunk).replace(b"\r\n", b"\n") - while b"\n\n" in buffer: - raw_event, buffer = buffer.split(b"\n\n", 1) + splitter = SSEEventSplitter() + async for chunk in guarded_chunks: + for raw_event in splitter.feed(chunk): for out in _process_event(raw_event): yield out + buffer = splitter.flush() if buffer.strip(): for out in _process_event(buffer, final=True): yield out @@ -1825,7 +1887,10 @@ class BaseUpstreamProvider: if usage_chunk_data is None: usage_chunk_data = { - "type": "response.completed", + "type": "response.failed" + if guarded_chunks.timed_out + else "response.completed", + "provider": provider_seen, "response": { "model": last_model_seen or "unknown", "usage": { @@ -1846,6 +1911,14 @@ class BaseUpstreamProvider: + cost_data.get("output_tokens", 0), }, } + if guarded_chunks.timed_out: + usage_chunk_data["type"] = "response.failed" + response_data = usage_chunk_data.get("response") + if isinstance(response_data, dict): + response_data["error"] = { + "code": "UPSTREAM_TIMEOUT", + "message": "Upstream stream stalled", + } try: self.inject_cost_metadata( @@ -1861,9 +1934,22 @@ class BaseUpstreamProvider: yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() - if done_seen: + if guarded_chunks.timed_out and ( + usage_chunk_data is None + or usage_chunk_data.get("type") != "response.failed" + ): + yield b'data: {"error":{"code":"UPSTREAM_TIMEOUT","message":"Upstream stream stalled"}}\n\n' + if done_seen and not guarded_chunks.timed_out: yield b"data: [DONE]\n\n" + except httpx.RemoteProtocolError as stream_error: + logger.warning( + "Upstream Responses API stream ended before the response was complete", + extra={ + "error": str(stream_error), + "key_hash": key.hashed_key[:8] + "...", + }, + ) except Exception as stream_error: logger.warning( "Responses API streaming interrupted; finalizing before closing upstream", @@ -1875,23 +1961,16 @@ class BaseUpstreamProvider: ) raise finally: - # Shielded so a client disconnect cannot cancel billing - # finalization or leak the upstream connection. - await asyncio.shield( - _finalize_and_close_stream( - None if usage_finalized else finalize_db_only, - response, - client, - ) - ) + await stream_finalizer.run() # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) response_headers.pop("content-encoding", None) response_headers.pop("content-length", None) - return StreamingResponse( + return ClosingStreamingResponse( stream_with_responses_cost(max_cost_for_model), + finalizer=stream_finalizer, status_code=response.status_code, headers=response_headers, ) @@ -2055,12 +2134,12 @@ class BaseUpstreamProvider: provider_fee: float | None, reservation_snapshot: ReservationSnapshot, ) -> None: - """Background task to finalize payment for generic streaming requests.""" + """Finalize payment for a generic streaming request.""" async with create_session() as session: key = await session.get(ApiKey, key_hash) if not key: logger.warning( - "Key not found during background payment finalization", + "Key not found during generic streaming payment finalization", extra={"key_hash": key_hash[:8] + "..."}, ) return @@ -2079,7 +2158,7 @@ class BaseUpstreamProvider: reservation_snapshot=reservation_snapshot, ) logger.debug( - "Finalized generic streaming payment in background", + "Finalized generic streaming payment", extra={ "path": path, "key_hash": key_hash[:8] + "...", @@ -2087,7 +2166,7 @@ class BaseUpstreamProvider: ) except Exception as e: logger.error( - "Error finalizing generic streaming payment in background", + "Error finalizing generic streaming payment", extra={ "error": str(e), "key_hash": key_hash[:8] + "...", @@ -2095,6 +2174,91 @@ class BaseUpstreamProvider: }, ) + async def _stream_generic_with_settlement( + self, + response: httpx.Response, + key_hash: str, + max_cost: int, + path: str, + model_obj: Model | None, + provider_fee: float | None, + reservation_snapshot: ReservationSnapshot, + finalizer: PersistentStreamFinalizer | None = None, + guarded_chunks: GuardedStream | None = None, + ) -> AsyncGenerator[bytes, None]: + """Relay an opaque stream and settle it even if the caller disconnects.""" + if finalizer is None: + finalizer = PersistentStreamFinalizer( + lambda: finalize_and_close_stream( + lambda: self._finalize_generic_streaming_payment( + key_hash, + max_cost, + path, + model_obj, + provider_fee, + reservation_snapshot, + ), + response, + ) + ) + try: + if guarded_chunks is None: + guarded_chunks = await self._guard_stream( + response, model_obj, sse=False + ) + async for chunk in guarded_chunks: + yield chunk + if guarded_chunks.timed_out: + raise UpstreamError( + "Upstream stream stalled", + status_code=UPSTREAM_ERROR_STATUS, + code="UPSTREAM_TIMEOUT", + ) + finally: + await finalizer.run() + + async def _generic_streaming_response( + self, + response: httpx.Response, + key_hash: str, + max_cost: int, + path: str, + model_obj: Model | None, + provider_fee: float | None, + reservation_snapshot: ReservationSnapshot, + ) -> ClosingStreamingResponse: + guarded_chunks = await self._guard_stream(response, model_obj, sse=False) + finalizer = PersistentStreamFinalizer( + lambda: finalize_and_close_stream( + lambda: self._finalize_generic_streaming_payment( + key_hash, + max_cost, + path, + model_obj, + provider_fee, + reservation_snapshot, + ), + response, + ) + ) + stream = self._stream_generic_with_settlement( + response, + key_hash, + max_cost, + path, + model_obj, + provider_fee, + reservation_snapshot, + finalizer, + guarded_chunks, + ) + return ClosingStreamingResponse( + stream, + finalizer=finalizer, + status_code=response.status_code, + headers=dict(response.headers), + ) + async def handle_streaming_messages_completion( self, response: httpx.Response, @@ -2105,14 +2269,63 @@ class BaseUpstreamProvider: reservation_snapshot: ReservationSnapshot | None = None, request_body: bytes | None = None, ) -> StreamingResponse: + guarded_chunks = await self._guard_stream(response, model_obj, sse=True) + usage_estimator = MissingUsageEstimator(request_body, model_obj) + usage_finalized = False + last_model_seen: str | None = None + provider_seen: str | None = None + + async def finalize_without_usage() -> bytes | None: + nonlocal usage_finalized + if usage_finalized: + return None + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + usage_finalized = True + return None + try: + cost_data = await adjust_payment_for_tokens( + fresh_key, + usage_estimator.response_data(last_model_seen), + new_session, + max_cost_for_model, + model_obj, + self.provider_fee, + reservation_snapshot, + ) + usage_finalized = True + return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() + except BaseException as e: + logger.critical( + "Error during Messages API usage finalization — CRITICAL", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + exc_info=True, + ) + usage_finalized = await self._release_failed_streaming_reservation( + fresh_key, + new_session, + reservation_snapshot, + ) + raise + + async def finalize_db_only() -> None: + if not usage_finalized: + await finalize_without_usage() + + stream_finalizer = PersistentStreamFinalizer( + lambda: finalize_and_close_stream(finalize_db_only, response) + ) async def stream_with_cost( max_cost_for_model: int, ) -> AsyncGenerator[bytes, None]: + nonlocal usage_finalized, last_model_seen, provider_seen stored_chunks: list[bytes] = [] - usage_finalized: bool = False - last_model_seen: str | None = None input_tokens: int = 0 output_tokens: int = 0 cache_read_input_tokens: int = 0 @@ -2150,47 +2363,8 @@ class BaseUpstreamProvider: for field in ("total_cost", "cost"): total_cost = max(total_cost, _coerce_usd(usage_or_root.get(field))) - async def finalize_without_usage() -> bytes | None: - nonlocal usage_finalized - if usage_finalized: - return None - async with create_session() as new_session: - fresh_key = await new_session.get(key.__class__, key.hashed_key) - if not fresh_key: - usage_finalized = True - return None - try: - cost_data = await adjust_payment_for_tokens( - fresh_key, - usage_estimator.response_data(last_model_seen), - new_session, - max_cost_for_model, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - usage_finalized = True - return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() - except BaseException as e: - logger.critical( - "Error during Messages API usage finalization — CRITICAL", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "error": str(e), - }, - exc_info=True, - ) - usage_finalized = ( - await self._release_failed_streaming_reservation( - fresh_key, - new_session, - reservation_snapshot, - ) - ) - raise - try: - async for chunk in response.aiter_bytes(): + async for chunk in guarded_chunks: stored_chunks.append(chunk) try: decoded_chunk = chunk.decode("utf-8", errors="ignore") @@ -2207,7 +2381,9 @@ class BaseUpstreamProvider: last_model_seen = str(msg.get("model")) provider_added = "provider" not in data - self._apply_provider_field(data) + provider_seen = self._stamp_streamed_provider( + data, provider_seen + ) if requested_model: # Apply requested_model override @@ -2325,6 +2501,7 @@ class BaseUpstreamProvider: try: combined_data = { "model": last_model_seen or "unknown", + "provider": provider_seen, "usage": usage_data, } cost_data = await adjust_payment_for_tokens( @@ -2366,6 +2543,8 @@ class BaseUpstreamProvider: maybe_cost_event = await finalize_without_usage() if maybe_cost_event is not None: yield maybe_cost_event + if guarded_chunks.timed_out: + yield b'event: error\ndata: {"error":{"code":"UPSTREAM_TIMEOUT","message":"Upstream stream stalled"}}\n\n' except httpx.ReadError: if not usage_finalized: @@ -2376,15 +2555,15 @@ class BaseUpstreamProvider: await finalize_without_usage() raise finally: - if not usage_finalized: - await finalize_without_usage() + await stream_finalizer.run() response_headers = dict(response.headers) response_headers.pop("content-encoding", None) response_headers.pop("content-length", None) - return StreamingResponse( + return ClosingStreamingResponse( stream_with_cost(max_cost_for_model), + finalizer=stream_finalizer, status_code=response.status_code, headers=response_headers, ) @@ -2485,6 +2664,21 @@ class BaseUpstreamProvider: ) -> dict: return await messages_dispatch.aggregate_anthropic_events_to_message(iterator) + def transform_messages_stream( + self, stream: AsyncIterator[Any] + ) -> AsyncIterator[Any]: + return stream + + def adapt_messages_request(self, body: dict, model_obj: Model) -> str: + """Rewrite an allowlisted /v1/messages body for this upstream. + + Returns a suffix appended to the upstream model name, empty when the + provider needs none. Subclasses override this to express an Anthropic + feature the upstream spells differently; the base forwards the body + untouched. + """ + return "" + async def _dispatch_anthropic_messages( self, request_body: bytes | None, @@ -2499,6 +2693,8 @@ class BaseUpstreamProvider: api_key=self.api_key, provider_prefix=self.get_litellm_provider_prefix(), transform_model_name=self.transform_model_name, + adapt_request=lambda body: self.adapt_messages_request(body, model_obj), + transform_stream=self.transform_messages_stream, log_extra=log_extra, ) @@ -2662,10 +2858,71 @@ class BaseUpstreamProvider: with cost reconciliation appended at end of stream.""" usage_estimator = MissingUsageEstimator(request_body, model_obj) + usage_finalized = False + last_model_seen: str | None = None + + async def finalize_without_usage() -> bytes | None: + nonlocal usage_finalized + if usage_finalized: + return None + logger.warning( + "Finalizing /v1/messages stream with locally estimated " + "usage because the upstream omitted `usage` from SSE. " + "Check that the upstream emits a final usage chunk; the " + "reservation ceiling will not be used as the charge.", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "model": last_model_seen or "unknown", + "provider": self.provider_type or self.base_url, + "max_cost_msats": max_cost_for_model, + }, + ) + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + usage_finalized = True + return None + try: + cost_data = await adjust_payment_for_tokens( + fresh_key, + usage_estimator.response_data(last_model_seen), + new_session, + max_cost_for_model, + model_obj, + self.provider_fee, + reservation_snapshot, + ) + usage_finalized = True + return ( + f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n" + ).encode() + except BaseException as e: + logger.critical( + "Error during LiteLLM Messages usage finalization — CRITICAL", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + exc_info=True, + ) + usage_finalized = await self._release_failed_streaming_reservation( + fresh_key, + new_session, + reservation_snapshot, + ) + raise + + async def finalize_stream() -> None: + try: + if not usage_finalized: + await finalize_without_usage() + finally: + await aclose_if_needed(iterator) + + stream_finalizer = PersistentStreamFinalizer(finalize_stream) async def stream_with_cost() -> AsyncGenerator[bytes, None]: - usage_finalized = False - last_model_seen: str | None = None + nonlocal usage_finalized, last_model_seen input_tokens = 0 output_tokens = 0 cache_read_input_tokens = 0 @@ -2674,59 +2931,6 @@ class BaseUpstreamProvider: input_cost = 0.0 output_cost = 0.0 - async def finalize_without_usage() -> bytes | None: - nonlocal usage_finalized - if usage_finalized: - return None - logger.warning( - "Finalizing /v1/messages stream with locally estimated " - "usage because the upstream omitted `usage` from SSE. " - "Check that the upstream emits a final usage chunk; the " - "reservation ceiling will not be used as the charge.", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "model": last_model_seen or "unknown", - "provider": self.provider_type or self.base_url, - "max_cost_msats": max_cost_for_model, - }, - ) - async with create_session() as new_session: - fresh_key = await new_session.get(key.__class__, key.hashed_key) - if not fresh_key: - usage_finalized = True - return None - try: - cost_data = await adjust_payment_for_tokens( - fresh_key, - usage_estimator.response_data(last_model_seen), - new_session, - max_cost_for_model, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - usage_finalized = True - return ( - f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n" - ).encode() - except BaseException as e: - logger.critical( - "Error during LiteLLM Messages usage finalization — CRITICAL", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "error": str(e), - }, - exc_info=True, - ) - usage_finalized = ( - await self._release_failed_streaming_reservation( - fresh_key, - new_session, - reservation_snapshot, - ) - ) - raise - try: async for annotated in messages_dispatch.stream_annotated_events( iterator, requested_model @@ -2829,11 +3033,11 @@ class BaseUpstreamProvider: await finalize_without_usage() raise finally: - if not usage_finalized: - await finalize_without_usage() + await stream_finalizer.run() - return StreamingResponse( + return ClosingStreamingResponse( stream_with_cost(), + finalizer=stream_finalizer, media_type="text/event-stream", headers={"Cache-Control": "no-cache", "Connection": "keep-alive"}, ) @@ -2872,25 +3076,40 @@ class BaseUpstreamProvider: input_cost = 0.0 output_cost = 0.0 - async for annotated in messages_dispatch.stream_annotated_events( - iterator, requested_model - ): - if annotated.model: - last_model_seen = annotated.model - # See _stream_litellm_messages for why this is max() not +=. - input_tokens = max(input_tokens, annotated.input_tokens) - output_tokens = max(output_tokens, annotated.output_tokens) - cache_read_input_tokens = max( - cache_read_input_tokens, annotated.cache_read_input_tokens + try: + annotated_events = messages_dispatch.stream_annotated_events( + iterator, requested_model ) - cache_creation_input_tokens = max( - cache_creation_input_tokens, - annotated.cache_creation_input_tokens, - ) - total_cost = max(total_cost, annotated.total_cost) - input_cost = max(input_cost, annotated.input_cost) - output_cost = max(output_cost, annotated.output_cost) - buffered.append(annotated) + async for annotated in annotated_events: + if annotated.model: + last_model_seen = annotated.model + # See _stream_litellm_messages for why this is max() not +=. + input_tokens = max(input_tokens, annotated.input_tokens) + output_tokens = max(output_tokens, annotated.output_tokens) + cache_read_input_tokens = max( + cache_read_input_tokens, annotated.cache_read_input_tokens + ) + cache_creation_input_tokens = max( + cache_creation_input_tokens, + annotated.cache_creation_input_tokens, + ) + total_cost = max(total_cost, annotated.total_cost) + input_cost = max(input_cost, annotated.input_cost) + output_cost = max(output_cost, annotated.output_cost) + buffered.append(annotated) + except Exception as exc: + # Buffering lets us return an HTTP error before sending headers. + if messages_dispatch.is_provider_exception(exc): + raise messages_dispatch.upstream_error_from_exception( + exc, + log_message="Upstream stream failed mid-flight", + log_extra={ + "model": last_model_seen or requested_model or "unknown", + "provider": self.provider_type or self.base_url, + "request_id": request_id, + }, + ) from exc + raise response_headers: dict[str, str] = { "Cache-Control": "no-cache", @@ -2996,7 +3215,7 @@ class BaseUpstreamProvider: for annotated in buffered: yield annotated.sse_bytes - return StreamingResponse( + return ClosingStreamingResponse( replay(), media_type="text/event-stream", headers=response_headers, @@ -3075,12 +3294,11 @@ class BaseUpstreamProvider: }, ) - client = httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, - ) + response: httpx.Response | None = None + response_handoff = ResponseHandoff() try: + client = acquire_upstream_http_client(url) if transformed_body is not None: response = await client.send( client.build_request( @@ -3103,6 +3321,7 @@ class BaseUpstreamProvider: ), stream=True, ) + response_handoff.acquire(response) if response.status_code != 200: if response.status_code >= 500: @@ -3137,8 +3356,7 @@ class BaseUpstreamProvider: "body_preview": body_preview, }, ) - await response.aclose() - await client.aclose() + await response_handoff.close() raise UpstreamError( f"Upstream {self.provider_type} returned {response.status_code} " f"for model {original_model_id or 'unknown'}: " @@ -3146,6 +3364,7 @@ class BaseUpstreamProvider: status_code=response.status_code, code=rate_limit.code if rate_limit else None, details=rate_limit.as_details() if rate_limit else None, + from_upstream_response=True, ) try: @@ -3153,8 +3372,7 @@ class BaseUpstreamProvider: request, path, response, model_id=original_model_id ) finally: - await response.aclose() - await client.aclose() + await response_handoff.close() return mapped_error if ( @@ -3187,10 +3405,7 @@ class BaseUpstreamProvider: reservation_snapshot=reservation_snapshot, request_body=request_body, ) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - result.background = background_tasks + response_handoff.handoff() return result if response.status_code == 200: @@ -3207,8 +3422,7 @@ class BaseUpstreamProvider: request_body=request_body, ) finally: - await response.aclose() - await client.aclose() + await response_handoff.close() if path.endswith("messages/count_tokens"): if response.status_code == 200: @@ -3225,8 +3439,7 @@ class BaseUpstreamProvider: request_body=request_body, ) finally: - await response.aclose() - await client.aclose() + await response_handoff.close() if completion_path is not None: client_wants_streaming = False @@ -3263,19 +3476,18 @@ class BaseUpstreamProvider: ) if is_streaming and response.status_code == 200: - background_tasks = BackgroundTasks() - return await self.handle_streaming_chat_completion( + result = await self.handle_streaming_chat_completion( response, key, max_cost_for_model, - background_tasks, requested_model=original_model_id, model_obj=model_obj, reservation_snapshot=reservation_snapshot, - client=client, request_body=request_body, legacy_completion=completion_path == "completions", ) + response_handoff.handoff() + return result # Handle both non-streaming chat completions and embeddings if response.status_code == 200: @@ -3292,25 +3504,11 @@ class BaseUpstreamProvider: legacy_completion=completion_path == "completions", ) finally: - await response.aclose() - await client.aclose() + await response_handoff.close() if reservation_snapshot is None: reservation_snapshot = await get_reservation_snapshot(key, session) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - background_tasks.add_task( - self._finalize_generic_streaming_payment, - key.hashed_key, - max_cost_for_model, - path, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - logger.debug( "Streaming non-chat response", extra={ @@ -3320,18 +3518,24 @@ class BaseUpstreamProvider: }, ) - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - background=background_tasks, + result = await self._generic_streaming_response( + response, + key.hashed_key, + max_cost_for_model, + path, + model_obj, + self.provider_fee, + reservation_snapshot, ) + response_handoff.handoff() + return result except UpstreamError: + await response_handoff.close() raise except httpx.RequestError as exc: - await client.aclose() + await response_handoff.close() error_type = type(exc).__name__ error_details = str(exc) @@ -3349,19 +3553,26 @@ class BaseUpstreamProvider: ) # Don't revert here — proxy.py owns payment revert to avoid double-revert - if isinstance(exc, httpx.ConnectError): + if isinstance(exc, httpx.PoolTimeout): + error_message = "Upstream connection pool is busy" + status_code = 503 + elif isinstance(exc, httpx.ConnectError): error_message = "Unable to connect to upstream service" + status_code = 502 elif isinstance(exc, httpx.TimeoutException): error_message = "Upstream service request timed out" + status_code = 502 elif isinstance(exc, httpx.NetworkError): error_message = "Network error while connecting to upstream service" + status_code = 502 else: error_message = f"Error connecting to upstream service: {error_type}" + status_code = 502 - raise UpstreamError(error_message, status_code=502) + raise UpstreamError(error_message, status_code=status_code) except Exception as exc: - await client.aclose() + await response_handoff.close() tb = traceback.format_exc() logger.error( @@ -3379,7 +3590,15 @@ class BaseUpstreamProvider: ) # Don't revert here — proxy.py owns payment revert to avoid double-revert - raise UpstreamError("An unexpected server error occurred", status_code=500) + raise UpstreamError( + "An unexpected server error occurred", + status_code=500, + scope=ERROR_SCOPE_NODE, + ) + + except BaseException: + await response_handoff.close(suppress_errors=True) + raise supports_ehbp: bool = False @@ -3451,12 +3670,11 @@ class BaseUpstreamProvider: }, ) - client = httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, - ) + response: httpx.Response | None = None + response_handoff = ResponseHandoff() try: + client = acquire_upstream_http_client(url) if transformed_body is not None: response = await client.send( client.build_request( @@ -3479,6 +3697,7 @@ class BaseUpstreamProvider: ), stream=True, ) + response_handoff.acquire(response) if response.status_code != 200: if response.status_code >= 500: @@ -3512,8 +3731,7 @@ class BaseUpstreamProvider: "body_preview": body_preview, }, ) - await response.aclose() - await client.aclose() + await response_handoff.close() raise UpstreamError( f"Upstream {self.provider_type} returned {response.status_code} " f"for model {original_model_id or 'unknown'}: " @@ -3521,6 +3739,7 @@ class BaseUpstreamProvider: status_code=response.status_code, code=rate_limit.code if rate_limit else None, details=rate_limit.as_details() if rate_limit else None, + from_upstream_response=True, ) try: @@ -3528,8 +3747,7 @@ class BaseUpstreamProvider: request, path, response, model_id=original_model_id ) finally: - await response.aclose() - await client.aclose() + await response_handoff.close() return mapped_error if path.startswith("responses"): @@ -3546,16 +3764,17 @@ class BaseUpstreamProvider: ) if is_streaming and response.status_code == 200: - return await self.handle_streaming_responses_completion( + result = await self.handle_streaming_responses_completion( response, key, max_cost_for_model, requested_model=original_model_id, model_obj=model_obj, reservation_snapshot=reservation_snapshot, - client=client, request_body=transformed_body, ) + response_handoff.handoff() + return result if response.status_code == 200: try: @@ -3570,25 +3789,11 @@ class BaseUpstreamProvider: request_body=transformed_body, ) finally: - await response.aclose() - await client.aclose() + await response_handoff.close() if reservation_snapshot is None: reservation_snapshot = await get_reservation_snapshot(key, session) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - background_tasks.add_task( - self._finalize_generic_streaming_payment, - key.hashed_key, - max_cost_for_model, - path, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - logger.debug( "Streaming non-Responses API response", extra={ @@ -3598,18 +3803,24 @@ class BaseUpstreamProvider: }, ) - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - background=background_tasks, + result = await self._generic_streaming_response( + response, + key.hashed_key, + max_cost_for_model, + path, + model_obj, + self.provider_fee, + reservation_snapshot, ) + response_handoff.handoff() + return result except UpstreamError: + await response_handoff.close() raise except httpx.RequestError as exc: - await client.aclose() + await response_handoff.close() error_type = type(exc).__name__ error_details = str(exc) @@ -3627,19 +3838,26 @@ class BaseUpstreamProvider: ) # Don't revert here — proxy.py owns payment revert to avoid double-revert - if isinstance(exc, httpx.ConnectError): + if isinstance(exc, httpx.PoolTimeout): + error_message = "Upstream connection pool is busy" + status_code = 503 + elif isinstance(exc, httpx.ConnectError): error_message = "Unable to connect to upstream service" + status_code = 502 elif isinstance(exc, httpx.TimeoutException): error_message = "Upstream service request timed out" + status_code = 502 elif isinstance(exc, httpx.NetworkError): error_message = "Network error while connecting to upstream service" + status_code = 502 else: error_message = f"Error connecting to upstream service: {error_type}" + status_code = 502 - raise UpstreamError(error_message, status_code=502) + raise UpstreamError(error_message, status_code=status_code) except Exception as exc: - await client.aclose() + await response_handoff.close() tb = traceback.format_exc() logger.error( @@ -3657,7 +3875,15 @@ class BaseUpstreamProvider: ) # Don't revert here — proxy.py owns payment revert to avoid double-revert - raise UpstreamError("An unexpected server error occurred", status_code=500) + raise UpstreamError( + "An unexpected server error occurred", + status_code=500, + scope=ERROR_SCOPE_NODE, + ) + + except BaseException: + await response_handoff.close(suppress_errors=True) + raise async def forward_get_request( self, @@ -3688,66 +3914,91 @@ class BaseUpstreamProvider: }, ) - async with httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, - ) as client: - try: - response = await client.send( - client.build_request( - request.method, - url, - headers=headers, - content=request.stream(), - params=self.prepare_params(path, request.query_params), - ), + response: httpx.Response | None = None + try: + client = acquire_upstream_http_client(url) + response = await client.send( + client.build_request( + request.method, + url, + headers=headers, + content=request.stream(), + params=self.prepare_params(path, request.query_params), + ), + ) + + logger.debug( + "GET request forwarded", + extra={ + "path": path, + "status_code": response.status_code, + "provider": self.provider_type, + }, + ) + if response.status_code != 200: + return await self.forward_upstream_error_response( + request, path, response ) - logger.debug( - "GET request forwarded", - extra={ - "path": path, - "status_code": response.status_code, - "provider": self.provider_type, - }, - ) - if response.status_code != 200: - try: - mapped = await self.forward_upstream_error_response( - request, path, response - ) - finally: - await response.aclose() - return mapped - - response_headers = dict(response.headers) - response_headers.pop("content-encoding", None) - response_headers.pop("content-length", None) - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=response_headers, - ) - except Exception as exc: - tb = traceback.format_exc() - logger.error( - "Error forwarding GET request", - extra={ - "error": str(exc), - "error_type": type(exc).__name__, - "method": request.method, - "url": url, - "path": path, - "query_params": dict(request.query_params), - "traceback": tb, - }, - ) - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, - ) + response_headers = dict(response.headers) + response_headers.pop("content-encoding", None) + response_headers.pop("content-length", None) + return Response( + content=response.content, + status_code=response.status_code, + headers=response_headers, + ) + except UpstreamError: + raise + except httpx.PoolTimeout: + logger.warning( + "Upstream connection pool exhausted on GET", + extra={"path": path, "url": url, "provider": self.provider_type}, + ) + return create_error_response( + "service_unavailable", + "Upstream connection pool is busy", + 503, + request=request, + ) + except httpx.RequestError as exc: + logger.warning( + "Upstream request error on GET", + extra={ + "error": str(exc), + "error_type": type(exc).__name__, + "path": path, + "url": url, + "provider": self.provider_type, + }, + ) + return create_error_response( + "upstream_error", + "Unable to reach upstream service", + 502, + request=request, + ) + except Exception as exc: + logger.error( + "Error forwarding GET request", + extra={ + "error": str(exc), + "error_type": type(exc).__name__, + "method": request.method, + "url": url, + "path": path, + "query_params": dict(request.query_params), + "traceback": traceback.format_exc(), + }, + ) + return create_error_response( + "internal_error", + "An unexpected server error occurred", + 500, + request=request, + ) + finally: + await aclose_if_needed(response) async def get_x_cashu_cost( self, @@ -4052,6 +4303,7 @@ class BaseUpstreamProvider: }, ) + provider_seen: str | None = None for i, line in enumerate(lines): if line.startswith("data: "): try: @@ -4059,7 +4311,9 @@ class BaseUpstreamProvider: if not isinstance(data_json, dict): continue provider_before = data_json.get("provider") - self._apply_provider_field(data_json) + provider_seen = self._stamp_streamed_provider( + data_json, provider_seen + ) changed = data_json.get("provider") != provider_before if cost_data and "usage" in data_json and data_json["usage"]: _inject_cost_into_usage(data_json, cost_data) @@ -4073,7 +4327,7 @@ class BaseUpstreamProvider: for line in lines: yield (line + "\n").encode("utf-8") - return StreamingResponse( + return ClosingStreamingResponse( generate(), status_code=response.status_code, headers=response_headers, @@ -4329,7 +4583,7 @@ class BaseUpstreamProvider: "unit": unit, }, ) - return StreamingResponse( + return ClosingStreamingResponse( response.aiter_bytes(), status_code=response.status_code, headers=dict(response.headers), @@ -4417,146 +4671,150 @@ class BaseUpstreamProvider: }, ) - async with httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, - ) as client: - try: - response = await client.send( - client.build_request( - request.method, - url, - headers=headers, - content=transformed_body if transformed_body else request_body, - params=self.prepare_params(path, request.query_params), - ), - stream=True, - ) + client = build_x_cashu_client() + response: httpx.Response | None = None + try: + response = await client.send( + client.build_request( + request.method, + url, + headers=headers, + content=transformed_body if transformed_body else request_body, + params=self.prepare_params(path, request.query_params), + ), + stream=True, + ) - if response.status_code != 200: - logger.error( - "Received upstream response", - extra={ - "reason_phrase": response.reason_phrase, - "status_code": response.status_code, - "path": path, - "response_headers": dict(response.headers), - }, - ) - else: - logger.debug( - "Received upstream response", - extra={ - "status_code": response.status_code, - "path": path, - "response_headers": dict(response.headers), - }, - ) - - if response.status_code != 200: - logger.warning( - "Upstream request failed, processing refund", - extra={ - "status_code": response.status_code, - "path": path, - "amount": amount, - "unit": unit, - }, - ) - - refund_token = await self.send_refund( - amount, - unit, - mint, - request_id=getattr(request.state, "request_id", None), - ) - - logger.info( - "Refund processed for failed upstream request", - extra={ - "status_code": response.status_code, - "refund_amount": amount, - "unit": unit, - "refund_token_preview": refund_token[:20] + "..." - if len(refund_token) > 20 - else refund_token, - }, - ) - - error_response = Response( - content=json.dumps( - { - "error": { - "message": "Error forwarding request to upstream", - "type": "upstream_error", - "code": response.status_code, - "refund_token": refund_token, - } - } - ), - status_code=response.status_code, - media_type="application/json", - ) - error_response.headers["X-Cashu"] = refund_token - return error_response - - if _x_cashu_path_has_settlement_handler(path): - logger.debug( - "Processing completion/embeddings/messages response", - extra={"path": path, "amount": amount, "unit": unit}, - ) - - result = await self.handle_x_cashu_chat_completion( - response, - amount, - unit, - max_cost_for_model, - mint, - request_id=getattr(request.state, "request_id", None), - model_obj=model_obj, - request_body=request_body, - ) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - result.background = background_tasks - return result - - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - - logger.debug( - "Streaming non-chat response", - extra={"path": path, "status_code": response.status_code}, - ) - - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - background=background_tasks, - ) - except Exception as exc: - tb = traceback.format_exc() + if response.status_code != 200: logger.error( - "Unexpected error in upstream forwarding", + "Received upstream response", extra={ - "error": str(exc), - "error_type": type(exc).__name__, - "method": request.method, - "url": url, + "reason_phrase": response.reason_phrase, + "status_code": response.status_code, "path": path, - "query_params": dict(request.query_params), - "traceback": tb, + "response_headers": dict(response.headers), }, ) - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, + else: + logger.debug( + "Received upstream response", + extra={ + "status_code": response.status_code, + "path": path, + "response_headers": dict(response.headers), + }, ) + if response.status_code != 200: + logger.warning( + "Upstream request failed, processing refund", + extra={ + "status_code": response.status_code, + "path": path, + "amount": amount, + "unit": unit, + }, + ) + + refund_token = await self.send_refund( + amount, + unit, + mint, + request_id=getattr(request.state, "request_id", None), + ) + + logger.info( + "Refund processed for failed upstream request", + extra={ + "status_code": response.status_code, + "refund_amount": amount, + "unit": unit, + "refund_token_preview": refund_token[:20] + "..." + if len(refund_token) > 20 + else refund_token, + }, + ) + + error_response = Response( + content=json.dumps( + { + "error": { + "message": "Error forwarding request to upstream", + "type": "upstream_error", + # Pass the status as the code so a provider + # 4xx keeps the legacy numeric ``code``. + "code": client_code_for_upstream_error( + response.status_code, response.status_code + ), + "upstream_status": response.status_code, + "refund_token": refund_token, + } + } + ), + status_code=client_status_for_upstream_error(response.status_code), + media_type="application/json", + ) + error_response.headers["X-Cashu"] = refund_token + error_response.headers[ERROR_SCOPE_HEADER] = ERROR_SCOPE_UPSTREAM + await close_upstream_exchange(response, client) + return error_response + + if _x_cashu_path_has_settlement_handler(path): + logger.debug( + "Processing completion/embeddings/messages response", + extra={"path": path, "amount": amount, "unit": unit}, + ) + + result = await self.handle_x_cashu_chat_completion( + response, + amount, + unit, + max_cost_for_model, + mint, + request_id=getattr(request.state, "request_id", None), + model_obj=model_obj, + request_body=request_body, + ) + if isinstance(result, StreamingResponse) and not response.is_closed: + return attach_upstream_stream_owner(result, response, client) + await close_upstream_exchange(response, client) + return result + + logger.debug( + "Streaming non-chat response", + extra={"path": path, "status_code": response.status_code}, + ) + + return ClosingStreamingResponse( + OwnedUpstreamStream(response.aiter_bytes(), response, client), + status_code=response.status_code, + headers=dict(response.headers), + ) + except asyncio.CancelledError: + await close_upstream_exchange(response, client) + raise + except Exception as exc: + await close_upstream_exchange(response, client) + tb = traceback.format_exc() + logger.error( + "Unexpected error in upstream forwarding", + extra={ + "error": str(exc), + "error_type": type(exc).__name__, + "method": request.method, + "url": url, + "path": path, + "query_params": dict(request.query_params), + "traceback": tb, + }, + ) + return create_error_response( + "internal_error", + "An unexpected server error occurred", + 500, + request=request, + ) + async def handle_x_cashu_responses( self, request: Request, @@ -4647,12 +4905,16 @@ class BaseUpstreamProvider: # Post-redemption the token is spent; a forwarding failure must not # be reported as a retryable redemption error (see handle_x_cashu). if redeemed: + upstream_status = getattr(e, "status_code", None) + upstream_code = getattr(e, "code", None) return create_error_response( "upstream_error", "Payment succeeded but the upstream request failed", - 502, + client_status_for_upstream_error(upstream_status, upstream_code), request=request, - code="upstream_request_failed", + code=client_code_for_upstream_error(upstream_status, upstream_code), + details=upstream_status_details(None, upstream_status), + error_scope=ERROR_SCOPE_UPSTREAM, ) classified = classify_redemption_error(e) @@ -4725,134 +4987,138 @@ class BaseUpstreamProvider: }, ) - async with httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, - ) as client: - try: - response = await client.send( - client.build_request( - request.method, - url, - headers=headers, - content=transformed_body if transformed_body else request_body, - params=self.prepare_params(path, request.query_params), - ), - stream=True, - ) + client = build_x_cashu_client() + response: httpx.Response | None = None + try: + response = await client.send( + client.build_request( + request.method, + url, + headers=headers, + content=transformed_body if transformed_body else request_body, + params=self.prepare_params(path, request.query_params), + ), + stream=True, + ) - logger.debug( - "Received upstream Responses API response", + logger.debug( + "Received upstream Responses API response", + extra={ + "status_code": response.status_code, + "path": path, + "response_headers": dict(response.headers), + }, + ) + + if response.status_code != 200: + logger.warning( + "Upstream Responses API request failed, processing refund", extra={ "status_code": response.status_code, "path": path, - "response_headers": dict(response.headers), + "amount": amount, + "unit": unit, }, ) - if response.status_code != 200: - logger.warning( - "Upstream Responses API request failed, processing refund", - extra={ - "status_code": response.status_code, - "path": path, - "amount": amount, - "unit": unit, - }, - ) - - refund_token = await self.send_refund( - amount, - unit, - mint, - request_id=getattr(request.state, "request_id", None), - ) - - logger.info( - "Refund processed for failed upstream Responses API request", - extra={ - "status_code": response.status_code, - "refund_amount": amount, - "unit": unit, - "refund_token_preview": refund_token[:20] + "..." - if len(refund_token) > 20 - else refund_token, - }, - ) - - error_response = Response( - content=json.dumps( - { - "error": { - "message": "Error forwarding Responses API request to upstream", - "type": "upstream_error", - "code": response.status_code, - "refund_token": refund_token, - } - } - ), - status_code=response.status_code, - media_type="application/json", - ) - error_response.headers["X-Cashu"] = refund_token - return error_response - - if path.startswith("responses"): - logger.debug( - "Processing Responses API response", - extra={"path": path, "amount": amount, "unit": unit}, - ) - - result = await self.handle_x_cashu_responses_completion( - response, - amount, - unit, - max_cost_for_model, - mint, - request_id=getattr(request.state, "request_id", None), - model_obj=model_obj, - request_body=request_body, - ) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - result.background = background_tasks - return result - - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - - logger.debug( - "Streaming non-responses response", - extra={"path": path, "status_code": response.status_code}, + refund_token = await self.send_refund( + amount, + unit, + mint, + request_id=getattr(request.state, "request_id", None), ) - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - background=background_tasks, - ) - except Exception as exc: - tb = traceback.format_exc() - logger.error( - "Unexpected error in upstream Responses API forwarding", + logger.info( + "Refund processed for failed upstream Responses API request", extra={ - "error": str(exc), - "error_type": type(exc).__name__, - "method": request.method, - "url": url, - "path": path, - "query_params": dict(request.query_params), - "traceback": tb, + "status_code": response.status_code, + "refund_amount": amount, + "unit": unit, + "refund_token_preview": refund_token[:20] + "..." + if len(refund_token) > 20 + else refund_token, }, ) - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, + + error_response = Response( + content=json.dumps( + { + "error": { + "message": "Error forwarding Responses API request to upstream", + "type": "upstream_error", + # Pass the status as the code so a provider + # 4xx keeps the legacy numeric ``code``. + "code": client_code_for_upstream_error( + response.status_code, response.status_code + ), + "upstream_status": response.status_code, + "refund_token": refund_token, + } + } + ), + status_code=client_status_for_upstream_error(response.status_code), + media_type="application/json", ) + error_response.headers["X-Cashu"] = refund_token + error_response.headers[ERROR_SCOPE_HEADER] = ERROR_SCOPE_UPSTREAM + await close_upstream_exchange(response, client) + return error_response + + if path.startswith("responses"): + logger.debug( + "Processing Responses API response", + extra={"path": path, "amount": amount, "unit": unit}, + ) + + result = await self.handle_x_cashu_responses_completion( + response, + amount, + unit, + max_cost_for_model, + mint, + request_id=getattr(request.state, "request_id", None), + model_obj=model_obj, + request_body=request_body, + ) + if isinstance(result, StreamingResponse) and not response.is_closed: + return attach_upstream_stream_owner(result, response, client) + await close_upstream_exchange(response, client) + return result + + logger.debug( + "Streaming non-responses response", + extra={"path": path, "status_code": response.status_code}, + ) + + return ClosingStreamingResponse( + OwnedUpstreamStream(response.aiter_bytes(), response, client), + status_code=response.status_code, + headers=dict(response.headers), + ) + except asyncio.CancelledError: + await close_upstream_exchange(response, client) + raise + except Exception as exc: + await close_upstream_exchange(response, client) + tb = traceback.format_exc() + logger.error( + "Unexpected error in upstream Responses API forwarding", + extra={ + "error": str(exc), + "error_type": type(exc).__name__, + "method": request.method, + "url": url, + "path": path, + "query_params": dict(request.query_params), + "traceback": tb, + }, + ) + return create_error_response( + "internal_error", + "An unexpected server error occurred", + 500, + request=request, + ) async def handle_x_cashu_responses_completion( self, @@ -4936,7 +5202,7 @@ class BaseUpstreamProvider: "unit": unit, }, ) - return StreamingResponse( + return ClosingStreamingResponse( response.aiter_bytes(), status_code=response.status_code, headers=dict(response.headers), @@ -5108,6 +5374,7 @@ class BaseUpstreamProvider: }, ) + provider_seen: str | None = None for i, (fields, data) in enumerate(events): if data.strip() == "[DONE]": continue @@ -5118,7 +5385,7 @@ class BaseUpstreamProvider: if not isinstance(data_json, dict): continue provider_before = data_json.get("provider") - self._apply_provider_field(data_json) + provider_seen = self._stamp_streamed_provider(data_json, provider_seen) changed = data_json.get("provider") != provider_before payload = _responses_usage_payload(data_json) if cost_data and isinstance(payload.get("usage"), dict): @@ -5131,7 +5398,7 @@ class BaseUpstreamProvider: for fields, data in events: yield _render_sse_event(fields, data).encode("utf-8") - return StreamingResponse( + return ClosingStreamingResponse( generate(), status_code=response.status_code, headers=response_headers, @@ -5398,12 +5665,16 @@ class BaseUpstreamProvider: # must not surface as a retryable mint_unreachable (spent-token retry # bait). Redemption classification only applies while not redeemed. if redeemed: + upstream_status = getattr(e, "status_code", None) + upstream_code = getattr(e, "code", None) return create_error_response( "upstream_error", "Payment succeeded but the upstream request failed", - 502, + client_status_for_upstream_error(upstream_status, upstream_code), request=request, - code="upstream_request_failed", + code=client_code_for_upstream_error(upstream_status, upstream_code), + details=upstream_status_details(None, upstream_status), + error_scope=ERROR_SCOPE_UPSTREAM, ) classified = classify_redemption_error(e) diff --git a/routstr/upstream/cooldown.py b/routstr/upstream/cooldown.py new file mode 100644 index 00000000..7031817c --- /dev/null +++ b/routstr/upstream/cooldown.py @@ -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() diff --git a/routstr/upstream/deepseek.py b/routstr/upstream/deepseek.py new file mode 100644 index 00000000..a4eed7b3 --- /dev/null +++ b/routstr/upstream/deepseek.py @@ -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, + ) diff --git a/routstr/upstream/deepseek_v4_pricing_shim.py b/routstr/upstream/deepseek_v4_pricing_shim.py deleted file mode 100644 index ba0c392d..00000000 --- a/routstr/upstream/deepseek_v4_pricing_shim.py +++ /dev/null @@ -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}, - ) diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index 90b92719..c1ee8cfc 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -5,7 +5,7 @@ import math import time import traceback from dataclasses import dataclass, field -from typing import AsyncIterator, Mapping +from typing import AsyncIterator, Awaitable, Mapping from urllib.parse import urlsplit, urlunsplit from fastapi import Request @@ -31,6 +31,14 @@ from ..core.db import ( from ..core.db import ( 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.settings import settings 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( key: ApiKey, session: AsyncSession, @@ -879,6 +918,7 @@ async def forward_ehbp_request( f"EHBP upstream {provider_type} returned {resp.status_code} " f"for model {model_obj.id}: {body_preview[:200] or ''}", status_code=resp.status_code, + from_upstream_response=True, ) # 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 computed_msats = int(cost_info["total_msats"]) - charged_msats = await finalize_ehbp_actual_cost_payment( - key, - session, - max_cost_for_model, - billing_model, - cost_info, - reservation_snapshot, + charged_msats = await _record_ehbp_settlement( + finalize_ehbp_actual_cost_payment( + key, + session, + max_cost_for_model, + billing_model, + cost_info, + reservation_snapshot, + ), + key=key, + model_id=billing_model, + settlement_type="ehbp_usage", ) cost_data = { **cost_info, @@ -954,12 +999,17 @@ async def forward_ehbp_request( "key_hash": key.hashed_key[:8] + "...", }, ) - charged_msats = await finalize_ehbp_max_cost_payment( - key, - session, - max_cost_for_model, - model_obj.id, - reservation_snapshot, + charged_msats = await _record_ehbp_settlement( + finalize_ehbp_max_cost_payment( + key, + session, + max_cost_for_model, + model_obj.id, + reservation_snapshot, + ), + key=key, + model_id=model_obj.id, + settlement_type="ehbp_unmeasured_release", ) cost_data = { "total_msats": charged_msats, @@ -1031,7 +1081,11 @@ async def forward_ehbp_request( "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( @@ -1127,15 +1181,21 @@ async def forward_ehbp_x_cashu_request( "error": { "message": "Error forwarding EHBP request to upstream", "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, } } ), - status_code=resp.status_code, + status_code=client_status_for_upstream_error(resp.status_code), media_type="application/json", ) error_response.headers["X-Cashu"] = refund_token + error_response.headers[ERROR_SCOPE_HEADER] = ERROR_SCOPE_UPSTREAM return error_response # 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( "upstream_timeout", str(e), - 504, + UPSTREAM_ERROR_STATUS, request=request, code="UPSTREAM_TIMEOUT", + error_scope=ERROR_SCOPE_UPSTREAM, ) error_response.headers["X-Cashu"] = refund_token return error_response @@ -1259,9 +1320,10 @@ async def forward_ehbp_x_cashu_request( return create_error_response( "upstream_timeout", str(e), - 504, + UPSTREAM_ERROR_STATUS, request=request, code="UPSTREAM_TIMEOUT", + error_scope=ERROR_SCOPE_UPSTREAM, ) except Exception as e: @@ -1283,8 +1345,9 @@ async def forward_ehbp_x_cashu_request( error_response = create_error_response( "upstream_error", "EHBP request failed after token redemption; refunded token", - 502, + UPSTREAM_ERROR_STATUS, request=request, + error_scope=ERROR_SCOPE_UPSTREAM, ) error_response.headers["X-Cashu"] = refund_token return error_response @@ -1351,7 +1414,8 @@ async def forward_ehbp_x_cashu_request( return create_error_response( "cashu_error" if not redeemed else "upstream_error", f"EHBP X-Cashu request failed: {error_message}", - 400 if not redeemed else 502, + 400 if not redeemed else UPSTREAM_ERROR_STATUS, request=request, token=x_cashu_token if not redeemed else None, + error_scope=None if not redeemed else ERROR_SCOPE_UPSTREAM, ) diff --git a/routstr/upstream/gemini_messages.py b/routstr/upstream/gemini_messages.py index 11440b91..f2c454e0 100644 --- a/routstr/upstream/gemini_messages.py +++ b/routstr/upstream/gemini_messages.py @@ -44,6 +44,7 @@ Pipeline from __future__ import annotations +import asyncio import json import uuid from collections.abc import AsyncGenerator, AsyncIterator @@ -52,8 +53,10 @@ from typing import Any, Callable import httpx from ..core import get_logger +from ..core.error_scope import ERROR_SCOPE_NODE from ..core.exceptions import UpstreamError from ..payment.models import Model +from .http_client import acquire_upstream_http_client from .messages_dispatch import ( ANTHROPIC_ONLY_FIELDS, aggregate_anthropic_events_to_message, @@ -63,6 +66,46 @@ logger = get_logger(__name__) 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 _FINISH_TO_STOP = { "stop": "end_turn", @@ -112,6 +155,7 @@ def _translate_anthropic_to_openai(body: dict, model: str) -> dict: raise UpstreamError( "Failed to translate Anthropic body to OpenAI format", status_code=500, + scope=ERROR_SCOPE_NODE, ) return dict(translated) @@ -297,17 +341,21 @@ async def _openai_chunks_to_anthropic_events( yield _sse_event("message_stop", {"type": "message_stop"}) +GEMINI_STREAM_READ_TIMEOUT_SECONDS = 120.0 + + async def _post_and_stream( base_url: str, api_key: str, payload: dict, log_extra: dict[str, Any] | None, -) -> tuple[httpx.AsyncClient, httpx.Response]: - """POST to upstream chat-completions and return (client, response) for - streaming. Caller is responsible for closing both.""" +) -> httpx.Response: + """POST to upstream chat-completions and return a streaming response.""" url = f"{base_url.rstrip('/')}/chat/completions" - client = httpx.AsyncClient(timeout=httpx.Timeout(120.0, read=120.0)) try: + client = acquire_upstream_http_client(url) + # HTTPX replaces rather than merges per-request timeout settings. + client_timeout = client.timeout request = client.build_request( "POST", url, @@ -317,10 +365,25 @@ async def _post_and_stream( "Content-Type": "application/json", "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) + 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: - await client.aclose() logger.error( "Gemini messages dispatch HTTP error", extra={"error": str(exc), "url": url, **(log_extra or {})}, @@ -334,7 +397,6 @@ async def _post_and_stream( body_bytes = await response.aread() finally: await response.aclose() - await client.aclose() body_text = body_bytes.decode("utf-8", errors="replace") logger.error( "Gemini messages dispatch upstream error", @@ -348,9 +410,10 @@ async def _post_and_stream( raise UpstreamError( f"Upstream error via gemini compat: {body_text}", status_code=response.status_code, + from_upstream_response=True, ) - return client, response + return response async def dispatch_gemini_messages( @@ -371,9 +434,7 @@ async def dispatch_gemini_messages( aggregates). """ if not request_body: - raise UpstreamError( - "Missing request body for /v1/messages", status_code=400 - ) + raise UpstreamError("Missing request body for /v1/messages", status_code=400) try: body: dict = json.loads(request_body) @@ -441,9 +502,7 @@ async def dispatch_gemini_messages( }, ) - http_client, response = await _post_and_stream( - base_url, api_key, openai_kwargs, log_extra - ) + response = await _post_and_stream(base_url, api_key, openai_kwargs, log_extra) async def line_iter() -> AsyncGenerator[str, None]: try: @@ -451,10 +510,9 @@ async def dispatch_gemini_messages( yield line finally: await response.aclose() - await http_client.aclose() - anthropic_event_iter = _openai_chunks_to_anthropic_events( - line_iter(), requested_model + anthropic_event_iter = _ResponseOwnedIterator( + _openai_chunks_to_anthropic_events(line_iter(), requested_model), response ) if not client_stream: @@ -475,6 +533,8 @@ async def dispatch_gemini_messages( f"Failed to aggregate upstream stream: {exc}", status_code=502, ) from exc + finally: + await anthropic_event_iter.aclose() return client_stream, aggregated, requested_model return client_stream, anthropic_event_iter, requested_model diff --git a/routstr/upstream/generic.py b/routstr/upstream/generic.py index 63b6ffc5..1032e85d 100644 --- a/routstr/upstream/generic.py +++ b/routstr/upstream/generic.py @@ -1,10 +1,12 @@ from __future__ import annotations from typing import TYPE_CHECKING +from urllib.parse import urlparse import httpx -from .base import BaseUpstreamProvider +from .base import BaseUpstreamProvider, _reported_provider +from .model_paths import public_provider_url from .pricing_resolver import ( FallbackPricingResolver, ResolvedPricing, @@ -26,7 +28,11 @@ class GenericUpstreamProvider(BaseUpstreamProvider): provider_type = "generic" 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__( self, @@ -50,6 +56,21 @@ class GenericUpstreamProvider(BaseUpstreamProvider): provider_fee=provider_fee, ) + def _apply_provider_field(self, response_json: object) -> None: + """Stamp ``"generic:"`` 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 def _build_from_row( cls, provider_row: "UpstreamProviderRow" @@ -145,7 +166,7 @@ class GenericUpstreamProvider(BaseUpstreamProvider): model_spec = model_data.get("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) if resolved is None: diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index dc544b3a..c94733bc 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -272,6 +272,7 @@ async def _seed_providers_from_settings( ("PERPLEXITY_API_KEY", "perplexity", None, None), ("FIREWORKS_API_KEY", "fireworks", None, None), ("XAI_API_KEY", "xai", None, None), + ("DEEPSEEK_API_KEY", "deepseek", None, None), ("TINFOIL_API_KEY", "tinfoil", None, None), ("TYPESAFE_API_KEY", "typesafe", None, None), ] diff --git a/routstr/upstream/http_client.py b/routstr/upstream/http_client.py new file mode 100644 index 00000000..fb8cf7f9 --- /dev/null +++ b/routstr/upstream/http_client.py @@ -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, + ), + ) diff --git a/routstr/upstream/json_codec.py b/routstr/upstream/json_codec.py new file mode 100644 index 00000000..81570a71 --- /dev/null +++ b/routstr/upstream/json_codec.py @@ -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() diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index c7e698e3..da7c6f9f 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -36,6 +36,12 @@ from .reasoning_effort import adapt_messages_body_for_litellm 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 # Chat Completions. ``litellm.drop_params`` only filters *known* # 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: """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")), ) - 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( event, - sse_bytes, + encode_sse(event), in_tokens, out_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( iterator: AsyncIterator[Any], 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}") +_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 = "" + 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( *, request_body: bytes | None, @@ -458,6 +563,8 @@ async def dispatch_anthropic_messages( api_key: str, provider_prefix: 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, ) -> tuple[bool, Any, str | None]: """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` 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: 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} + 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; # `forwarded_model_id` is the public alias the internal API exposes # and echoes back to the client. requested_model = ( (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}" kwargs: dict = { "model": litellm_model, "api_base": base_url, - "api_key": api_key, + "api_key": api_key or KEYLESS_UPSTREAM_API_KEY, "stream": upstream_stream, **body, } @@ -530,45 +676,15 @@ async def dispatch_anthropic_messages( try: result = await litellm.anthropic.messages.acreate(**kwargs) except Exception as exc: - raw_message = getattr(exc, "message", None) or str(exc) or repr(exc) - # Redact provider account identifiers before the message reaches logs - # or the surfaced error. - exc_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 = "" - 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, + raise upstream_error_from_exception( + exc, + log_message="litellm dispatch failed", + log_extra={"model": litellm_model, "api_base": base_url}, ) 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__"): # Client asked for a non-streaming response but we always stream # from upstream — drain the events into a single Anthropic Message @@ -581,6 +697,13 @@ async def dispatch_anthropic_messages( cast(AsyncIterator[Any], result) ) 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( "Failed to aggregate streamed events into message", extra={ diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 0ac4e09d..db110e53 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -17,6 +17,7 @@ provider they named. from __future__ import annotations import asyncio +import functools import ipaddress import json import random @@ -102,6 +103,8 @@ class ProviderPathSnapshot: 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: """Mask private IP addresses and URLs with explicit ports.""" parsed = urlsplit(base_url) diff --git a/routstr/upstream/openai.py b/routstr/upstream/openai.py index f2f03cc5..b8d7c923 100644 --- a/routstr/upstream/openai.py +++ b/routstr/upstream/openai.py @@ -1,11 +1,23 @@ +import json 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 .base import BaseUpstreamProvider if TYPE_CHECKING: 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): """Upstream provider specifically configured for OpenAI API.""" @@ -42,6 +54,33 @@ class OpenAIUpstreamProvider(BaseUpstreamProvider): """Strip 'openai/' prefix for OpenAI API compatibility.""" 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]: """Fetch OpenAI models from OpenRouter API filtered by openai source.""" models_data = await async_fetch_openrouter_models(source_filter="openai") diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index 1caeaa5c..3ff13d90 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -2,12 +2,27 @@ from typing import TYPE_CHECKING import httpx +from ..core.logging import get_logger 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: 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): """Upstream provider specifically configured for OpenRouter API.""" @@ -26,22 +41,37 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): - Real upstream sub-provider (e.g. ``"GMICloud"``) -> ``"openrouter:GMICloud"``. - 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:..."``; the ``openrouter:`` prefix appears at most once. """ if not isinstance(response_json, dict): return + response_json["provider_url"] = public_provider_url(self.base_url) provider_type = (self.provider_type or "").strip() - existing = response_json.get("provider") - sub = existing.strip() if isinstance(existing, str) else "" + sub = _reported_provider(response_json) or "" # Strip any already-applied "openrouter:" prefixes (idempotency). prefix = f"{provider_type}:" while sub.lower().startswith(prefix.lower()): 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. 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 response_json["provider"] = f"{provider_type}:{sub}" diff --git a/routstr/upstream/request_correction.py b/routstr/upstream/request_correction.py index d959a015..6bb73632 100644 --- a/routstr/upstream/request_correction.py +++ b/routstr/upstream/request_correction.py @@ -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[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[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 # 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. @@ -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: """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. -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( diff --git a/routstr/upstream/sse_splitter.py b/routstr/upstream/sse_splitter.py new file mode 100644 index 00000000..f58a96a3 --- /dev/null +++ b/routstr/upstream/sse_splitter.py @@ -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 diff --git a/routstr/upstream/stream_ownership.py b/routstr/upstream/stream_ownership.py new file mode 100644 index 00000000..ea06e400 --- /dev/null +++ b/routstr/upstream/stream_ownership.py @@ -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() diff --git a/routstr/upstream/stream_timeout.py b/routstr/upstream/stream_timeout.py new file mode 100644 index 00000000..d6256da4 --- /dev/null +++ b/routstr/upstream/stream_timeout.py @@ -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 diff --git a/routstr/upstream/tinfoil.py b/routstr/upstream/tinfoil.py index 0928cbc9..63467fa4 100644 --- a/routstr/upstream/tinfoil.py +++ b/routstr/upstream/tinfoil.py @@ -1,5 +1,6 @@ from __future__ import annotations +import json from typing import TYPE_CHECKING, Optional import httpx @@ -7,6 +8,12 @@ from fastapi import Request from fastapi.responses import Response, StreamingResponse 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.logging import get_logger from ..payment.models import Architecture, Model, Pricing @@ -138,6 +145,30 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider): response_headers = dict(resp.headers) response_headers.pop("content-encoding", 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( content=resp.content, status_code=resp.status_code, diff --git a/routstr/upstream/tinfoil_trailer.py b/routstr/upstream/tinfoil_trailer.py index 1b4bb60d..2246a250 100644 --- a/routstr/upstream/tinfoil_trailer.py +++ b/routstr/upstream/tinfoil_trailer.py @@ -20,7 +20,7 @@ from urllib.parse import urlsplit import h11 from ..core import get_logger -from ..core.exceptions import EhbpTimeoutError +from ..core.exceptions import EhbpConnectionError, EhbpTimeoutError logger = get_logger(__name__) @@ -110,6 +110,21 @@ async def forward_with_trailer( raise EhbpTimeoutError( f"EHBP upstream {host} timed out after {timeout_seconds:g}s connecting" ) 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: # Build HTTP/1.1 request diff --git a/routstr/upstream/venice.py b/routstr/upstream/venice.py new file mode 100644 index 00000000..15a476fa --- /dev/null +++ b/routstr/upstream/venice.py @@ -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, + ) diff --git a/routstr/wallet.py b/routstr/wallet.py index beecde99..1f2a9e0e 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -15,12 +15,14 @@ import httpx from cashu.core.base import MeltQuote, Proof, Token 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_proofs as get_cashu_proofs from cashu.wallet.helpers import deserialize_token_from_string from cashu.wallet.wallet import Wallet as _CashuWallet from pydantic_core import PydanticUndefined from sqlmodel import col, select, update from .cashu_compat import install_cashu_httpx_shim +from .checkstate import filter_unspent_proofs from .core import db, get_logger from .core.db import store_cashu_transaction_with_retry as store_cashu_transaction from .core.settings import settings @@ -35,7 +37,7 @@ from .mint import ( mint_cooldown_remaining, 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 # 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]: """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: mint_urls.append(settings.primary_mint) return mint_urls @@ -143,6 +145,12 @@ class Wallet(_CashuWallet): request=resp.request, 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: response_data = resp.json() except json.JSONDecodeError: @@ -177,9 +185,12 @@ class Wallet(_CashuWallet): pass await self.load_mint_keysets(force_old_keysets) - await self.activate_keyset(keyset_id) 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() + await self.activate_keyset(keyset_id) class MintConnectionError(Exception): @@ -694,18 +705,64 @@ class Bolt11PaymentPlan: 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( mint_url: str, unit: str, proofs_balance: 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: - # Refund mint is a preference, not funding provenance. Mirror payout's - # conservative rule and protect the full liability at every mint. - user_liability = await db.total_user_liability(session) - # API-key balances are stored in msats. Cashu ``sat`` proofs are not. - if unit == "sat": - user_liability = _msats_to_sats_ceil(user_liability) - return max(0, proofs_balance - user_liability) + mint_liability = await db.user_liability_for_mint_and_unit( + session, mint_url, unit + ) + total_liability = await db.total_user_liability(session) + proofs_msats = _to_msats(proofs_balance, unit) + surplus_msats = min( + 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: @@ -780,9 +837,7 @@ async def _prepare_bolt11_payment(invoice: str) -> Bolt11PaymentPlan: ) if owner_balance < required: continue - owner_balance_msats = ( - owner_balance * 1000 if unit == "sat" else owner_balance - ) + owner_balance_msats = _to_msats(owner_balance, unit) candidates.append( (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. _wallet_last_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] = {} @@ -1165,6 +1223,7 @@ async def get_wallet( retry_on_rate_limit: bool = True, force_reload: bool = False, load_proofs: bool = True, + force_reload_proofs: bool = False, ) -> Wallet: global _wallets, _wallet_last_load, _wallet_last_mint_load, _wallet_load_locks id = f"{mint_url}_{unit}" @@ -1181,22 +1240,41 @@ async def get_wallet( or last_mint_load is None or now - last_mint_load >= _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS ): - 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, - ) + cached_error = _wallet_mint_load_errors.get(id) + if ( + not force_reload + and cached_error is not None + and now - cached_error[0] < _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS + ): + raise cached_error[1] + try: + 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() if load_proofs: last_proof_load = _wallet_last_load.get(id) if ( force_reload + or force_reload_proofs or last_proof_load is None or now - last_proof_load >= _WALLET_PROOF_RELOAD_MIN_INTERVAL_SECONDS @@ -1231,29 +1309,9 @@ async def slow_filter_spend_proofs( *, retry_on_rate_limit: bool = True, ) -> list[Proof]: - if not proofs: - return [] - _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 + return await filter_unspent_proofs( + proofs, wallet, retry_on_rate_limit=retry_on_rate_limit + ) 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: """Send only conservatively proven owner funds for one wallet.""" try: # Runs under wallet_operation_guard; a cached wallet may carry a proof # snapshot up to 30s stale from another process's reservation, so the - # cross-process lock is only safe with a fresh reload. - wallet = await get_wallet(mint_url, unit, force_reload=True) + # cross-process lock is only safe with fresh local proofs, not a + # 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) - if not proofs: - # Nothing to pay out, so skip the settle delay rather than hold the - # cross-process guard (and block credits) for a wallet with no funds. + min_amount = ( + settings.min_payout_sat + if unit == "sat" + else _sats_to_msats(settings.min_payout_sat) + ) + if sum(proof.amount for proof in proofs) <= min_amount: return proofs = await slow_filter_spend_proofs(proofs, wallet) await asyncio.sleep(5) @@ -1561,15 +1719,13 @@ async def _payout_mint_and_unit(mint_url: str, unit: str) -> None: ) return - # Fetch liability after the proofs snapshot and settle delay while the - # wallet operation guard excludes concurrent proof mutation and crediting. + # Read liabilities and the other wallets' proofs after this wallet's proofs + # snapshot and settle delay, while the wallet operation guard excludes + # concurrent proof mutation and crediting. try: - async with db.create_session() as session: - # ApiKey stores a refund preference, not funding provenance. Until - # 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) + available_balance = await _owner_balance_for_mint_and_unit( + mint_url, unit, sum(proof.amount for proof in proofs) + ) except Exception as e: logger.error( 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 try: - if unit == "sat": - user_balance = _msats_to_sats_ceil(user_balance) - proofs_balance = sum(proof.amount for proof in proofs) - available_balance = proofs_balance - user_balance - min_amount = ( - settings.min_payout_sat + max_amount = ( + settings.max_payout_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: - amount_received = await raw_send_to_lnurl( - wallet, - proofs, - settings.receive_ln_address, - unit, - amount=available_balance, - ) + payout_amount = min(available_balance, max_amount) + payout_quote_id: str | None = None + + async def record_payout(quote_id: str, bolt11: str) -> None: + nonlocal payout_quote_id + 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( "Payout sent successfully", extra={ "mint_url": mint_url, "unit": unit, "balance": available_balance, + "amount": payout_amount, "amount_received": amount_received, }, ) @@ -1640,6 +1830,7 @@ async def periodic_payout() -> None: # Proof mutation, liability observation, and sending are one # cross-process critical section. Credits take the same lock. async with wallet_operation_guard(): + await _reconcile_stale_payout_history(mint_url, unit) await _payout_mint_and_unit(mint_url, unit) except Exception as e: logger.error( @@ -1866,6 +2057,7 @@ async def periodic_routstr_fee_payout() -> None: payout_unit, ) if completed: + await _settle_payout_history(payout_quote_id, status="paid") logger.info( "Routstr fee payout reconciled as paid", extra={"payout_quote_id": payout_quote_id}, @@ -1880,6 +2072,9 @@ async def periodic_routstr_fee_payout() -> None: payout_unit, ) if restored: + await _settle_payout_history( + payout_quote_id, status="failed" + ) logger.warning( "Routstr fee payout reconciled as unpaid and restored for retry", extra={"payout_quote_id": payout_quote_id}, @@ -1910,7 +2105,7 @@ async def periodic_routstr_fee_payout() -> 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 async with db.create_session() as session: checkpointed = await db.reset_routstr_fee( @@ -1923,6 +2118,13 @@ async def periodic_routstr_fee_payout() -> None: if not checkpointed: raise _RoutstrFeePayoutAlreadyClaimed 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: 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}, exc_info=isinstance(e, Exception), ) + await _settle_payout_history( + attempt_quote_id, status="reconciliation_required" + ) if not isinstance(e, Exception): raise continue @@ -1969,6 +2174,9 @@ async def periodic_routstr_fee_payout() -> None: extra={"payout_in_progress_msats": paid_msats}, exc_info=isinstance(e, Exception), ) + await _settle_payout_history( + attempt_quote_id, status="reconciliation_required" + ) if not isinstance(e, Exception): raise continue @@ -1977,8 +2185,17 @@ async def periodic_routstr_fee_payout() -> None: "Routstr fee payout sent but checkpoint was not completed; awaiting quote reconciliation", extra={"payout_in_progress_msats": paid_msats}, ) + await _settle_payout_history( + attempt_quote_id, status="reconciliation_required" + ) continue + await _settle_payout_history( + attempt_quote_id, + status="paid", + amount_sats=_msats_to_sats(amount_received), + ) + logger.info( "Routstr fee payout sent", extra={ @@ -1995,8 +2212,8 @@ async def periodic_routstr_fee_payout() -> None: def _quote_callback( notify: Callable[[str, str], Awaitable[None]], mint: str -) -> Callable[[str], Awaitable[None]]: - async def callback(quote_id: str) -> None: +) -> Callable[[str, str], Awaitable[None]]: + async def callback(quote_id: str, _bolt11: str) -> None: await notify(quote_id, mint) return callback diff --git a/tests/conftest.py b/tests/conftest.py index d1bfa919..0e0ccbb3 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -31,3 +31,13 @@ def _isolate_redemption_negative_cache() -> Iterator[None]: redemption_negative_cache.clear() yield 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() diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index f8babb17..badaf549 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -365,6 +365,10 @@ async def test_database_url(tmp_path: Any) -> str: @pytest_asyncio.fixture async def integration_engine(test_database_url: str) -> AsyncGenerator[Any, None]: """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( test_database_url, echo=False, @@ -372,6 +376,7 @@ async def integration_engine(test_database_url: str) -> AsyncGenerator[Any, None pool_pre_ping=True, pool_size=5, max_overflow=10, + connect_args={"timeout": settings.database_busy_timeout}, ) # Initialize database schema diff --git a/tests/integration/test_failover_billing.py b/tests/integration/test_failover_billing.py index d67dc32b..f82b3248 100644 --- a/tests/integration/test_failover_billing.py +++ b/tests/integration/test_failover_billing.py @@ -27,6 +27,12 @@ EXPENSIVE_BASE_URL = "https://expensive.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( model_id: str, prompt_sats: float, @@ -191,14 +197,15 @@ async def test_failover_serve_billed_at_serving_providers_rate( assert response.status_code == 200 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] == [ + "cheap.example.com", "cheap.example.com", "expensive.example.com", ] # The fallback must be asked for ITS OWN model spelling, not the winner's. - forwarded_body = json.loads(sent_requests[1].content) + forwarded_body = json.loads(sent_requests[2].content) assert forwarded_body["model"] == "provb/dual-model" # 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 + # The winner's 502 is retried in place before failover. assert [r.url.host for r in sent_requests] == [ + "cheap.example.com", "cheap.example.com", "expensive.example.com", ] @@ -459,7 +468,9 @@ async def test_usd_cost_serve_carries_serving_providers_fee( ) assert response.status_code == 200 + # The winner's 502 is retried in place before failover. assert [r.url.host for r in sent_requests] == [ + "cheap.example.com", "cheap.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 # fallback must be rejected before its upstream is ever contacted. assert response.status_code == 402 - assert [r.url.host for r in sent_requests] == ["cheap.example.com"] + # 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 async def raised_envelope_provider_maps( patched_db_engine: None, @@ -593,7 +608,9 @@ async def test_failover_reserves_serving_candidates_envelope( ) assert response.status_code == 200 + # The winner's 502 is retried in place before failover. assert [r.url.host for r in sent_requests] == [ + "cheap.example.com", "cheap.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") assert charged.reserved_msats > released.reserved_msats 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", + ] diff --git a/tests/integration/test_lightning_invoice_constraints.py b/tests/integration/test_lightning_invoice_constraints.py index 7370ebe4..f715e67d 100644 --- a/tests/integration/test_lightning_invoice_constraints.py +++ b/tests/integration/test_lightning_invoice_constraints.py @@ -17,8 +17,10 @@ import pytest from cashu.core.base import Proof from sqlalchemy import inspect from sqlalchemy.ext.asyncio import AsyncEngine +from sqlmodel import select from sqlmodel.ext.asyncio.session import AsyncSession +from routstr.core import db from routstr.core.db import ApiKey, LightningInvoice from routstr.lightning import _create_api_key_record @@ -81,6 +83,41 @@ async def test_invoice_persists_validity_date( 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 # --------------------------------------------------------------------------- diff --git a/tests/integration/test_lightning_settlement.py b/tests/integration/test_lightning_settlement.py index 28803315..1a9a28e4 100644 --- a/tests/integration/test_lightning_settlement.py +++ b/tests/integration/test_lightning_settlement.py @@ -5,6 +5,7 @@ from unittest.mock import AsyncMock, Mock, patch import pytest from cashu.core.base import Proof +from fastapi import HTTPException from sqlalchemy.ext.asyncio import AsyncEngine from sqlmodel import col, update from sqlmodel.ext.asyncio.session import AsyncSession @@ -13,12 +14,15 @@ from routstr.core.db import ApiKey, LightningInvoice from routstr.lightning import ( INVOICE_EXPIRY_GRACE_SECONDS, INVOICE_WATCH_BATCH_LIMIT, + InvoiceRecoverRequest, _expire_invoice_if_authoritatively_unpaid, _expire_overdue_invoices, _finalize_invoice_settlement, _InvoiceSettlement, _process_invoice_watch_batch, 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 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) @@ -191,9 +196,7 @@ async def test_failed_final_commit_rolls_back_claim_and_credit_for_retry( assert unchanged.balance == 100_000 async with AsyncSession(integration_engine, expire_on_commit=False) as retry: - settled, _ = await _finalize_invoice_settlement( - snapshot, retry, 1_700_000_001 - ) + settled, _ = await _finalize_invoice_settlement(snapshot, retry, 1_700_000_001) assert settled 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 await paid.commit() - expired = await _expire_invoice_if_authoritatively_unpaid( - stale, caller, True - ) + expired = await _expire_invoice_if_authoritatively_unpaid(stale, caller, True) assert expired is False assert stale.status == "paid" @@ -381,8 +382,9 @@ async def test_sweep_expires_only_overdue_pending_invoices( overdue = _lightning_invoice(expires_at=now - 1) fresh = _lightning_invoice(expires_at=now + 3600) 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: - seed.add_all([overdue, fresh, settling]) + seed.add_all([overdue, fresh, settling, outgoing]) await seed.commit() await _expire_overdue_invoices(now) @@ -392,6 +394,7 @@ async def test_sweep_expires_only_overdue_pending_invoices( (overdue, "expired"), (fresh, "pending"), (settling, "settlement_pending"), + (outgoing, "pending"), ): stored = await verify.get(LightningInvoice, invoice.id) assert stored is not None @@ -409,8 +412,11 @@ async def test_watch_batch_expires_overdue_invoices_and_keeps_settling_rows( settling = _lightning_invoice( 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: - seed.add_all([overdue, fresh, settling]) + seed.add_all([overdue, fresh, settling, outgoing]) await seed.commit() 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 settling.id in polled + assert outgoing.id not in polled async with AsyncSession(integration_engine, expire_on_commit=False) as verify: 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 {inv.id for inv in settling} <= 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" diff --git a/tests/integration/test_proxy_session_lifecycle.py b/tests/integration/test_proxy_session_lifecycle.py index 9a70f294..4233b653 100644 --- a/tests/integration/test_proxy_session_lifecycle.py +++ b/tests/integration/test_proxy_session_lifecycle.py @@ -33,7 +33,6 @@ async def test_authenticated_proxy_releases_db_connection_before_upstream_header request = MagicMock() request.method = "POST" request.headers = {"authorization": "Bearer test-key"} - request.body = AsyncMock(return_value=json.dumps({"model": "test-model"}).encode()) request.url.path = "/v1/chat/completions" 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)), ): 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 diff --git a/tests/integration/test_reservation_lifecycle.py b/tests/integration/test_reservation_lifecycle.py index 1f60ce9d..d3fb9f9f 100644 --- a/tests/integration/test_reservation_lifecycle.py +++ b/tests/integration/test_reservation_lifecycle.py @@ -55,9 +55,13 @@ async def test_reserve_increases_reserved_balance( cost = 100 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) + 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.balance == 500 # balance column is NOT decremented on reserve assert key.total_balance == 500 - cost # available = balance - reserved @@ -75,11 +79,11 @@ async def test_revert_releases_reservation( cost = 150 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) 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) assert key.reserved_balance == 0 diff --git a/tests/integration/test_secret_bootstrap.py b/tests/integration/test_secret_bootstrap.py index 12f0d9a1..6dcdb990 100644 --- a/tests/integration/test_secret_bootstrap.py +++ b/tests/integration/test_secret_bootstrap.py @@ -470,7 +470,6 @@ async def test_startup_runs_bootstrap_before_settings_initialize( return None monkeypatch.setattr(main, "configure_litellm", lambda: None) - monkeypatch.setattr(main, "register_deepseek_v4_pricing", lambda: None) monkeypatch.setattr(main, "run_migrations", lambda: None) monkeypatch.setattr(main, "init_db", noop_init_db) monkeypatch.setattr(main, "create_session", fake_create_session) diff --git a/tests/integration/test_venice_web_search_wire.py b/tests/integration/test_venice_web_search_wire.py new file mode 100644 index 00000000..2636af1f --- /dev/null +++ b/tests/integration/test_venice_web_search_wire.py @@ -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"] diff --git a/tests/unit/proxy_test_utils.py b/tests/unit/proxy_test_utils.py new file mode 100644 index 00000000..a4fe4d3b --- /dev/null +++ b/tests/unit/proxy_test_utils.py @@ -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) diff --git a/tests/unit/test_bounded_request_body.py b/tests/unit/test_bounded_request_body.py new file mode 100644 index 00000000..8b6cd06b --- /dev/null +++ b/tests/unit/test_bounded_request_body.py @@ -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 diff --git a/tests/unit/test_cache_pricing.py b/tests/unit/test_cache_pricing.py index cc4a4dbc..cfb7e9a6 100644 --- a/tests/unit/test_cache_pricing.py +++ b/tests/unit/test_cache_pricing.py @@ -31,15 +31,6 @@ from routstr.payment.models import ( backfill_cache_pricing, ) 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: diff --git a/tests/unit/test_checkstate.py b/tests/unit/test_checkstate.py new file mode 100644 index 00000000..9b4685b8 --- /dev/null +++ b/tests/unit/test_checkstate.py @@ -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"}, "error"]) +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() diff --git a/tests/unit/test_ehbp_timeout.py b/tests/unit/test_ehbp_timeout.py index 7ae78cc4..8ec300c2 100644 --- a/tests/unit/test_ehbp_timeout.py +++ b/tests/unit/test_ehbp_timeout.py @@ -1,14 +1,24 @@ from __future__ import annotations +import json from unittest.mock import AsyncMock, MagicMock 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 # --------------------------------------------------------------------------- -# 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 -async def test_x_cashu_timeout_refunds_and_returns_504( +async def test_x_cashu_timeout_refunds_and_returns_424( monkeypatch: pytest.MonkeyPatch, ) -> None: monkeypatch.setattr( @@ -78,27 +88,32 @@ async def test_x_cashu_timeout_refunds_and_returns_504( 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" + 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") # --------------------------------------------------------------------------- # 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 -async def test_bearer_timeout_propagates_504( +async def test_bearer_timeout_propagates_424( monkeypatch: pytest.MonkeyPatch, ) -> None: """A timed-out bearer request must not be rewritten to a 500. ``forward_ehbp_request`` ends in a bare ``except Exception`` that turns any error into ``UpstreamError(..., status_code=500)``. The ``except - UpstreamError: raise`` above it is the only thing preserving the 504 that - ``proxy.py`` returns to the client, so this test pins that handler. + UpstreamError: raise`` above it is the only thing preserving the upstream + timeout status that ``proxy.py`` returns to the client, so this test pins + that handler. """ monkeypatch.setattr( ehbp_module, @@ -126,6 +141,51 @@ async def test_bearer_timeout_propagates_504( 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.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) diff --git a/tests/unit/test_fee_payout_crash_safety.py b/tests/unit/test_fee_payout_crash_safety.py index c6fc3073..b60f42a0 100644 --- a/tests/unit/test_fee_payout_crash_safety.py +++ b/tests/unit/test_fee_payout_crash_safety.py @@ -1,5 +1,5 @@ import asyncio -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, Generator from contextlib import asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, Mock, patch @@ -29,6 +29,21 @@ def _session_context(session: Mock) -> _SessionContext: 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 async def test_fee_payout_checkpoint_is_atomic_and_durable() -> None: 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: 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") 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: 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() return 5 @@ -289,7 +304,9 @@ async def test_fee_payout_lost_checkpoint_race_does_not_send() -> None: @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() fee = SimpleNamespace( accumulated_msats=10_000, @@ -331,10 +348,14 @@ async def test_fee_payout_finalizes_a_paid_unresolved_quote_without_resending() ) restore.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 -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() unresolved_fee = SimpleNamespace( 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: - 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 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" ) 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 @@ -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: 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") 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: 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 with ( @@ -666,6 +691,8 @@ async def test_fee_payout_completion_failures_use_sent_checkpoint_alert( side_effect=[ _session_context(session), _session_context(session), + _session_context(session), + 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: 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 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: assert engine.pool.checkedout() == 0 # type: ignore[attr-defined] 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] return 5 diff --git a/tests/unit/test_lightning_settlement.py b/tests/unit/test_lightning_settlement.py index 8d7c96a9..205f3db7 100644 --- a/tests/unit/test_lightning_settlement.py +++ b/tests/unit/test_lightning_settlement.py @@ -37,6 +37,7 @@ def _invoice(**overrides: object) -> SimpleNamespace: "payment_hash": "quote-1", "amount_sats": 100, "purpose": "create", + "direction": "in", "status": "pending", "paid_at": None, "api_key_hash": None, diff --git a/tests/unit/test_litellm_routing.py b/tests/unit/test_litellm_routing.py index 264a142d..077a765b 100644 --- a/tests/unit/test_litellm_routing.py +++ b/tests/unit/test_litellm_routing.py @@ -91,3 +91,24 @@ def test_detect_litellm_prefix_custom_default() -> None: assert detect_litellm_prefix("https://example.com", default="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} diff --git a/tests/unit/test_lnurl_amount_and_destination.py b/tests/unit/test_lnurl_amount_and_destination.py index 72986500..3a4f033e 100644 --- a/tests/unit/test_lnurl_amount_and_destination.py +++ b/tests/unit/test_lnurl_amount_and_destination.py @@ -192,7 +192,7 @@ async def test_raw_send_to_lnurl_requotes_for_exact_input_fees_without_recursion assert paid == 485_000 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() selected = wallet.melt.await_args.kwargs["proofs"] assert sum(proof.amount for proof in selected) == 500 @@ -358,9 +358,7 @@ def _patch_getaddrinfo(ip: str) -> Any: loop = MagicMock() loop.getaddrinfo = fake_getaddrinfo - return patch.object( - lnurl_module.asyncio, "get_running_loop", return_value=loop - ) + return patch.object(lnurl_module.asyncio, "get_running_loop", return_value=loop) @pytest.mark.asyncio diff --git a/tests/unit/test_lnurl_change.py b/tests/unit/test_lnurl_change.py new file mode 100644 index 00000000..90ea1a9b --- /dev/null +++ b/tests/unit/test_lnurl_change.py @@ -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() diff --git a/tests/unit/test_lnurl_melt_timeout.py b/tests/unit/test_lnurl_melt_timeout.py index ede3f5af..4c9ea306 100644 --- a/tests/unit/test_lnurl_melt_timeout.py +++ b/tests/unit/test_lnurl_melt_timeout.py @@ -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 async def test_raw_send_to_lnurl_direct_unpaid_is_retry_safe() -> None: wallet, proofs = _wallet() @@ -319,8 +346,9 @@ async def test_raw_send_to_lnurl_checkpoints_quote_before_melt_dispatch() -> Non wallet, proofs = _wallet() events: list[str] = [] - async def checkpoint(quote_id: str) -> None: + async def checkpoint(quote_id: str, bolt11: str) -> None: assert quote_id == "q" + assert bolt11 == "lnbc1..." events.append("checkpoint") async def melt(**_kwargs: object) -> MagicMock: diff --git a/tests/unit/test_log_model_provider_attribution.py b/tests/unit/test_log_model_provider_attribution.py new file mode 100644 index 00000000..76529af1 --- /dev/null +++ b/tests/unit/test_log_model_provider_attribution.py @@ -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" diff --git a/tests/unit/test_log_secret_redaction.py b/tests/unit/test_log_secret_redaction.py index 0f660dda..3b51644d 100644 --- a/tests/unit/test_log_secret_redaction.py +++ b/tests/unit/test_log_secret_redaction.py @@ -42,7 +42,6 @@ def log_dir(tmp_path: Path) -> Path: @pytest.fixture def handler(log_dir: Path) -> Iterator[DailyRotatingFileHandler]: - """A file handler configured exactly like the production ``file`` handler.""" handler = DailyRotatingFileHandler( str(log_dir / "app.log"), when="midnight", diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index 294f5c0c..47b3ccd9 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -23,6 +23,9 @@ from routstr.core.db import ApiKey # noqa: E402 from routstr.payment.cost_calculation import CostData # noqa: E402 from routstr.payment.models import Architecture, Model, Pricing # noqa: E402 from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 +from routstr.upstream.messages_dispatch import ( # noqa: E402 + prune_blank_system_blocks, +) 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: out = BaseUpstreamProvider._coerce_litellm_payload({"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 ) -> None: """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() model = _make_model() request = _make_request() @@ -1623,9 +1659,11 @@ async def test_x_cashu_transport_error_after_redemption_is_not_retryable( 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)) assert body["error"]["type"] == "upstream_error" + assert body["error"]["code"] == "UPSTREAM_UNAVAILABLE" 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" # Spent-to-zero token must not be echoed back for retry. 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" diff --git a/tests/unit/test_messages_upstream_errors.py b/tests/unit/test_messages_upstream_errors.py new file mode 100644 index 00000000..62699f69 --- /dev/null +++ b/tests/unit/test_messages_upstream_errors.py @@ -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) diff --git a/tests/unit/test_model_path_routing.py b/tests/unit/test_model_path_routing.py index 4dfa2ef8..cb764289 100644 --- a/tests/unit/test_model_path_routing.py +++ b/tests/unit/test_model_path_routing.py @@ -9,8 +9,16 @@ import pytest from routstr import proxy as proxy_module from routstr.auth import ReservationSnapshot 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 .proxy_test_utils import mock_request_stream, patch_proxy_session + MODEL_ID = "test-model" @@ -32,7 +40,7 @@ def _make_request(headers: dict[str, str], body: bytes) -> MagicMock: request = MagicMock() request.method = "POST" request.headers = headers - request.body = AsyncMock(return_value=body) + mock_request_stream(request, body) request.state = MagicMock() request.state.request_id = "req-model-path" return request @@ -62,15 +70,13 @@ async def _run_proxy( ), patch.object(proxy_module, "check_token_balance", MagicMock()), 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( - proxy_module, - "get_reservation_snapshot", - AsyncMock(return_value=reservation), + proxy_module, "pay_for_request", AsyncMock(return_value=reservation) ), 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: @@ -391,8 +397,13 @@ def test_model_path_header_is_not_forwarded() -> None: @pytest.mark.asyncio @pytest.mark.parametrize("path", ["v1/chat/completions", "v1/responses"]) -@pytest.mark.parametrize("status_code", [200, 429, 502]) -async def test_cashu_pin_reaches_http_transport(path: str, status_code: int) -> None: +@pytest.mark.parametrize( + "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 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( 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() assert len(sent) == 1 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( 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() fallback.forward_request.assert_not_awaited() @@ -514,8 +532,9 @@ async def test_unsupported_endpoint_pins_fail_before_payment( patch.object( 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 json.loads(response.body)["error"]["type"] == "unsupported_request" payment.assert_not_called() @@ -545,7 +564,8 @@ async def test_ehbp_pin_does_not_fall_back(cashu: bool) -> None: response = await _run_proxy( 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() assert forward.await_args is not None 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 selected.forward_request.assert_awaited_once() 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 diff --git a/tests/unit/test_openai_output_cap.py b/tests/unit/test_openai_output_cap.py new file mode 100644 index 00000000..84c7e234 --- /dev/null +++ b/tests/unit/test_openai_output_cap.py @@ -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 diff --git a/tests/unit/test_openrouter_models_fetch_retry.py b/tests/unit/test_openrouter_models_fetch_retry.py new file mode 100644 index 00000000..862bfffe --- /dev/null +++ b/tests/unit/test_openrouter_models_fetch_retry.py @@ -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"] diff --git a/tests/unit/test_payment_settlement_timing.py b/tests/unit/test_payment_settlement_timing.py new file mode 100644 index 00000000..4893fc7d --- /dev/null +++ b/tests/unit/test_payment_settlement_timing.py @@ -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 diff --git a/tests/unit/test_payout_liability_bounds.py b/tests/unit/test_payout_liability_bounds.py new file mode 100644 index 00000000..2fb8ae0d --- /dev/null +++ b/tests/unit/test_payout_liability_bounds.py @@ -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 diff --git a/tests/unit/test_payout_limits.py b/tests/unit/test_payout_limits.py new file mode 100644 index 00000000..87403074 --- /dev/null +++ b/tests/unit/test_payout_limits.py @@ -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() diff --git a/tests/unit/test_periodic_payout.py b/tests/unit/test_periodic_payout.py index fb105cd5..4286bb40 100644 --- a/tests/unit/test_periodic_payout.py +++ b/tests/unit/test_periodic_payout.py @@ -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). """ -from collections.abc import Callable, Coroutine +from collections.abc import Callable, Coroutine, Iterator from contextlib import asynccontextmanager +from pathlib import Path from typing import Any -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import ANY, AsyncMock, MagicMock, patch 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 # exactly one full cycle. @@ -53,11 +77,20 @@ def _one_cycle_sleep() -> Callable[[float], Coroutine[Any, Any, None]]: @pytest.mark.asyncio async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> None: - """primary_mint absent from cashu_mints is still paid out.""" + """primary_mint absent from cashu_mints is paid out and recorded.""" from routstr.core.settings import settings 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 ( 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", 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), ): 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} assert processed == {"http://primary:3338"} 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 @@ -142,6 +195,10 @@ async def test_periodic_payout_releases_session_before_slow_mint_send() -> None: "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", AsyncMock(side_effect=raw_send)), ): 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.""" from routstr.core.settings import settings - async def _get_wallet( - mint_url: str, unit: str, force_reload: bool = False - ) -> MagicMock: + async def _get_wallet(mint_url: str, unit: str, **_: object) -> MagicMock: if mint_url == "http://bad:3338": raise RuntimeError("mint unreachable") return MagicMock() @@ -190,6 +245,10 @@ async def test_periodic_payout_isolates_failing_mint() -> None: "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): @@ -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 # still reached and paid out for both units — failures are isolated. - good_calls = [ - c for c in get_wallet.await_args_list if c.args[0] == "http://good:3338" + good_reloads = [ + 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 @@ -219,6 +281,7 @@ async def test_periodic_payout_handles_session_creation_failure() -> None: patch.object(settings, "payout_interval_seconds", _INTERVAL), patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()), patch("routstr.wallet.db.create_session", create_session), + patch.object(settings, "min_payout_sat", 10), patch( "routstr.wallet._get_supported_mint_units", AsyncMock(return_value=["sat", "msat"]), @@ -237,11 +300,12 @@ async def test_periodic_payout_handles_session_creation_failure() -> None: with pytest.raises(_LoopBreak): await periodic_payout() - # The liability session is opened per mint/unit (sat + msat), and each - # DB failure retains the cycle-specific alert wording while remaining - # isolated to its own iteration. - assert create_session.call_count == 2 - assert logger.error.call_count == 2 + # Per mint/unit (sat + msat) a session is opened twice: once by the stale + # payout-history sweep and once for the liability read. Each DB failure is + # logged and isolated to its own step; the liability error keeps the + # cycle-specific alert wording. + assert create_session.call_count == 4 + assert logger.error.call_count == 4 message = logger.error.call_args.args[0] extra = logger.error.call_args.kwargs["extra"] 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"]), ): 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}), + ] diff --git a/tests/unit/test_pre_handoff_stream_ownership.py b/tests/unit/test_pre_handoff_stream_ownership.py new file mode 100644 index 00000000..894798e0 --- /dev/null +++ b/tests/unit/test_pre_handoff_stream_ownership.py @@ -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() diff --git a/tests/unit/test_provider_field_injection.py b/tests/unit/test_provider_field_injection.py index bf2813e0..6caea063 100644 --- a/tests/unit/test_provider_field_injection.py +++ b/tests/unit/test_provider_field_injection.py @@ -1,5 +1,8 @@ +from unittest.mock import patch + from routstr.upstream.anthropic import AnthropicUpstreamProvider from routstr.upstream.base import BaseUpstreamProvider +from routstr.upstream.generic import GenericUpstreamProvider 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: - """If OpenRouter omits the provider field, the real serving provider is - unknown — a bare ``openrouter`` value carries no information.""" + """If OpenRouter omits the provider field, the serving provider is + unknown but the router is not.""" p = _make_provider(OpenRouterUpstreamProvider, "openrouter") data: dict = {"id": "gen-abc"} p._apply_provider_field(data) - assert data["provider"] == "unknown" + assert data["provider"] == "openrouter:unknown" 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") data: dict = {"provider": "openrouter"} 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: @@ -78,14 +109,41 @@ def test_apply_provider_field_blank_upstream_treated_as_missing() -> None: p = _make_provider(OpenRouterUpstreamProvider, "openrouter") data: dict = {"provider": " "} 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: p = _make_provider(OpenRouterUpstreamProvider, "openrouter") data: dict = {"provider": 42} 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: @@ -127,3 +185,53 @@ def test_inject_cost_metadata_sets_provider() -> None: p.inject_cost_metadata(response_json, cost_data, key) 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:``.""" + 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" diff --git a/tests/unit/test_proxy_path_allowlist.py b/tests/unit/test_proxy_path_allowlist.py index 4cd1c67f..1019dbad 100644 --- a/tests/unit/test_proxy_path_allowlist.py +++ b/tests/unit/test_proxy_path_allowlist.py @@ -51,6 +51,8 @@ def test_ambiguous_paths_are_rejected(path: str) -> None: "v1/chat/completions", "chat/completions", "v1/responses", + "v1/messages", + "v1/messages/count_tokens", "v1/embeddings", "models", "v1/models/gpt-4", @@ -138,6 +140,7 @@ def test_known_prefix_does_not_carry_an_unknown_endpoint(path: str) -> None: ("completions", "POST"), ("v1/responses", "POST"), ("v1/messages", "POST"), + ("v1/messages/count_tokens", "POST"), ("v1/embeddings", "POST"), ("models", "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 +@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: parsed = _parse_extra_allowed_endpoints("POST:v1/rerank, GET:batches ,post:audio/x") assert parsed == { diff --git a/tests/unit/test_proxy_session_lifecycle.py b/tests/unit/test_proxy_session_lifecycle.py index 5d0416d5..7cc01b03 100644 --- a/tests/unit/test_proxy_session_lifecycle.py +++ b/tests/unit/test_proxy_session_lifecycle.py @@ -6,6 +6,8 @@ from fastapi.responses import StreamingResponse from routstr import proxy as proxy_module +from .proxy_test_utils import mock_request_stream, patch_proxy_session + @pytest.mark.asyncio 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.url.path = "/not-an-api-route" request.state.request_id = "test-request" + mock_request_stream(request, b"") 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 session.close.assert_awaited_once() @@ -26,6 +30,8 @@ async def test_proxy_closes_request_session_before_returning_response() -> None: @pytest.mark.asyncio async def test_proxy_session_is_closed_before_first_stream_chunk() -> None: request = MagicMock() + request.headers = {} + mock_request_stream(request, b"") session = AsyncMock() async def stream() -> AsyncIterator[bytes]: @@ -33,10 +39,11 @@ async def test_proxy_session_is_closed_before_first_stream_chunk() -> None: yield b"chunk" upstream_response = StreamingResponse(stream()) - with patch("routstr.proxy._proxy", AsyncMock(return_value=upstream_response)): - response = await proxy_module.proxy( - request, "v1/chat/completions", session=session - ) + with ( + patch("routstr.proxy._proxy", AsyncMock(return_value=upstream_response)), + patch_proxy_session(session), + ): + response = await proxy_module.proxy(request, "v1/chat/completions") assert isinstance(response, StreamingResponse) chunks = [chunk async for chunk in response.body_iterator] diff --git a/tests/unit/test_proxy_tinfoil_attestation_routing.py b/tests/unit/test_proxy_tinfoil_attestation_routing.py index c367a049..7480d3ac 100644 --- a/tests/unit/test_proxy_tinfoil_attestation_routing.py +++ b/tests/unit/test_proxy_tinfoil_attestation_routing.py @@ -1,13 +1,21 @@ from __future__ import annotations +import json from unittest.mock import AsyncMock, MagicMock +import httpx import pytest from fastapi import FastAPI from fastapi.responses import Response from httpx import ASGITransport, AsyncClient 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 @@ -150,3 +158,141 @@ def test_attestation_upstream_selection_is_tinfoil_only() -> None: assert proxy_module._select_unauthenticated_get_upstreams( "attestationjunk", [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" diff --git a/tests/unit/test_queued_logging.py b/tests/unit/test_queued_logging.py new file mode 100644 index 00000000..138bce35 --- /dev/null +++ b/tests/unit/test_queued_logging.py @@ -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) diff --git a/tests/unit/test_request_correction.py b/tests/unit/test_request_correction.py index 3b1110a2..56a3423c 100644 --- a/tests/unit/test_request_correction.py +++ b/tests/unit/test_request_correction.py @@ -16,9 +16,15 @@ from routstr.upstream.request_correction import ( Correction, correct_request, extract_error_message, + rename_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: return json.dumps(kwargs).encode() @@ -139,6 +145,190 @@ class TestStripUnsupportedParam: 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: def test_extracts_nested_error_message(self) -> None: resp = Response( diff --git a/tests/unit/test_request_lifecycle.py b/tests/unit/test_request_lifecycle.py new file mode 100644 index 00000000..1573f074 --- /dev/null +++ b/tests/unit/test_request_lifecycle.py @@ -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() diff --git a/tests/unit/test_request_stage_timing.py b/tests/unit/test_request_stage_timing.py new file mode 100644 index 00000000..38845c3a --- /dev/null +++ b/tests/unit/test_request_stage_timing.py @@ -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 diff --git a/tests/unit/test_settings.py b/tests/unit/test_settings.py index 23834665..61f657ac 100644 --- a/tests/unit/test_settings.py +++ b/tests/unit/test_settings.py @@ -1,5 +1,6 @@ import json import os +from pathlib import Path import pytest from pydantic.v1 import ValidationError @@ -7,7 +8,7 @@ from sqlalchemy.ext.asyncio import create_async_engine from sqlmodel import text 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 @@ -72,6 +73,20 @@ def test_database_pool_defaults_provide_concurrency_headroom() -> None: 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( ("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. await session.exec( # type: ignore - text( - "UPDATE settings SET data = :data WHERE id = 1" - ).bindparams( + text("UPDATE settings SET data = :data WHERE id = 1").bindparams( 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 session.exec( # type: ignore 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() diff --git a/tests/unit/test_sse_splitter.py b/tests/unit/test_sse_splitter.py new file mode 100644 index 00000000..0952500a --- /dev/null +++ b/tests/unit/test_sse_splitter.py @@ -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"" diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py index 498b57f5..61a3da5a 100644 --- a/tests/unit/test_stale_reservations.py +++ b/tests/unit/test_stale_reservations.py @@ -9,6 +9,7 @@ Covers: """ import asyncio +import math import time from typing import AsyncGenerator from unittest.mock import AsyncMock, MagicMock, patch @@ -16,17 +17,21 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from sqlalchemy.pool import StaticPool -from sqlmodel import SQLModel +from sqlmodel import SQLModel, select from sqlmodel.ext.asyncio.session import AsyncSession +import routstr.auth as auth_module from routstr.auth import pay_for_request from routstr.balance import refund_wallet_endpoint from routstr.core.db import ( ApiKey, + ReservationRelease, release_stale_reservations, reset_all_reserved_balances, ) +from .proxy_test_utils import mock_request_stream, patch_proxy_session + def _make_engine() -> AsyncEngine: return create_async_engine( @@ -55,11 +60,16 @@ async def session() -> "AsyncGenerator[AsyncSession, None]": @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) session.add(key) 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()) 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_at is not None 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 @@ -328,7 +417,7 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None: request = MagicMock() request.method = "POST" 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.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, "get_bearer_token_key", AsyncMock(return_value=key)), - patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)), patch.object( proxy_module, - "get_reservation_snapshot", + "pay_for_request", AsyncMock(return_value=reservation_snapshot), ), patch.object(proxy_module, "revert_pay_for_request", revert_mock), + patch_proxy_session(session), ): 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) + + +@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) diff --git a/tests/unit/test_stream_id_injection.py b/tests/unit/test_stream_id_injection.py index 2d682bc5..29912a4f 100644 --- a/tests/unit/test_stream_id_injection.py +++ b/tests/unit/test_stream_id_injection.py @@ -42,8 +42,6 @@ async def test_stream_with_id_injection() -> None: key.hashed_key = "test_hash" key.balance = 1000 - background_tasks = MagicMock() - # We need to mock adjust_payment_for_tokens since it's called at the end with MagicMock(): from routstr.upstream import base @@ -66,7 +64,6 @@ async def test_stream_with_id_injection() -> None: response=mock_response, key=key, max_cost_for_model=100, - background_tasks=background_tasks, requested_model="test-model", reservation_snapshot=ReservationSnapshot( release_id="test-release", diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index e70804c1..eef0c44f 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -6,13 +6,13 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest -from fastapi import BackgroundTasks from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from sqlmodel import SQLModel from sqlmodel.ext.asyncio.session import AsyncSession import routstr.auth as auth_module +import routstr.upstream.gemini_messages as gemini_messages from routstr.auth import ( ReservationSnapshot, adjust_payment_for_tokens, @@ -160,7 +160,7 @@ async def test_post_commit_failure_cannot_release_charged_reservation() -> None: @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() provider = BaseUpstreamProvider( 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() +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 -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( 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) reservation_snapshot = MagicMock() reservation_snapshot.reserved_msats = 500 - background_tasks = MagicMock() with ( patch( @@ -255,7 +531,6 @@ async def test_streaming_release_is_terminal_and_suppresses_background_charge() response=upstream_response, key=key, max_cost_for_model=500, - background_tasks=background_tasks, ) 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() release.assert_awaited_once_with(reservation_snapshot, session, 500) - background_tasks.add_task.assert_not_called() @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.aclose = AsyncMock() - client = MagicMock() - client.aclose = AsyncMock() key = MagicMock(spec=ApiKey) key.hashed_key = f"{api}-partial" key.balance = 10_000 @@ -391,9 +663,7 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once( response=upstream_response, key=key, max_cost_for_model=500, - background_tasks=BackgroundTasks(), reservation_snapshot=snapshot, - client=client, ) else: response = await provider.handle_streaming_responses_completion( @@ -401,14 +671,10 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once( key=key, max_cost_for_model=500, reservation_snapshot=snapshot, - client=client, ) emitted = bytearray() - with pytest.raises(httpx.RemoteProtocolError): - async for chunk in response.body_iterator: - emitted.extend( - chunk.encode() if isinstance(chunk, str) else bytes(chunk) - ) + async for chunk in response.body_iterator: + emitted.extend(chunk.encode() if isinstance(chunk, str) else bytes(chunk)) adjust.assert_awaited_once() if finalization_fails: @@ -417,13 +683,12 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once( else: release.assert_not_awaited() upstream_response.aclose.assert_awaited_once() - client.aclose.assert_awaited_once() assert b"[DONE]" not in emitted @pytest.mark.asyncio @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, ) -> None: 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.aclose = AsyncMock() - client = MagicMock() - client.aclose = AsyncMock() key = MagicMock(spec=ApiKey) key.hashed_key = f"{api}-database-down" key.balance = 10_000 @@ -464,9 +727,7 @@ async def test_partial_stream_preserves_transport_error_when_billing_db_is_down( response=upstream_response, key=key, max_cost_for_model=500, - background_tasks=BackgroundTasks(), reservation_snapshot=snapshot, - client=client, ) else: 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, max_cost_for_model=500, reservation_snapshot=snapshot, - client=client, ) - with pytest.raises(httpx.RemoteProtocolError, match="incomplete chunked read"): - async for _ in response.body_iterator: - pass + async for _ in response.body_iterator: + pass upstream_response.aclose.assert_awaited_once() - client.aclose.assert_awaited_once() @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) +@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 async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() -> None: 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"}]} ).encode() - background_tasks = BackgroundTasks() try: with ( patch( @@ -746,7 +1097,6 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat() response=upstream_response, key=key, max_cost_for_model=500, - background_tasks=background_tasks, model_obj=model, reservation_snapshot=snapshot, 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) await iterator.__anext__() # first chunk reaches the client 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: await auth_module._stop_reservation_heartbeat(snapshot.release_id) diff --git a/tests/unit/test_streaming_sse_providers.py b/tests/unit/test_streaming_sse_providers.py index ffb7e266..afbf22ba 100644 --- a/tests/unit/test_streaming_sse_providers.py +++ b/tests/unit/test_streaming_sse_providers.py @@ -42,7 +42,9 @@ def _make_response(chunks: list[bytes]) -> MagicMock: 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.""" provider = BaseUpstreamProvider( 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), key=key, max_cost_for_model=100, - background_tasks=MagicMock(), requested_model=requested_model, reservation_snapshot=ReservationSnapshot( release_id="test-release", diff --git a/tests/unit/test_tinfoil_integration.py b/tests/unit/test_tinfoil_integration.py index 36e76765..14715f58 100644 --- a/tests/unit/test_tinfoil_integration.py +++ b/tests/unit/test_tinfoil_integration.py @@ -31,6 +31,8 @@ from routstr.upstream.tinfoil import ( ) from routstr.upstream.tinfoil_trailer import TrailerResponse +from .proxy_test_utils import patch_proxy_session + # --------------------------------------------------------------------------- # 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, "get_bearer_token_key", AsyncMock(return_value=key)), - patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)), patch.object( proxy_module, - "get_reservation_snapshot", + "pay_for_request", AsyncMock(return_value=reservation_snapshot), ), 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", AsyncMock(return_value=upstream_resp), ), + patch_proxy_session(session), ): - response = await proxy_module.proxy( - request, "v1/chat/completions", session=session - ) + response = await proxy_module.proxy(request, "v1/chat/completions") # The reservation was released despite the early passthrough return. revert_mock.assert_awaited_once_with(key, session, 1_000, reservation_snapshot) diff --git a/tests/unit/test_tinfoil_trailer.py b/tests/unit/test_tinfoil_trailer.py index 3e4d3e0f..12107aca 100644 --- a/tests/unit/test_tinfoil_trailer.py +++ b/tests/unit/test_tinfoil_trailer.py @@ -1,11 +1,13 @@ from __future__ import annotations import asyncio +import ssl from unittest.mock import AsyncMock, MagicMock 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 @@ -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 async def test_forward_with_trailer_read_timeout_raises_ehbp_timeout( 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: exc = EhbpTimeoutError("boom") - assert exc.status_code == 504 + assert exc.status_code == UPSTREAM_ERROR_STATUS assert exc.code == "UPSTREAM_TIMEOUT" assert exc.details is None + assert exc.scope == ERROR_SCOPE_UPSTREAM 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.""" exc = EhbpTimeoutError("boom", 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" + + +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" diff --git a/tests/unit/test_upstream_deepseek.py b/tests/unit/test_upstream_deepseek.py new file mode 100644 index 00000000..96ab14f9 --- /dev/null +++ b/tests/unit/test_upstream_deepseek.py @@ -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" diff --git a/tests/unit/test_upstream_error_response.py b/tests/unit/test_upstream_error_response.py index 62b9a32c..ff2be223 100644 --- a/tests/unit/test_upstream_error_response.py +++ b/tests/unit/test_upstream_error_response.py @@ -14,7 +14,19 @@ from unittest.mock import Mock import httpx 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.rate_limit import UPSTREAM_RATE_LIMIT 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 ) - 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" payload = json.loads(bytes(response.body)) assert payload["error"]["message"] == "Service Unavailable" + assert payload["error"]["code"] == UPSTREAM_UNAVAILABLE + assert payload["error"]["upstream_status"] == 503 @pytest.mark.asyncio @@ -123,10 +138,13 @@ async def test_empty_body_with_non_json_content_type_normalizes( _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" payload = json.loads(bytes(response.body)) 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 @@ -148,3 +166,211 @@ async def test_json_error_body_is_passed_through_unchanged( assert response.status_code == 400 assert bytes(response.body) == json_body 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"bad gateway", 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 + ) diff --git a/tests/unit/test_upstream_gemini.py b/tests/unit/test_upstream_gemini.py index 13683887..b9cf0b73 100644 --- a/tests/unit/test_upstream_gemini.py +++ b/tests/unit/test_upstream_gemini.py @@ -14,15 +14,21 @@ These tests cover the two pure helpers that drive the dispatcher: from __future__ import annotations +import asyncio import json from collections.abc import AsyncGenerator from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest +import routstr.upstream.gemini_messages as gemini_messages +from routstr.core.exceptions import UpstreamError from routstr.upstream.gemini_messages import ( DUMMY_THOUGHT_SIGNATURE, _openai_chunks_to_anthropic_events, + _ResponseOwnedIterator, inject_thought_signatures, ) @@ -81,9 +87,7 @@ def test_inject_thought_signatures_preserves_existing_signature() -> None: inject_thought_signatures(messages) assert ( - messages[0]["tool_calls"][0]["extra_content"]["google"][ - "thought_signature" - ] + messages[0]["tool_calls"][0]["extra_content"]["google"]["thought_signature"] == "real-signature" ) @@ -138,6 +142,46 @@ async def _lines(*chunks: dict | str) -> AsyncGenerator[str, None]: yield c +class _TrackingStream(httpx.AsyncByteStream): + def __init__( + self, + *chunks: bytes, + error: Exception | None = None, + started: asyncio.Event | None = None, + ) -> None: + self._chunks = chunks + self._error = error + self._started = started + self.close_count = 0 + + async def __aiter__(self) -> AsyncGenerator[bytes, None]: + if self._started is not None: + self._started.set() + await asyncio.Event().wait() + for chunk in self._chunks: + yield chunk + if self._error is not None: + raise self._error + + async def aclose(self) -> None: + self.close_count += 1 + + +def _owned_events( + response: httpx.Response, +) -> _ResponseOwnedIterator: + async def line_iter() -> AsyncGenerator[str, None]: + try: + async for line in response.aiter_lines(): + yield line + finally: + await response.aclose() + + return _ResponseOwnedIterator( + _openai_chunks_to_anthropic_events(line_iter(), "gemini-test"), response + ) + + def _parse_anthropic_sse(blocks: list[bytes]) -> list[dict]: """Flatten a list of Anthropic SSE byte chunks into event dicts.""" events: list[dict] = [] @@ -150,6 +194,58 @@ def _parse_anthropic_sse(blocks: list[bytes]) -> list[dict]: return events +@pytest.mark.asyncio +async def test_response_owner_closes_once_after_normal_completion() -> None: + stream = _TrackingStream( + b'data: {"model":"gemini-test","choices":[{"delta":{"content":"ok"},"finish_reason":"stop"}]}\n\n' + ) + response = httpx.Response( + 200, + request=httpx.Request("POST", "https://gemini.example/chat/completions"), + stream=stream, + ) + + assert [event async for event in _owned_events(response)] + assert response.is_closed + assert stream.close_count == 1 + + +@pytest.mark.asyncio +async def test_response_owner_closes_once_after_body_failure() -> None: + stream = _TrackingStream(error=RuntimeError("upstream body failed")) + response = httpx.Response( + 200, + request=httpx.Request("POST", "https://gemini.example/chat/completions"), + stream=stream, + ) + + with pytest.raises(RuntimeError, match="upstream body failed"): + await _owned_events(response).__anext__() + + assert response.is_closed + assert stream.close_count == 1 + + +@pytest.mark.asyncio +async def test_response_owner_closes_once_after_cancellation() -> None: + started = asyncio.Event() + stream = _TrackingStream(started=started) + response = httpx.Response( + 200, + request=httpx.Request("POST", "https://gemini.example/chat/completions"), + stream=stream, + ) + task = asyncio.create_task(_owned_events(response).__anext__()) + await started.wait() + + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert response.is_closed + assert stream.close_count == 1 + + @pytest.mark.asyncio async def test_translator_emits_text_only_response() -> None: """Plain text response: message_start → content_block_* (text) → @@ -187,9 +283,7 @@ async def test_translator_emits_text_only_response() -> None: ] # Text deltas concatenate to "Hello, world". text_deltas = [ - e["delta"]["text"] - for e in events - if e["type"] == "content_block_delta" + e["delta"]["text"] for e in events if e["type"] == "content_block_delta" ] assert "".join(text_deltas) == "Hello, world" # Stop reason was mapped from openai's "stop". @@ -270,9 +364,7 @@ async def test_translator_emits_tool_use_block() -> None: # Argument deltas were forwarded as input_json_delta partials. deltas = [e for e in events if e["type"] == "content_block_delta"] assert all(d["delta"]["type"] == "input_json_delta" for d in deltas) - assert "".join(d["delta"]["partial_json"] for d in deltas) == ( - '{"cmd": "ls"}' - ) + assert "".join(d["delta"]["partial_json"] for d in deltas) == ('{"cmd": "ls"}') # tool_calls finish_reason → tool_use stop_reason. msg_delta = next(e for e in events if e["type"] == "message_delta") assert msg_delta["delta"]["stop_reason"] == "tool_use" @@ -306,8 +398,36 @@ async def test_translator_handles_done_sentinel_and_blank_lines() -> None: assert events[0]["type"] == "message_start" assert events[-1]["type"] == "message_stop" text = "".join( - e["delta"]["text"] - for e in events - if e["type"] == "content_block_delta" + e["delta"]["text"] for e in events if e["type"] == "content_block_delta" ) assert text == "ok" + + +@pytest.mark.asyncio +async def test_post_and_stream_maps_pool_timeout_to_503() -> None: + client = MagicMock() + client.timeout = httpx.Timeout(10.0) + client.build_request = MagicMock(return_value=MagicMock()) + client.send = AsyncMock(side_effect=httpx.PoolTimeout("pool busy")) + with patch( + "routstr.upstream.gemini_messages.acquire_upstream_http_client", + return_value=client, + ): + with pytest.raises(UpstreamError) as exc_info: + await gemini_messages._post_and_stream( + "https://gemini.example", "key", {"model": "m"}, None + ) + assert exc_info.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_post_and_stream_surfaces_shutdown_as_503() -> None: + with patch( + "routstr.upstream.gemini_messages.acquire_upstream_http_client", + side_effect=UpstreamError("shutting down", status_code=503), + ): + with pytest.raises(UpstreamError) as exc_info: + await gemini_messages._post_and_stream( + "https://gemini.example", "key", {"model": "m"}, None + ) + assert exc_info.value.status_code == 503 diff --git a/tests/unit/test_upstream_http_client.py b/tests/unit/test_upstream_http_client.py new file mode 100644 index 00000000..3de90aa4 --- /dev/null +++ b/tests/unit/test_upstream_http_client.py @@ -0,0 +1,751 @@ +import asyncio +import concurrent.futures +import threading +from collections.abc import Callable +from typing import Any, cast +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +import routstr.upstream.http_client as http_client_module +from routstr.core.exceptions import UpstreamError +from routstr.core.settings import settings +from routstr.upstream.http_client import ( + acquire_upstream_http_client, + build_x_cashu_client, + close_upstream_http_client, + get_upstream_http_client, + upstream_origin_key, +) + + +@pytest.mark.asyncio +async def test_x_cashu_client_reuses_the_process_ssl_context() -> None: + """A per-request client must not reload the CA bundle on every call.""" + pooled = get_upstream_http_client("https://api.example.com/v1/chat") + owned = build_x_cashu_client() + try: + pooled_transport = cast(Any, pooled)._transport + owned_transport = cast(Any, owned)._transport + assert owned_transport._pool._ssl_context is pooled_transport._pool._ssl_context + finally: + await owned.aclose() + await close_upstream_http_client() + + +@pytest.mark.asyncio +async def test_upstream_http_client_is_reused_until_shutdown() -> None: + first = get_upstream_http_client("https://api.example.com/v1/chat") + second = get_upstream_http_client("https://api.example.com/v1/models") + + assert second is first + assert not first.is_closed + + await close_upstream_http_client() + assert first.is_closed + + replacement = get_upstream_http_client("https://api.example.com/v1/chat") + try: + assert replacement is not first + assert not replacement.is_closed + finally: + await close_upstream_http_client() + + +@pytest.mark.asyncio +async def test_upstream_http_client_is_isolated_per_origin() -> None: + try: + first = get_upstream_http_client("https://one.example.com/v1/chat") + second = get_upstream_http_client("https://two.example.com/v1/chat") + other_port = get_upstream_http_client("https://one.example.com:8443/v1/chat") + + assert first is not second + assert first is not other_port + finally: + await close_upstream_http_client() + + +@pytest.mark.parametrize( + ("url", "expected"), + [ + ("https://api.example.com/v1/chat?x=1", "https://api.example.com"), + ("HTTPS://API.EXAMPLE.COM:443/v1/chat", "https://api.example.com"), + ("http://API.EXAMPLE.COM:80/v1/chat", "http://api.example.com"), + ("http://api.example.com:8080/v1/chat", "http://api.example.com:8080"), + ("https://bücher.example/v1/chat", "https://xn--bcher-kva.example"), + ("https://xn--bcher-kva.example/v1/chat", "https://xn--bcher-kva.example"), + ("https://[2001:db8::1]/v1/chat", "https://[2001:db8::1]"), + ( + "https://[2001:0DB8:0:0:0:0:0:1]:443/v1/chat", + "https://[2001:db8::1]", + ), + ("https://[2001:db8::1]:8443/v1/chat", "https://[2001:db8::1]:8443"), + ], +) +def test_upstream_origin_key_returns_http_origin(url: str, expected: str) -> None: + assert upstream_origin_key(url) == expected + + +@pytest.mark.parametrize( + "url", + [ + "", + "/v1/chat", + "ftp://api.example.com", + "https://:443", + "https://user@", + "https://example.com:", + "https://example.com:not-a-port", + "https://example.com:65536", + "https://[2001:db8::1", + "https://exa mple.com", + "https://user@example.com", + "https://user:secret@example.com", + "https://:secret@example.com", + "https://@example.com", + "https://exa\u200bmple.com", + None, + ], +) +def test_upstream_origin_key_rejects_invalid_urls(url: object) -> None: + with pytest.raises(ValueError, match="absolute HTTP") as exc_info: + upstream_origin_key(url) # type: ignore[arg-type] + assert "secret" not in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_acquire_maps_invalid_provider_url_to_502() -> None: + with pytest.raises(UpstreamError) as exc_info: + acquire_upstream_http_client("ftp://api.example.com") + assert exc_info.value.status_code == 502 + + +@pytest.mark.asyncio +async def test_acquire_maps_shutdown_to_503() -> None: + with patch.object( + http_client_module, + "get_upstream_http_client", + side_effect=RuntimeError("Upstream HTTP client is shutting down"), + ): + with pytest.raises(UpstreamError) as exc_info: + acquire_upstream_http_client("https://api.example.com") + assert exc_info.value.status_code == 503 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("first_url", "second_url"), + [ + ("https://EXAMPLE.com:443/v1/chat", "https://example.com/v1/models"), + ( + "https://bücher.example/v1/chat", + "https://xn--bcher-kva.example/v1/models", + ), + ], +) +async def test_equivalent_origins_share_one_client( + first_url: str, second_url: str +) -> None: + try: + first = get_upstream_http_client(first_url) + second = get_upstream_http_client(second_url) + assert second is first + finally: + await close_upstream_http_client() + + +@pytest.mark.asyncio +async def test_upstream_http_client_applies_configured_pool_bounds() -> None: + with ( + patch.object( + http_client_module.httpx, + "Limits", + wraps=httpx.Limits, + ) as build_limits, + patch.object( + http_client_module.httpx, + "AsyncHTTPTransport", + wraps=httpx.AsyncHTTPTransport, + ) as build_transport, + ): + client = get_upstream_http_client("https://api.example.com") + + try: + assert client.timeout.pool == settings.upstream_pool_timeout + assert client.timeout.read == settings.upstream_read_timeout + assert client.timeout.connect == http_client_module.UPSTREAM_CONNECT_TIMEOUT + assert client.timeout.write == http_client_module.UPSTREAM_WRITE_TIMEOUT + build_limits.assert_called_once_with( + max_connections=settings.upstream_max_connections, + max_keepalive_connections=http_client_module.UPSTREAM_MAX_KEEPALIVE_CONNECTIONS, + keepalive_expiry=http_client_module.UPSTREAM_KEEPALIVE_EXPIRY, + ) + build_transport.assert_called_once() + assert ( + build_transport.call_args.kwargs["retries"] + == http_client_module.UPSTREAM_CONNECT_RETRIES + ) + finally: + await close_upstream_http_client() + + +@pytest.mark.asyncio +async def test_upstream_http_client_does_not_share_cookies() -> None: + client = get_upstream_http_client("https://example.com") + try: + first = client.build_request("GET", "https://example.com/test") + response = httpx.Response( + 200, + headers={"set-cookie": "sticky=upstream; Path=/"}, + request=first, + ) + client.cookies.extract_cookies(response) + + later = client.build_request("GET", "https://example.com/test") + explicit = client.build_request( + "GET", "https://example.com/test", headers={"cookie": "user=provided"} + ) + + assert "cookie" not in later.headers + assert explicit.headers["cookie"] == "user=provided" + finally: + await close_upstream_http_client() + + +@pytest.mark.asyncio +async def test_shutdown_closes_foreign_client_on_its_owner_loop( + monkeypatch: pytest.MonkeyPatch, +) -> None: + foreign_loop = asyncio.new_event_loop() + loop_ready = threading.Event() + close_finished = threading.Event() + close_loops: list[asyncio.AbstractEventLoop] = [] + + def run_foreign_loop() -> None: + asyncio.set_event_loop(foreign_loop) + loop_ready.set() + foreign_loop.run_forever() + + thread = threading.Thread(target=run_foreign_loop) + thread.start() + assert loop_ready.wait(timeout=10) + + async def make_client() -> httpx.AsyncClient: + client = get_upstream_http_client("https://example.com") + original_close = client.aclose + + async def tracked_close() -> None: + close_loops.append(asyncio.get_running_loop()) + await original_close() + close_finished.set() + + monkeypatch.setattr(client, "aclose", tracked_close) + return client + + client_future = asyncio.run_coroutine_threadsafe(make_client(), foreign_loop) + client = await asyncio.to_thread(client_future.result, 10) + try: + await close_upstream_http_client() + assert await asyncio.to_thread(close_finished.wait, 10) + assert client.is_closed + assert close_loops == [foreign_loop] + + await close_upstream_http_client() + assert not http_client_module._pending_closes + finally: + foreign_loop.call_soon_threadsafe(foreign_loop.stop) + await asyncio.to_thread(thread.join, 10) + assert not thread.is_alive() + foreign_loop.close() + + +@pytest.mark.asyncio +async def test_shutdown_rehomes_queued_close_when_owner_loop_stops( + monkeypatch: pytest.MonkeyPatch, +) -> None: + foreign_loop = asyncio.new_event_loop() + loop_ready = threading.Event() + blocker_started = threading.Event() + allow_stop = threading.Event() + + def run_foreign_loop() -> None: + asyncio.set_event_loop(foreign_loop) + loop_ready.set() + foreign_loop.run_forever() + + thread = threading.Thread(target=run_foreign_loop) + thread.start() + assert loop_ready.wait(timeout=10) + + async def make_client() -> httpx.AsyncClient: + return get_upstream_http_client("https://example.com") + + client_future = asyncio.run_coroutine_threadsafe(make_client(), foreign_loop) + client = await asyncio.to_thread(client_future.result, 10) + original_close = client.aclose + close_loops: list[asyncio.AbstractEventLoop] = [] + + async def tracked_close() -> None: + close_loops.append(asyncio.get_running_loop()) + await original_close() + + monkeypatch.setattr(client, "aclose", tracked_close) + + def stop_before_next_iteration() -> None: + blocker_started.set() + assert allow_stop.wait(timeout=10) + foreign_loop.stop() + + foreign_loop.call_soon_threadsafe(stop_before_next_iteration) + assert blocker_started.wait(timeout=10) + + try: + closing = asyncio.create_task(close_upstream_http_client()) + while not http_client_module._pending_closes: + await asyncio.sleep(0) + assert not client.is_closed + + allow_stop.set() + await asyncio.to_thread(thread.join, 10) + assert not thread.is_alive() + + await closing + assert client.is_closed + assert close_loops == [asyncio.get_running_loop()] + assert not http_client_module._pending_closes + finally: + allow_stop.set() + if thread.is_alive(): + foreign_loop.call_soon_threadsafe(foreign_loop.stop) + await asyncio.to_thread(thread.join, 10) + foreign_loop.close() + + +@pytest.mark.asyncio +async def test_shutdown_finishes_transport_close_on_stopped_owner_loop( + monkeypatch: pytest.MonkeyPatch, +) -> None: + foreign_loop = asyncio.new_event_loop() + loop_ready = threading.Event() + transport_started = threading.Event() + allow_transport_close = threading.Event() + transport_finished = threading.Event() + + class BlockingTransport(httpx.AsyncBaseTransport): + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + return httpx.Response(200, request=request) + + async def aclose(self) -> None: + transport_started.set() + while not allow_transport_close.is_set(): + await asyncio.sleep(0) + transport_finished.set() + + client = httpx.AsyncClient(transport=BlockingTransport()) + monkeypatch.setattr(http_client_module, "_build_client", lambda: client) + + def run_foreign_loop() -> None: + asyncio.set_event_loop(foreign_loop) + loop_ready.set() + foreign_loop.run_forever() + + thread = threading.Thread(target=run_foreign_loop) + thread.start() + assert loop_ready.wait(timeout=10) + + async def register_client() -> None: + assert get_upstream_http_client("https://example.com") is client + + registered = asyncio.run_coroutine_threadsafe(register_client(), foreign_loop) + await asyncio.to_thread(registered.result, 10) + + try: + closing = asyncio.create_task(close_upstream_http_client()) + assert await asyncio.to_thread(transport_started.wait, 10) + + foreign_loop.call_soon_threadsafe(foreign_loop.stop) + await asyncio.to_thread(thread.join, 10) + assert not thread.is_alive() + + allow_transport_close.set() + await closing + + assert transport_finished.is_set() + assert client.is_closed + assert not http_client_module._pending_closes + finally: + allow_transport_close.set() + if thread.is_alive(): + foreign_loop.call_soon_threadsafe(foreign_loop.stop) + await asyncio.to_thread(thread.join, 10) + if not foreign_loop.is_closed(): + foreign_loop.close() + + +@pytest.mark.asyncio +async def test_started_close_on_closed_owner_loop_retries_transport() -> None: + class CountingTransport(httpx.AsyncBaseTransport): + def __init__(self) -> None: + self.close_count = 0 + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + return httpx.Response(200, request=request) + + async def aclose(self) -> None: + self.close_count += 1 + + owner_loop = MagicMock(spec=asyncio.AbstractEventLoop) + owner_loop.is_closed.return_value = True + owner_loop.is_running.return_value = False + transport = CountingTransport() + client = httpx.AsyncClient(transport=transport) + completion: concurrent.futures.Future[None] = concurrent.futures.Future() + completion.set_running_or_notify_cancel() + task = MagicMock(spec=asyncio.Task) + task.done.return_value = False + submission = http_client_module._CloseSubmission( + client=client, + completion=completion, + task=task, + ) + http_client_module._pending_closes[owner_loop] = {completion: submission} + + http_client_module._collect_completed_closes() + await http_client_module._drain_pending_closes() + + assert submission.retired + assert transport.close_count == 1 + assert client.is_closed + assert not http_client_module._pending_closes + + +@pytest.mark.asyncio +async def test_shutdown_closes_client_after_owner_loop_stopped( + monkeypatch: pytest.MonkeyPatch, +) -> None: + created: list[httpx.AsyncClient] = [] + owner_loops: list[asyncio.AbstractEventLoop] = [] + + def create_on_stopped_loop() -> None: + owner_loop = asyncio.new_event_loop() + asyncio.set_event_loop(owner_loop) + owner_loops.append(owner_loop) + + async def make_client() -> None: + created.append(get_upstream_http_client("https://example.com")) + + owner_loop.run_until_complete(make_client()) + + thread = threading.Thread(target=create_on_stopped_loop) + thread.start() + await asyncio.to_thread(thread.join, 10) + assert not thread.is_alive() + + client = created[0] + owner_loop = owner_loops[0] + close_loops: list[asyncio.AbstractEventLoop] = [] + original_close = client.aclose + + async def tracked_close() -> None: + close_loops.append(asyncio.get_running_loop()) + await original_close() + + monkeypatch.setattr(client, "aclose", tracked_close) + try: + await close_upstream_http_client() + assert client.is_closed + assert close_loops == [asyncio.get_running_loop()] + assert not http_client_module._pending_closes + finally: + owner_loop.close() + + +@pytest.mark.asyncio +async def test_shutdown_closes_client_after_owner_loop_closed() -> None: + created: list[httpx.AsyncClient] = [] + + def create_and_close_loop() -> None: + owner_loop = asyncio.new_event_loop() + asyncio.set_event_loop(owner_loop) + + async def make_client() -> None: + created.append(get_upstream_http_client("https://example.com")) + + owner_loop.run_until_complete(make_client()) + owner_loop.close() + + thread = threading.Thread(target=create_and_close_loop) + thread.start() + await asyncio.to_thread(thread.join, 10) + assert not thread.is_alive() + + client = created[0] + await close_upstream_http_client() + assert client.is_closed + + +@pytest.mark.asyncio +async def test_shutdown_retries_failed_client_close( + monkeypatch: pytest.MonkeyPatch, +) -> None: + client = get_upstream_http_client("https://example.com") + original_close = client.aclose + attempts = 0 + + async def flaky_close() -> None: + nonlocal attempts + attempts += 1 + if attempts == 1: + raise RuntimeError("close failed") + await original_close() + + monkeypatch.setattr(client, "aclose", flaky_close) + + await close_upstream_http_client() + assert not client.is_closed + assert any( + client in failed for failed in http_client_module._failed_closes.values() + ) + + await close_upstream_http_client() + assert client.is_closed + assert attempts == 2 + assert not http_client_module._failed_closes + + +@pytest.mark.asyncio +async def test_shutdown_prunes_externally_closed_failed_client( + monkeypatch: pytest.MonkeyPatch, +) -> None: + client = get_upstream_http_client("https://example.com") + original_close = client.aclose + + async def fail_close() -> None: + raise RuntimeError("close failed") + + monkeypatch.setattr(client, "aclose", fail_close) + await close_upstream_http_client() + assert http_client_module._failed_closes + + await original_close() + await close_upstream_http_client() + + assert client.is_closed + assert not http_client_module._failed_closes + assert not http_client_module._pending_closes + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", [RuntimeError("failed"), asyncio.CancelledError()]) +async def test_shutdown_retries_transport_after_httpx_marks_client_closed( + failure: BaseException, + monkeypatch: pytest.MonkeyPatch, +) -> None: + class FailOnceTransport(httpx.AsyncBaseTransport): + def __init__(self) -> None: + self.attempts = 0 + self.completed = False + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + return httpx.Response(200, request=request) + + async def aclose(self) -> None: + self.attempts += 1 + if self.attempts == 1: + raise failure + self.completed = True + + transport = FailOnceTransport() + client = httpx.AsyncClient(transport=transport) + monkeypatch.setattr(http_client_module, "_build_client", lambda: client) + assert get_upstream_http_client("https://example.com") is client + + await close_upstream_http_client() + assert client.is_closed + assert transport.attempts == 1 + assert not transport.completed + assert any( + client in failed for failed in http_client_module._failed_closes.values() + ) + + await close_upstream_http_client() + assert transport.attempts == 2 + assert transport.completed + assert not http_client_module._failed_closes + assert not http_client_module._pending_closes + + +@pytest.mark.asyncio +async def test_shutdown_collects_done_task_before_owner_loop_callback( + monkeypatch: pytest.MonkeyPatch, +) -> None: + foreign_loop = asyncio.new_event_loop() + loop_ready = threading.Event() + transport_finished = threading.Event() + + class StopAfterCloseTransport(httpx.AsyncBaseTransport): + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + return httpx.Response(200, request=request) + + async def aclose(self) -> None: + transport_finished.set() + asyncio.get_running_loop().stop() + + client = httpx.AsyncClient(transport=StopAfterCloseTransport()) + monkeypatch.setattr(http_client_module, "_build_client", lambda: client) + + def run_foreign_loop() -> None: + asyncio.set_event_loop(foreign_loop) + loop_ready.set() + foreign_loop.run_forever() + + thread = threading.Thread(target=run_foreign_loop) + thread.start() + assert loop_ready.wait(timeout=10) + + async def register_client() -> None: + assert get_upstream_http_client("https://example.com") is client + + registered_client = asyncio.run_coroutine_threadsafe( + register_client(), foreign_loop + ) + await asyncio.to_thread(registered_client.result, 10) + + try: + await asyncio.wait_for(close_upstream_http_client(), timeout=1) + assert transport_finished.is_set() + assert client.is_closed + assert not http_client_module._pending_closes + assert not http_client_module._failed_closes + finally: + if thread.is_alive(): + foreign_loop.call_soon_threadsafe(foreign_loop.stop) + await asyncio.to_thread(thread.join, 10) + foreign_loop.close() + + +@pytest.mark.asyncio +async def test_close_submission_settlement_is_atomic_across_threads( + monkeypatch: pytest.MonkeyPatch, +) -> None: + task = asyncio.create_task(asyncio.sleep(0)) + await task + + client = httpx.AsyncClient() + completion: concurrent.futures.Future[None] = concurrent.futures.Future() + completion.set_running_or_notify_cancel() + submission = http_client_module._CloseSubmission( + client=client, + completion=completion, + task=task, + ) + loop = asyncio.get_running_loop() + http_client_module._pending_closes[loop] = {completion: submission} + + barrier = threading.Barrier(2) + errors: list[BaseException] = [] + original_settle = http_client_module._settle_close_submission + + def synchronized_settle( + close_submission: http_client_module._CloseSubmission, + completed: asyncio.Task[None], + ) -> None: + barrier.wait(timeout=10) + original_settle(close_submission, completed) + + def run(action: Callable[[], None]) -> None: + try: + action() + except BaseException as exc: + errors.append(exc) + + monkeypatch.setattr( + http_client_module, + "_settle_close_submission", + synchronized_settle, + ) + collector = threading.Thread( + target=run, + args=(lambda: http_client_module._settle_submission_from_task(submission),), + ) + callback = threading.Thread( + target=run, + args=(lambda: http_client_module._finish_close_submission(submission, task),), + ) + + try: + collector.start() + callback.start() + collector.join(timeout=10) + callback.join(timeout=10) + + assert not collector.is_alive() + assert not callback.is_alive() + assert errors == [] + assert completion.result() is None + + monkeypatch.setattr( + http_client_module, + "_settle_close_submission", + original_settle, + ) + http_client_module._collect_completed_closes() + + assert http_client_module._close_completed.get(client) is True + assert not http_client_module._pending_closes + assert not http_client_module._failed_closes + finally: + http_client_module._pending_closes.pop(loop, None) + http_client_module._close_completed.pop(client, None) + await client.aclose() + + +@pytest.mark.asyncio +async def test_close_submission_rejects_conflicting_outcomes() -> None: + succeeded = asyncio.create_task(asyncio.sleep(0)) + + async def fail() -> None: + raise RuntimeError("different outcome") + + failed = asyncio.create_task(fail()) + await succeeded + with pytest.raises(RuntimeError, match="different outcome"): + await failed + + client = httpx.AsyncClient() + completion: concurrent.futures.Future[None] = concurrent.futures.Future() + completion.set_running_or_notify_cancel() + submission = http_client_module._CloseSubmission(client, completion) + + try: + http_client_module._settle_close_submission(submission, succeeded) + with pytest.raises(RuntimeError, match="conflicting outcomes"): + http_client_module._settle_close_submission(submission, failed) + finally: + await client.aclose() + + +@pytest.mark.asyncio +async def test_upstream_http_client_cannot_reopen_during_shutdown( + monkeypatch: pytest.MonkeyPatch, +) -> None: + client = get_upstream_http_client("https://example.com") + close_started = asyncio.Event() + allow_close = asyncio.Event() + original_close = client.aclose + + async def delayed_close() -> None: + close_started.set() + await allow_close.wait() + await original_close() + + monkeypatch.setattr(client, "aclose", delayed_close) + closing = asyncio.create_task(close_upstream_http_client()) + await close_started.wait() + + with pytest.raises(RuntimeError, match="shutting down"): + get_upstream_http_client("https://example.com") + + allow_close.set() + await closing diff --git a/tests/unit/test_upstream_rate_limit.py b/tests/unit/test_upstream_rate_limit.py index 495f1e57..0b73199c 100644 --- a/tests/unit/test_upstream_rate_limit.py +++ b/tests/unit/test_upstream_rate_limit.py @@ -15,6 +15,11 @@ from unittest.mock import AsyncMock, MagicMock, Mock, patch import httpx import pytest +from routstr.core.error_scope import ( + ERROR_SCOPE_HEADER, + ERROR_SCOPE_UPSTREAM, + UPSTREAM_UNAVAILABLE, +) from routstr.core.redaction import redact_org_ids from routstr.upstream.base import BaseUpstreamProvider from routstr.upstream.rate_limit import ( @@ -23,6 +28,8 @@ from routstr.upstream.rate_limit import ( classify_rate_limit, ) +from .proxy_test_utils import mock_request_stream, patch_proxy_session + # The exact scenario from the issue, with a realistic (fake) org identifier. RAW_ORG_ID = "org-abc123XYZ456def" RATE_LIMIT_MESSAGE = ( @@ -231,7 +238,8 @@ def test_create_upstream_error_response_preserves_structure() -> None: assert "org-[REDACTED]" in serialized -def test_generic_upstream_error_still_defaults_to_502() -> None: +def test_generic_upstream_error_reports_424() -> None: + """An upstream-attributable failure is reported as 424, not 502.""" from routstr.core.exceptions import UpstreamError from routstr.payment.helpers import create_upstream_error_response @@ -239,11 +247,12 @@ def test_generic_upstream_error_still_defaults_to_502() -> None: response = create_upstream_error_response(err, _make_request()) - assert response.status_code == 502 + assert response.status_code == 424 + 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"] == 502 - assert "details" not in payload["error"] + assert payload["error"]["code"] == UPSTREAM_UNAVAILABLE + assert payload["error"]["details"]["upstream_status"] == 502 # --------------------------------------------------------------------------- # @@ -307,7 +316,8 @@ async def test_5xx_wrapped_rate_limit_is_classified( provider: BaseUpstreamProvider, ) -> None: # Some providers wrap a rate-limit in a 5xx envelope; classification must - # key off the message marker, not only the 429 status. + # key off the message marker, not only the 429 status. The retry hint wins + # over the 424 mapping: a caller must still see a retryable 429. body = json.dumps({"error": {"message": RATE_LIMIT_MESSAGE}}).encode() upstream = _make_upstream_response(body=body, status_code=500) @@ -315,9 +325,11 @@ async def test_5xx_wrapped_rate_limit_is_classified( _make_request(), "v1/chat/completions", upstream ) - assert response.status_code == 500 + assert response.status_code == 429 payload: dict[str, Any] = json.loads(bytes(response.body)) assert payload["error"]["code"] == UPSTREAM_RATE_LIMIT + assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM + assert payload["error"]["upstream_status"] == 500 serialized = json.dumps(payload) assert RAW_ORG_ID not in serialized assert "org-[REDACTED]" in serialized @@ -343,7 +355,7 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None: request = MagicMock() request.method = "POST" request.headers = {"authorization": "Bearer sk-rlkey"} - request.body = AsyncMock(return_value=b'{"model": "test-model"}') + mock_request_stream(request, b'{"model": "test-model"}') request.state = MagicMock() request.state.request_id = "req-rl" @@ -384,17 +396,15 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None: ), patch.object(proxy_module, "check_token_balance", MagicMock()), 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( proxy_module, - "get_reservation_snapshot", + "pay_for_request", AsyncMock(return_value=reservation), ), patch.object(proxy_module, "revert_pay_for_request", revert_mock), + patch_proxy_session(session), ): - response = await proxy_module.proxy( - request, "v1/chat/completions", session=session - ) + response = await proxy_module.proxy(request, "v1/chat/completions") # Original 429 status and the stable code/details survive to the client. assert response.status_code == 429 diff --git a/tests/unit/test_upstream_stream_timeout.py b/tests/unit/test_upstream_stream_timeout.py new file mode 100644 index 00000000..7942337c --- /dev/null +++ b/tests/unit/test_upstream_stream_timeout.py @@ -0,0 +1,506 @@ +"""First-token / idle stream guards and the per-(provider, model) cooldown.""" + +import asyncio +import json +from collections.abc import AsyncIterator +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +from routstr.core.error_scope import ( + ERROR_SCOPE_HEADER, + ERROR_SCOPE_NODE, + ERROR_SCOPE_UPSTREAM, +) +from routstr.core.exceptions import UpstreamError +from routstr.core.settings import Settings, settings +from routstr.upstream.base import BaseUpstreamProvider +from routstr.upstream.cooldown import is_cooling_down, record_failure +from routstr.upstream.stream_timeout import open_guarded_stream + + +def _response(chunks: AsyncIterator[bytes]) -> MagicMock: + response = MagicMock(spec=httpx.Response) + response.aiter_bytes = MagicMock(return_value=chunks) + response.aclose = AsyncMock() + return response + + +async def _never() -> AsyncIterator[bytes]: + await asyncio.sleep(10) + yield b"late" + + +async def _stalls_after_first() -> AsyncIterator[bytes]: + yield b"first" + await asyncio.sleep(10) + yield b"never delivered" + + +async def _heartbeat_only(frame: bytes = b": keepalive\n\n") -> AsyncIterator[bytes]: + while True: + yield frame + await asyncio.sleep(0.002) + + +@pytest.fixture +def fast_timeouts(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(settings, "upstream_first_token_timeout_seconds", 0.01) + monkeypatch.setattr(settings, "upstream_stream_idle_timeout_seconds", 0.01) + + +@pytest.mark.asyncio +async def test_first_token_timeout_closes_response_and_raises( + fast_timeouts: None, +) -> None: + response = _response(_never()) + + with pytest.raises(UpstreamError) as exc_info: + await open_guarded_stream(response, "test") + + assert exc_info.value.code == "UPSTREAM_TIMEOUT" + assert exc_info.value.from_upstream_response is False + response.aclose.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_generic_stream_times_out_before_response_is_handed_off( + fast_timeouts: None, +) -> None: + provider = BaseUpstreamProvider(base_url="https://slow.example", api_key="test") + response = _response(_never()) + + with pytest.raises(UpstreamError, match="no first chunk"): + await provider._generic_streaming_response( + response, "key-hash", 100, "audio/speech", None, None, MagicMock() + ) + + response.aclose.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_generic_stream_idle_abort_settles_without_clean_completion( + fast_timeouts: None, +) -> None: + provider = BaseUpstreamProvider(base_url="https://slow.example", api_key="test") + finalize = AsyncMock() + provider._finalize_generic_streaming_payment = finalize # type: ignore[method-assign] + upstream = _response(_stalls_after_first()) + upstream.status_code = 200 + upstream.headers = {} + response = await provider._generic_streaming_response( + upstream, "key-hash", 100, "audio/speech", None, None, MagicMock() + ) + chunks = [] + with pytest.raises(UpstreamError, match="stream stalled"): + async for chunk in response.body_iterator: + chunks.append(chunk) + + assert chunks == [b"first"] + finalize.assert_awaited_once() + upstream.aclose.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("frame", [b": keepalive\n\n", b"data: \n\n"]) +async def test_sse_heartbeats_do_not_satisfy_first_token_timeout( + fast_timeouts: None, frame: bytes +) -> None: + response = _response(_heartbeat_only(frame)) + with pytest.raises(UpstreamError, match="no first chunk"): + await open_guarded_stream(response, "test", sse=True) + response.aclose.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("frame", [b": keepalive\n\n", b"data: \n\n"]) +async def test_sse_heartbeats_do_not_reset_idle_timeout( + fast_timeouts: None, frame: bytes +) -> None: + async def chunks() -> AsyncIterator[bytes]: + yield b'data: {"delta":"first"}\n\n' + async for chunk in _heartbeat_only(frame): + yield chunk + + failures = MagicMock() + stream = await open_guarded_stream( + _response(chunks()), "test", sse=True, on_idle_timeout=failures + ) + assert [chunk async for chunk in stream] == [b'data: {"delta":"first"}\n\n'] + assert stream.timed_out is True + failures.assert_called_once() + + +@pytest.mark.asyncio +async def test_zero_first_token_timeout_disables_the_guard( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "upstream_first_token_timeout_seconds", 0) + monkeypatch.setattr(settings, "upstream_stream_idle_timeout_seconds", 0) + + async def _slow() -> AsyncIterator[bytes]: + await asyncio.sleep(0.02) + yield b"first" + + stream = await open_guarded_stream(_response(_slow()), "test") + + assert [chunk async for chunk in stream] == [b"first"] + + +def test_stream_guards_are_off_by_default() -> None: + # Reasoning models can think silently for minutes; on by default, the + # guards would fail requests that succeed without them. + fields = Settings.__fields__ + assert fields["upstream_first_token_timeout_seconds"].default == 0 + assert fields["upstream_stream_idle_timeout_seconds"].default == 0 + + +@pytest.mark.asyncio +async def test_idle_timeout_cools_down_the_serving_provider( + fast_timeouts: None, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(settings, "upstream_allowed_fails", 1) + provider = BaseUpstreamProvider(base_url="https://slow.example", api_key="test") + provider.db_id = 17 + model = MagicMock(id="test-model") + guarded = await provider._guard_stream( + _response(_stalls_after_first()), model, sse=False + ) + + assert [chunk async for chunk in guarded] == [b"first"] + assert is_cooling_down("db:17", "test-model") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("terminal_before_stall", [False, True]) +async def test_responses_idle_timeout_does_not_emit_completed( + fast_timeouts: None, terminal_before_stall: bool +) -> None: + async def chunks() -> AsyncIterator[bytes]: + event = ( + b'data: {"type":"response.completed","response":{"model":"test","usage":{"input_tokens":0,"output_tokens":1}}}\n\n' + if terminal_before_stall + else b'data: {"type":"response.created","response":{"model":"test"}}\n\n' + ) + yield event + await asyncio.sleep(10) + + response = _response(chunks()) + response.status_code = 200 + response.headers = {"content-type": "text/event-stream"} + key = MagicMock() + key.hashed_key = "test-key" + key.balance = 1000 + 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) + provider = BaseUpstreamProvider(base_url="https://slow.example", api_key="test") + + with ( + patch("routstr.upstream.base.create_session", return_value=session_context), + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + AsyncMock(return_value={"input_tokens": 0, "output_tokens": 1}), + ), + ): + result = await provider.handle_streaming_responses_completion( + response, key, 100, reservation_snapshot=MagicMock() + ) + emitted = b"".join( + [ + chunk.encode() if isinstance(chunk, str) else bytes(chunk) + async for chunk in result.body_iterator + ] + ) + + assert b'"type": "response.failed"' in emitted + assert b'"code": "UPSTREAM_TIMEOUT"' in emitted + assert b'"type": "response.completed"' not in emitted + response.aclose.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_guarded_stream_passes_every_chunk_through() -> None: + async def _chunks() -> AsyncIterator[bytes]: + yield b"a" + yield b"b" + yield b"c" + + stream = await open_guarded_stream(_response(_chunks()), "test") + + assert [chunk async for chunk in stream] == [b"a", b"b", b"c"] + + +def test_cooldown_opens_after_allowed_fails_and_expires( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "upstream_allowed_fails", 3) + monkeypatch.setattr(settings, "upstream_cooldown_seconds", 30) + + for _ in range(2): + record_failure("https://a.example", "m") + assert is_cooling_down("https://a.example", "m") is False + + record_failure("https://a.example", "m") + assert is_cooling_down("https://a.example", "m") is True + # Scoped to the exact pair. + assert is_cooling_down("https://b.example", "m") is False + assert is_cooling_down("https://a.example", "other") is False + + with patch("routstr.upstream.cooldown.time.monotonic", return_value=1e6): + assert is_cooling_down("https://a.example", "m") is False + + +def test_zero_cooldown_disables_skipping(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(settings, "upstream_allowed_fails", 1) + monkeypatch.setattr(settings, "upstream_cooldown_seconds", 0) + + record_failure("https://a.example", "m") + + assert is_cooling_down("https://a.example", "m") is False + + +def _proxy_request() -> MagicMock: + request = MagicMock() + request.method = "POST" + request.headers = {"authorization": "Bearer sk-key"} + request.body = AsyncMock(return_value=b'{"model": "test-model", "stream": true}') + request.state = MagicMock() + request.state.request_id = "req-1" + return request + + +def _upstream(base_url: str, forward: AsyncMock) -> MagicMock: + upstream = MagicMock() + upstream.provider_type = "test" + upstream.base_url = base_url + upstream.db_id = None + upstream.prepare_headers = MagicMock(side_effect=lambda h: h) + upstream.forward_request = forward + return upstream + + +async def _run_proxy( + candidates: list[tuple[MagicMock, MagicMock]], + revert_mock: AsyncMock, + request: MagicMock | None = None, +) -> Any: + from routstr import proxy as proxy_module + from routstr.auth import ReservationSnapshot + from routstr.core.db import ApiKey + + key = ApiKey(hashed_key="streamkey", balance=10_000) + reservation = ReservationSnapshot( + release_id="release", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=1_000, + ) + + with ( + patch.object(proxy_module, "get_candidates", return_value=candidates), + patch.object( + proxy_module, "get_max_cost_for_model", AsyncMock(return_value=1_000) + ), + patch.object( + proxy_module, + "calculate_discounted_max_cost", + AsyncMock(return_value=1_000), + ), + patch.object(proxy_module, "check_token_balance", MagicMock()), + patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)), + patch.object( + proxy_module, "pay_for_request", AsyncMock(return_value=reservation) + ), + patch.object(proxy_module, "revert_pay_for_request", revert_mock), + ): + request = request or _proxy_request() + return await proxy_module._proxy( + request, "v1/chat/completions", MagicMock(), await request.body() + ) + + +@pytest.mark.asyncio +async def test_first_token_timeout_fails_over_to_the_next_candidate( + fast_timeouts: None, +) -> None: + async def _timing_out(*args: Any, **kwargs: Any) -> Any: + return await open_guarded_stream(_response(_never()), "test") + + served = MagicMock() + served.status_code = 200 + slow = _upstream("https://slow.example", AsyncMock(side_effect=_timing_out)) + fast = _upstream("https://fast.example", AsyncMock(return_value=served)) + revert_mock = AsyncMock(return_value=True) + + response = await _run_proxy([(MagicMock(), slow), (MagicMock(), fast)], revert_mock) + + assert response is served + fast.forward_request.assert_awaited_once() + # The reservation carries over to the candidate that served the request. + revert_mock.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_first_token_timeout_on_last_candidate_reverts_reservation( + fast_timeouts: None, +) -> None: + async def _timing_out(*args: Any, **kwargs: Any) -> Any: + return await open_guarded_stream(_response(_never()), "test") + + slow = _upstream("https://slow.example", AsyncMock(side_effect=_timing_out)) + revert_mock = AsyncMock(return_value=True) + + response = await _run_proxy([(MagicMock(), slow)], revert_mock) + + assert response.status_code == 424 + assert json.loads(bytes(response.body))["error"]["code"] == "UPSTREAM_TIMEOUT" + revert_mock.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_cooling_down_candidate_is_skipped_then_recovers( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "upstream_allowed_fails", 1) + monkeypatch.setattr(settings, "upstream_cooldown_seconds", 30) + + sick_response = MagicMock() + sick_response.status_code = 200 + healthy_response = MagicMock() + healthy_response.status_code = 200 + sick = _upstream("https://sick.example", AsyncMock(return_value=sick_response)) + healthy = _upstream("https://ok.example", AsyncMock(return_value=healthy_response)) + candidates = [(MagicMock(), sick), (MagicMock(), healthy)] + + record_failure("test|https://sick.example", "test-model") + assert await _run_proxy(candidates, AsyncMock()) is healthy_response + sick.forward_request.assert_not_awaited() + + with patch("routstr.upstream.cooldown.time.monotonic", return_value=1e6): + assert await _run_proxy(candidates, AsyncMock()) is sick_response + + +@pytest.mark.asyncio +async def test_cooldown_never_empties_the_candidate_list( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "upstream_allowed_fails", 1) + monkeypatch.setattr(settings, "upstream_cooldown_seconds", 30) + + only_response = MagicMock() + only_response.status_code = 200 + only = _upstream("https://only.example", AsyncMock(return_value=only_response)) + + record_failure("test|https://only.example", "test-model") + + assert await _run_proxy([(MagicMock(), only)], AsyncMock()) is only_response + + +@pytest.mark.asyncio +async def test_cooldown_distinguishes_credentials_at_same_url( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "upstream_allowed_fails", 1) + bad = _upstream("https://same.example", AsyncMock()) + bad.db_id = 1 + good_response = MagicMock(status_code=200) + good = _upstream("https://same.example", AsyncMock(return_value=good_response)) + good.db_id = 2 + other = _upstream("https://other.example", AsyncMock()) + record_failure("db:1", "test-model") + + assert ( + await _run_proxy( + [(MagicMock(), bad), (MagicMock(), good), (MagicMock(), other)], AsyncMock() + ) + is good_response + ) + bad.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_cooldown_normalizes_model_spelling( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "upstream_allowed_fails", 1) + bad = _upstream("https://bad.example", AsyncMock()) + good_response = MagicMock(status_code=200) + good = _upstream("https://good.example", AsyncMock(return_value=good_response)) + record_failure("test|https://bad.example", "test-model") + request = _proxy_request() + request.body = AsyncMock( + return_value=b'{"model":"TEST-MODEL-20251222","stream":true}' + ) + + assert ( + await _run_proxy( + [(MagicMock(id="test-model"), bad), (MagicMock(id="test-model"), good)], + AsyncMock(), + request, + ) + is good_response + ) + bad.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_x_cashu_upstream_failure_opens_cooldown( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "upstream_allowed_fails", 1) + upstream = _upstream("https://cashu.example", AsyncMock()) + upstream.handle_x_cashu = AsyncMock( + return_value=MagicMock( + status_code=503, headers={ERROR_SCOPE_HEADER: ERROR_SCOPE_UPSTREAM} + ) + ) + request = _proxy_request() + request.headers = {"x-cashu": "token"} + + response = await _run_proxy([(MagicMock(), upstream)], AsyncMock(), request) + + assert response.status_code == 503 + assert is_cooling_down("test|https://cashu.example", "test-model") + + +@pytest.mark.asyncio +async def test_x_cashu_local_mint_failure_does_not_cool_provider( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "upstream_allowed_fails", 1) + upstream = _upstream("https://cashu.example", AsyncMock()) + upstream.handle_x_cashu = AsyncMock( + return_value=MagicMock(status_code=503, headers={}) + ) + request = _proxy_request() + request.headers = {"x-cashu": "token"} + + response = await _run_proxy([(MagicMock(), upstream)], AsyncMock(), request) + + assert response.status_code == 503 + assert not is_cooling_down("test|https://cashu.example", "test-model") + + +@pytest.mark.asyncio +async def test_node_scoped_upstream_exception_does_not_cool_provider( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "upstream_allowed_fails", 1) + upstream = _upstream( + "https://healthy.example", + AsyncMock( + side_effect=UpstreamError( + "local fault", status_code=500, scope=ERROR_SCOPE_NODE + ) + ), + ) + + response = await _run_proxy([(MagicMock(), upstream)], AsyncMock()) + + assert response.status_code == 500 + assert not is_cooling_down("test|https://healthy.example", "test-model") diff --git a/tests/unit/test_upstream_venice.py b/tests/unit/test_upstream_venice.py new file mode 100644 index 00000000..742c38ae --- /dev/null +++ b/tests/unit/test_upstream_venice.py @@ -0,0 +1,293 @@ +"""Unit tests for ``VeniceUpstreamProvider.fetch_models``. + +Venice answers ``/models`` with only its text catalog unless ``type`` is +passed, which is why the same account configured as a generic upstream sees a +different catalog. These tests pin that query parameter, the per-token pricing +shape, and the families dropped as unpriceable. +""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import patch + +import pytest + +from routstr.upstream.venice import VeniceUpstreamProvider + + +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, *_: object) -> None: + return None + + async def get( + self, + url: str, + params: dict[str, Any] | None = None, + headers: dict[str, str] | None = None, + ) -> _FakeResponse: + self._calls.append({"url": url, "params": params, "headers": headers}) + return _FakeResponse(self._payload) + + +CATALOG: dict[str, Any] = { + "object": "list", + "data": [ + { + "id": "venice-uncensored-1-2", + "type": "text", + "created": 1727966436, + "model_spec": { + "name": "Venice Uncensored 1.2", + "availableContextTokens": 128000, + "maxCompletionTokens": 8192, + "capabilities": {"supportsVision": True}, + "pricing": { + "input": {"usd": 0.2, "diem": 0.2}, + "output": {"usd": 0.9, "diem": 0.9}, + "cache_input": {"usd": 0.02, "diem": 0.02}, + "cache_write": {"usd": 0.25, "diem": 0.25}, + }, + }, + }, + { + "id": "text-embedding-bge-m3", + "type": "embedding", + "created": 1727966436, + "model_spec": { + "name": "BGE m3", + "availableContextTokens": 8192, + "pricing": {"input": {"usd": 0.01, "diem": 0.01}}, + }, + }, + { + "id": "unpriced-text", + "type": "text", + "created": 1727966436, + "model_spec": {"name": "Unpriced", "pricing": {}}, + }, + { + "id": "offline-model", + "type": "text", + "created": 1727966436, + "model_spec": { + "name": "Offline", + "offline": True, + "pricing": {"input": {"usd": 0.2, "diem": 0.2}}, + }, + }, + { + "id": "venice-sd35", + "type": "image", + "created": 1727966436, + "model_spec": { + "name": "Venice SD35", + "pricing": {"generation": {"usd": 0.01, "diem": 0.01}}, + }, + }, + { + "id": "flux-2-max-edit", + "type": "inpaint", + "created": 1727966436, + "model_spec": { + "name": "FLUX.2 Max Edit", + "pricing": {"inpaint": {"usd": 0.12, "diem": 0.12}}, + }, + }, + { + "id": "tts-kokoro", + "type": "tts", + "created": 1727966436, + "model_spec": { + "name": "Kokoro", + "pricing": {"input": {"usd": 3.5, "diem": 3.5}}, + }, + }, + { + "id": "unpriced-video", + "type": "video", + "created": 1727966436, + "model_spec": {"name": "Video"}, + }, + ], +} + + +def _fetch(payload: dict[str, Any] = CATALOG) -> tuple[list[Any], list[dict[str, Any]]]: + import asyncio + + calls: list[dict[str, Any]] = [] + provider = VeniceUpstreamProvider(api_key="sk-test") + with patch( + "routstr.upstream.venice.httpx.AsyncClient", + lambda *a, **kw: _FakeAsyncClient(payload, calls), + ): + models = asyncio.run(provider.fetch_models()) + return models, calls + + +def test_requests_every_model_family() -> None: + _, calls = _fetch() + assert calls[0]["params"] == {"type": "all"} + assert calls[0]["url"] == "https://api.venice.ai/api/v1/models" + assert calls[0]["headers"] == {"Authorization": "Bearer sk-test"} + + +def test_text_pricing_is_per_token() -> None: + models, _ = _fetch() + model = next(m for m in models if m.id == "venice-uncensored-1-2") + assert model.pricing.prompt == pytest.approx(0.2 / 1_000_000) + assert model.pricing.completion == pytest.approx(0.9 / 1_000_000) + assert model.pricing.input_cache_read == pytest.approx(0.02 / 1_000_000) + assert model.pricing.input_cache_write == pytest.approx(0.25 / 1_000_000) + assert model.context_length == 128000 + assert model.top_provider is not None + assert model.top_provider.max_completion_tokens == 8192 + assert model.architecture.input_modalities == ["text", "image"] + assert model.architecture.modality == "text+image->text" + + +def test_embedding_models_are_listed() -> None: + models, _ = _fetch() + model = next(m for m in models if m.id == "text-embedding-bge-m3") + assert model.architecture.output_modalities == ["embedding"] + assert model.pricing.prompt == pytest.approx(0.01 / 1_000_000) + assert model.pricing.completion == 0.0 + + +def test_families_billed_per_clip_are_dropped() -> None: + """Image, audio and video return no usage to settle against, so listing + them here would hand out inference this provider cannot price.""" + models, _ = _fetch() + ids = {m.id for m in models} + assert "venice-sd35" not in ids + assert "flux-2-max-edit" not in ids + assert "tts-kokoro" not in ids + assert "unpriced-video" not in ids + + +def test_offline_and_unpriced_models_are_dropped() -> None: + models, _ = _fetch() + ids = {m.id for m in models} + assert "offline-model" not in ids + assert "unpriced-text" not in ids + + +def _priced_entry(model_id: str, model_type: str, pricing: dict[str, Any]) -> dict: + return { + "id": model_id, + "type": model_type, + "created": 1727966436, + "model_spec": {"name": model_id, "pricing": pricing}, + } + + +@pytest.mark.parametrize( + "pricing", + [ + pytest.param({"input": {"usd": 0.2, "diem": 0.2}}, id="missing-output"), + pytest.param( + {"input": {"usd": 0.0, "diem": 0.0}, "output": {"usd": 0.0, "diem": 0.0}}, + id="both-zero", + ), + pytest.param( + {"input": {"usd": -0.2, "diem": 0.2}, "output": {"usd": 0.9, "diem": 0.9}}, + id="negative-input", + ), + pytest.param( + {"input": {"usd": 0.2, "diem": 0.2}, "output": {"usd": -0.9, "diem": 0.9}}, + id="negative-output", + ), + ], +) +def test_text_models_that_would_bill_free_or_negative_are_dropped( + pricing: dict[str, Any], +) -> None: + models, _ = _fetch( + {"object": "list", "data": [_priced_entry("bad-text", "text", pricing)]} + ) + assert models == [] + + +def test_embedding_with_only_an_input_price_is_listed() -> None: + models, _ = _fetch( + { + "object": "list", + "data": [ + _priced_entry("emb", "embedding", {"input": {"usd": 0.05, "diem": 0}}) + ], + } + ) + assert [m.id for m in models] == ["emb"] + assert models[0].pricing.prompt == pytest.approx(0.05 / 1_000_000) + assert models[0].pricing.completion == 0.0 + + +def test_embedding_with_a_negative_price_is_dropped() -> None: + models, _ = _fetch( + { + "object": "list", + "data": [ + _priced_entry("emb", "embedding", {"input": {"usd": -0.05, "diem": 0}}) + ], + } + ) + assert models == [] + + +def test_text_model_with_one_zero_price_is_listed() -> None: + """Only both-zero is free; a free prompt with a paid completion is priced.""" + pricing = {"input": {"usd": 0.0, "diem": 0}, "output": {"usd": 0.9, "diem": 0}} + models, _ = _fetch( + {"object": "list", "data": [_priced_entry("t", "text", pricing)]} + ) + assert [m.id for m in models] == ["t"] + assert models[0].pricing.completion == pytest.approx(0.9 / 1_000_000) + + +def test_model_name_drops_the_venice_prefix() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + assert provider.transform_model_name("venice/venice-uncensored-1-2") == ( + "venice-uncensored-1-2" + ) + assert provider.transform_model_name("venice-uncensored-1-2") == ( + "venice-uncensored-1-2" + ) + + +def test_provider_metadata_pins_the_base_url() -> None: + metadata = VeniceUpstreamProvider.get_provider_metadata() + assert metadata["id"] == "venice" + assert metadata["default_base_url"] == "https://api.venice.ai/api/v1" + assert metadata["fixed_base_url"] is True + + +def test_fetch_returns_empty_on_upstream_failure() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + + with patch.object( + VeniceUpstreamProvider, + "_fetch_provider_models", + side_effect=RuntimeError("boom"), + ): + import asyncio + + assert asyncio.run(provider.fetch_models()) == [] diff --git a/tests/unit/test_user_liability_for_mint_and_unit.py b/tests/unit/test_user_liability_for_mint_and_unit.py new file mode 100644 index 00000000..7d8abbc1 --- /dev/null +++ b/tests/unit/test_user_liability_for_mint_and_unit.py @@ -0,0 +1,131 @@ +"""Real-DB coverage for the per-mint liability query that bounds owner payout.""" + +from typing import AsyncGenerator + +import pytest +from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine +from sqlalchemy.pool import StaticPool +from sqlmodel import SQLModel +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ApiKey, Refund, user_liability_for_mint_and_unit + +MINT = "http://m1" + + +def _make_engine() -> AsyncEngine: + return create_async_engine( + "sqlite+aiosqlite://", + poolclass=StaticPool, + connect_args={"check_same_thread": False}, + ) + + +@pytest.fixture +async def session() -> "AsyncGenerator[AsyncSession, None]": + engine = _make_engine() + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + db_session = AsyncSession(engine, expire_on_commit=False) + try: + yield db_session + finally: + await db_session.close() + await engine.dispose() + + +async def _add_key( + session: AsyncSession, + hashed_key: str, + balance: int, + mint_url: str | None = MINT, + currency: str | None = "sat", +) -> None: + session.add( + ApiKey( + hashed_key=hashed_key, + balance=balance, + refund_mint_url=mint_url, + refund_currency=currency, + ) + ) + await session.commit() + + +async def _add_refund( + session: AsyncSession, + hashed_key: str, + amount_msats: int, + status: str, + mint_url: str = MINT, + unit: str = "sat", +) -> None: + session.add( + Refund( + api_key_hashed_key=hashed_key, + method="lightning", + amount_msats=amount_msats, + unit=unit, + mint_url=mint_url, + status=status, + ) + ) + await session.commit() + + +@pytest.mark.asyncio +async def test_sums_key_balances_for_the_mint_and_unit(session: AsyncSession) -> None: + await _add_key(session, "a", 1000) + await _add_key(session, "b", 500) + + assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 1500 + + +@pytest.mark.asyncio +async def test_adds_unresolved_refunds_to_key_balances(session: AsyncSession) -> None: + # Only one pending/ambiguous claim per key is allowed. + await _add_key(session, "a", 1000) + await _add_refund(session, "a", 300, "pending") + await _add_key(session, "b", 0) + await _add_refund(session, "b", 40, "ambiguous") + await _add_key(session, "c", 0) + await _add_refund(session, "c", 7, "stuck") + + assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 1347 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status", ["paid", "failed"]) +async def test_excludes_resolved_refunds(session: AsyncSession, status: str) -> None: + await _add_key(session, "a", 0) + await _add_refund(session, "a", 900, status) + + assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 0 + + +@pytest.mark.asyncio +async def test_excludes_other_mints_and_units(session: AsyncSession) -> None: + await _add_key(session, "a", 1000) + await _add_key(session, "other-mint", 111, mint_url="http://m2") + await _add_key(session, "other-unit", 222, currency="msat") + await _add_refund(session, "a", 300, "pending") + await _add_refund(session, "other-mint", 444, "pending", mint_url="http://m2") + await _add_refund(session, "other-unit", 555, "pending", unit="msat") + + assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 1300 + + +@pytest.mark.asyncio +async def test_excludes_keys_without_a_refund_mint(session: AsyncSession) -> None: + await _add_key(session, "a", 1000) + await _add_key(session, "unattributed", 4242, mint_url=None, currency=None) + + assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 1000 + + +@pytest.mark.asyncio +async def test_unknown_mint_has_no_liability(session: AsyncSession) -> None: + await _add_key(session, "a", 1000) + await _add_refund(session, "a", 300, "pending") + + assert await user_liability_for_mint_and_unit(session, "http://missing", "sat") == 0 diff --git a/tests/unit/test_venice_encrypted_reasoning.py b/tests/unit/test_venice_encrypted_reasoning.py new file mode 100644 index 00000000..43af7614 --- /dev/null +++ b/tests/unit/test_venice_encrypted_reasoning.py @@ -0,0 +1,210 @@ +import json +from collections.abc import AsyncIterator +from typing import Any +from unittest.mock import AsyncMock, patch + +import pytest + +from routstr.upstream import messages_dispatch +from routstr.upstream.base import BaseUpstreamProvider +from routstr.upstream.venice import VeniceUpstreamProvider, _drop_encrypted_reasoning + +from .test_venice_web_search import _model + +ENCRYPTED = "__ENCRYPTED_REASONING__id=rs_0b04\ngAAAAABqvDJD" + + +def _block(index: int, block: dict, deltas: list[dict]) -> list[dict]: + return [ + {"type": "content_block_start", "index": index, "content_block": block}, + *({"type": "content_block_delta", "index": index, "delta": d} for d in deltas), + {"type": "content_block_stop", "index": index}, + ] + + +def _thinking(index: int, text: str) -> list[dict]: + return _block( + index, + {"type": "thinking", "thinking": "", "signature": ""}, + [{"type": "thinking_delta", "thinking": text}], + ) + + +def _text(index: int, text: str) -> list[dict]: + return _block( + index, + {"type": "text", "text": ""}, + [{"type": "text_delta", "text": text}], + ) + + +def _tool(index: int) -> list[dict]: + return _block( + index, + {"type": "tool_use", "id": "call_1", "name": "Bash", "input": {}}, + [{"type": "input_json_delta", "partial_json": '{"command":"ls"}'}], + ) + + +def _message(blocks: list[dict], stop_reason: str = "end_turn") -> list[dict]: + return [ + { + "type": "message_start", + "message": {"id": "msg_1", "role": "assistant", "content": []}, + }, + *blocks, + {"type": "message_delta", "delta": {"stop_reason": stop_reason}}, + {"type": "message_stop"}, + ] + + +async def _upstream(events: list[dict], *, split: bool = False) -> AsyncIterator[Any]: + payload = b"".join(messages_dispatch.encode_sse(e) for e in events) + if split: + for i in range(0, len(payload), 7): + yield payload[i : i + 7] + else: + yield payload + + +async def _filtered(events: list[dict], **kwargs: Any) -> list[dict]: + buffer = b"" + out: list[dict] = [] + async for chunk in _drop_encrypted_reasoning(_upstream(events, **kwargs)): + parsed, buffer = messages_dispatch.events_from_chunk(chunk, buffer) + out.extend(parsed) + return out + + +def _starts(events: list[dict]) -> list[tuple[int, str]]: + return [ + (e["index"], e["content_block"]["type"]) + for e in events + if e["type"] == "content_block_start" + ] + + +@pytest.mark.asyncio +async def test_trailing_encrypted_reasoning_is_dropped() -> None: + events = _message([*_text(0, "a.txt contains: hello"), *_thinking(1, ENCRYPTED)]) + + out = await _filtered(events) + + assert _starts(out) == [(0, "text")] + assert all(ENCRYPTED not in json.dumps(e) for e in out) + assert out[-2]["delta"]["stop_reason"] == "end_turn" + + +@pytest.mark.asyncio +async def test_leading_encrypted_reasoning_closes_index_gap() -> None: + events = _message( + [*_thinking(0, ENCRYPTED), *_text(1, "hi"), *_tool(2)], "tool_use" + ) + + out = await _filtered(events, split=True) + + assert _starts(out) == [(0, "text"), (1, "tool_use")] + assert {e["index"] for e in out if "index" in e} == {0, 1} + + +@pytest.mark.asyncio +async def test_plaintext_thinking_is_kept_in_order() -> None: + events = _message([*_thinking(0, "Let me list files."), *_tool(1)], "tool_use") + + out = await _filtered(events) + + assert out == events + + +@pytest.mark.asyncio +async def test_thinking_start_without_delta_is_flushed() -> None: + events = _message( + [ + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "thinking", "thinking": "", "signature": ""}, + }, + {"type": "content_block_stop", "index": 0}, + *_text(1, "ok"), + ] + ) + + out = await _filtered(events) + + assert out == events + + +@pytest.mark.asyncio +async def test_aggregated_message_ends_with_answer_text() -> None: + events = _message([*_text(0, "hello"), *_thinking(1, ENCRYPTED)]) + + message = await messages_dispatch.aggregate_anthropic_events_to_message( + _drop_encrypted_reasoning(_upstream(events)) + ) + + assert [b["type"] for b in message["content"]] == ["text"] + assert message["content"][0]["text"] == "hello" + + +async def _dispatched_blocks( + provider: BaseUpstreamProvider, *, stream: bool +) -> list[str]: + events = _message([*_text(0, "hello"), *_thinking(1, ENCRYPTED)]) + with patch( + "litellm.anthropic.messages.acreate", + new=AsyncMock(return_value=_upstream(events)), + ): + _, result, _ = await provider._dispatch_anthropic_messages( + request_body=json.dumps( + { + "model": "x", + "stream": stream, + "max_tokens": 64, + "messages": [{"role": "user", "content": "hi"}], + } + ).encode(), + model_obj=_model(), + ) + if not stream: + return [b["type"] for b in result["content"]] + buffer = b"" + out: list[dict] = [] + async for chunk in result: + parsed, buffer = messages_dispatch.events_from_chunk(chunk, buffer) + out.extend(parsed) + return [t for _, t in _starts(out)] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", [True, False]) +async def test_venice_dispatch_drops_encrypted_reasoning(stream: bool) -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + + assert await _dispatched_blocks(provider, stream=stream) == ["text"] + + +@pytest.mark.asyncio +async def test_other_providers_keep_thinking_blocks() -> None: + provider = BaseUpstreamProvider(base_url="https://example.com/v1", api_key="k") + + assert await _dispatched_blocks(provider, stream=True) == ["text", "thinking"] + + +@pytest.mark.asyncio +async def test_closing_the_filter_closes_upstream() -> None: + closed = False + + async def upstream() -> AsyncIterator[bytes]: + nonlocal closed + try: + for event in _message(_text(0, "hello")): + yield messages_dispatch.encode_sse(event) + finally: + closed = True + + filtered = _drop_encrypted_reasoning(upstream()) + await filtered.__anext__() + await filtered.aclose() + + assert closed diff --git a/tests/unit/test_venice_system_cache.py b/tests/unit/test_venice_system_cache.py new file mode 100644 index 00000000..6b8cc67f --- /dev/null +++ b/tests/unit/test_venice_system_cache.py @@ -0,0 +1,86 @@ +from __future__ import annotations + +import pytest + +from routstr.upstream.venice import VeniceUpstreamProvider + +from .test_venice_web_search import _body, _dispatch + +EPHEMERAL = {"type": "ephemeral"} + +CLAUDE_CODE_SYSTEM = [ + { + "type": "text", + "text": "x-anthropic-billing-header: cc_version=2.1.281; cc_entrypoint=cli;", + }, + {"type": "text", "text": "You are a Claude agent.", "cache_control": EPHEMERAL}, + { + "type": "text", + "text": "\nYou are an interactive agent.", + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + }, +] + + +@pytest.mark.asyncio +async def test_cache_marked_multi_block_system_is_merged_into_one_block() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + + kwargs = await _dispatch(provider, _body(system=CLAUDE_CODE_SYSTEM)) + + assert kwargs["system"] == [ + { + "type": "text", + "text": ( + "x-anthropic-billing-header: cc_version=2.1.281; cc_entrypoint=cli;" + "\n\nYou are a Claude agent.\n\n\nYou are an interactive agent." + ), + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + } + ] + + +@pytest.mark.asyncio +async def test_unmarked_multi_block_system_is_untouched() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + system = [{"type": "text", "text": "A."}, {"type": "text", "text": "B."}] + + kwargs = await _dispatch(provider, _body(system=system)) + + assert kwargs["system"] == system + + +@pytest.mark.asyncio +async def test_single_marked_block_and_string_system_are_untouched() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + single = [{"type": "text", "text": "A.", "cache_control": EPHEMERAL}] + + assert (await _dispatch(provider, _body(system=single)))["system"] == single + assert (await _dispatch(provider, _body(system="A.")))["system"] == "A." + + +@pytest.mark.asyncio +async def test_message_and_tool_cache_markers_are_kept() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + messages = [ + { + "role": "user", + "content": [{"type": "text", "text": "hi", "cache_control": EPHEMERAL}], + } + ] + tools = [ + { + "name": "Bash", + "description": "Run a command", + "input_schema": {"type": "object", "properties": {}}, + "cache_control": EPHEMERAL, + } + ] + + kwargs = await _dispatch( + provider, + _body(system=CLAUDE_CODE_SYSTEM, messages=messages, tools=tools), + ) + + assert kwargs["messages"] == messages + assert kwargs["tools"] == tools diff --git a/tests/unit/test_venice_web_search.py b/tests/unit/test_venice_web_search.py new file mode 100644 index 00000000..380ccc33 --- /dev/null +++ b/tests/unit/test_venice_web_search.py @@ -0,0 +1,303 @@ +"""Venice web search over ``/v1/messages``. + +litellm's Anthropic adapter rewrites an Anthropic server-side web-search tool +into a top-level ``web_search_options``, which Venice rejects with +``400 Unrecognized key(s) in object: 'web_search_options'``. These tests pin +the trade: the tool is lifted out of the body and the same intent re-expressed +as a Venice model feature suffix. +""" + +from __future__ import annotations + +import json +from typing import Any, AsyncIterator +from unittest.mock import AsyncMock, patch + +import pytest + +from routstr.core.exceptions import UpstreamError +from routstr.payment.models import Architecture, Model, Pricing +from routstr.upstream.base import BaseUpstreamProvider +from routstr.upstream.venice import VeniceUpstreamProvider + +WEB_SEARCH_TOOL = {"type": "web_search_20250305", "name": "web_search"} +FUNCTION_TOOL = { + "name": "lookup", + "description": "Look something up", + "input_schema": {"type": "object", "properties": {}}, +} + + +def _model(model_id: str = "deepseek-v4-flash-0731") -> Model: + return Model( + id=model_id, + name=model_id, + 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), + ) + + +def _body(**extra: Any) -> dict[str, Any]: + return { + "messages": [{"role": "user", "content": "what shipped today?"}], + "max_tokens": 64, + **extra, + } + + +async def _dispatch(provider: BaseUpstreamProvider, body: dict[str, Any]) -> dict: + """Run the real dispatcher, capturing the kwargs litellm would receive.""" + captured: dict[str, Any] = {} + + async def empty_iter() -> AsyncIterator[dict]: + if False: + yield {} + + async def fake_acreate(**kwargs: Any) -> AsyncIterator[dict]: + captured.update(kwargs) + return empty_iter() + + with patch( + "litellm.anthropic.messages.acreate", + new=AsyncMock(side_effect=fake_acreate), + ): + await provider._dispatch_anthropic_messages( + request_body=json.dumps( + {"model": "venice/x", "stream": True, **body} + ).encode(), + model_obj=_model(), + ) + return captured + + +@pytest.mark.asyncio +async def test_web_search_tool_never_reaches_venice_as_web_search_options() -> None: + """The reported 400: the derived parameter must not be sent at all.""" + provider = VeniceUpstreamProvider(api_key="sk-test") + + kwargs = await _dispatch(provider, _body(tools=[WEB_SEARCH_TOOL])) + + assert "web_search_options" not in kwargs + assert "tools" not in kwargs + assert kwargs["model"] == ( + "openai/deepseek-v4-flash-0731:enable_web_search=auto&enable_web_citations=true" + ) + assert kwargs["api_base"] == "https://api.venice.ai/api/v1" + + +@pytest.mark.asyncio +async def test_function_tools_survive_alongside_web_search() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + + kwargs = await _dispatch( + provider, + _body( + tools=[WEB_SEARCH_TOOL, FUNCTION_TOOL], + tool_choice={"type": "tool", "name": "lookup"}, + ), + ) + + assert kwargs["tools"] == [FUNCTION_TOOL] + assert kwargs["tool_choice"] == {"type": "tool", "name": "lookup"} + assert "web_search_options" not in kwargs + assert kwargs["model"].endswith(":enable_web_search=auto&enable_web_citations=true") + + +@pytest.mark.asyncio +async def test_requests_without_web_search_are_untouched() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + + kwargs = await _dispatch(provider, _body(tools=[FUNCTION_TOOL])) + + assert kwargs["model"] == "openai/deepseek-v4-flash-0731" + assert kwargs["tools"] == [FUNCTION_TOOL] + + +@pytest.mark.asyncio +async def test_generic_openai_upstream_rejects_untranslatable_web_search() -> None: + """Do not let LiteLLM send unsupported web_search_options to a generic API.""" + provider = BaseUpstreamProvider(base_url="http://test", api_key="k") + + with pytest.raises(UpstreamError) as excinfo: + await _dispatch(provider, _body(tools=[WEB_SEARCH_TOOL, FUNCTION_TOOL])) + + assert excinfo.value.status_code == 400 + assert excinfo.value.code == "UNSUPPORTED_WEB_SEARCH" + + +@pytest.mark.asyncio +async def test_generic_openai_upstream_still_accepts_function_tools() -> None: + provider = BaseUpstreamProvider(base_url="http://test", api_key="k") + + kwargs = await _dispatch(provider, _body(tools=[FUNCTION_TOOL])) + + assert kwargs["model"] == "openai/deepseek-v4-flash-0731" + assert kwargs["tools"] == [FUNCTION_TOOL] + + +@pytest.mark.asyncio +async def test_non_openai_adapter_can_still_handle_search_tool() -> None: + provider = BaseUpstreamProvider( + base_url="https://openrouter.ai/api/v1", api_key="k" + ) + + kwargs = await _dispatch(provider, _body(tools=[WEB_SEARCH_TOOL])) + + assert kwargs["tools"] == [WEB_SEARCH_TOOL] + + +@pytest.mark.parametrize( + "tool", + [ + { + "type": "web_search_20250305", + "name": "web_search", + "allowed_domains": ["example.com"], + }, + {"type": "web_search_20250305", "name": "web_search", "blocked_domains": ["x"]}, + { + "type": "web_search_20250305", + "name": "web_search", + "user_location": {"type": "approximate", "country": "DE"}, + }, + ], +) +def test_constraints_venice_cannot_enforce_are_refused(tool: dict[str, Any]) -> None: + """Better an explicit 400 than a search that quietly ignored the limit.""" + provider = VeniceUpstreamProvider(api_key="sk-test") + + with pytest.raises(UpstreamError) as excinfo: + provider.adapt_messages_request(_body(tools=[tool]), _model()) + + assert excinfo.value.status_code == 400 + assert excinfo.value.code == "UNSUPPORTED_WEB_SEARCH_OPTION" + + +@pytest.mark.parametrize( + "tool", + [ + {"type": "web_search_20250305", "name": "web_search", "max_uses": None}, + {"type": "web_search_20250305", "name": "web_search", "allowed_domains": []}, + ], +) +def test_constraint_keys_stating_nothing_are_read_as_absent( + tool: dict[str, Any], +) -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + + assert provider.adapt_messages_request(_body(tools=[tool]), _model()) != "" + + +def test_forcing_web_search_through_tool_choice_is_refused() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + body = _body( + tools=[WEB_SEARCH_TOOL], + tool_choice={"type": "tool", "name": "web_search"}, + ) + + with pytest.raises(UpstreamError) as excinfo: + provider.adapt_messages_request(body, _model()) + + assert excinfo.value.status_code == 400 + assert excinfo.value.code == "UNSUPPORTED_WEB_SEARCH_OPTION" + assert excinfo.value.details == {"unsupported_options": ["tool_choice"]} + + +def test_web_search_only_request_drops_tool_choice() -> None: + """Without tools left, a surviving tool_choice is rejected upstream.""" + provider = VeniceUpstreamProvider(api_key="sk-test") + body = _body(tools=[WEB_SEARCH_TOOL], tool_choice={"type": "auto"}) + + provider.adapt_messages_request(body, _model()) + + assert "tools" not in body + assert "tool_choice" not in body + + +@pytest.mark.asyncio +async def test_claude_code_web_search_tool_is_accepted() -> None: + """Claude Code always sends ``max_uses: 8``; Venice's single ``auto`` + search already stays under any cap of one or more.""" + provider = VeniceUpstreamProvider(api_key="sk-test") + tool = { + "type": "web_search_20250305", + "name": "web_search", + "allowed_domains": None, + "blocked_domains": None, + "max_uses": 8, + } + + kwargs = await _dispatch(provider, _body(tools=[tool])) + + assert "web_search_options" not in kwargs + assert "tools" not in kwargs + assert kwargs["model"] == ( + "openai/deepseek-v4-flash-0731:enable_web_search=auto&enable_web_citations=true" + ) + + +@pytest.mark.parametrize("max_uses", [1, None]) +def test_max_uses_of_one_or_absent_is_accepted(max_uses: Any) -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + tool = {"type": "web_search_20250305", "name": "web_search", "max_uses": max_uses} + + assert provider.adapt_messages_request(_body(tools=[tool]), _model()) != "" + + +@pytest.mark.parametrize("max_uses", [0, -1, 1.5, True, "0", "8"]) +def test_max_uses_other_than_a_positive_integer_is_refused(max_uses: Any) -> None: + """``auto`` may still search, so a cap below one cannot be met, and a + malformed cap cannot be shown to be met.""" + provider = VeniceUpstreamProvider(api_key="sk-test") + tool = {"type": "web_search_20250305", "name": "web_search", "max_uses": max_uses} + + with pytest.raises(UpstreamError) as excinfo: + provider.adapt_messages_request(_body(tools=[tool]), _model()) + + assert excinfo.value.status_code == 400 + assert excinfo.value.code == "UNSUPPORTED_WEB_SEARCH_OPTION" + assert excinfo.value.details == {"unsupported_options": ["max_uses"]} + + +def test_tool_named_web_search_without_the_type_marker_is_caught() -> None: + """litellm matches on either marker, so this one would also be rewritten.""" + provider = VeniceUpstreamProvider(api_key="sk-test") + body = _body(tools=[{"name": "web_search"}]) + + assert provider.adapt_messages_request(body, _model()) != "" + assert "tools" not in body + + +def test_litellm_adapter_derives_no_web_search_options_from_the_adapted_body() -> None: + """The fix at its cause: run the real litellm translation over the body + this provider produces and assert the rejected key is never derived.""" + from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import ( # noqa: E501 + LiteLLMAnthropicMessagesAdapter, + ) + + provider = VeniceUpstreamProvider(api_key="sk-test") + adapter = LiteLLMAnthropicMessagesAdapter() # type: ignore[no-untyped-call] + body = _body(tools=[WEB_SEARCH_TOOL, FUNCTION_TOOL]) + + def translate(request: dict[str, Any]) -> dict: + # litellm types the request as a TypedDict; these bodies are built + # from client JSON, so they are plain dicts at this seam. + translated, _ = adapter.translate_anthropic_to_openai(request) # type: ignore[arg-type] + return dict(translated) + + # Unadapted, litellm derives the parameter Venice rejects. + before = translate({"model": "m", **_body(tools=[WEB_SEARCH_TOOL])}) + assert "web_search_options" in before + + provider.adapt_messages_request(body, _model()) + + assert "web_search_options" not in translate({"model": "m", **body}) diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 86c90b63..30b71f09 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -4,6 +4,7 @@ import json import socket from collections.abc import AsyncIterator, Generator from contextlib import asynccontextmanager +from pathlib import Path from unittest.mock import AsyncMock, MagicMock, Mock, patch import httpx @@ -31,17 +32,23 @@ from routstr.wallet import ( @pytest.fixture(autouse=True) -def isolate_wallet_runtime_state() -> Generator[None, None, None]: +def isolate_wallet_runtime_state( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> Generator[None, None, None]: """Keep production limiter/wallet caches from leaking across unit tests.""" from routstr import wallet as wallet_module from routstr.core.settings import settings + monkeypatch.setattr( + wallet_module, "_WALLET_OPERATION_LOCK", tmp_path / "wallet.lock" + ) original_concurrency = settings.mint_max_concurrency settings.mint_max_concurrency = 0 wallet_module._MintRateGuard._guards.clear() wallet_module._wallets.clear() wallet_module._wallet_last_load.clear() wallet_module._wallet_last_mint_load.clear() + wallet_module._wallet_mint_load_errors.clear() wallet_module._wallet_load_locks.clear() wallet_module._mint_metadata_last_load.clear() wallet_module._mint_metadata_load_locks.clear() @@ -51,6 +58,7 @@ def isolate_wallet_runtime_state() -> Generator[None, None, None]: wallet_module._wallets.clear() wallet_module._wallet_last_load.clear() wallet_module._wallet_last_mint_load.clear() + wallet_module._wallet_mint_load_errors.clear() wallet_module._wallet_load_locks.clear() wallet_module._mint_metadata_last_load.clear() wallet_module._mint_metadata_load_locks.clear() @@ -143,6 +151,70 @@ async def test_get_wallet_force_reload_bypasses_reload_interval() -> None: assert mock_wallet.load_proofs.await_count == 2 +@pytest.mark.asyncio +async def test_unservable_mint_load_is_not_retried_every_call() -> None: + """Retrying it per call refetched keysets and got the node rate-limited.""" + from routstr.wallet import get_wallet + + failure = Exception("No active keyset found for unit msat.") + mock_wallet = Mock( + load_mint=AsyncMock(side_effect=failure), load_proofs=AsyncMock() + ) + with patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet)): + for _ in range(3): + with pytest.raises(Exception, match="No active keyset"): + await get_wallet("http://mint:3338", "msat") + + assert mock_wallet.load_mint.await_count == 1 + + +@pytest.mark.asyncio +async def test_unreachable_mint_load_stays_retryable() -> None: + """Transport failures are the rate guard's job, not the metadata throttle's.""" + from routstr.wallet import get_wallet + + failure = httpx.ConnectError("mint unreachable") + mock_wallet = Mock( + load_mint=AsyncMock(side_effect=failure), load_proofs=AsyncMock() + ) + with patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet)): + for _ in range(2): + with pytest.raises(Exception): + await get_wallet("http://mint:3338", "sat") + + assert mock_wallet.load_mint.await_count == 2 + + +@pytest.mark.asyncio +async def test_force_reload_retries_an_unservable_mint_load() -> None: + from routstr.wallet import get_wallet + + failure = Exception("No active keyset found for unit msat.") + mock_wallet = Mock( + load_mint=AsyncMock(side_effect=failure), load_proofs=AsyncMock() + ) + with patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet)): + with pytest.raises(Exception, match="No active keyset"): + await get_wallet("http://mint:3338", "msat") + with pytest.raises(Exception, match="No active keyset"): + await get_wallet("http://mint:3338", "msat", force_reload=True) + + assert mock_wallet.load_mint.await_count == 2 + + +@pytest.mark.asyncio +async def test_get_wallet_force_reload_proofs_keeps_cached_keysets() -> None: + from routstr.wallet import get_wallet + + mock_wallet = Mock(load_mint=AsyncMock(), load_proofs=AsyncMock()) + with patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet)): + await get_wallet("http://mint:3338", "sat") + await get_wallet("http://mint:3338", "sat", force_reload_proofs=True) + + assert mock_wallet.load_mint.await_count == 1 + assert mock_wallet.load_proofs.await_count == 2 + + @pytest.mark.asyncio async def test_public_recieve_token_holds_wallet_operation_guard() -> None: inside_guard = False @@ -1235,9 +1307,14 @@ async def test_prepare_bolt11_payment_does_not_spend_user_liabilities() -> None: @pytest.mark.asyncio -async def test_prepare_bolt11_payment_rounds_user_liability_up_to_whole_sats() -> None: +async def test_prepare_bolt11_payment_floors_fractional_owner_surplus() -> None: + """A sub-sat surplus is not enough to fund a 1 sat invoice.""" from routstr.core.settings import settings + @asynccontextmanager + async def session() -> AsyncIterator[MagicMock]: + yield MagicMock() + wallet = MagicMock() wallet.proofs = [MagicMock(amount=100)] wallet.melt_quote = AsyncMock( @@ -1262,10 +1339,15 @@ async def test_prepare_bolt11_payment_rounds_user_liability_up_to_whole_sats() - "routstr.wallet.slow_filter_spend_proofs", side_effect=lambda proofs, wallet: proofs, ), + patch("routstr.wallet.db.create_session", session), patch( "routstr.wallet.db.total_user_liability", AsyncMock(return_value=99_999), ), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(return_value=99_999), + ), pytest.raises(ValueError, match="user liabilities"), ): await prepare_bolt11_payment("lnbc-invoice") @@ -1301,10 +1383,12 @@ async def test_execute_bolt11_payment_rereserves_when_cancelled() -> None: @pytest.mark.asyncio async def test_balance_proof_check_uses_large_batches_to_avoid_rate_limit() -> None: """Balance reads must not turn a few hundred proofs into many mint requests.""" + from cashu.core.base import ProofSpentState + from routstr.wallet import slow_filter_spend_proofs - proofs = [Mock() for _ in range(250)] - states = [Mock(state="UNSPENT") for _ in proofs] + proofs = [Mock(Y=str(i)) for i in range(250)] + states = [Mock(Y=proof.Y, state=ProofSpentState.unspent) for proof in proofs] wallet = Mock() wallet.url = "http://mint:3338" wallet.check_proof_state = AsyncMock(return_value=Mock(states=states)) @@ -2017,7 +2101,7 @@ async def test_payout_reloads_wallet_snapshot_under_guard() -> None: await _payout_mint_and_unit("https://mint.example.com", "sat") mock_get_wallet.assert_awaited_once_with( - "https://mint.example.com", "sat", force_reload=True + "https://mint.example.com", "sat", force_reload_proofs=True ) diff --git a/tests/unit/test_x_cashu_provider_path.py b/tests/unit/test_x_cashu_provider_path.py index 11b47d6f..6b4fdec2 100644 --- a/tests/unit/test_x_cashu_provider_path.py +++ b/tests/unit/test_x_cashu_provider_path.py @@ -52,3 +52,50 @@ async def test_x_cashu_responses_stream_reports_complete_provider_path() -> None payload = json.loads((await _body(response)).decode().removeprefix("data: ")) assert payload["provider"] == "openrouter:z.ai" + + +@pytest.mark.asyncio +async def test_x_cashu_messages_stream_carries_provider_to_later_events() -> None: + provider = OpenRouterUpstreamProvider(api_key="test-key") + events = [ + {"type": "message_start", "message": {"provider": "Anthropic"}}, + {"type": "content_block_delta", "delta": {"text": "hi"}}, + ] + content = "".join(f"data: {json.dumps(e)}\n" for e in events) + + response = await provider.handle_x_cashu_streaming_response( + content, + httpx.Response(200, headers={"content-type": "text/event-stream"}), + amount=1, + unit="sat", + max_cost_for_model=1, + ) + + lines = (await _body(response)).decode().splitlines() + stamped = [json.loads(line.removeprefix("data: ")) for line in lines if line] + assert [e["provider"] for e in stamped] == ["openrouter:Anthropic"] * 2 + + +@pytest.mark.asyncio +async def test_x_cashu_responses_stream_carries_nested_provider() -> None: + provider = OpenRouterUpstreamProvider(api_key="test-key") + events = [ + {"type": "response.created", "response": {"provider": "OpenAI"}}, + {"type": "response.output_text.delta", "delta": "hi"}, + ] + content = "".join(f"data: {json.dumps(e)}\n\n" for e in events) + + with patch.object( + provider, "get_x_cashu_cost", new=AsyncMock(return_value=None) + ): + response = await provider.handle_x_cashu_streaming_responses_response( + content, + httpx.Response(200, headers={"content-type": "text/event-stream"}), + amount=1, + unit="sat", + max_cost_for_model=1, + ) + + lines = (await _body(response)).decode().splitlines() + stamped = [json.loads(line.removeprefix("data: ")) for line in lines if line] + assert [e["provider"] for e in stamped] == ["openrouter:OpenAI"] * 2 diff --git a/tests/unit/test_x_cashu_stream_ownership.py b/tests/unit/test_x_cashu_stream_ownership.py new file mode 100644 index 00000000..4751bd54 --- /dev/null +++ b/tests/unit/test_x_cashu_stream_ownership.py @@ -0,0 +1,330 @@ +import asyncio +from collections.abc import AsyncIterator +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from fastapi import Request +from fastapi.responses import Response, StreamingResponse +from starlette.types import Message, Send + +from routstr.upstream.base import BaseUpstreamProvider +from routstr.upstream.stream_ownership import OwnedUpstreamStream + + +class _CountingStream(httpx.AsyncByteStream): + def __init__(self, payload: bytes) -> None: + self.payload = payload + self.close_count = 0 + + async def __aiter__(self) -> AsyncIterator[bytes]: + yield self.payload + + async def aclose(self) -> None: + self.close_count += 1 + + +class _CountingTransport(httpx.AsyncBaseTransport): + def __init__(self, payload: bytes) -> None: + self.stream = _CountingStream(payload) + self.close_count = 0 + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + return httpx.Response(200, request=request, stream=self.stream) + + async def aclose(self) -> None: + self.close_count += 1 + + +class _CountingClient(httpx.AsyncClient): + def __init__(self, transport: _CountingTransport) -> None: + super().__init__(transport=transport) + self.close_count = 0 + + async def aclose(self) -> None: + self.close_count += 1 + await super().aclose() + + +def _request() -> Request: + sent = False + + async def receive() -> dict[str, object]: + nonlocal sent + if sent: + return {"type": "http.disconnect"} + sent = True + return {"type": "http.request", "body": b"{}", "more_body": False} + + return Request( + { + "type": "http", + "asgi": {"version": "3.0", "spec_version": "2.4"}, + "method": "POST", + "scheme": "http", + "path": "/v1/audio/speech", + "raw_path": b"/v1/audio/speech", + "query_string": b"", + "headers": [], + "client": ("test", 1), + "server": ("test", 80), + }, + receive, + ) + + +async def _forward( + provider: BaseUpstreamProvider, + method_name: str, +) -> tuple[StreamingResponse, _CountingClient, _CountingTransport]: + transport = _CountingTransport(b"live-stream") + client = _CountingClient(transport) + model = MagicMock() + + with patch("routstr.upstream.base.httpx.AsyncClient", return_value=client): + result = await getattr(provider, method_name)( + request=_request(), + path="v1/audio/speech", + headers={}, + amount=10, + unit="sat", + max_cost_for_model=10_000, + model_obj=model, + ) + + assert isinstance(result, StreamingResponse) + return result, client, transport + + +async def _run_asgi_response( + response: StreamingResponse, + send: Send, +) -> None: + async def receive() -> dict[str, str]: + return {"type": "http.disconnect"} + + await response( + { + "type": "http", + "asgi": {"version": "3.0", "spec_version": "2.4"}, + }, + receive, + send, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "method_name", + ["forward_x_cashu_request", "forward_x_cashu_responses_request"], +) +async def test_x_cashu_opaque_stream_owns_client_until_normal_completion( + method_name: str, +) -> None: + provider = BaseUpstreamProvider(base_url="http://upstream", api_key="test") + response, client, transport = await _forward(provider, method_name) + messages: list[dict[str, Any]] = [] + + assert client.close_count == 0 + assert transport.stream.close_count == 0 + + async def send(message: Message) -> None: + messages.append(dict(message)) + + await _run_asgi_response(response, send) + + assert ( + b"".join( + message.get("body", b"") + for message in messages + if message["type"] == "http.response.body" + ) + == b"live-stream" + ) + assert transport.stream.close_count == 1 + assert client.close_count == 1 + assert transport.close_count == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "method_name", + ["forward_x_cashu_request", "forward_x_cashu_responses_request"], +) +@pytest.mark.parametrize( + "failure", + [RuntimeError("downstream send failed"), asyncio.CancelledError()], +) +async def test_x_cashu_opaque_stream_closes_client_when_send_fails( + method_name: str, + failure: BaseException, +) -> None: + provider = BaseUpstreamProvider(base_url="http://upstream", api_key="test") + response, client, transport = await _forward(provider, method_name) + + async def send(message: Message) -> None: + if message["type"] == "http.response.body" and message.get("body"): + raise failure + + with pytest.raises(type(failure)): + await _run_asgi_response(response, send) + + assert transport.stream.close_count == 1 + assert client.close_count == 1 + assert transport.close_count == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("method_name", "path", "payload"), + [ + ( + "forward_x_cashu_request", + "v1/chat/completions", + b'data: {"model":"m","usage":{"prompt_tokens":1,"completion_tokens":1}}\n\ndata: [DONE]\n\n', + ), + ( + "forward_x_cashu_responses_request", + "v1/responses", + b'data: {"type":"response.completed","response":{"model":"m","usage":{"input_tokens":1,"output_tokens":1}}}\n\ndata: [DONE]\n\n', + ), + ], +) +@pytest.mark.parametrize( + "failure", [None, RuntimeError("send failed"), asyncio.CancelledError()] +) +async def test_x_cashu_real_processed_stream_releases_buffered_upstream_promptly( + method_name: str, + path: str, + payload: bytes, + failure: BaseException | None, +) -> None: + provider = BaseUpstreamProvider(base_url="http://upstream", api_key="test") + transport = _CountingTransport(payload) + client = _CountingClient(transport) + + with ( + patch("routstr.upstream.base.httpx.AsyncClient", return_value=client), + patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=None)), + ): + result = await getattr(provider, method_name)( + request=_request(), + path=path, + headers={}, + amount=10, + unit="sat", + max_cost_for_model=10_000, + model_obj=MagicMock(), + ) + + assert isinstance(result, StreamingResponse) + assert transport.stream.close_count == 1 + assert client.close_count == 1 + assert transport.close_count == 1 + + messages: list[Message] = [] + + async def send(message: Message) -> None: + messages.append(message) + if failure is not None and message["type"] == "http.response.body": + if message.get("body"): + raise failure + + if failure is None: + await _run_asgi_response(result, send) + assert any(message.get("body") for message in messages) + else: + with pytest.raises(type(failure)): + await _run_asgi_response(result, send) + + assert transport.stream.close_count == 1 + assert client.close_count == 1 + assert transport.close_count == 1 + + +@pytest.mark.asyncio +async def test_owned_upstream_cleanup_survives_caller_cancellation() -> None: + cleanup_started = asyncio.Event() + allow_cleanup = asyncio.Event() + cleanup_finished = asyncio.Event() + client_close_count = 0 + + async def body() -> AsyncIterator[bytes]: + yield b"body" + + response = MagicMock(spec=httpx.Response) + response.aclose = AsyncMock() + client = MagicMock(spec=httpx.AsyncClient) + + async def close_client() -> None: + nonlocal client_close_count + client_close_count += 1 + cleanup_started.set() + await allow_cleanup.wait() + cleanup_finished.set() + + client.aclose = close_client + owned = OwnedUpstreamStream(body(), response, client) + + first_close = asyncio.create_task(owned.aclose()) + await cleanup_started.wait() + first_close.cancel() + with pytest.raises(asyncio.CancelledError): + await first_close + + allow_cleanup.set() + await asyncio.wait_for(cleanup_finished.wait(), timeout=1) + await owned.aclose() + + assert response.aclose.await_count == 1 + assert client_close_count == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("method_name", "path", "handler_name"), + [ + ( + "forward_x_cashu_request", + "v1/chat/completions", + "handle_x_cashu_chat_completion", + ), + ( + "forward_x_cashu_responses_request", + "v1/responses", + "handle_x_cashu_responses_completion", + ), + ], +) +async def test_x_cashu_non_streaming_result_closes_upstream_promptly( + method_name: str, + path: str, + handler_name: str, +) -> None: + provider = BaseUpstreamProvider(base_url="http://upstream", api_key="test") + transport = _CountingTransport(b"{}") + client = _CountingClient(transport) + + with ( + patch("routstr.upstream.base.httpx.AsyncClient", return_value=client), + patch.object( + provider, + handler_name, + new=AsyncMock(return_value=Response(b"done")), + ), + ): + result = await getattr(provider, method_name)( + request=_request(), + path=path, + headers={}, + amount=10, + unit="sat", + max_cost_for_model=10_000, + model_obj=MagicMock(), + ) + + assert not isinstance(result, StreamingResponse) + assert transport.stream.close_count == 1 + assert client.close_count == 1 + assert transport.close_count == 1 diff --git a/ui/app/model/loading.tsx b/ui/app/model/loading.tsx new file mode 100644 index 00000000..390be4d8 --- /dev/null +++ b/ui/app/model/loading.tsx @@ -0,0 +1,23 @@ +import { AppPageShell } from '@/components/app-page-shell'; +import { PageHeader } from '@/components/page-header'; +import { Skeleton } from '@/components/ui/skeleton'; + +/** + * Route-level fallback so clicking "Models" lands on the page immediately + * instead of holding the previous route until this one's chunk is parsed. + */ +export default function ModelPageLoading() { + return ( + +
+ + + + +
+
+ ); +} diff --git a/ui/app/transactions/page.tsx b/ui/app/transactions/page.tsx index 56bc1b74..0e3eca85 100644 --- a/ui/app/transactions/page.tsx +++ b/ui/app/transactions/page.tsx @@ -264,7 +264,8 @@ function LightningInvoiceTable({ No invoices found - Lightning invoices created via /lightning/invoice will show here. + Lightning invoices created via /lightning/invoice and payouts sent + to your Lightning address will show here.
@@ -290,6 +291,33 @@ function LightningInvoiceTable({ Expired ); + if (status === 'failed') + return ( + + Failed + + ); + if (status === 'settlement_pending') + return ( + + Settling + + ); + if (status === 'reconciliation_required') + return ( + + Reconciling + + ); if (status === 'cancelled') return ( + Direction Purpose Amount Status @@ -328,6 +357,11 @@ function LightningInvoiceTable({ {invoices.map((inv) => ( + + + {inv.direction === 'out' ? 'Sent' : 'Received'} + + {inv.purpose} @@ -514,15 +548,27 @@ export default function TransactionsPage() { placeholderData: keepPreviousData, }); - const LIGHTNING_STATUSES = ['pending', 'paid', 'expired', 'cancelled']; + const LIGHTNING_STATUSES = [ + 'pending', + 'settlement_pending', + 'paid', + 'failed', + 'expired', + 'cancelled', + 'reconciliation_required', + ]; const lightningStatusParam = LIGHTNING_STATUSES.includes(status) ? status : undefined; + const lightningDirectionParam = ['in', 'out'].includes(type) + ? type + : undefined; const lightningQuery = useQuery({ queryKey: [ 'lightning-invoices', lightningStatusParam, + lightningDirectionParam, searchParam, lightningPage, ], @@ -530,6 +576,7 @@ export default function TransactionsPage() { AdminService.getLightningInvoices( lightningStatusParam, undefined, + lightningDirectionParam, searchParam, PAGE_SIZE, lightningPage * PAGE_SIZE @@ -710,7 +757,9 @@ export default function TransactionsPage() { All Types Incoming (Payments) - Outgoing (Refunds) + + Outgoing (Refunds & Payouts) + @@ -727,6 +776,13 @@ export default function TransactionsPage() { Collected Swept Paid (Lightning) + + Settling (Lightning) + + Failed (Lightning) + + Reconciling (Lightning) + Expired (Lightning) Cancelled (Lightning) diff --git a/ui/components/model-provider-section.tsx b/ui/components/model-provider-section.tsx index 9f441678..3acd132e 100644 --- a/ui/components/model-provider-section.tsx +++ b/ui/components/model-provider-section.tsx @@ -1,5 +1,6 @@ import { useMemo } from 'react'; import type { Model } from '@/lib/api/schemas/models'; +import { useProgressiveList } from '@/lib/hooks/use-progressive-list'; import type { AdminModelGroup } from '@/lib/api/services/admin'; import type { DisplayUnit } from '@/lib/types/units'; import { ModelItemCard } from '@/components/model-item-card'; @@ -24,6 +25,7 @@ import { Edit3, Globe, Key, + Loader2, MoreVertical, RefreshCw, } from 'lucide-react'; @@ -103,10 +105,21 @@ export function ModelProviderSection({ }); }, [provider, providerModels]); + const { visibleItems: visibleProviderModels, hiddenCount } = + useProgressiveList(keyedProviderModels); + + const pendingRowsNotice = + hiddenCount > 0 ? ( +
+ + Rendering {hiddenCount} more model{hiddenCount === 1 ? '' : 's'}… +
+ ) : null; + if (filterProvider) { return (
- {keyedProviderModels.map(({ model, renderKey }) => ( + {visibleProviderModels.map(({ model, renderKey }) => ( onDeleteModel(model.id)} /> ))} + {pendingRowsNotice}
); } @@ -217,7 +231,7 @@ export function ModelProviderSection({
- {keyedProviderModels.map(({ model, renderKey }) => ( + {visibleProviderModels.map(({ model, renderKey }) => ( onDeleteModel(model.id)} /> ))} + {pendingRowsNotice}
diff --git a/ui/components/model-selector.tsx b/ui/components/model-selector.tsx index 8636fb36..54a345b8 100644 --- a/ui/components/model-selector.tsx +++ b/ui/components/model-selector.tsx @@ -1,7 +1,7 @@ 'use client'; import React, { useState, useMemo } from 'react'; -import { useQuery, useMutation, useQueryClient } from '@tanstack/react-query'; +import { useMutation, useQueryClient } from '@tanstack/react-query'; import { type Model, type GroupSettings } from '@/lib/api/schemas/models'; import { AdminService, @@ -13,6 +13,7 @@ import { AddProviderModelDialog } from '@/components/add-provider-model-dialog'; import { EditGroupForm } from '@/components/edit-group-form'; import { ModelProviderSection } from '@/components/model-provider-section'; import { useDisplayCurrency } from '@/lib/hooks/use-display-currency'; +import { useModelsWithProviders } from '@/lib/hooks/use-models-with-providers'; import { Button } from '@/components/ui/button'; import { Checkbox } from '@/components/ui/checkbox'; import { Skeleton } from '@/components/ui/skeleton'; @@ -131,19 +132,14 @@ export function ModelSelector({ const queryClient = useQueryClient(); - // Fetch models and groups + // Shared with the page shell, so mounting this panel costs no extra fetch. const { - data: modelsData, + models, + groups, isLoading: isLoadingModels, error: modelsError, refetch: refetchModels, - } = useQuery({ - queryKey: ['models-with-providers'], - queryFn: () => AdminService.getModelsWithProviders(), - refetchOnWindowFocus: false, - }); - - const { models = [], groups = [] } = modelsData || {}; + } = useModelsWithProviders(); const allOverrideModels = useMemo( () => models.filter(isOverrideModel), [models] diff --git a/ui/components/models-page.tsx b/ui/components/models-page.tsx index bff3a379..2bfc2149 100644 --- a/ui/components/models-page.tsx +++ b/ui/components/models-page.tsx @@ -1,16 +1,14 @@ 'use client'; import { useMemo, useState } from 'react'; -import { useQuery } from '@tanstack/react-query'; +import dynamic from 'next/dynamic'; import { AlertCircle } from 'lucide-react'; import type { Model } from '@/lib/api/schemas/models'; -import { AdminService } from '@/lib/api/services/admin'; +import { useModelsWithProviders } from '@/lib/hooks/use-models-with-providers'; import { groupAndSortModelsByProvider } from '@/lib/utils/model-sort'; import { AppPageShell } from '@/components/app-page-shell'; import { PageHeader } from '@/components/page-header'; import { ModelSelector } from '@/components/model-selector'; -import { ModelTester } from '@/components/model-tester'; -import { ApiEndpointTester } from '@/components/api-endpoint-tester'; import { ModelSearchFilter } from '@/components/model-search-filter'; import { Alert, AlertDescription } from '@/components/ui/alert'; import { @@ -23,6 +21,19 @@ import { import { Skeleton } from '@/components/ui/skeleton'; import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs'; +// The testing tabs are never the landing view, so keeping them out of this +// route's chunk is what lets the navigation itself resolve quickly. +const ModelTester = dynamic( + () => import('@/components/model-tester').then((m) => m.ModelTester), + { loading: () => , ssr: false } +); + +const ApiEndpointTester = dynamic( + () => + import('@/components/api-endpoint-tester').then((m) => m.ApiEndpointTester), + { loading: () => , ssr: false } +); + export function ModelsPage() { const [filteredModels, setFilteredModels] = useState( undefined @@ -31,16 +42,11 @@ export function ModelsPage() { useState('all'); const { - data: modelsData, + models, + groups, isLoading: isLoadingModels, error: modelsError, - } = useQuery({ - queryKey: ['admin-models-with-providers'], - queryFn: () => AdminService.getModelsWithProviders(), - refetchOnWindowFocus: false, - }); - - const { models = [], groups = [] } = modelsData || {}; + } = useModelsWithProviders(); const groupedModels = useMemo( () => groupAndSortModelsByProvider(models), diff --git a/ui/lib/api/services/admin.ts b/ui/lib/api/services/admin.ts index b3781f6e..52ec852f 100644 --- a/ui/lib/api/services/admin.ts +++ b/ui/lib/api/services/admin.ts @@ -501,14 +501,50 @@ export class AdminService { const allModels: AdminModelAsModel[] = []; const seenModelIds = new Set(); - for (const provider of providers) { - try { - const providerModels = await this.getProviderModels(provider.id); + // One provider's catalog never depends on another's, and each miss costs an + // upstream round trip, so the whole fan-out happens in a single wave. + const providerResults = await Promise.all( + providers.map(async (provider) => { + try { + return { + provider, + models: await this.getProviderModels(provider.id), + }; + } catch (error) { + console.error( + `Failed to fetch models for provider ${provider.id}:`, + error + ); + return null; + } + }) + ); - providerModels.db_models.forEach((dbModel) => { - seenModelIds.add(dbModel.id); + for (const result of providerResults) { + if (!result) { + continue; + } + const { provider, models: providerModels } = result; + providerModels.db_models.forEach((dbModel) => { + seenModelIds.add(dbModel.id); + const modelWithProvider = { + ...dbModel, + upstream_provider_id: provider.id, + }; + allModels.push({ + ...this.transformAdminModelToModel( + modelWithProvider, + provider.provider_type + ), + has_own_api_key: false, + api_key_type: 'group', + }); + }); + + providerModels.remote_models.forEach((remoteModel) => { + if (!seenModelIds.has(remoteModel.id)) { const modelWithProvider = { - ...dbModel, + ...remoteModel, upstream_provider_id: provider.id, }; allModels.push({ @@ -517,33 +553,11 @@ export class AdminService { provider.provider_type ), has_own_api_key: false, - api_key_type: 'group', + api_key_type: 'remote', + soft_deleted: false, }); - }); - - providerModels.remote_models.forEach((remoteModel) => { - if (!seenModelIds.has(remoteModel.id)) { - const modelWithProvider = { - ...remoteModel, - upstream_provider_id: provider.id, - }; - allModels.push({ - ...this.transformAdminModelToModel( - modelWithProvider, - provider.provider_type - ), - has_own_api_key: false, - api_key_type: 'remote', - soft_deleted: false, - }); - } - }); - } catch (error) { - console.error( - `Failed to fetch models for provider ${provider.id}:`, - error - ); - } + } + }); } return { models: allModels, groups }; @@ -993,6 +1007,7 @@ export class AdminService { static async getLightningInvoices( status?: string, purpose?: string, + direction?: string, search?: string, limit: number = 50, offset: number = 0 @@ -1000,6 +1015,7 @@ export class AdminService { const params = new URLSearchParams(); if (status) params.append('status', status); if (purpose) params.append('purpose', purpose); + if (direction) params.append('direction', direction); if (search) params.append('search', search); params.append('limit', limit.toString()); params.append('offset', offset.toString()); @@ -1359,9 +1375,17 @@ export interface LightningInvoice { amount_sats: number; description: string; payment_hash: string; - status: 'pending' | 'paid' | 'expired' | 'cancelled'; + status: + | 'pending' + | 'settlement_pending' + | 'paid' + | 'failed' + | 'expired' + | 'cancelled' + | 'reconciliation_required'; api_key_hash: string | null; - purpose: 'create' | 'topup'; + direction: 'in' | 'out'; + purpose: 'create' | 'topup' | 'payout'; created_at: number; expires_at: number; paid_at: number | null; diff --git a/ui/lib/hooks/use-models-with-providers.ts b/ui/lib/hooks/use-models-with-providers.ts new file mode 100644 index 00000000..6aeec5cf --- /dev/null +++ b/ui/lib/hooks/use-models-with-providers.ts @@ -0,0 +1,26 @@ +'use client'; + +import { useQuery } from '@tanstack/react-query'; +import { AdminService } from '@/lib/api/services/admin'; + +export const modelsWithProvidersQueryKey = ['models-with-providers'] as const; + +/** + * Shared catalog read for every models view, so the page shell and the + * selector panel share one request instead of each fanning out to providers. + */ +export function useModelsWithProviders() { + const query = useQuery({ + queryKey: modelsWithProvidersQueryKey, + queryFn: () => AdminService.getModelsWithProviders(), + refetchOnWindowFocus: false, + }); + + return { + models: query.data?.models ?? [], + groups: query.data?.groups ?? [], + isLoading: query.isLoading, + error: query.error, + refetch: query.refetch, + }; +} diff --git a/ui/lib/hooks/use-progressive-list.ts b/ui/lib/hooks/use-progressive-list.ts new file mode 100644 index 00000000..8b55ffc4 --- /dev/null +++ b/ui/lib/hooks/use-progressive-list.ts @@ -0,0 +1,46 @@ +'use client'; + +import { useEffect, useState } from 'react'; + +/** + * Reveal a long list in frame-sized batches. + * + * A provider catalog can hold thousands of rows, and mounting them in one + * commit blocks the main thread long enough that the page looks frozen right + * after navigation. Each batch yields back to the browser, so the first rows + * paint immediately and the rest fill in without freezing input. + */ +export function useProgressiveList( + items: T[], + initialCount = 40, + step = 80 +): { visibleItems: T[]; hiddenCount: number } { + const [count, setCount] = useState(initialCount); + const [trackedItems, setTrackedItems] = useState(items); + + // Reset during render, not in an effect: an effect would first commit the new + // list at the old (possibly full) count, which is the freeze this avoids. + if (trackedItems !== items) { + setTrackedItems(items); + setCount(initialCount); + } + + useEffect(() => { + if (count >= items.length) { + return; + } + + const frame = requestAnimationFrame(() => { + setCount((current) => Math.min(items.length, current + step)); + }); + + return () => cancelAnimationFrame(frame); + }, [count, items.length, step]); + + const visibleCount = Math.min(count, items.length); + + return { + visibleItems: items.slice(0, visibleCount), + hiddenCount: items.length - visibleCount, + }; +} diff --git a/ui/package.json b/ui/package.json index 1d34f131..209c4fad 100644 --- a/ui/package.json +++ b/ui/package.json @@ -39,7 +39,7 @@ "@radix-ui/react-toggle-group": "^1.1.11", "@radix-ui/react-tooltip": "^1.2.8", "@tanstack/react-query": "^5.90.21", - "axios": "^1.16.0", + "axios": "^1.20.0", "class-variance-authority": "^0.7.1", "clsx": "^2.1.1", "cmdk": "^1.1.1", @@ -48,7 +48,7 @@ "geist": "^1.7.0", "input-otp": "^1.4.2", "lucide-react": "^0.575.0", - "next": "16.3.4", + "next": "16.3.6", "next-themes": "^0.4.6", "qrcode": "^1.5.4", "radix-ui": "^1.4.3", @@ -74,7 +74,7 @@ "@types/react": "^19.2.14", "@types/react-dom": "^19.2.3", "eslint": "^9.7.0", - "eslint-config-next": "16.3.4", + "eslint-config-next": "16.3.6", "eslint-config-prettier": "^10.1.8", "eslint-plugin-prettier": "^5.5.5", "eslint-plugin-react": "^7.37.5", diff --git a/ui/pnpm-lock.yaml b/ui/pnpm-lock.yaml index 097f3490..836a06f6 100644 --- a/ui/pnpm-lock.yaml +++ b/ui/pnpm-lock.yaml @@ -7,8 +7,8 @@ settings: overrides: '@babel/core': 7.29.6 ajv@6: 6.14.0 - brace-expansion@1: 1.1.18 - brace-expansion@5: 5.0.9 + brace-expansion@1: 1.1.21 + brace-expansion@5: 5.0.12 flatted: 3.4.2 follow-redirects: 1.16.0 form-data: 4.0.6 @@ -103,8 +103,8 @@ importers: specifier: ^5.90.21 version: 5.90.21(react@19.2.4) axios: - specifier: ^1.16.0 - version: 1.18.1 + specifier: ^1.20.0 + version: 1.20.0 class-variance-authority: specifier: ^0.7.1 version: 0.7.1 @@ -122,7 +122,7 @@ importers: version: 8.6.0(react@19.2.4) geist: specifier: ^1.7.0 - version: 1.7.0(next@16.3.4(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4)) + version: 1.7.0(next@16.3.6(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4)) input-otp: specifier: ^1.4.2 version: 1.4.2(react-dom@19.2.4(react@19.2.4))(react@19.2.4) @@ -130,8 +130,8 @@ importers: specifier: ^0.575.0 version: 0.575.0(react@19.2.4) next: - specifier: 16.3.4 - version: 16.3.4(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4) + specifier: 16.3.6 + version: 16.3.6(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4) next-themes: specifier: ^0.4.6 version: 0.4.6(react-dom@19.2.4(react@19.2.4))(react@19.2.4) @@ -203,8 +203,8 @@ importers: specifier: ^9.7.0 version: 9.38.0(jiti@2.6.1) eslint-config-next: - specifier: 16.3.4 - version: 16.3.4(@typescript-eslint/parser@8.57.0(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3))(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3) + specifier: 16.3.6 + version: 16.3.6(@typescript-eslint/parser@8.57.0(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3))(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3) eslint-config-prettier: specifier: ^10.1.8 version: 10.1.8(eslint@9.38.0(jiti@2.6.1)) @@ -580,56 +580,56 @@ packages: '@napi-rs/wasm-runtime@0.2.12': resolution: {integrity: sha512-ZVWUcfwY4E/yPitQJl481FjFo3K22D6qF0DuFH6Y/nbnE11GY5uguDxZMGXPQ8WQ0128MXQD7TnfHyK4oWoIJQ==} - '@next/env@16.3.4': - resolution: {integrity: sha512-cjWZnUUa6jZq2kFaNe/ZyJdZonOZ/QoN0Zka2nz/FLOrfx14pQuM9c5RaSVkWMqgdt4ksgPAMWPyHSs/CyV48Q==} + '@next/env@16.3.6': + resolution: {integrity: sha512-x9Vblze1EbtltQYnNH38xCPWU3TVfBd1eXqA3+w9+BTpedkkdNpAaltXlGQ/nsc1+E0mVTNrtcbX3GoO09zeLQ==} - '@next/eslint-plugin-next@16.3.4': - resolution: {integrity: sha512-szW9y2Aumu4z88YXfTzcFsgUAg2k64uzbtcO5L9f1AKS4w/GUKJcbFllRflROVyNPgJtGOnvNxiyp3v6b+prIA==} + '@next/eslint-plugin-next@16.3.6': + resolution: {integrity: sha512-jowwDX+7DOlDIjJLgTMxudw+k37QnWu1JkZLkSi9MaJBfDYcfhAPMKBhXL0idYzFN/AGg//axnOR4cLkHX/Rng==} - '@next/swc-darwin-arm64@16.3.4': - resolution: {integrity: sha512-iBr3I5LZNk5/bgl5//iTgD2tcym14MX0Xo7fD//u9dYAEgGzza1y9oywluPtf74YnOswVdH1908aK9xVz7zQTw==} + '@next/swc-darwin-arm64@16.3.6': + resolution: {integrity: sha512-E/7GEqaUkt8mk/T8v9lAnrhzR06kdq1ZBkC12F8tAMkdIadwNp3H1KqHynDHrpcTlGCUdq/qu6vUL2aYVyYBdw==} engines: {node: '>= 10'} cpu: [arm64] os: [darwin] - '@next/swc-darwin-x64@16.3.4': - resolution: {integrity: sha512-2dpiSyl2Jw/NrBPaU2MAKGSa+2MR82pJIn4Sm5Rjr+gxAeuh0z158Su3Z2O8zn7UNNq+ej4bToed6RcRN/Lydg==} + '@next/swc-darwin-x64@16.3.6': + resolution: {integrity: sha512-yBE893/nDWTlaiBD1p+qgt7NUen4U5R6FXyH0s67Npq1S3E0cVSef1WIXC2xBRgQvwAvJq6DnS6Y6PrY0cy4Ew==} engines: {node: '>= 10'} cpu: [x64] os: [darwin] - '@next/swc-linux-arm64-gnu@16.3.4': - resolution: {integrity: sha512-+t+U8HZT+fApePCS5h89CSH3datz29MkzyfCn+6fpsZBG/oiEOhINcb9rtkv6sdpToLGFn2e6146NzaKCXkqrA==} + '@next/swc-linux-arm64-gnu@16.3.6': + resolution: {integrity: sha512-KJDpjBqBPYlvkivmyrp+Qys6k/7ksbqGQvRVc6ZEGfR+cjQxx+nUkJaWmNZJsmoOrqYNbaXByF8wa0lBwDhB3Q==} engines: {node: '>= 10'} cpu: [arm64] os: [linux] - '@next/swc-linux-arm64-musl@16.3.4': - resolution: {integrity: sha512-mx03GNs1ocQA5JQ4FxDMmIsNkdrZh8cuezKCrId28e5/gIPU/l7Kcy2+vmCCzdjnnmXJy+iOAu+7K0QppO6Urg==} + '@next/swc-linux-arm64-musl@16.3.6': + resolution: {integrity: sha512-mqNg2K+hvWskSRb/QM+Ix412DvBsuSF0XV+frTSw5vmoucNnIlynFwKYew8D01bfATErMOM7Bujrf0BA5DRKFA==} engines: {node: '>= 10'} cpu: [arm64] os: [linux] - '@next/swc-linux-x64-gnu@16.3.4': - resolution: {integrity: sha512-YIhGY6fSMfha52bnVxnzc9zaVBzJg+cqQTOD8tXIBSx4fuv0pVMxQTE0PaS59YhnMOiYiG09IMwxJAf/CFm/Dw==} + '@next/swc-linux-x64-gnu@16.3.6': + resolution: {integrity: sha512-nFncBNGAYouRHjRVaITs9beZRfhX4ssVwpnvPIAbkZVH6LtGoAVlH4bJ8Cnf9SOo9bsXgPFer/GdHtEE3JNOkw==} engines: {node: '>= 10'} cpu: [x64] os: [linux] - '@next/swc-linux-x64-musl@16.3.4': - resolution: {integrity: sha512-+eaaX6axpDb0yF1GCpiERe6njplvdC+nks/fKfcHu3XPGRrald8P3/X7yv7QLdjA51knnxwl9pxdIJsg+w1L+Q==} + '@next/swc-linux-x64-musl@16.3.6': + resolution: {integrity: sha512-5Mf3cHDGR/Iz0ng2Bj3zUR3p5QS9YK3Hn2QiAfavFmyF48zwThAjpFoiTKNIcOHLYS4zEk+gzyJ/9deQ2ZB8yQ==} engines: {node: '>= 10'} cpu: [x64] os: [linux] - '@next/swc-win32-arm64-msvc@16.3.4': - resolution: {integrity: sha512-0jcXW7Xs/uzICrmgV3MhDYDeRy++1CqnpDIerlPIqYO4bhzB4WNbX/aRnQclustsAyTkFKB0z6rbcjmNg5tR8A==} + '@next/swc-win32-arm64-msvc@16.3.6': + resolution: {integrity: sha512-0jkJy0C2kbrJWTk4YLa3xk80pVBpx8FCHJym7CnUfDAXe/FWv5qT7SQJbR0KuemyxaEDlEx5WT4VQJoTW+/9Qw==} engines: {node: '>= 10'} cpu: [arm64] os: [win32] - '@next/swc-win32-x64-msvc@16.3.4': - resolution: {integrity: sha512-vvBzwu1pYQCp92maZCFCIw/XgOTMR5tur9GjakwIo2cmwRTMKajRZZDS9+e4KsUZWKu1E007WUeAFXRRjZeuzw==} + '@next/swc-win32-x64-msvc@16.3.6': + resolution: {integrity: sha512-/YXjI1e5OXcZ7YpxRwgP/1jAV/SBKTzeVKqN2mk7mLpcICsyn3Gl5+dIfDTJp70M0ccMhyMMRso4v6mPDCGepg==} engines: {node: '>= 10'} cpu: [x64] os: [win32] @@ -1842,8 +1842,8 @@ packages: resolution: {integrity: sha512-BASOg+YwO2C+346x3LZOeoovTIoTrRqEsqMa6fmfAV0P+U9mFr9NsyOEpiYvFjbc64NMrSswhV50WdXzdb/Z5A==} engines: {node: '>=4'} - axios@1.18.1: - resolution: {integrity: sha512-3nTvFlvpn9Zu/RkHUqtc7/+al4UpRW5az71ap5zccp6e8RAYEzhMTecX8Dz1wWDYrPpUoB1HAQEGEAEvUr7S9g==} + axios@1.20.0: + resolution: {integrity: sha512-r8aOh8j9cGKpgQAqpzrUHnSIc6a59Y3Xf/cv8sy1DrHCkZHzQGEuoq1tARk6qSyDdtQGSDgpb9kFlruzPvrgwg==} axobject-query@4.1.0: resolution: {integrity: sha512-qIj0G9wZbMGNLjLmg1PT6v2mE9AH2zlnADJD/2tC6E00hgmhUOfEB6greHPAfLRSufHqROIUTkw6E+M3lH0PTQ==} @@ -1861,11 +1861,11 @@ packages: engines: {node: '>=6.0.0'} hasBin: true - brace-expansion@1.1.18: - resolution: {integrity: sha512-Edep/X9fGqVNmzKBVsDYIOtD+z1tuezV70LBjdCst9Tqu76lsnvRiZ6oTic1n+/BIwX6QDGAO94PN4N2SADvtw==} + brace-expansion@1.1.21: + resolution: {integrity: sha512-9zeA+KLZNNzglF2TPKRQEDyx6Yby7daAkuy8MiPzpXPsYDWi/DRM8jmwUDxokQjYqBpv5DgPiwD4h4ZZSy1Ujw==} - brace-expansion@5.0.9: - resolution: {integrity: sha512-ScQ4IuvIEF1TMlP7Zt+vjJ//9zlPb2SDcxWxM3bk8s6t6GGdJ7KO1dCcTidOPJKePW30LE/2cT7wCyPho9/Wxg==} + brace-expansion@5.0.12: + resolution: {integrity: sha512-YovQ3rzhaLMIrDjNDMkNS01tea93qhEhG5xy8f6+R0l+dw3Ki+5sCoIoI942iuLZTHWogWktgwVDhU09iNEimQ==} engines: {node: 20 || >=22} braces@3.0.3: @@ -2142,8 +2142,8 @@ packages: resolution: {integrity: sha512-TtpcNJ3XAzx3Gq8sWRzJaVajRs0uVxA2YAkdb1jm2YkPz4G6egUFAyA3n5vtEIZefPk5Wa4UXbKuS5fKkJWdgA==} engines: {node: '>=10'} - eslint-config-next@16.3.4: - resolution: {integrity: sha512-35/8RM10huEL9vlr8hUZMERMENHBrnyHN3ZZkF9efSgzGaqK34jIqry44A956//zriUhUAUW0XSkcolhrryqAA==} + eslint-config-next@16.3.6: + resolution: {integrity: sha512-1Upt3U7BDwU+ilpe2byZjAfts9oNq4d4fv/zXEvs8/4yS+cwOQW/WCxUNy8gCDquX67SzeehDvKblVC6ZBMocQ==} peerDependencies: eslint: '>=9.0.0' typescript: '>=3.3.1' @@ -2821,8 +2821,8 @@ packages: react: ^16.8 || ^17 || ^18 || ^19 || ^19.0.0-rc react-dom: ^16.8 || ^17 || ^18 || ^19 || ^19.0.0-rc - next@16.3.4: - resolution: {integrity: sha512-/Ztf6CeRH+ejEXUrYtqI4gkS66eFIHuSwqi60RgcpWKodxFZx2/dqVCMKBwILfAHXQ+F1b1vAudgj3mnxqtoIA==} + next@16.3.6: + resolution: {integrity: sha512-L+otWM/aQbYTx98aZhgEoMb4bZAXx1YVW4UMA/vuCyCoWG5HJyZUili8QAkqzrcC+5///tsz3s0M+SlyB5bLMw==} engines: {node: '>=20.9.0'} hasBin: true peerDependencies: @@ -3899,37 +3899,37 @@ snapshots: '@tybys/wasm-util': 0.10.1 optional: true - '@next/env@16.3.4': {} + '@next/env@16.3.6': {} - '@next/eslint-plugin-next@16.3.4(eslint@9.38.0(jiti@2.6.1))': + '@next/eslint-plugin-next@16.3.6(eslint@9.38.0(jiti@2.6.1))': dependencies: '@eslint-community/eslint-utils': 4.9.1(eslint@9.38.0(jiti@2.6.1)) fast-glob: 3.3.1 transitivePeerDependencies: - eslint - '@next/swc-darwin-arm64@16.3.4': + '@next/swc-darwin-arm64@16.3.6': optional: true - '@next/swc-darwin-x64@16.3.4': + '@next/swc-darwin-x64@16.3.6': optional: true - '@next/swc-linux-arm64-gnu@16.3.4': + '@next/swc-linux-arm64-gnu@16.3.6': optional: true - '@next/swc-linux-arm64-musl@16.3.4': + '@next/swc-linux-arm64-musl@16.3.6': optional: true - '@next/swc-linux-x64-gnu@16.3.4': + '@next/swc-linux-x64-gnu@16.3.6': optional: true - '@next/swc-linux-x64-musl@16.3.4': + '@next/swc-linux-x64-musl@16.3.6': optional: true - '@next/swc-win32-arm64-msvc@16.3.4': + '@next/swc-win32-arm64-msvc@16.3.6': optional: true - '@next/swc-win32-x64-msvc@16.3.4': + '@next/swc-win32-x64-msvc@16.3.6': optional: true '@nodelib/fs.scandir@2.1.5': @@ -5168,7 +5168,7 @@ snapshots: axe-core@4.11.1: {} - axios@1.18.1: + axios@1.20.0: dependencies: follow-redirects: 1.16.0 form-data: 4.0.6 @@ -5186,12 +5186,12 @@ snapshots: baseline-browser-mapping@2.11.21: {} - brace-expansion@1.1.18: + brace-expansion@1.1.21: dependencies: balanced-match: 1.0.2 concat-map: 0.0.1 - brace-expansion@5.0.9: + brace-expansion@5.0.12: dependencies: balanced-match: 4.0.4 @@ -5576,9 +5576,9 @@ snapshots: escape-string-regexp@4.0.0: {} - eslint-config-next@16.3.4(@typescript-eslint/parser@8.57.0(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3))(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3): + eslint-config-next@16.3.6(@typescript-eslint/parser@8.57.0(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3))(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3): dependencies: - '@next/eslint-plugin-next': 16.3.4(eslint@9.38.0(jiti@2.6.1)) + '@next/eslint-plugin-next': 16.3.6(eslint@9.38.0(jiti@2.6.1)) eslint: 9.38.0(jiti@2.6.1) eslint-import-resolver-node: 0.3.9 eslint-import-resolver-typescript: 3.10.1(eslint-plugin-import@2.32.0)(eslint@9.38.0(jiti@2.6.1)) @@ -5872,9 +5872,9 @@ snapshots: functions-have-names@1.2.3: {} - geist@1.7.0(next@16.3.4(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4)): + geist@1.7.0(next@16.3.6(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4)): dependencies: - next: 16.3.4(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4) + next: 16.3.6(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4) generator-function@2.0.1: {} @@ -6263,11 +6263,11 @@ snapshots: minimatch@10.2.4: dependencies: - brace-expansion: 5.0.9 + brace-expansion: 5.0.12 minimatch@3.1.4: dependencies: - brace-expansion: 1.1.18 + brace-expansion: 1.1.21 minimist@1.2.8: {} @@ -6284,9 +6284,9 @@ snapshots: react: 19.2.4 react-dom: 19.2.4(react@19.2.4) - next@16.3.4(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4): + next@16.3.6(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4): dependencies: - '@next/env': 16.3.4 + '@next/env': 16.3.6 '@swc/helpers': 0.5.23 baseline-browser-mapping: 2.11.21 caniuse-lite: 1.0.30001810 @@ -6295,14 +6295,14 @@ snapshots: react-dom: 19.2.4(react@19.2.4) styled-jsx: 5.1.6(@babel/core@7.29.6)(react@19.2.4) optionalDependencies: - '@next/swc-darwin-arm64': 16.3.4 - '@next/swc-darwin-x64': 16.3.4 - '@next/swc-linux-arm64-gnu': 16.3.4 - '@next/swc-linux-arm64-musl': 16.3.4 - '@next/swc-linux-x64-gnu': 16.3.4 - '@next/swc-linux-x64-musl': 16.3.4 - '@next/swc-win32-arm64-msvc': 16.3.4 - '@next/swc-win32-x64-msvc': 16.3.4 + '@next/swc-darwin-arm64': 16.3.6 + '@next/swc-darwin-x64': 16.3.6 + '@next/swc-linux-arm64-gnu': 16.3.6 + '@next/swc-linux-arm64-musl': 16.3.6 + '@next/swc-linux-x64-gnu': 16.3.6 + '@next/swc-linux-x64-musl': 16.3.6 + '@next/swc-win32-arm64-msvc': 16.3.6 + '@next/swc-win32-x64-msvc': 16.3.6 sharp: 0.35.4(@types/node@25.4.0) transitivePeerDependencies: - '@babel/core' diff --git a/ui/pnpm-workspace.yaml b/ui/pnpm-workspace.yaml index 317f28e0..5eb49756 100644 --- a/ui/pnpm-workspace.yaml +++ b/ui/pnpm-workspace.yaml @@ -5,8 +5,8 @@ onlyBuiltDependencies: overrides: '@babel/core': 7.29.6 ajv@6: 6.14.0 - brace-expansion@1: 1.1.18 - brace-expansion@5: 5.0.9 + brace-expansion@1: 1.1.21 + brace-expansion@5: 5.0.12 flatted: 3.4.2 follow-redirects: 1.16.0 form-data: 4.0.6 diff --git a/uv.lock b/uv.lock index 1da5b968..6abbc6ab 100644 --- a/uv.lock +++ b/uv.lock @@ -9,10 +9,13 @@ resolution-markers = [ [manifest] constraints = [ + { name = "anyio", specifier = ">=4.14.2" }, { name = "grpcio", specifier = ">=1.76.0,<2.0.0" }, { name = "grpcio-tools", specifier = ">=1.76.0,<2.0.0" }, { name = "httpcore", specifier = ">=1.0.9" }, + { name = "pyjwt", specifier = ">=2.15.0" }, { name = "starlette", specifier = ">=1.3.1" }, + { name = "urllib3", specifier = ">=2.8.0" }, ] overrides = [ { name = "cryptography", specifier = ">=49.0.0" }, @@ -211,16 +214,15 @@ wheels = [ [[package]] name = "anyio" -version = "4.9.0" +version = "4.14.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "idna" }, - { name = "sniffio" }, { name = "typing-extensions", marker = "python_full_version < '3.13'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/95/7d/4c1bd541d4dffa1b52bd83fb8527089e097a106fc90b467a7313b105f840/anyio-4.9.0.tar.gz", hash = "sha256:673c0c244e15788651a4ff38710fea9675823028a6f08a5eda409e0c9840a028", size = 190949, upload-time = "2025-03-17T00:02:54.77Z" } +sdist = { url = "https://files.pythonhosted.org/packages/61/cc/a381afa6efea9f496eff839d4a6a1aed3bfafc7b3ab4b0d1b243a12573dd/anyio-4.14.2.tar.gz", hash = "sha256:cfa139f3ed1a23ee8f88a145ddb5ac7605b8bbfd8592baacd7ce3d8bb4313c7f", size = 260176, upload-time = "2026-07-12T20:29:07.082Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/a1/ee/48ca1a7c89ffec8b6a0c5d02b89c305671d5ffd8d3c94acf8b8c408575bb/anyio-4.9.0-py3-none-any.whl", hash = "sha256:9f76d541cad6e36af7beb62e978876f3b41e3e04f2c1fbf0884604c0a9c4d93c", size = 100916, upload-time = "2025-03-17T00:02:52.713Z" }, + { url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" }, ] [[package]] @@ -282,6 +284,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/77/06/bb80f5f86020c4551da315d78b3ab75e8228f89f0162f2c3a819e407941a/attrs-25.3.0-py3-none-any.whl", hash = "sha256:427318ce031701fea540783410126f03899a97ffc6f61596ad581ac2e40e3bc3", size = 63815, upload-time = "2025-03-13T11:10:21.14Z" }, ] +[[package]] +name = "backoff" +version = "2.2.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/47/d7/5bbeb12c44d7c4f2fb5b56abce497eb5ed9f34d85701de869acedd602619/backoff-2.2.1.tar.gz", hash = "sha256:03f829f5bb1923180821643f8753b0502c3b682293992485b0eef2807afa5cba", size = 17001, upload-time = "2022-10-05T19:19:32.061Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/df/73/b6e24bd22e6720ca8ee9a85a0c4a2971af8497d8f3193fa05390cbd46e09/backoff-2.2.1-py3-none-any.whl", hash = "sha256:63579f9a0628e06278f7e47b7d7d5b6ce20dc65c5e96a6f3ca99a6adca0396e8", size = 15148, upload-time = "2022-10-05T19:19:30.546Z" }, +] + [[package]] name = "base58" version = "2.1.1" @@ -337,6 +348,34 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/9d/9e/78e59887cbf94116bdc890af7726ae264d55df14f1c777724c656e8a35fe/bolt11-2.1.1-py3-none-any.whl", hash = "sha256:fd4edb9e73e27bf5e017f47c97f7c6827b523fcf9cab152b123961ca78323e2d", size = 17102, upload-time = "2025-03-12T13:33:08.142Z" }, ] +[[package]] +name = "boto3" +version = "1.43.105" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "botocore" }, + { name = "jmespath" }, + { name = "s3transfer" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/75/46/d8c87ada70a7647fb3d206c7f19eafca3580a0ae4c06d62da539a1ee1207/boto3-1.43.105.tar.gz", hash = "sha256:e51260aed9cc1474778b5488bc6f97ad28f27a0a7002f4bbaaf8191aff1422ea", size = 112682, upload-time = "2026-09-29T19:37:40.784Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/bc/8e/0310a37ff609529dab9153cbc9fd0b66c685364d741bb6d1ae31134b728e/boto3-1.43.105-py3-none-any.whl", hash = "sha256:b8b6236ae7fe2724eee608c9b0649afbb86f0ec98158f39f67e64e678ec47499", size = 140042, upload-time = "2026-09-29T19:37:39.415Z" }, +] + +[[package]] +name = "botocore" +version = "1.43.105" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "jmespath" }, + { name = "python-dateutil" }, + { name = "urllib3" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/2b/30/668f3c0533a440787e212cf56404cb6ec234ae8e6baf97fe17329d512d88/botocore-1.43.105.tar.gz", hash = "sha256:afb3e7706b123ab069d1c34571ca1fdf82528a48425574fe4693df3d039d503f", size = 16263910, upload-time = "2026-09-29T19:37:36.456Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/f3/94/50923cd46840e4d2b56cad1dcf5008fb20c099b04b2d93d32231f2d0bfaa/botocore-1.43.105-py3-none-any.whl", hash = "sha256:7abd19e1ef2c5e4a0314ca493fa7cebefabe33e559d7dd570fe2432a5431e6ec", size = 15958067, upload-time = "2026-09-29T19:37:33.373Z" }, +] + [[package]] name = "brotli" version = "1.2.0" @@ -1509,6 +1548,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/b3/4a/4175a563579e884192ba6e81725fc0448b042024419be8d83aa8a80a3f44/jiter-0.10.0-cp314-cp314t-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:3aa96f2abba33dc77f79b4cf791840230375f9534e5fac927ccceb58c5e604a5", size = 354213, upload-time = "2025-05-18T19:04:41.894Z" }, ] +[[package]] +name = "jmespath" +version = "1.1.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/d3/59/322338183ecda247fb5d1763a6cbe46eff7222eaeebafd9fa65d4bf5cb11/jmespath-1.1.0.tar.gz", hash = "sha256:472c87d80f36026ae83c6ddd0f1d05d4e510134ed462851fd5f754c8c3cbb88d", size = 27377, upload-time = "2026-01-22T16:35:26.279Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/14/2f/967ba146e6d58cf6a652da73885f52fc68001525b4197effc174321d70b4/jmespath-1.1.0-py3-none-any.whl", hash = "sha256:a5663118de4908c91729bea0acadca56526eb2698e83de10cd116ae0f4e97c64", size = 20419, upload-time = "2026-01-22T16:35:24.919Z" }, +] + [[package]] name = "jsonschema" version = "4.26.0" @@ -1552,10 +1600,11 @@ wheels = [ [[package]] name = "litellm" -version = "1.93.2" +version = "1.101.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "aiohttp" }, + { name = "boto3" }, { name = "click" }, { name = "fastuuid" }, { name = "httpx", extra = ["socks"] }, @@ -1564,40 +1613,20 @@ dependencies = [ { name = "jsonschema" }, { name = "openai" }, { name = "pydantic" }, + { name = "pydantic-settings" }, { name = "python-dotenv" }, { name = "tiktoken" }, { name = "tokenizers" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/97/dd/28024c0e4cf2dc6ab1bad59b8357af7f460e952c69526eae28f12ac4ee5e/litellm-1.93.2.tar.gz", hash = "sha256:c5d5223ef07f36e0886397fb45cc9db4150f86a0c6f6835cee1d5524cab69dfd", size = 15955441, upload-time = "2026-08-09T02:17:49.646Z" } +sdist = { url = "https://files.pythonhosted.org/packages/26/c9/cb2730c6c763233e322fe7c5b2f53783eb10893cea9304e5474f1f20c306/litellm-1.101.2.tar.gz", hash = "sha256:790adf4ce19116d7bf4342492b1be5a90dd56e08d89795979bf6c1c3446a9670", size = 17493188, upload-time = "2026-09-24T00:04:22.712Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/64/c7/cb3f49dc60d57dda7fe368310fd5da2a94ec9b6a746bcf343a61e10bdeda/litellm-1.93.2-cp311-cp311-macosx_10_12_x86_64.whl", hash = "sha256:1bd0690efc94357e559de97927fd98437555cd5b5dd832544cfcca87297ccb80", size = 19938326, upload-time = "2026-08-09T02:16:38.041Z" }, - { url = "https://files.pythonhosted.org/packages/0c/bd/d77184fdaaf57d67d65da91dcfc61c7f656703e7ce4f950e07523e7de4e3/litellm-1.93.2-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:845ececc628737909b1422d1af18bd19ae453727a66244aa9da3ca37a3773111", size = 19862606, upload-time = "2026-08-09T02:16:40.653Z" }, - { url = "https://files.pythonhosted.org/packages/53/99/d8dd58b6840754a13cc2e1111b283aa28cbfc0ccc653a8725050916bb08e/litellm-1.93.2-cp311-cp311-manylinux_2_28_aarch64.whl", hash = "sha256:498f9878ea773305e0638b6159d7e1ef27bb0b9a4292538d6634312d18a4e781", size = 20168532, upload-time = "2026-08-09T02:16:42.997Z" }, - { url = "https://files.pythonhosted.org/packages/d7/ca/559ca0f5e0b99b9f641086ae924c782f8d521d09384fbe9abbe0bddb6e61/litellm-1.93.2-cp311-cp311-manylinux_2_28_x86_64.whl", hash = "sha256:1e5618ef495b2e02299b376ca84ffb2647837aafee478cf3a1be17d47a8f0f73", size = 20162696, upload-time = "2026-08-09T02:16:45.283Z" }, - { url = "https://files.pythonhosted.org/packages/92/3e/18c31b27c7d1271b43bdc8ffbef01bfba68d90248bbe60bb2130dd17e43c/litellm-1.93.2-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:c2da463d70c9fffbea9532fd000e035328f5266b399a2fb4c6c76b3470478337", size = 20233518, upload-time = "2026-08-09T02:16:47.87Z" }, - { url = "https://files.pythonhosted.org/packages/d9/98/a6bae7c52f09cd03487a040f98eeedb899b3cf3fc541b87c6d051ee92e0d/litellm-1.93.2-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:2cf122399f84f8f04621ed6ef8f276dd6d61f4fab108932ce0e30368de34dd42", size = 20291180, upload-time = "2026-08-09T02:16:50.549Z" }, - { url = "https://files.pythonhosted.org/packages/77/2d/81d974f2533cf039afda7e3e0f769dc73dc692c75ec867cf29ec6f41c06f/litellm-1.93.2-cp311-cp311-win_amd64.whl", hash = "sha256:8eaaf780fab9a19234735ef94225172179d15bc28b67ddbec125194249a504b7", size = 19775654, upload-time = "2026-08-09T02:16:53.162Z" }, - { url = "https://files.pythonhosted.org/packages/d0/05/72fd8051f0f2f3c84b90986e6f4551db7c8b190ba3300f111461b7701689/litellm-1.93.2-cp312-cp312-macosx_10_12_x86_64.whl", hash = "sha256:3bf532c164ad7cb1b76f2c62afefdcc656b9b296374d075a4150e2ce10bb74c3", size = 19937403, upload-time = "2026-08-09T02:16:55.545Z" }, - { url = "https://files.pythonhosted.org/packages/9e/4d/5081b39bdb73cab04f8a86294a4534a029cf0434ac6932c7ae8049d55723/litellm-1.93.2-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:526b7afc037f79dfdd5c607f5085ac597c7fd301a6dedabea40baae899b27f19", size = 19853652, upload-time = "2026-08-09T02:16:57.977Z" }, - { url = "https://files.pythonhosted.org/packages/70/3f/fb70691266a7fd08c202406abea0153e82fa17f134cd9d58e4029cc741db/litellm-1.93.2-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:294ad19f356f821ce97a5428d09439be5f38d22b218c73008d8a49e3e42eb145", size = 20165680, upload-time = "2026-08-09T02:17:00.65Z" }, - { url = "https://files.pythonhosted.org/packages/81/91/84424ce2a25595463e5d24e9cf8949877cd4ce93c0fcbf6486ecd685094f/litellm-1.93.2-cp312-cp312-manylinux_2_28_x86_64.whl", hash = "sha256:6f6a5e3907f0a1c9d8ff8d71a6cbac8a592e47a40da3f97167074947b5ba7d11", size = 20157772, upload-time = "2026-08-09T02:17:03.027Z" }, - { url = "https://files.pythonhosted.org/packages/8f/8d/b0eac7ee6d174564f820565c8c9a726ae83dbb8c4d3522daf175b95da002/litellm-1.93.2-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:8541f1b7fd5c437ad249ad68d0a11f68e5e2866b0649da5fa7d63b595e9b8b22", size = 20229256, upload-time = "2026-08-09T02:17:05.271Z" }, - { url = "https://files.pythonhosted.org/packages/ee/6d/03e931c1cb2d1e1b7a968de21aa9e4db853928200da856c35c940ee6faa9/litellm-1.93.2-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:712c9387419d7b06a10df59973f5e530592d61b2314102b0fa3142f3743f9a9e", size = 20287257, upload-time = "2026-08-09T02:17:08.175Z" }, - { url = "https://files.pythonhosted.org/packages/16/05/6c0fe2fcf31c260474c55fabe4ecb0e9e1343c9b9132e28589391b2ad33e/litellm-1.93.2-cp312-cp312-win_amd64.whl", hash = "sha256:cc0d58ccabd22ef7ef44a9e6f7247deb54ae42f5e126e6f00360c2b28b41bc2b", size = 19772580, upload-time = "2026-08-09T02:17:11.254Z" }, - { url = "https://files.pythonhosted.org/packages/70/74/e9046cffa69b32b710452480598e418b26a29896ece680c80ec23997fd16/litellm-1.93.2-cp313-cp313-macosx_10_12_x86_64.whl", hash = "sha256:f4071bef03e4c2942cd2ddc752727345b85447d6a7fee1ff5a4f8b92187966b0", size = 19938095, upload-time = "2026-08-09T02:17:13.929Z" }, - { url = "https://files.pythonhosted.org/packages/fa/db/6ef38a7a2f73d5cc507423954fa535a8546ead375c4c71265c093bdb4e9e/litellm-1.93.2-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:8a99ac7c0c1b78acd6bfd1959e9f203dca71fdbceb5f0c8691c2ad8eee450d7d", size = 19854187, upload-time = "2026-08-09T02:17:16.588Z" }, - { url = "https://files.pythonhosted.org/packages/cb/b3/80ee0143b88e2921f8c8f24c7331478258a8bf25a3d4d4450bd96043403e/litellm-1.93.2-cp313-cp313-manylinux_2_28_aarch64.whl", hash = "sha256:a81ceff44c58ef504ab8bd787d03b82618765b9cfd530942386ae6d23c58be94", size = 20166307, upload-time = "2026-08-09T02:17:19.078Z" }, - { url = "https://files.pythonhosted.org/packages/98/60/cb326e1094f7042f28f9e21543d9f367a8aa25af6915bf4253b77da5c2a2/litellm-1.93.2-cp313-cp313-manylinux_2_28_x86_64.whl", hash = "sha256:dee1b02b7f52a5a408bf7c8d499f0834e49194651743758a511dcdd926c0b692", size = 20158336, upload-time = "2026-08-09T02:17:21.507Z" }, - { url = "https://files.pythonhosted.org/packages/b1/87/bad75146863531172c9dbae189486c7f4425b56a6641b55ab20745316048/litellm-1.93.2-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:d2edfa14b99bce706b35981703692e3ee631f9b87bf6dc28fb53b574f6480b20", size = 20229711, upload-time = "2026-08-09T02:17:24.073Z" }, - { url = "https://files.pythonhosted.org/packages/df/28/040b1853021ed8fd57be19eb2affb024d168951fe7e7abdbad91da3f6f3f/litellm-1.93.2-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:ae75a61c9abc827aa3131b7e640c952367a450830bb7c531b426b4ec2bb45f85", size = 20287584, upload-time = "2026-08-09T02:17:26.542Z" }, - { url = "https://files.pythonhosted.org/packages/d9/0b/4208815b0d666636cbf7afbd571eec3004d3a15d3150a23a9009fc2ce930/litellm-1.93.2-cp313-cp313-win_amd64.whl", hash = "sha256:c54a09ab20f94120a9d60a30d9970439dcefa00d2565d190505ff006a80c7a69", size = 19772641, upload-time = "2026-08-09T02:17:29.308Z" }, - { url = "https://files.pythonhosted.org/packages/09/4a/ff7a9c000519d2bab362318bf744a24c2500228e5fceaa6ac23acab96fa0/litellm-1.93.2-cp314-cp314-macosx_10_12_x86_64.whl", hash = "sha256:204cb0763fff9285bc87eb2dc0fc59b591999e5d94863d0964f806424d3c0cd6", size = 19943639, upload-time = "2026-08-09T02:17:31.811Z" }, - { url = "https://files.pythonhosted.org/packages/c4/26/29e9276ce4aa8ed133d9fd5ecc07375017d2228215547c6bbb17ccbc59b4/litellm-1.93.2-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:3126c84361606b9fb07fde7d57eccd8a1747304d64c4143e2e5e40ae6e7693fb", size = 19855435, upload-time = "2026-08-09T02:17:34.376Z" }, - { url = "https://files.pythonhosted.org/packages/f8/20/2c9c818248ae019b2d496ca41900a9a5651ab05e2400794cd8dc8b89b6d2/litellm-1.93.2-cp314-cp314-manylinux_2_28_aarch64.whl", hash = "sha256:1c84f7c4acb4e926a79b93145ab23231b300fc687bde7172ef884fc52d6011e0", size = 20166947, upload-time = "2026-08-09T02:17:36.828Z" }, - { url = "https://files.pythonhosted.org/packages/50/af/4016682be48350407837941ad1a1ae8185cca65b102e04e89eee2a2abccb/litellm-1.93.2-cp314-cp314-manylinux_2_28_x86_64.whl", hash = "sha256:cacf35cf703b12c54516fc6464a3e08c6dbb1dcfb97239e1f629294fe36a1cba", size = 20160055, upload-time = "2026-08-09T02:17:39.674Z" }, - { url = "https://files.pythonhosted.org/packages/5b/b5/c25d7fbe08490d8211bd6b69af23f3a922b68ad8c87c776480b0de64a505/litellm-1.93.2-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:0a7f3e5138e307e429bd8fa29cc0c48bb1e2b827792e8f7799ca4c8cff736103", size = 20230910, upload-time = "2026-08-09T02:17:42.159Z" }, - { url = "https://files.pythonhosted.org/packages/21/27/341b18a40d4d98a2ac09025c248a3a7edddaf15ce4096ac4a783ff2f70db/litellm-1.93.2-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:d8684629be3f7b5f8e2b6e5fe5ea27ff957c63a8d525d81c1460d8436e2e1857", size = 20288903, upload-time = "2026-08-09T02:17:44.433Z" }, - { url = "https://files.pythonhosted.org/packages/8d/45/dd9ef72075a83854f852b1bf9a97ec7029a2be9fb4e338fc6623eb09fc90/litellm-1.93.2-cp314-cp314-win_amd64.whl", hash = "sha256:a783b8b18ed68cb6a3b79d2b00273ec21aef92442e9b2712a50036cb84bfe583", size = 19772974, upload-time = "2026-08-09T02:17:46.972Z" }, + { url = "https://files.pythonhosted.org/packages/44/a7/4bccec0ac9cb1b2e94e391b666458d07480d342039c66383ac191819e8e7/litellm-1.101.2-cp310-abi3-macosx_10_12_x86_64.whl", hash = "sha256:48c42c2c2cf9d4b0d75f4e1670b1b64b9e0513488d0b737fe057fda0bc716551", size = 23827328, upload-time = "2026-09-24T00:04:03.196Z" }, + { url = "https://files.pythonhosted.org/packages/3d/3d/faf394e5ac5a1469de5cbd3939c9e330f744e648351934b981078ccc40d9/litellm-1.101.2-cp310-abi3-macosx_11_0_arm64.whl", hash = "sha256:77195c8ed502c052bb31d4c3887356a308ac2e8c3b0b30c97e2b04b0be10dd44", size = 23484770, upload-time = "2026-09-24T00:04:06.248Z" }, + { url = "https://files.pythonhosted.org/packages/ba/84/60f70aa2683666626c4abe7aa44b52acca52cad911ee870ea244fd3b0796/litellm-1.101.2-cp310-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:abb7b3ac04f56ced46e53cca2369a5dd29539cab9fdcd6cc9a94987aa55a38d1", size = 23618000, upload-time = "2026-09-24T00:04:08.827Z" }, + { url = "https://files.pythonhosted.org/packages/e2/8e/c57a4e157f97b1bcef9b410d51e17507047bbb11c676f5c81b22e7190c7c/litellm-1.101.2-cp310-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:210c89194225778759aa6649f5c0d605572bf14708ec85712462ace479f47f04", size = 23994795, upload-time = "2026-09-24T00:04:12.292Z" }, + { url = "https://files.pythonhosted.org/packages/b3/49/8737aee5a5a15cac7eb8a972b800529923e2837bbbadf0617f56111344ab/litellm-1.101.2-cp310-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:144dd8d1ead7174a718d1748deffcda7438cf7dfe8dc7c20761b72c35117e9a3", size = 23693332, upload-time = "2026-09-24T00:04:14.886Z" }, + { url = "https://files.pythonhosted.org/packages/04/50/4e711caa0374309d6aaf5696549449c2078f0225dd36a22ee0ca44dd068f/litellm-1.101.2-cp310-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:ae95e7ef15e109472f2cec69da6028a874b416e56a704cd9b797362993fc63cf", size = 24092655, upload-time = "2026-09-24T00:04:17.677Z" }, + { url = "https://files.pythonhosted.org/packages/c1/7d/32d391ddcb30d4d5d08fddd0abe918e836f9b3f753237c2b12ecb3d7425a/litellm-1.101.2-cp310-abi3-win_amd64.whl", hash = "sha256:0f5ee6daf9082b7efca1dc851c10c0d4884a2f1e0d7ea410bd508961b5a2cdae", size = 23894930, upload-time = "2026-09-24T00:04:20.432Z" }, ] [[package]] @@ -1895,6 +1924,69 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/1e/c1/d6e64ccd0536bf616556f0cad2b6d94a8125f508d25cfd814b1d2db4e2f1/openai-2.32.0-py3-none-any.whl", hash = "sha256:4dcc9badeb4bf54ad0d187453742f290226d30150890b7890711bda4f32f192f", size = 1162570, upload-time = "2026-04-15T22:28:17.714Z" }, ] +[[package]] +name = "orjson" +version = "3.12.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/0f/f3/742fb1f62b825f2c010697eaf4e828004bc2a81e7e806666989c132c7c42/orjson-3.12.0.tar.gz", hash = "sha256:d14203fb1aae2ad9b3d52f8a0e82aeb10197ef1c9bc61da7f358bd70b00123d5", size = 4142915, upload-time = "2026-08-14T16:13:30.607Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/75/1a/a7075a8e8b0d3f5097d17ac3099017104b6b7b42012041147995d5b2da05/orjson-3.12.0-cp311-cp311-macosx_10_15_x86_64.macosx_11_0_arm64.macosx_10_15_universal2.whl", hash = "sha256:a94f0f0c6fcbb2b5bd9734c57a489c7584a732bbdf04a39e8c83b861e9d03e92", size = 223409, upload-time = "2026-08-14T16:12:12.654Z" }, + { url = "https://files.pythonhosted.org/packages/05/34/c2eb3b2900e5597db7841a4c6416ac2d90081bd956b02d4dd1833fa2b96b/orjson-3.12.0-cp311-cp311-macosx_15_0_arm64.whl", hash = "sha256:a696529ec96a90d9a5f9570207efe403c8b08f8e4aa2783ee3403511e2fdfa10", size = 124015, upload-time = "2026-08-14T16:12:14.025Z" }, + { url = "https://files.pythonhosted.org/packages/1c/df/b49081766a75b6a37b3d33bdc0a39e492abab8441dd25e3e1998e7b83fcb/orjson-3.12.0-cp311-cp311-manylinux2014_armv7l.manylinux_2_17_armv7l.whl", hash = "sha256:e4ac5059baab4b3acbd99485de019ff8cda0fdf34b61fa74f7197a53db78bfe8", size = 113471, upload-time = "2026-08-14T16:12:15.81Z" }, + { url = "https://files.pythonhosted.org/packages/48/d4/58ea28eeef95c2a27358ed927380a621162cf20bd740bbccf9c3f09a200a/orjson-3.12.0-cp311-cp311-manylinux2014_i686.manylinux_2_17_i686.whl", hash = "sha256:8e29957429c35bbb5a185a119c523aa2428b7bbf1a293724c7b9375ed8f892a3", size = 129998, upload-time = "2026-08-14T16:12:17.503Z" }, + { url = "https://files.pythonhosted.org/packages/e2/f4/1e82aa2efc9916422d804697876ce433c907a1abd7c7e5c6d3d48565e5f9/orjson-3.12.0-cp311-cp311-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:dce0166feb0a737ab84f598c9a338cbc0b764a036617aa686194f53c7eba0c3e", size = 130891, upload-time = "2026-08-14T16:12:18.762Z" }, + { url = "https://files.pythonhosted.org/packages/5b/e1/15169e9d22b59a406264f99d6db387c0b0b12b6357a8a0169917c2a713eb/orjson-3.12.0-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:9caf3d09f47c3c70c4451ada20ef9bc4a4cdffa26f49862cf0a253b329aae2d5", size = 131285, upload-time = "2026-08-14T16:12:20.251Z" }, + { url = "https://files.pythonhosted.org/packages/a4/3a/763dbd426290d044ec3e615a05e70adb6d8b6f95bf17dc355c0081a5e8b6/orjson-3.12.0-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:b9dca132b1fda5565088e65a6b6e742285e0aeceb6fae549fa8863e16c7d3998", size = 135707, upload-time = "2026-08-14T16:12:21.652Z" }, + { url = "https://files.pythonhosted.org/packages/04/d1/3b2038ed168d22e14182ed715d6963f9c073a83a2ba43cfe918a4fc43c64/orjson-3.12.0-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:a791f793b287bbc135b8e87c34e35c8bfc693e2a8a620fab1ae682b925f9a32e", size = 127669, upload-time = "2026-08-14T16:12:22.926Z" }, + { url = "https://files.pythonhosted.org/packages/88/ae/b84b3d3e65f5629ada0edcb1d2bccc55d7c5f89d8b981537ecdc3d6f31ec/orjson-3.12.0-cp311-cp311-win32.whl", hash = "sha256:31ed278a36304390adc3eec5d7f6fd593a7c3e99e5a06cd07866396c4b1b4710", size = 128043, upload-time = "2026-08-14T16:12:24.367Z" }, + { url = "https://files.pythonhosted.org/packages/35/24/2ed0e6f51ea3d0af45d807233a851175af75bec83ef5fd0d6a2601904ec0/orjson-3.12.0-cp311-cp311-win_amd64.whl", hash = "sha256:fb2539159dfe8d371914f354360fa50e4a577cc89222a3828b9650a5e5040252", size = 122084, upload-time = "2026-08-14T16:12:25.813Z" }, + { url = "https://files.pythonhosted.org/packages/21/dd/95d25fcfbc9471799ef6bb01c552d64ee5cde93ee40ba2f423dd3442c708/orjson-3.12.0-cp311-cp311-win_arm64.whl", hash = "sha256:61318b6de893c7a9d9f3e5ecbadccbfc26a7eb417ccc7bbf0771de3b4d72f868", size = 127035, upload-time = "2026-08-14T16:12:27.201Z" }, + { url = "https://files.pythonhosted.org/packages/be/4a/295da39c651c2faac8bd351a2a346f0fdedd9d50b847ee9dfc27d2207ef6/orjson-3.12.0-cp312-cp312-macosx_10_15_x86_64.macosx_11_0_arm64.macosx_10_15_universal2.whl", hash = "sha256:aa3e43a6846e91d7bde3d5a9c66090fcd8744f569a9b6cffc5e1ca38f6a461c0", size = 223427, upload-time = "2026-08-14T16:12:28.525Z" }, + { url = "https://files.pythonhosted.org/packages/29/98/758cf90fbeaaafb7f8141bfac75a432099959f3a2f5db93a412e876415d8/orjson-3.12.0-cp312-cp312-macosx_15_0_arm64.whl", hash = "sha256:11edb4660a6680abee9788a3a9072208a2c96538cc1322bd79542065229d8e54", size = 123725, upload-time = "2026-08-14T16:12:30.013Z" }, + { url = "https://files.pythonhosted.org/packages/32/b5/5b934d251f8651f7e41df180ad0c57a6e1cabe15c7bd331638413a50ebc9/orjson-3.12.0-cp312-cp312-manylinux2014_armv7l.manylinux_2_17_armv7l.whl", hash = "sha256:2d3a9da945a4d96ae758fdaaca56742e6b73b6fd554c5d8876f252a6dad70b83", size = 113375, upload-time = "2026-08-14T16:12:31.209Z" }, + { url = "https://files.pythonhosted.org/packages/cd/d2/37efb5b12a176ce3ced29f4144f20da57d02757f78ce549637dc1b4e1fc8/orjson-3.12.0-cp312-cp312-manylinux2014_i686.manylinux_2_17_i686.whl", hash = "sha256:92ffc09e07233a6ab6d4e067f7841edcbcc134cb4812155cf171ea5255a421d7", size = 129983, upload-time = "2026-08-14T16:12:32.721Z" }, + { url = "https://files.pythonhosted.org/packages/50/22/0644b87c73f13e0092df8f35a1fe280d991e5e90072087411e0dd7e44e0c/orjson-3.12.0-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:bf44e374aadde77b1f6109f1030be51433eb61984379852766b6f4e187db7b1e", size = 130629, upload-time = "2026-08-14T16:12:34.084Z" }, + { url = "https://files.pythonhosted.org/packages/8c/57/80b986ebfecd9c6a177ddf1c2319717f0cd8feffb2b78946595a18a2fc88/orjson-3.12.0-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:1192a7021b6d071aaf909864f6e924d6a2675ca360485b972b8401749311750b", size = 131245, upload-time = "2026-08-14T16:12:35.713Z" }, + { url = "https://files.pythonhosted.org/packages/80/3d/75c5ac5a69161f44492a68fbdde66f4cc4ce48cd5e1fb05918e46f0c8848/orjson-3.12.0-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:53c0c474a9d9aff9aebfc0c88de1f28f843d940e6e3a80729abdf6a20274356f", size = 135397, upload-time = "2026-08-14T16:12:37.128Z" }, + { url = "https://files.pythonhosted.org/packages/71/93/4d71f2df314a97ff0d27a4559bf5888fc8406e3c6dec90e92291e3511215/orjson-3.12.0-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:532ff8cd4bd59a327a953a7dcde922c7fc25b85e29721bb8633265430d3a3873", size = 127693, upload-time = "2026-08-14T16:12:38.627Z" }, + { url = "https://files.pythonhosted.org/packages/bc/1d/0dbc6be5adfd1730491072fb60beb6bcdf5d7b2596ee41b7fc2e298bfc09/orjson-3.12.0-cp312-cp312-win32.whl", hash = "sha256:a6cf4b18e7de173f209f2084ffbd736dd72389a396326ee80a7022168be232e5", size = 128000, upload-time = "2026-08-14T16:12:39.954Z" }, + { url = "https://files.pythonhosted.org/packages/2d/c9/97b1ce0112ebf5e949c775ed5b1755e562233179f3584579673cc24d6378/orjson-3.12.0-cp312-cp312-win_amd64.whl", hash = "sha256:010811c1b69773450a01cef97727a67b223242f350b77d4ca000e59a9ef2155a", size = 122106, upload-time = "2026-08-14T16:12:41.324Z" }, + { url = "https://files.pythonhosted.org/packages/a8/6a/facd8b312e4a0d3a7fa978c7e15821f74a336adf1d65529faec33b48e18b/orjson-3.12.0-cp312-cp312-win_arm64.whl", hash = "sha256:ad29eece0c601737f2a60edc2752a84e7a0785df3efb62e3012834700a5afe0d", size = 126869, upload-time = "2026-08-14T16:12:42.651Z" }, + { url = "https://files.pythonhosted.org/packages/54/cb/d7b78218a987eb8a8ce4eeae0286b1bb679333eb631ea0eeaf6371680bfc/orjson-3.12.0-cp313-cp313-macosx_10_15_x86_64.macosx_11_0_arm64.macosx_10_15_universal2.whl", hash = "sha256:9a36ec60f1796f9a3f13e3b98390295e17a1c7c10155b448d264098bf9ee5900", size = 223397, upload-time = "2026-08-14T16:12:44.003Z" }, + { url = "https://files.pythonhosted.org/packages/f8/4a/bc87c45e7ec639d35ebefd62618e01939531ac8e171426606a01bda05914/orjson-3.12.0-cp313-cp313-macosx_15_0_arm64.whl", hash = "sha256:ad0422b92d5195443a39f80c3bcf731cc2e00f153bd32063a47b73b057bd0f03", size = 123662, upload-time = "2026-08-14T16:12:45.433Z" }, + { url = "https://files.pythonhosted.org/packages/94/ee/c9a4ff3f2dbedbbe9e635d0fa72c8866adede09b6335ef9644f53752f0d8/orjson-3.12.0-cp313-cp313-manylinux2014_armv7l.manylinux_2_17_armv7l.whl", hash = "sha256:5a0fdbc216388f653d3752ff310e710f59253bd4ed6a2bfb3f4f06b84714bbd8", size = 113374, upload-time = "2026-08-14T16:12:46.755Z" }, + { url = "https://files.pythonhosted.org/packages/75/09/3f330a026a796c8b4c97a6f429652a5e912e7065039bf96ed25e42aa7b25/orjson-3.12.0-cp313-cp313-manylinux2014_i686.manylinux_2_17_i686.whl", hash = "sha256:2eb5c56e534127b2b8fa38d2363c8b1b8190367ee0d1d16c041517d880843b94", size = 130029, upload-time = "2026-08-14T16:12:48.06Z" }, + { url = "https://files.pythonhosted.org/packages/7d/40/094cc53126a3d22f76cdf83b6ea67338bed01d774037621a785aa8e6e5ea/orjson-3.12.0-cp313-cp313-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:784106539f4b9d4b930e0b4eb8d45168507dae001945e71b4675a367f1e5e806", size = 130528, upload-time = "2026-08-14T16:12:49.362Z" }, + { url = "https://files.pythonhosted.org/packages/bc/74/89bb236deb9565f99434b13052bb40ddfcce4adf3afbfa3132ee7e421468/orjson-3.12.0-cp313-cp313-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:1c680706fc8396d95e7c4c1f9482563f552137aef91b57237a3ad5aaf64629df", size = 131075, upload-time = "2026-08-14T16:12:50.692Z" }, + { url = "https://files.pythonhosted.org/packages/0c/ac/1176360d762c01b5bd34acd56fc098e936c491363d8b6b397ad4aa475547/orjson-3.12.0-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:83445adc40cba26d6d621185a45128ce455b766af368cad2ab64b970603a7978", size = 135321, upload-time = "2026-08-14T16:12:52.114Z" }, + { url = "https://files.pythonhosted.org/packages/7a/02/bbd881c8b9276d50b998de38b4e97de8ace1aac940b0ee545aedbf65ed00/orjson-3.12.0-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:644d005bc82f917337a95ce270c9f6f92f9834c2bed7b1477572f8db00784222", size = 127472, upload-time = "2026-08-14T16:12:53.517Z" }, + { url = "https://files.pythonhosted.org/packages/8e/02/a0934d7503e6dcbedd6afac3e7f3f8597fd09389949ad94d0f7540e9dbca/orjson-3.12.0-cp313-cp313-win32.whl", hash = "sha256:d8e78d3d93705e3d27cc17cdb209e44d7a8ea203010cac6ce9c7ffc1ae1996f1", size = 128000, upload-time = "2026-08-14T16:12:55.14Z" }, + { url = "https://files.pythonhosted.org/packages/52/87/69f98f8d40faff103a965a5fbb83f08241b01beaf92badb5413fbc9358cc/orjson-3.12.0-cp313-cp313-win_amd64.whl", hash = "sha256:b85931be5b6763c31283805c9bdaae1ca03ad9f6f12a15f1cbf6745b907932c2", size = 121841, upload-time = "2026-08-14T16:12:56.507Z" }, + { url = "https://files.pythonhosted.org/packages/e6/07/b83046a4e3cadcc0987d0f160696107c4af706a619b56e4ad01940cadadf/orjson-3.12.0-cp313-cp313-win_arm64.whl", hash = "sha256:6a31348d7dfa64cd9c78bd1f510ff44c48fe64d71094e6b90e364dba3b55949e", size = 126765, upload-time = "2026-08-14T16:12:57.806Z" }, + { url = "https://files.pythonhosted.org/packages/12/9d/3931253e6f3148abf2cbe14830367042a4806b362ea520df2303db188fb9/orjson-3.12.0-cp314-cp314-macosx_10_15_x86_64.macosx_11_0_arm64.macosx_10_15_universal2.whl", hash = "sha256:9e6fee342a48760e854d743e7a81534d8e2925a6f46e09f750cf56b50fd1de5d", size = 223391, upload-time = "2026-08-14T16:12:59.184Z" }, + { url = "https://files.pythonhosted.org/packages/8a/0e/b4a4f1e305367245877b967a0bad70fcf001d77c54ac4339a120b66fdae4/orjson-3.12.0-cp314-cp314-macosx_15_0_arm64.whl", hash = "sha256:8c3bb86dd10f39b3fbf434b7d5dc7cac77d6fc8ac572ae30a10731ede2c4b647", size = 123659, upload-time = "2026-08-14T16:13:00.548Z" }, + { url = "https://files.pythonhosted.org/packages/96/f3/6782c6fa85e2702bc66be183c3b421486167dcf266ee4dc1403fe3824870/orjson-3.12.0-cp314-cp314-manylinux2014_armv7l.manylinux_2_17_armv7l.whl", hash = "sha256:2bb3ce43203936072dd8b4917b01d3aecfc02329bfb42510cb7cfb24708adc9c", size = 113337, upload-time = "2026-08-14T16:13:02.009Z" }, + { url = "https://files.pythonhosted.org/packages/bf/79/b32ab64bacda9d0fa4942ef483bd03cabf0eaf2be819ca9fb7ff610c559d/orjson-3.12.0-cp314-cp314-manylinux2014_i686.manylinux_2_17_i686.whl", hash = "sha256:6a2a79c89984dc719817d388c8709e0efc2a2795a934eaa746b4882eb6045adc", size = 130112, upload-time = "2026-08-14T16:13:03.404Z" }, + { url = "https://files.pythonhosted.org/packages/ee/49/6e6142999ca01509219be5e5a9c338a3e5ea011f63e91ff473fbbf3734ed/orjson-3.12.0-cp314-cp314-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:f06dd838d1e07d9b1de0932ec0485ec92c4d5f5d1ad4817a656268c3e88be1e1", size = 130520, upload-time = "2026-08-14T16:13:04.798Z" }, + { url = "https://files.pythonhosted.org/packages/49/d0/3745af0a4cc9867784f29722929cec4d10bd1c877cd754b01ba6d96eb21a/orjson-3.12.0-cp314-cp314-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:c6b11be792c3d2c6a4be2af4ebf97a68d0bf5f580aca6e86a418a354f6cc846a", size = 131053, upload-time = "2026-08-14T16:13:06.14Z" }, + { url = "https://files.pythonhosted.org/packages/c3/f4/6fe5a22fa478fffb190e65c338c84df5c311ef597b363150a17cc57063c0/orjson-3.12.0-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:477ecaf6b9f88f873341b91fcc736119ca81b5e002a9f7f308ff5b4f2ce2a70e", size = 135321, upload-time = "2026-08-14T16:13:07.544Z" }, + { url = "https://files.pythonhosted.org/packages/ff/41/b1b0ec30289646a81a76e2dbaae2686b96fcccb7cb0323dc1dd78cbc7875/orjson-3.12.0-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:f3c0683136acdc29afdf88a5bc2f7d3d0e34087788d1d63c0144b805a87a196f", size = 127485, upload-time = "2026-08-14T16:13:08.88Z" }, + { url = "https://files.pythonhosted.org/packages/bf/2b/277404bdcc21c93b112b963655b76443ebfe828f8a3ff1de7d90f8850eb3/orjson-3.12.0-cp314-cp314-win32.whl", hash = "sha256:d39f3f5c3927e2dc0913fe5bbc1a2f6b1b9d1bba1de6358340d0ad0d0c00ca92", size = 128048, upload-time = "2026-08-14T16:13:10.305Z" }, + { url = "https://files.pythonhosted.org/packages/41/2b/395b36fa2b4ce7af70b651d715e88f80d884b2c2b14a6b53e84d554fb5f0/orjson-3.12.0-cp314-cp314-win_amd64.whl", hash = "sha256:0b1ac5bf6609b2716c7954011c5fef6254922df029f45d032ee4ebf5d363cbed", size = 121858, upload-time = "2026-08-14T16:13:11.634Z" }, + { url = "https://files.pythonhosted.org/packages/ea/a3/833e895ff452859eebe75093d26691fe9108f1a7a6a08435d7a5780ea652/orjson-3.12.0-cp314-cp314-win_arm64.whl", hash = "sha256:50fae885cb073eac7556353ff3df93312b0d5137b0a5056b2bb63f97ed9a93c7", size = 126749, upload-time = "2026-08-14T16:13:13.117Z" }, + { url = "https://files.pythonhosted.org/packages/58/64/99c8947ece10c17176af9aae85c4948f1d109da77440ec14d87239efaf73/orjson-3.12.0-cp315-cp315-macosx_10_15_x86_64.macosx_11_0_arm64.macosx_10_15_universal2.whl", hash = "sha256:01efac2074fffb4cb1ea3fab7861e9d0f2a26913854a972f5ac760525dbdaf6e", size = 223398, upload-time = "2026-08-14T16:13:14.694Z" }, + { url = "https://files.pythonhosted.org/packages/3e/30/cf983fe09f2731420fda097a9f7ef4343f47fa216c228961ad8f6da44f3d/orjson-3.12.0-cp315-cp315-macosx_15_0_arm64.whl", hash = "sha256:ed4ca42bd55955aa34deedcfdfd0e0c31abf51143aae158ae2bc3520b626e517", size = 123655, upload-time = "2026-08-14T16:13:16.221Z" }, + { url = "https://files.pythonhosted.org/packages/11/50/9cb8ae73fa4749dbbc20f617004213b5ff01c20aaeec34c3f31124f2c1d8/orjson-3.12.0-cp315-cp315-manylinux_2_39_aarch64.whl", hash = "sha256:40f92192227505acca4e2533ce565f8e6b9535f7d0d09b0968452f18b7376b38", size = 130515, upload-time = "2026-08-14T16:13:17.601Z" }, + { url = "https://files.pythonhosted.org/packages/9f/0a/adb6ce1a5b5fbf9cb1790f9961bb668a0dd5429aadaf6cee044724681795/orjson-3.12.0-cp315-cp315-manylinux_2_39_armv7l.whl", hash = "sha256:33efefcf5d88eaf400b47e2eba02f91f319bb9951be61ca500b7d536d3f2079d", size = 113327, upload-time = "2026-08-14T16:13:18.927Z" }, + { url = "https://files.pythonhosted.org/packages/51/5c/d17f61581d8dbdde7048f87a330fa24915edec38db4d72b381fec14fbb56/orjson-3.12.0-cp315-cp315-manylinux_2_39_i686.whl", hash = "sha256:8e386b0bc0ddd7cd2056f884b5a0af33592bd01ac66a7ca4b42a65a7e7774a13", size = 130105, upload-time = "2026-08-14T16:13:20.317Z" }, + { url = "https://files.pythonhosted.org/packages/9f/b7/938befcf33bee4704a92ecec6a2731224c539d939bf9429fd39396d28931/orjson-3.12.0-cp315-cp315-manylinux_2_39_x86_64.whl", hash = "sha256:58c58e1de0006ffb580368d6793c36c7b0b021db066479cf281bf5061e732328", size = 131049, upload-time = "2026-08-14T16:13:21.719Z" }, + { url = "https://files.pythonhosted.org/packages/b0/15/cfa2021d64d5aa8bb5c9f604ef375e00ec8b657651b5dd650b1b7ad13df1/orjson-3.12.0-cp315-cp315-musllinux_1_2_aarch64.whl", hash = "sha256:08231552159be266a7269555bd9f7c016aee7d9ad6dab06eb58796c5ccb7101c", size = 135320, upload-time = "2026-08-14T16:13:23.415Z" }, + { url = "https://files.pythonhosted.org/packages/1a/50/3e75dfe357c1e8f9e287c7a5740260ef15bd23a5299eae8d0835dcad5375/orjson-3.12.0-cp315-cp315-musllinux_1_2_x86_64.whl", hash = "sha256:a15f9a891bce5f5cc5d210e3ad8614d4d1b489a56448c099d6d2a7168b2d954a", size = 127488, upload-time = "2026-08-14T16:13:24.791Z" }, + { url = "https://files.pythonhosted.org/packages/11/a6/79aed402eb3ab284dc5b4791a7ad62c5875127de01b8e3f04bd92d551298/orjson-3.12.0-cp315-cp315-win32.whl", hash = "sha256:03091c8a64db4be38746597ceea68f33c238e27acd9bfe99fb59420224ae7a55", size = 128048, upload-time = "2026-08-14T16:13:26.217Z" }, + { url = "https://files.pythonhosted.org/packages/64/f7/2723e264aab7248c1ed6ecaad8e5d0cb866c0cffde75442102ffa7491aba/orjson-3.12.0-cp315-cp315-win_amd64.whl", hash = "sha256:2b7bcefb9f40fa242fa6b06377232c048e655747790829609168c01162f60578", size = 121860, upload-time = "2026-08-14T16:13:27.577Z" }, + { url = "https://files.pythonhosted.org/packages/82/56/630c9113ec8996778f1f0304b364b091b9a9db5fef5fdc17cca622f5ea24/orjson-3.12.0-cp315-cp315-win_arm64.whl", hash = "sha256:859fc4196855890150bb08e649b30d2c93b249b3e3edd0d3bb2231abf8aa8adc", size = 126754, upload-time = "2026-08-14T16:13:28.962Z" }, +] + [[package]] name = "packaging" version = "25.0" @@ -2313,11 +2405,11 @@ wheels = [ [[package]] name = "pyjwt" -version = "2.13.0" +version = "2.15.1" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/3b/81/58d0ac84e1ef3a3843791d6954d94c0b33d526c75eeb1efbce9d0a4c4077/pyjwt-2.13.0.tar.gz", hash = "sha256:41571c89ca91598c79e8ef18a2d07367d4810fbbd6f637794879baf1b7703423", size = 107515, upload-time = "2026-05-21T19:54:36.618Z" } +sdist = { url = "https://files.pythonhosted.org/packages/43/ea/5194e52748b0da83d71e082d75496eaec6e58f419f5e184786ded517e6a9/pyjwt-2.15.1.tar.gz", hash = "sha256:4f259e80cdfb6b3fc18a7de51fd1ef9ec79652f25019bae68975ca2468a34df8", size = 121252, upload-time = "2026-09-28T18:40:42.598Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/a3/5e/ecf12fdb62546d64385c158514e9b2b671f7832108ef2ecd2020ce0af2d1/pyjwt-2.13.0-py3-none-any.whl", hash = "sha256:66adcc2aff09b3f1bbd95fc1e1577df8ac8723c978552fd43304c8a290ac5728", size = 31274, upload-time = "2026-05-21T19:54:35.362Z" }, + { url = "https://files.pythonhosted.org/packages/50/ca/44de4e75f8aadc457f0634be3b542815078ded46dca30efb960edeecad6e/pyjwt-2.15.1-py3-none-any.whl", hash = "sha256:42d59d631f7768a1028a64c7ff581a9bf7519804daf91fc5b6c56e30eec5e193", size = 33860, upload-time = "2026-09-28T18:40:41.429Z" }, ] [[package]] @@ -2376,6 +2468,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/bc/16/4ea354101abb1287856baa4af2732be351c7bee728065aed451b678153fd/pytest_cov-6.2.1-py3-none-any.whl", hash = "sha256:f5bc4c23f42f1cdd23c70b1dab1bbaef4fc505ba950d53e0081d0730dd7e86d5", size = 24644, upload-time = "2025-06-12T10:47:45.932Z" }, ] +[[package]] +name = "python-dateutil" +version = "2.9.0.post0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "six" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/66/c0/0c8b6ad9f17a802ee498c46e004a0eb49bc148f2fd230864601a86dcf6db/python-dateutil-2.9.0.post0.tar.gz", hash = "sha256:37dd54208da7e1cd875388217d5e00ebd4179249f90fb72437e91a35459a0ad3", size = 342432, upload-time = "2024-03-01T18:36:20.211Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/ec/57/56b9bcc3c9c6a792fcbaf139543cee77261f3651ca9da0c93f5c1221264b/python_dateutil-2.9.0.post0-py2.py3-none-any.whl", hash = "sha256:a8b2bc7bffae282281c8140a97d3aa9c14da0b136dfe83f850eea9a5f7470427", size = 229892, upload-time = "2024-03-01T18:36:18.57Z" }, +] + [[package]] name = "python-dotenv" version = "1.2.2" @@ -2630,6 +2734,7 @@ source = { editable = "." } dependencies = [ { name = "aiosqlite" }, { name = "alembic" }, + { name = "backoff" }, { name = "cashu" }, { name = "fastapi", extra = ["standard-no-fastapi-cloud-cli"] }, { name = "greenlet" }, @@ -2640,6 +2745,7 @@ dependencies = [ { name = "mdurl" }, { name = "nostr-sdk" }, { name = "openai" }, + { name = "orjson" }, { name = "pillow" }, { name = "python-json-logger" }, { name = "sqlmodel" }, @@ -2666,16 +2772,18 @@ dev = [ requires-dist = [ { name = "aiosqlite", specifier = ">=0.20" }, { name = "alembic", specifier = ">=1.13" }, + { name = "backoff", specifier = ">=2.2" }, { name = "cashu", specifier = ">=0.20" }, { name = "fastapi", extras = ["standard-no-fastapi-cloud-cli"], specifier = ">=0.141" }, { name = "greenlet", specifier = ">=3.2.1" }, { name = "h11", specifier = ">=0.16" }, { name = "httpx", extras = ["socks"], specifier = ">=0.28.1" }, - { name = "litellm", specifier = ">=1.93.0,<1.94" }, + { name = "litellm", specifier = ">=1.101.2,<1.102" }, { name = "marshmallow", specifier = ">=3.13,<4.0" }, { name = "mdurl", specifier = "==0.1.2" }, { name = "nostr-sdk", specifier = ">=0.45.1,<0.46" }, { name = "openai", specifier = ">=1.98.0" }, + { name = "orjson", specifier = ">=3.10" }, { name = "pillow", specifier = ">=10" }, { name = "python-json-logger", specifier = ">=2.0.0" }, { name = "sqlmodel", specifier = ">=0.0.42" }, @@ -2831,6 +2939,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/4c/9b/0b8aa09817b63e78d94b4977f18b1fcaead3165a5ee49251c5d5c245bb2d/ruff-0.12.7-py3-none-win_arm64.whl", hash = "sha256:dfce05101dbd11833a0776716d5d1578641b7fddb537fe7fa956ab85d1769b69", size = 11982083, upload-time = "2025-07-29T22:32:33.881Z" }, ] +[[package]] +name = "s3transfer" +version = "0.19.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "botocore" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/76/43/35e4d8aa320bffe8287fe8f65f578fa2d2db0a64212f0e710dce58267854/s3transfer-0.19.2.tar.gz", hash = "sha256:ba0309fd86be3c27dbf78cdd813c13c5e1df16e5874b99d2535ebbdfb9892993", size = 165592, upload-time = "2026-07-22T19:30:44.432Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/bc/e7/5c595c75e9f41a44f30e526eda465ea0b4eec93470e074e4a111b253f13a/s3transfer-0.19.2-py3-none-any.whl", hash = "sha256:d8168eccca828cbb2cd573675333f3bddd254313a9c42494b84c76b539e8ba25", size = 90216, upload-time = "2026-07-22T19:30:43.251Z" }, +] + [[package]] name = "setuptools" version = "84.0.0" @@ -3136,11 +3256,11 @@ wheels = [ [[package]] name = "urllib3" -version = "2.7.0" +version = "2.8.0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/53/0c/06f8b233b8fd13b9e5ee11424ef85419ba0d8ba0b3138bf360be2ff56953/urllib3-2.7.0.tar.gz", hash = "sha256:231e0ec3b63ceb14667c67be60f2f2c40a518cb38b03af60abc813da26505f4c", size = 433602, upload-time = "2026-05-07T16:13:18.596Z" } +sdist = { url = "https://files.pythonhosted.org/packages/e3/05/b17359e1cefb4f909b5e40b1b90a496d987258916dbbf88e842c729f510e/urllib3-2.8.0.tar.gz", hash = "sha256:63bf2ead4c879426ebf22ef2a781eeb4aa3b4ae798a0435506f8687fd5bb9b63", size = 458972, upload-time = "2026-09-15T19:29:36.253Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/7f/3e/5db95bcf282c52709639744ca2a8b149baccf648e39c8cc87553df9eae0c/urllib3-2.7.0-py3-none-any.whl", hash = "sha256:9fb4c81ebbb1ce9531cce37674bbc6f1360472bc18ca9a553ede278ef7276897", size = 131087, upload-time = "2026-05-07T16:13:17.151Z" }, + { url = "https://files.pythonhosted.org/packages/92/9d/c4e665119135114480843e7ab388fa94d8480650450e6f8e26b70d323a4c/urllib3-2.8.0-py3-none-any.whl", hash = "sha256:0cf3cae568d36aa9576b28dfb35f11328f1cb974ca7647d9475ebb86c75ac6e3", size = 135717, upload-time = "2026-09-15T19:29:34.577Z" }, ] [[package]]